diff --git a/config/config.default.yaml b/config/config.default.yaml index ca850515..22e91d63 100644 --- a/config/config.default.yaml +++ b/config/config.default.yaml @@ -5,7 +5,6 @@ # Edit the appropriate config file instead. debug: false # Run server with increased compatibility for breakpoint debugging. instance_env: dev # Instance environment. Used in Sentry, userAgent of subqueries, instance-appropriate behavior, etc. -instance_idx: 0 # Instance index. Use when multiple Retriever instances are run, so a leader can be determined. max_request_size: 2147483648 # Maximum request body size in bytes, post-decompression. Larger bodies are rejected with 413. log_level: DEBUG # Level of application logs to print/keep. host: 0.0.0.0 # Uvicorn listen host. @@ -55,6 +54,7 @@ redis: shutdown_timeout: 3 # Time in seconds to wait for batched tasks to finish before force-quitting. heartbeat_interval_seconds: 60 # Cadence at which workers / background / main re-register their process entry. process_ttl_seconds: 300 # TTL on each registered process entry; set to several heartbeat intervals so a single missed refresh doesn't drop the entry. + leader_lease_ttl_seconds: 180 # TTL on the build-leader lease, renewed each heartbeat interval; set to several intervals so a transient renewal failure doesn't trigger a leadership handoff. mongo: host: localhost port: 27017 diff --git a/src/retriever/background.py b/src/retriever/background.py index c6a6480e..78bbf148 100644 --- a/src/retriever/background.py +++ b/src/retriever/background.py @@ -16,6 +16,7 @@ from retriever.lookup.subclass import SubclassMapping from retriever.metadata.optable import OpTableManager from retriever.utils.general import tolerate_init +from retriever.utils.leader import LEADER_ELECTION from retriever.utils.logs import add_mongo_sink from retriever.utils.mongo import MongoClient, MongoQueue from retriever.utils.orphan_detection import periodically_mark_orphans @@ -50,15 +51,18 @@ async def _background_async() -> None: role_label="Background", ) - # The leader doesn't see query traffic, so opt into periodic backend pings. + # The builder doesn't see query traffic, so opt into periodic backend pings. tier_manager.enable_periodic_healthchecks() await tier_manager.initialize_drivers() metakg_manager = OpTableManager() - metakg_manager.promote_to_leader() + metakg_manager.promote_to_builder() await tolerate_init("OpTable build", metakg_manager.initialize()) subclass_manager = SubclassMapping() - subclass_manager.promote_to_leader() + subclass_manager.promote_to_builder() await tolerate_init("Subclass map build", subclass_manager.initialize()) + # Managers have registered their on_acquire hooks; contend for the build lease + # so exactly one instance drives the builds across the shared Redis. + await tolerate_init("Leader election", LEADER_ELECTION.start()) orphan_task = asyncio.create_task( periodically_mark_orphans(), name="orphan-detection" ) @@ -79,6 +83,9 @@ async def _background_async() -> None: with contextlib.suppress(asyncio.CancelledError): await orphan_task # Heartbeat task lives in RedisClient().tasks; cancelled in its wrapup. + # Relinquish the build lease first so a peer can take over promptly, and + # while Redis is still up to release it (rather than waiting on TTL). + await LEADER_ELECTION.stop() await SubclassMapping().wrapup() await metakg_manager.wrapup() await RedisClient().wrapup() diff --git a/src/retriever/config/general.py b/src/retriever/config/general.py index bbd93f19..4e41c689 100644 --- a/src/retriever/config/general.py +++ b/src/retriever/config/general.py @@ -74,6 +74,12 @@ class RedisSettings(BaseModel): description="TTL on each registered process entry; set to several heartbeat intervals so a single missed refresh doesn't drop the entry." ), ] = 300 + leader_lease_ttl_seconds: Annotated[ + int, + Field( + description="TTL on the build-leader lease, renewed each heartbeat interval. Set to several intervals to avoid touchy handoffs." + ), + ] = 180 class MongoSettings(BaseModel): @@ -346,12 +352,6 @@ class GeneralConfig(CommentedSettings): description="Instance environment. Used in Sentry, userAgent of subqueries, instance-appropriate behavior, etc." ), ] = "dev" - instance_idx: Annotated[ - int, - Field( - description="Instance index. Use when multiple Retriever instances are run, so a leader can be determined." - ), - ] = 0 max_request_size: Annotated[ int, Field( diff --git a/src/retriever/lookup/subclass.py b/src/retriever/lookup/subclass.py index fe59b73a..b70620e5 100644 --- a/src/retriever/lookup/subclass.py +++ b/src/retriever/lookup/subclass.py @@ -12,13 +12,14 @@ CURIE, ) from retriever.utils.general import BatchedAction +from retriever.utils.leader import LEADER_ELECTION from retriever.utils.redis import SUBCLASS_META_KEY, TIER_RECOVERED_CHANNEL, RedisClient REDIS_CLIENT = RedisClient() MAPPING_ID = "SubclassHashMap" MAPPING_BUILD_ID = f"{MAPPING_ID}:next" -"""Temp key the leader builds into, then atomically renames onto `MAPPING_ID`.""" +"""Temp key the builder builds into, then atomically renames onto `MAPPING_ID`.""" class SubclassMapping(BatchedAction): @@ -30,7 +31,7 @@ class SubclassMapping(BatchedAction): flush_time: float = 0 multibatch: bool = True - is_leader: bool = False + is_builder: bool = False subscriptions: dict[CURIE, list[Callable[[list[CURIE] | None], None]]] _refresh_lock: asyncio.Lock _pending_refresh: bool = False @@ -38,15 +39,15 @@ class SubclassMapping(BatchedAction): redis_setup_batch_size: int = 5000 def __init__(self) -> None: - """Initialize without leader role; call `promote_to_leader()` to flip the flag.""" + """Initialize without builder role; call `promote_to_builder()` to flip the flag.""" self.subscriptions = {} self._refresh_lock = asyncio.Lock() self._pending_refresh = False super().__init__() - def promote_to_leader(self) -> None: - """Flip this instance to leader mode. Must be called before `initialize()`.""" - self.is_leader = True + def promote_to_builder(self) -> None: + """Flip this process to builder mode. Must be called before `initialize()`.""" + self.is_builder = True @override async def initialize(self) -> None: @@ -55,12 +56,15 @@ async def initialize(self) -> None: logger.info("Implicit subclassing disabled, skipping initialization.") return await super().initialize() - if not self.is_leader: # Only need leader to update the redis setup + if not self.is_builder: # Only the builder updates the redis setup return await super().initialize() if self.initialized: return # rebuild loop already running + # Rebuild as soon as this instance wins the build lease. + LEADER_ELECTION.on_acquire(self.refresh) + try: await self.refresh() self.tasks.append(asyncio.create_task(self.rebuild())) @@ -72,7 +76,7 @@ async def initialize(self) -> None: REDIS_CLIENT.on_recover(self.refresh) tier_manager.get_driver(1).on_recover(self.refresh) # Also listen for worker-detected tier 1 recovery via Redis so - # the rebuild fires faster than the leader's own periodic ping. + # the rebuild fires faster than the builder's own periodic ping. with contextlib.suppress(Exception): await REDIS_CLIENT.subscribe( TIER_RECOVERED_CHANNEL, self._on_remote_tier_recover @@ -82,8 +86,8 @@ async def initialize(self) -> None: @override async def wrapup(self) -> None: - """Unsubscribe the leader's tier-recovery listener, then cancel the rebuild loop.""" - if self.is_leader: + """Unsubscribe the builder's tier-recovery listener, then cancel the rebuild loop.""" + if self.is_builder: with contextlib.suppress(Exception): await REDIS_CLIENT.unsubscribe( TIER_RECOVERED_CHANNEL, self._on_remote_tier_recover @@ -98,6 +102,10 @@ async def _on_remote_tier_recover(self, message: str) -> None: async def refresh(self) -> None: """Rebuild and publish the subclass mapping; concurrent calls collapse to one trailing rebuild.""" + # Only the leader writes the shared mapping, so instances don't race on + # the temp build key. + if not LEADER_ELECTION.is_leader: + return self._pending_refresh = True if self._refresh_lock.locked(): logger.debug( diff --git a/src/retriever/metadata/optable.py b/src/retriever/metadata/optable.py index 924c17ff..a5a5394e 100644 --- a/src/retriever/metadata/optable.py +++ b/src/retriever/metadata/optable.py @@ -37,6 +37,7 @@ from retriever.utils import biolink from retriever.utils.biolink import expand from retriever.utils.general import AsyncDaemon +from retriever.utils.leader import LEADER_ELECTION from retriever.utils.redis import ( OP_TABLE_KEY, OP_TABLE_META_KEY, @@ -97,30 +98,30 @@ class OpTableManager(AsyncDaemon): update_lock: asyncio.Lock _refresh_lock: asyncio.Lock _pending_refresh: bool = False - is_leader: bool = False + is_builder: bool = False def __init__(self) -> None: - """Initialize without leader role; call `promote_to_leader()` to flip the flag.""" + """Initialize without builder role; call `promote_to_builder()` to flip the flag.""" self.update_lock = asyncio.Lock() self._refresh_lock = asyncio.Lock() self._pending_refresh = False super().__init__() - def promote_to_leader(self) -> None: - """Flip this instance to leader mode. Must be called before `initialize()`.""" - self.is_leader = True + def promote_to_builder(self) -> None: + """Flip this process to builder mode. Must be called before `initialize()`.""" + self.is_builder = True @override def get_task_funcs(self) -> list[Callable[[], Coroutine[None, None, None]]]: tasks = list[Callable[[], Coroutine[None, None, None]]]() - if self.is_leader and CONFIG.job.metakg.build_time > -1: + if self.is_builder and CONFIG.job.metakg.build_time > -1: tasks.append(self.periodic_build_op_table) return tasks @override async def initialize(self) -> None: """Start the appropriate tasks for a given process.""" - if self.is_leader: + if self.is_builder: # Register hooks before the initial refresh so a startup # against a down dependency still recovers later. REDIS_CLIENT.on_recover(self.refresh) @@ -130,6 +131,8 @@ async def initialize(self) -> None: await REDIS_CLIENT.subscribe( TIER_RECOVERED_CHANNEL, self._on_remote_tier_recover ) + # Rebuild as soon as this instance wins the build lease. + LEADER_ELECTION.on_acquire(self.refresh) try: await self.refresh() except Exception: @@ -147,7 +150,7 @@ async def initialize(self) -> None: for tier_idx in range(0, 2): driver = tier_manager.get_driver(tier_idx) driver.on_recover(self._on_tier_recover) - # Tell the leader so it rebuilds without waiting on its periodic ping. + # Tell the builder so it rebuilds without waiting on its periodic ping. driver.on_recover(self._make_remote_publisher(tier_idx)) return await super().initialize() @@ -165,7 +168,7 @@ async def _publish() -> None: return _publish async def _on_remote_tier_recover(self, _message: str) -> None: - """Leader-side subscriber callback for cross-process tier recovery.""" + """Builder-side subscriber callback for cross-process tier recovery.""" await self.refresh() async def refresh(self) -> None: @@ -217,7 +220,7 @@ async def degraded_local_build(self) -> None: @override async def wrapup(self) -> None: """Cancel running tasks so connections can close.""" - if self.is_leader: + if self.is_builder: with contextlib.suppress(Exception): await REDIS_CLIENT.unsubscribe( TIER_RECOVERED_CHANNEL, self._on_remote_tier_recover @@ -374,7 +377,9 @@ async def _collect_tier_ops(self, *, bypass_cache: bool = False) -> OperationTab async def build_operation_table(self) -> None: """Build Retriever's internal OperationTable and store it to Redis.""" - if CONFIG.instance_idx != 0: + # Build+publish only when this builder's instance is the elected leader; + # workers (is_builder False) may still build on demand via get_op_table. + if self.is_builder and not LEADER_ELECTION.is_leader: return logger.info("Building Operation Table...") @@ -386,9 +391,9 @@ async def build_operation_table(self) -> None: logger.success( f"Built Operation Table containing {len(op_table.operations_flat)} operations / {len(op_table.nodes)} nodes." ) - # The leader never reads _operation_table back - it only exists to push to + # The builder never reads _operation_table back - it only exists to push to # Redis. Drop the reference so the snapshot doesn't sit in process memory. - if self.is_leader: + if self.is_builder: async with self.update_lock: self._operation_table = None @@ -428,7 +433,7 @@ async def get_op_table(self) -> OperationTable: op_table = self._operation_table if op_table is not None: return op_table - if not REDIS_CLIENT.up and not self.is_leader: + if not REDIS_CLIENT.up and not self.is_builder: # Worker can't pull the published copy; build from # available tiers and re-check. await self.degraded_local_build() diff --git a/src/retriever/status.py b/src/retriever/status.py index 119c0dc3..d59491bf 100644 --- a/src/retriever/status.py +++ b/src/retriever/status.py @@ -350,17 +350,26 @@ async def status_root() -> StatusSnapshot: redis_mem_r: object = None metakg_r: object = None subclass_r: object = None + leader_r: object = None + leader_elected_r: object = None if redis_client.up: + # gather()'s typed tuple overloads stop at 6 awaitables; past that it + # returns list[...], so route through object before the tuple cast. redis_results = cast( - tuple[object, object, object, object, object, object], - await asyncio.gather( - redis_client.list_main(), - redis_client.list_background(), - redis_client.list_workers(), - redis_client.used_memory_bytes(), - redis_client.metakg_freshness(), - redis_client.subclass_freshness(), - return_exceptions=True, + tuple[object, object, object, object, object, object, object, object], + cast( + object, + await asyncio.gather( + redis_client.list_main(), + redis_client.list_background(), + redis_client.list_workers(), + redis_client.used_memory_bytes(), + redis_client.metakg_freshness(), + redis_client.subclass_freshness(), + redis_client.get_leader(), + redis_client.get_leader_elected_at(), + return_exceptions=True, + ), ), ) ( @@ -370,6 +379,8 @@ async def status_root() -> StatusSnapshot: redis_mem_r, metakg_r, subclass_r, + leader_r, + leader_elected_r, ) = redis_results if any(isinstance(r, BaseException) for r in redis_results): redis_client.request_health_check() @@ -400,6 +411,8 @@ async def status_root() -> StatusSnapshot: stuck_job_count = _unwrap(stuck_r) metakg_record = cast(FreshnessRecord | None, _unwrap(metakg_r)) subclass_record = cast(FreshnessRecord | None, _unwrap(subclass_r)) + leader = cast(str | None, _unwrap(leader_r)) + leader_elected_at = cast(datetime | None, _unwrap(leader_elected_r)) # `registry_available` flips False when *any* of the three registry # reads failed (they all hit Redis; one failing means we can't trust @@ -498,6 +511,8 @@ async def status_root() -> StatusSnapshot: # True on their own snapshot - not currently wired. metakg=_metakg_row(metakg_record, self_reported=False, now=now), subclass_map=_subclass_map_row(subclass_record), + leader=leader, + leader_elected_at=leader_elected_at, ) diff --git a/src/retriever/types/status.py b/src/retriever/types/status.py index ddcd900e..15dd445d 100644 --- a/src/retriever/types/status.py +++ b/src/retriever/types/status.py @@ -309,3 +309,7 @@ class StatusSnapshot(TypedDict): tiers: list[StatusTier] metakg: StatusMetaKG subclass_map: StatusSubclassMap + leader: str | None + """Builder token of the current leader; None when the cluster is leaderless.""" + leader_elected_at: datetime | None + """Best-effort time the current leader won the lease; None if unknown/leaderless.""" diff --git a/src/retriever/utils/leader.py b/src/retriever/utils/leader.py new file mode 100644 index 00000000..ae00d484 --- /dev/null +++ b/src/retriever/utils/leader.py @@ -0,0 +1,161 @@ +import asyncio +import contextlib +from collections.abc import Awaitable, Callable +from datetime import datetime +from uuid import uuid4 + +from loguru import logger + +from retriever.config.general import CONFIG +from retriever.utils.general import Singleton +from retriever.utils.redis import LEADER_ELECTED_KEY, LEADER_LEASE_KEY, RedisClient + +REDIS_CLIENT = RedisClient() + +# Consecutive lease failures (while Redis reads still pass) before escalating from +# a debug line to a warning that the cluster may be leaderless. +LEASE_FAILURE_WARN_THRESHOLD = 3 + + +class LeaderElection(metaclass=Singleton): + """Elects one instance as the leader across instances via a single Redis lease. + + Each instance's *builder* (background process) attempts to acquire one Redis + lease. The instance whose builder holds the lease is the *leader*. It + relinquishes leadership the instant Redis goes down, so a stale build never + overwrites live state. Only the builder calls `start()`; workers import the + singleton but leave `is_leader` False. + """ + + is_leader: bool = False + """Whether this instance is the leader (its builder holds the lease).""" + + def __init__(self) -> None: + """Set up contention state; asyncio objects are deferred to `start()`.""" + self.token: str = uuid4().hex + self.is_leader = False + self._acquire_callbacks: list[Callable[[], Awaitable[None]]] = [] + # When this leadership episode began; published best-effort for /status. + self._elected_at: datetime | None = None + # Consecutive lease failures while Redis reads still succeed; drives escalation. + self._consecutive_failures: int = 0 + self._renew_task: asyncio.Task[None] | None = None + self._fired_tasks: set[asyncio.Task[None]] = set() + + def on_acquire(self, callback: Callable[[], Awaitable[None]]) -> None: + """Register a callback fired once each time this instance wins leadership.""" + if callback not in self._acquire_callbacks: + self._acquire_callbacks.append(callback) + + async def start(self) -> None: + """Begin contending for the lease; call once, from the builder process.""" + REDIS_CLIENT.on_outage(self._demote) + REDIS_CLIENT.on_recover(self._try_acquire) + self._renew_task = asyncio.create_task( + self._renew_loop(), name="leader-lease-renew" + ) + await self._attempt() + + async def stop(self) -> None: + """Stop contending, releasing the lease so a peer can take over promptly.""" + if self._renew_task is not None: + _ = self._renew_task.cancel() + with contextlib.suppress(asyncio.CancelledError): + await self._renew_task + self._renew_task = None + REDIS_CLIENT.deregister_callback("outage", self._demote) + REDIS_CLIENT.deregister_callback("recover", self._try_acquire) + if REDIS_CLIENT.up: + with contextlib.suppress(Exception): + await REDIS_CLIENT.release_lease(LEADER_LEASE_KEY, self.token) + self._set_leader(False) + + async def _renew_loop(self) -> None: + """Renew or reclaim the lease each heartbeat interval while Redis is up.""" + interval = CONFIG.redis.heartbeat_interval_seconds + try: + while True: + await asyncio.sleep(interval) + if REDIS_CLIENT.up: + await self._attempt() + except asyncio.CancelledError: + return + + async def _attempt(self) -> None: + """Renew or acquire the lease once, updating leader state from the result.""" + if not REDIS_CLIENT.up: + return + try: + held = await REDIS_CLIENT.renew_or_acquire_lease( + LEADER_LEASE_KEY, + self.token, + CONFIG.redis.leader_lease_ttl_seconds * 1000, + ) + except Exception: + # A write-rejecting Redis (memory limit / read-only replica) also fails + # the health probe's canary write, so an outage demotes us within a + # probe cycle. This path only catches lease-specific hiccups; warn if + # they persist since nobody may be leading. + self._consecutive_failures += 1 + if self._consecutive_failures == LEASE_FAILURE_WARN_THRESHOLD: + logger.warning( + f"Leader lease has failed {self._consecutive_failures} times while Redis reads succeed - the cluster may be leaderless. Check Redis write health (memory limit / read-only replica)." + ) + else: + logger.debug("Leader lease renew/acquire failed; will retry.") + REDIS_CLIENT.request_health_check() + return + self._consecutive_failures = 0 + # `held` is True only when our token holds the lease, so promoting while + # Redis is up always reflects reality; a concurrent outage leaves `up` False. + if held and REDIS_CLIENT.up: + self._set_leader(True) + elif not held: + self._set_leader(False) + # Best-effort: refresh when leadership began for /status, so it lives as + # long as the lease. Failures are non-fatal - it's only observability. + if self.is_leader and self._elected_at is not None: + with contextlib.suppress(Exception): + await REDIS_CLIENT.set( + LEADER_ELECTED_KEY, + self._elected_at.isoformat().encode(), + ttl=CONFIG.redis.leader_lease_ttl_seconds, + ) + + def _set_leader(self, value: bool) -> None: + """Flip leader state synchronously, firing `on_acquire` on a False->True edge. + + The flag is set before callbacks run so a callback that reads `is_leader` + (e.g. a gated `refresh`) sees leadership. Callbacks run as tracked + fire-and-forget tasks so a slow build never stalls the renew loop. + """ + was_leader = self.is_leader + self.is_leader = value + if value and not was_leader: + self._elected_at = datetime.now().astimezone() + logger.info("Won leadership; firing initial build.") + for callback in self._acquire_callbacks: + task = asyncio.create_task(self._run_callback(callback)) + self._fired_tasks.add(task) + task.add_done_callback(self._fired_tasks.discard) + elif was_leader and not value: + self._elected_at = None + logger.info("Relinquished leadership.") + + async def _run_callback(self, callback: Callable[[], Awaitable[None]]) -> None: + """Run an `on_acquire` callback, isolating its failures from the loop.""" + try: + await callback() + except Exception: + logger.exception("Leader on_acquire callback failed.") + + async def _demote(self) -> None: + """Relinquish leadership immediately when Redis goes down (on_outage hook).""" + self._set_leader(False) + + async def _try_acquire(self) -> None: + """Re-contend for the lease when Redis recovers (on_recover hook).""" + await self._attempt() + + +LEADER_ELECTION = LeaderElection() diff --git a/src/retriever/utils/redis.py b/src/retriever/utils/redis.py index 6360d783..f024c8bd 100644 --- a/src/retriever/utils/redis.py +++ b/src/retriever/utils/redis.py @@ -34,14 +34,54 @@ OP_TABLE_KEY = "op_table" OP_TABLE_UPDATE_CHANNEL = "op_table:update" -# Worker -> leader signal for a tier recovery; payload is the tier index as a string. +# Worker -> builder signal for a tier recovery; payload is the tier index as a string. TIER_RECOVERED_CHANNEL = "tier:recovered" +# Leader lease; holds the leading instance's builder-process token so only one +# instance builds and publishes shared artifacts at a time. +LEADER_LEASE_KEY = "leader:lease" + +# Best-effort ISO timestamp of when the current leader won, refreshed each renew +# so /status can show "since when"; expires with the lease. +LEADER_ELECTED_KEY = "leader:elected" + +# Extend the lease if we still own it, claim it if it's free, otherwise fail. +# The compare-and-set runs atomically so a lagging ex-leader can't overwrite the +# current holder's token. +_RENEW_OR_ACQUIRE_LEASE_LUA = """ +local current = redis.call('get', KEYS[1]) +if current == ARGV[1] then + redis.call('pexpire', KEYS[1], ARGV[2]) + return 1 +elseif current == false then + redis.call('set', KEYS[1], ARGV[1], 'px', ARGV[2]) + return 1 +else + return 0 +end +""" + +# Drop the lease only if this token still holds it, so a lagging release can't +# delete a successor's freshly-acquired lease. +_RELEASE_LEASE_LUA = """ +if redis.call('get', KEYS[1]) == ARGV[1] then + return redis.call('del', KEYS[1]) +else + return 0 +end +""" + # Timestamp keys written alongside published artifacts so /status # can show freshness without inferring from TTL. OP_TABLE_META_KEY = f"{PREFIX}op_table:meta" SUBCLASS_META_KEY = f"{PREFIX}SubclassHashMap:meta" +# Canary key the health probe writes so a read-only or out-of-memory Redis - +# where PING and reads still succeed but writes are rejected - registers as an +# outage instead of silently stranding writers. +HEALTH_CANARY_KEY = f"{PREFIX}health:canary" +HEALTH_CANARY_TTL_MS = 30_000 + # Process-registry hashes. WORKER_REGISTRY_KEY = f"{PREFIX}workers" BACKGROUND_REGISTRY_KEY = f"{PREFIX}background" @@ -374,8 +414,14 @@ def __init__(self) -> None: @override async def ping(self) -> None: - """Probe Redis. Raises on failure.""" + """Probe Redis with a read and a canary write; raises on failure. + + The write makes a read-only or out-of-memory Redis - where PING and + reads still succeed but writes are rejected - register as an outage, so + writers degrade gracefully instead of failing silently. + """ _ = await self.client.ping() + _ = await self.client.set(HEALTH_CANARY_KEY, b"1", px=HEALTH_CANARY_TTL_MS) @override def _build_client(self) -> None: @@ -658,6 +704,22 @@ async def used_memory_bytes(self) -> int: info = await self.client.info("memory") return int(info["used_memory"]) + async def renew_or_acquire_lease(self, key: str, token: str, ttl_ms: int) -> bool: + """Hold or claim a lease: extend if `token` owns `key`, acquire if free, else fail. + + Returns whether `token` holds the lease afterward. `eval` isn't on the + `AsyncRedis` protocol, so this reaches the raw client directly (as + `hexpire` does) and prefixes the key itself. + """ + script = self._client.register_script(_RENEW_OR_ACQUIRE_LEASE_LUA) + result = cast(int, await script(keys=[f"{PREFIX}{key}"], args=[token, ttl_ms])) + return result == 1 + + async def release_lease(self, key: str, token: str) -> None: + """Drop the lease only if `token` still holds it.""" + script = self._client.register_script(_RELEASE_LEASE_LUA) + _ = cast(int, await script(keys=[f"{PREFIX}{key}"], args=[token])) + async def _read_freshness(self, key: str) -> FreshnessRecord | None: """Read a JSON sidecar `{refreshed_at, count}` from the given key.""" data = await self.client.get(key) @@ -687,6 +749,21 @@ async def subclass_freshness(self) -> FreshnessRecord | None: """Return the last-refreshed timestamp + map size for the subclass map.""" return await self._read_freshness(SUBCLASS_META_KEY) + async def get_leader(self) -> str | None: + """Builder token of the current leader, or None when the cluster is leaderless.""" + holder = await self.get(LEADER_LEASE_KEY) + return holder.decode() if holder is not None else None + + async def get_leader_elected_at(self) -> datetime | None: + """When the current leader won the lease, or None if unknown/leaderless.""" + data = await self.get(LEADER_ELECTED_KEY) + if data is None: + return None + try: + return datetime.fromisoformat(data.decode()) + except ValueError: + return None + async def _register_process( self, registry_key: str, diff --git a/tests/test_utils_leader.py b/tests/test_utils_leader.py new file mode 100644 index 00000000..e3d646fa --- /dev/null +++ b/tests/test_utils_leader.py @@ -0,0 +1,203 @@ +"""Unit tests for the Redis-lease leader election state machine. + +These drive `LeaderElection` directly with a mocked `RedisClient` so the +promote/demote/failure behavior can be exercised without a live Redis. +""" + +from __future__ import annotations + +import asyncio +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from retriever.utils import leader +from retriever.utils.leader import LeaderElection + + +def _fresh() -> LeaderElection: + """A `LeaderElection` built outside the Singleton cache, for per-test isolation.""" + election = object.__new__(LeaderElection) + election.__init__() + return election + + +def _fake_redis(*, up: bool = True, held: bool = True) -> MagicMock: + """A stand-in `RedisClient` exposing just the lease/callback surface used here.""" + redis = MagicMock() + redis.up = up + redis.renew_or_acquire_lease = AsyncMock(return_value=held) + redis.release_lease = AsyncMock() + redis.set = AsyncMock() + redis.request_health_check = MagicMock() + redis.on_outage = MagicMock() + redis.on_recover = MagicMock() + redis.deregister_callback = MagicMock() + return redis + + +async def _drain(event: asyncio.Event) -> None: + """Wait for a fire-and-forget `on_acquire` callback to run.""" + await asyncio.wait_for(event.wait(), 1) + + +@pytest.mark.asyncio +async def test_win_on_start_promotes_and_fires_once(monkeypatch: pytest.MonkeyPatch): + redis = _fake_redis(up=True, held=True) + monkeypatch.setattr(leader, "REDIS_CLIENT", redis) + election = _fresh() + + fired = asyncio.Event() + seen_leader: list[bool] = [] + + async def callback() -> None: + # The flag must already be True when the callback runs (set-before-fire). + seen_leader.append(election.is_leader) + fired.set() + + election.on_acquire(callback) + + await election.start() + await _drain(fired) + + assert election.is_leader is True + assert seen_leader == [True] + redis.on_outage.assert_called_once() + redis.on_recover.assert_called_once() + + await election.stop() + redis.release_lease.assert_awaited_once() + assert election.is_leader is False + + +@pytest.mark.asyncio +async def test_loser_stays_follower(monkeypatch: pytest.MonkeyPatch): + redis = _fake_redis(up=True, held=False) + monkeypatch.setattr(leader, "REDIS_CLIENT", redis) + election = _fresh() + + fired = asyncio.Event() + + async def callback() -> None: + fired.set() + + election.on_acquire(callback) + + await election.start() + await asyncio.sleep(0) + + assert election.is_leader is False + assert not fired.is_set() + + await election.stop() + + +@pytest.mark.asyncio +async def test_down_at_start_then_recovers(monkeypatch: pytest.MonkeyPatch): + redis = _fake_redis(up=False, held=True) + monkeypatch.setattr(leader, "REDIS_CLIENT", redis) + election = _fresh() + + fired = asyncio.Event() + + async def callback() -> None: + fired.set() + + election.on_acquire(callback) + + await election.start() + redis.renew_or_acquire_lease.assert_not_awaited() + assert election.is_leader is False + + # Redis recovers; the registered on_recover hook drives `_try_acquire`. + redis.up = True + await election._try_acquire() + await _drain(fired) + + assert election.is_leader is True + + await election.stop() + + +@pytest.mark.asyncio +async def test_outage_demotes(monkeypatch: pytest.MonkeyPatch): + redis = _fake_redis(up=True, held=True) + monkeypatch.setattr(leader, "REDIS_CLIENT", redis) + election = _fresh() + + await election.start() + assert election.is_leader is True + + await election._demote() + + assert election.is_leader is False + assert election._elected_at is None + + await election.stop() + + +@pytest.mark.asyncio +async def test_win_publishes_elected_at(monkeypatch: pytest.MonkeyPatch): + redis = _fake_redis(up=True, held=True) + monkeypatch.setattr(leader, "REDIS_CLIENT", redis) + election = _fresh() + + await election.start() + + assert election.is_leader is True + assert election._elected_at is not None + redis.set.assert_awaited() + key_arg = redis.set.await_args.args[0] + assert key_arg == leader.LEADER_ELECTED_KEY + + await election.stop() + + +@pytest.mark.asyncio +async def test_persistent_lease_failure_stays_leaderless( + monkeypatch: pytest.MonkeyPatch, +): + redis = _fake_redis(up=True, held=True) + redis.renew_or_acquire_lease.side_effect = RuntimeError("write rejected") + monkeypatch.setattr(leader, "REDIS_CLIENT", redis) + election = _fresh() + + # Reads succeed (up stays True) but every lease write fails: nobody leads. + await election.start() + await election._attempt() + await election._attempt() + + assert election.is_leader is False + assert election._consecutive_failures >= leader.LEASE_FAILURE_WARN_THRESHOLD + redis.request_health_check.assert_called() + + # Writes recover: the next attempt promotes and clears the failure count. + redis.renew_or_acquire_lease.side_effect = None + redis.renew_or_acquire_lease.return_value = True + await election._attempt() + + assert election.is_leader is True + assert election._consecutive_failures == 0 + + await election.stop() + + +@pytest.mark.asyncio +async def test_transient_renew_failure_keeps_leadership( + monkeypatch: pytest.MonkeyPatch, +): + redis = _fake_redis(up=True, held=True) + monkeypatch.setattr(leader, "REDIS_CLIENT", redis) + election = _fresh() + + await election.start() + assert election.is_leader is True + + # A transient renew error must not drop leadership - the TTL still protects us. + redis.renew_or_acquire_lease.side_effect = RuntimeError("boom") + await election._attempt() + + assert election.is_leader is True + redis.request_health_check.assert_called() + + await election.stop() diff --git a/tests/test_utils_redis_health.py b/tests/test_utils_redis_health.py new file mode 100644 index 00000000..8aeff1e3 --- /dev/null +++ b/tests/test_utils_redis_health.py @@ -0,0 +1,42 @@ +"""Unit tests for the write-aware Redis health probe. + +`RedisClient.ping` issues a canary write so a read-only or out-of-memory Redis +(reads/PING succeed, writes rejected) registers as an outage rather than +silently stranding writers. +""" + +from __future__ import annotations + +from unittest.mock import AsyncMock + +import pytest + +from retriever.utils.redis import HEALTH_CANARY_KEY, RedisClient + + +@pytest.mark.asyncio +async def test_ping_reads_and_writes_a_canary(): + redis = RedisClient() + original = redis.client + try: + redis.client = AsyncMock() + await redis.ping() + redis.client.ping.assert_awaited_once() + redis.client.set.assert_awaited_once() + assert redis.client.set.await_args.args[0] == HEALTH_CANARY_KEY + finally: + redis.client = original + + +@pytest.mark.asyncio +async def test_ping_raises_when_writes_are_rejected(): + redis = RedisClient() + original = redis.client + try: + # Reads/PING fine, but the canary write is rejected (read-only / OOM). + redis.client = AsyncMock() + redis.client.set.side_effect = RuntimeError("READONLY") + with pytest.raises(RuntimeError): + await redis.ping() + finally: + redis.client = original