diff --git a/tensorrt_llm/serve/cluster_storage.py b/tensorrt_llm/serve/cluster_storage.py index bb4f1902b6f4..22daf14748fd 100644 --- a/tensorrt_llm/serve/cluster_storage.py +++ b/tensorrt_llm/serve/cluster_storage.py @@ -208,6 +208,9 @@ def key_time(): class HttpClusterStorageServer(ClusterStorage): + # Granularity at which the expiry sweep samples the event loop. + _OUTAGE_SAMPLE_SEC = 0.1 + def __init__(self, cluster_uri, cluster_name, @@ -219,9 +222,48 @@ def __init__(self, self._watch_lock = asyncio.Lock() self._check_expired_task = None self._check_expired_interval = 1 # in seconds + # Total time this event loop was unable to serve requests, excluded + # from the clock TTLs are measured on (see _service_now). + self._unserviceable_sec = 0.0 + self._outage_sample_deadline: Optional[float] = None if server: self.add_routes(server) + def _settle_outage_sample(self, now: float) -> None: + """Account for an overdue event-loop probe exactly once. + + Re-arm rather than disarm: the sampling task cannot re-arm until the + loop gives it a turn, and everything else already queued runs first. + Clearing the deadline here would leave that window unmeasured, so a + second handler blocking back-to-back with the first is charged to TTL + and expires a live worker. Arming the next sample from ``now`` keeps + the accounting continuous; ``_sleep_counting_outage`` overwrites this + deadline with its own the moment it does get a turn. + """ + deadline = self._outage_sample_deadline + if deadline is None or now < deadline: + return + self._unserviceable_sec += now - deadline + self._outage_sample_deadline = now + self._OUTAGE_SAMPLE_SEC + + def _service_now(self) -> float: + """Monotonic time minus however long this loop could not serve requests. + + Workers refresh their TTL over ``/expire`` on this very event loop, so + while the loop is blocked their refreshes sit unread rather than + arriving late. Charging that time to a key expires a live worker for + the storage's own outage: the coordinator evicts it from the routers + and requests routed there fail. A worker that really stopped refreshing + still expires, because the clock only pauses while nobody could have + been served. + """ + now = key_time() + # A queued request may run before the sleeping expiry task resumes. + # Settle its overdue probe here so no TTL operation can observe the + # stale service clock in that scheduling window. + self._settle_outage_sample(now) + return now - self._unserviceable_sec + def add_routes(self, server: FastAPI): server.add_api_route("/set", jsonify(self._set), methods=["POST"]) server.add_api_route("/get", jsonify(self.get), methods=["GET"]) @@ -242,6 +284,7 @@ async def stop(self): if self._check_expired_task: self._check_expired_task.cancel() self._check_expired_task = None + self._outage_sample_deadline = None async def set(self, key: str, @@ -259,7 +302,8 @@ async def _set(self, storage_item: StorageItem) -> bool: if storage_item.key in self._storage and not storage_item.overwrite_if_exists: return False if storage_item.expire_time < 0 and storage_item.ttl and storage_item.ttl > 0: - storage_item.expire_time = key_time() + storage_item.ttl + storage_item.expire_time = (self._service_now() + + storage_item.ttl) self._storage[storage_item.key] = storage_item await self._notify_watch_event(storage_item.key, storage_item, WatchEventType.SET) @@ -269,7 +313,8 @@ async def get(self, key: str) -> str: async with self._lock: if key in self._storage: item = self._storage[key] - if item.expire_time < 0 or item.expire_time > key_time(): + now = self._service_now() + if item.expire_time < 0 or item.expire_time > now: return item.value else: await self._notify_watch_event(key, item, @@ -280,7 +325,7 @@ async def get(self, key: str) -> str: async def expire(self, key: str, ttl: int) -> bool: async with self._lock: if key in self._storage: - self._storage[key].expire_time = key_time() + int(ttl) + self._storage[key].expire_time = self._service_now() + int(ttl) return True return False @@ -339,12 +384,30 @@ async def _notify_watch_event(self, key, storage_item: StorageItem, logger.info( f"Notified watch event for key {key} with type {event_type}") + async def _sleep_counting_outage(self, duration: float) -> None: + """Sleep, accumulating however long the loop overran its own timers. + + Sliced rather than one long sleep: a single sleep only reveals lateness + after its own deadline, so a block that started mid-sleep is credited + short by however far in it began. Crediting happens here, in the sweep's + own task, because a separate prober and the sweep would both be ready + when the loop resumes and the sweep could reap first. + """ + remaining = duration + while remaining > 0: + slice_sec = min(self._OUTAGE_SAMPLE_SEC, remaining) + before = key_time() + self._outage_sample_deadline = before + slice_sec + await asyncio.sleep(slice_sec) + self._settle_outage_sample(key_time()) + remaining -= slice_sec + async def _check_expired(self): while True: - await asyncio.sleep(self._check_expired_interval) + await self._sleep_counting_outage(self._check_expired_interval) try: before_len = len(self._storage) - current_time = key_time() + current_time = self._service_now() async with self._lock: kv_to_delete = { k: v diff --git a/tensorrt_llm/serve/disagg_auto_scaling.py b/tensorrt_llm/serve/disagg_auto_scaling.py index fcbeedbff97e..92d11361b31a 100644 --- a/tensorrt_llm/serve/disagg_auto_scaling.py +++ b/tensorrt_llm/serve/disagg_auto_scaling.py @@ -387,17 +387,25 @@ def _stamp_registration_expiry(self, ttl_applied_after: float) -> None: self._config.inactive_timeout_sec) async def _refresh_registration(self) -> bool: - """Refresh this worker's TTL, retrying while the window allows. + """Refresh this worker's TTL, retrying until the storage answers. A storage RPC is bounded only by the client's own timeout, which is as coarse as the heartbeat interval and half of inactive_timeout_sec, so - awaiting one stalled /expire burns the whole TTL window: a healthy - worker's registration expires, the coordinator evicts it from the - routers, and in-flight requests routed there fail. Bound each attempt - to a fraction of the time left and spend the rest of the window - retrying. Only a stall is retried; "not refreshed" is definitive and - returns immediately so the caller can re-register. + awaiting one stalled /expire burns the whole TTL window. Bound each + attempt to a fraction of the time left and spend the rest retrying. + + A timeout means "no answer yet", not "expired": the storage stamps the + new TTL when it *receives* the request and reaps lapsed keys lazily, so + _registration_expires_at is only a conservative local bound and the + registration normally outlives it. Treating it as definitive tears down + a healthy worker whose coordinator was merely busy for one interval -- + it re-registers, which the coordinator reads as a leave/join, evicting + it from the routers so requests routed there fail. So retry through a + grace period; only an explicit "not refreshed", or silence past the + grace, is definitive. """ + grace_deadline = (self._registration_expires_at + + self._config.inactive_timeout_sec) while not self._stop: attempt_start = key_time() remaining = self._registration_expires_at - attempt_start @@ -407,12 +415,15 @@ async def _refresh_registration(self) -> bool: self.worker_key, self._config.inactive_timeout_sec), timeout=max(self._MIN_REFRESH_TIMEOUT_SEC, remaining / 3)) except asyncio.TimeoutError: - if key_time() >= self._registration_expires_at: - return False + now = key_time() + lapsed = now >= grace_deadline + action = ("giving up past the grace period" + if lapsed else "retrying") logger.warning( f"Worker {self.worker_info.worker_id} heartbeat refresh " - f"stalled, retrying before the registration expires " - f"{key_time()}") + f"stalled, {action} {now}") + if lapsed: + return False continue if refreshed: self._stamp_registration_expiry(attempt_start) diff --git a/tensorrt_llm/serve/openai_server.py b/tensorrt_llm/serve/openai_server.py index 689ed7f341e2..b1165bdfbae6 100644 --- a/tensorrt_llm/serve/openai_server.py +++ b/tensorrt_llm/serve/openai_server.py @@ -3555,6 +3555,26 @@ async def update_weights(self, args=(request.weights, )) return JSONResponse(content={"status": "success"}) + async def _live_tokens_per_block(self) -> Optional[int]: + """Return the runtime's effective KV block size, or None if unknown. + + The executor layer already turns RPC failures into an empty dict, so + the only failure left to absorb here is ``encode_only``, which rejects + the call outright. Generators without a KV cache (VisualGen) have no + such method. Both mean "fall back to the configured value". + """ + get_capacity = getattr(self.generator, "get_kv_cache_capacity", None) + if get_capacity is None: + return None + try: + # Off-loop: the RPC blocks, and this worker's own heartbeat task + # shares this loop, so stalling it here can lapse its registration. + capacity = await asyncio.to_thread(get_capacity) + except RuntimeError as e: + logger.debug(f"Could not read live tokens_per_block: {e}") + return None + return capacity.get("tokensPerBlock") or None + async def get_server_info(self) -> JSONResponse: # Note: calling self.generator.disaggregated_params and startup_metrics below # may trigger an RPC sync call, blocking the server event loop. Since this server_info @@ -3573,6 +3593,14 @@ async def get_server_info(self) -> JSONResponse: if kv_cache_config.tokens_per_block is not None: content[ "tokens_per_block"] = kv_cache_config.tokens_per_block + # The runtime may override the configured block size (e.g. FlashMLA + # forces 64) in the worker process, so args.kv_cache_config still + # holds the pre-override value here. A kv-cache-aware router hashes + # prompts in whatever block size this endpoint publishes, so a + # stale value makes every block hash miss. Prefer the live value. + live_tokens_per_block = await self._live_tokens_per_block() + if live_tokens_per_block is not None: + content["tokens_per_block"] = live_tokens_per_block content["startup_metrics"] = getattr(self.generator, "startup_metrics", {}) return JSONResponse(content=content) diff --git a/tests/unittest/disaggregated/test_cluster_storage.py b/tests/unittest/disaggregated/test_cluster_storage.py index 98854f6957de..9c476d1e5302 100644 --- a/tests/unittest/disaggregated/test_cluster_storage.py +++ b/tests/unittest/disaggregated/test_cluster_storage.py @@ -13,7 +13,7 @@ from tensorrt_llm.serve.cluster_storage import ( HttpClusterStorageServer, StorageItem, WatchEvent, WatchEventType, create_cluster_storage, create_cluster_storage_client, is_loopback_host, - validate_http_cluster_storage_scope) + jsonify, validate_http_cluster_storage_scope) pytestmark = pytest.mark.cpu_only @@ -230,7 +230,7 @@ async def test_watch_set_and_delete(self, storage_server_client): ]) == {WatchEventType.DELETE, WatchEventType.SET} -def http_server_storage(port): +def http_server_storage(port, expire_block_sec=0): cluster_storage = HttpClusterStorageServer("", "") @contextlib.asynccontextmanager @@ -240,6 +240,15 @@ async def lifespan(app: FastAPI): await cluster_storage.stop() app = FastAPI(lifespan=lifespan) + if expire_block_sec > 0: + + async def blocked_expire(key: str, ttl: int) -> bool: + # Block inside the real HTTP request's event loop so its TTL read + # is ordered before the overdue expiry-sweep continuation. + time.sleep(expire_block_sec) + return await cluster_storage.expire(key, ttl) + + app.add_api_route("/expire", jsonify(blocked_expire), methods=["GET"]) cluster_storage.add_routes(app) server = Server( uvicorn.Config(app=app, host="localhost", port=port, log_level="info")) @@ -269,3 +278,137 @@ def storage_server(self): yield self.etcd, "etcd://localhost:2379" self.etcd.kill() self.etcd.wait() + + +@pytest.mark.asyncio +async def test_expiry_does_not_charge_the_storage_own_outage(): + """A worker refreshing its TTL must survive a block longer than that TTL. + + Workers refresh over ``/expire`` on the storage's own event loop, so while + that loop is blocked their refreshes sit unread rather than arriving late. + Charging the block to the key expires a live worker for the storage's own + outage, which evicts it from the routers and fails requests routed there + (https://nvbugs/6786712). Uses the tight functional-test timings + (ttl=2s, refresh every 1s) against a block that straddles the deadline. + """ + ttl, refresh_sec, block_sec = 2, 1, 2.5 + storage = HttpClusterStorageServer("", "") + await storage.start() + try: + key = gen_key("outage_key") + assert await storage.set(key, "worker", ttl=ttl) + + async def refresh_periodically(): + while True: + await asyncio.sleep(refresh_sec) + await storage.expire(key, ttl) + + refresher = asyncio.create_task(refresh_periodically()) + try: + await asyncio.sleep(refresh_sec + 0.2) # one clean refresh first + # Busy-wait, holding the loop exactly as a cold tokenizer build does. + block_end = time.monotonic() + block_sec + while time.monotonic() < block_end: + pass + # Let the loop resume so the expiry sweep runs at least once. + await asyncio.sleep(storage._check_expired_interval + 0.5) + assert await storage.get(key) == "worker", ( + "a live, refreshing worker was expired for time the storage " + "itself could not serve refreshes") + finally: + refresher.cancel() + with contextlib.suppress(asyncio.CancelledError): + await refresher + + # The clock must still expire a worker that really stopped refreshing, + # otherwise the fix above would keep dead workers registered forever. + await asyncio.sleep(ttl + storage._check_expired_interval + 0.5) + assert await storage.get(key) is None + finally: + await storage.stop() + + +@pytest.mark.asyncio +async def test_queued_http_refresh_settles_outage_before_ttl_operations( + unused_tcp_port): + """A queued HTTP refresh must settle a storage-loop outage first.""" + ttl, block_sec = 2, 2.5 + server, storage = http_server_storage(unused_tcp_port, + expire_block_sec=block_sec) + + with server.run_in_thread(): + client = create_cluster_storage_client( + f"http://localhost:{unused_tcp_port}", "test") + try: + key = gen_key("queued_outage_key") + assert await client.set(key, "worker", ttl=ttl) + + # The /expire handler blocks the uvicorn/storage loop past the + # current TTL, then refreshes before the queued sweep can resume. + assert await client.expire(key, ttl) + assert await client.get(key) == "worker" + + # A refresh made with the stale wall clock would extend this TTL + # by the outage duration a second time. + await asyncio.sleep(ttl + storage._check_expired_interval + 0.5) + assert await client.get(key) is None + finally: + await client._session.close() + + +@pytest.mark.asyncio +async def test_consecutive_outages_are_not_charged_to_ttl(): + """A stall right after a settled one must not expire a live worker. + + Any TTL operation settles the overdue outage sample, and the sampling task + cannot arm the next one until the loop gives it a turn -- which it cannot + while another ready handler is blocking. Leaving that window unmeasured + charges the second stall to the key and deletes a worker whose periodic + refresh is merely queued (https://nvbugs/6786712). Same tight + functional-test timings as above (ttl=2s, refresh every 1s). + """ + ttl, refresh_sec, block_sec = 2, 1, 2.5 + storage = HttpClusterStorageServer("", "") + await storage.start() + try: + key = gen_key("consecutive_outage_key") + assert await storage.set(key, "worker", ttl=ttl) + + async def refresh_periodically(): + while True: + await asyncio.sleep(refresh_sec) + await storage.expire(key, ttl) + + def block_the_loop(): + # Busy-wait, holding the loop exactly as a cold tokenizer build does. + block_end = time.monotonic() + block_sec + while time.monotonic() < block_end: + pass + + async def stall_then_read(): + # Settles the first stall's sample before the sweep can resume. + block_the_loop() + assert await storage.get(key) == "worker" + + async def stall_again(): + # Already queued, so it runs in the same batch as the handler above + # and blocks the loop again before the sweep gets a turn. + block_the_loop() + + refresher = asyncio.create_task(refresh_periodically()) + try: + await asyncio.sleep(refresh_sec + 0.2) # one clean refresh first + await asyncio.gather(stall_then_read(), stall_again()) + assert await storage.get(key) == "worker", ( + "the second consecutive stall was charged to the key's TTL, " + "expiring a live worker for the storage's own outage") + finally: + refresher.cancel() + with contextlib.suppress(asyncio.CancelledError): + await refresher + + # The clock must still expire a worker that really stopped refreshing. + await asyncio.sleep(ttl + storage._check_expired_interval + 0.5) + assert await storage.get(key) is None + finally: + await storage.stop() diff --git a/tests/unittest/disaggregated/test_disagg_cluster_manager_worker.py b/tests/unittest/disaggregated/test_disagg_cluster_manager_worker.py index f00ad888dadd..f93ed0de0acb 100644 --- a/tests/unittest/disaggregated/test_disagg_cluster_manager_worker.py +++ b/tests/unittest/disaggregated/test_disagg_cluster_manager_worker.py @@ -5,7 +5,7 @@ import subprocess import tempfile import time -from unittest.mock import ANY, AsyncMock +from unittest.mock import ANY, AsyncMock, patch import pytest import pytest_asyncio @@ -127,6 +127,22 @@ async def stalled_then_ok(*args, **kwargs): storage.set.assert_not_awaited() +@pytest.mark.asyncio +async def test_refresh_timeout_past_grace_returns_without_retry(): + config = worker_config() + storage = AsyncMock() + storage.expire.side_effect = asyncio.TimeoutError + worker = DisaggClusterWorker(ServerRole.CONTEXT, "127.0.0.1", 8001, config, + storage) + worker._registration_expires_at = 100 + + with patch("tensorrt_llm.serve.disagg_auto_scaling.key_time", + side_effect=[105, 105]): + assert not await worker._refresh_registration() + + assert storage.expire.await_count == 1 + + def get_uri(storage_type): if storage_type == "http": return f"http://localhost:18000" diff --git a/tests/unittest/disaggregated/test_openai_server_info.py b/tests/unittest/disaggregated/test_openai_server_info.py index 5f1aa142791f..114d6f10319b 100644 --- a/tests/unittest/disaggregated/test_openai_server_info.py +++ b/tests/unittest/disaggregated/test_openai_server_info.py @@ -14,6 +14,7 @@ import json from types import SimpleNamespace +from unittest.mock import Mock import pytest @@ -80,3 +81,29 @@ async def test_server_info_includes_tokens_per_block_from_kv_cache_config(): response = await server.get_server_info() content = json.loads(response.body) assert content["tokens_per_block"] == 64 + + +@pytest.mark.asyncio +async def test_server_info_prefers_live_tokens_per_block(): + server = _make_server(kv_cache_config=KvCacheConfig(tokens_per_block=32)) + server.generator.get_kv_cache_capacity = Mock(return_value={"tokensPerBlock": 64}) + + response = await server.get_server_info() + + content = json.loads(response.body) + assert content["tokens_per_block"] == 64 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("capacity", [{}, RuntimeError("encode only")]) +async def test_server_info_preserves_configured_tokens_per_block_on_live_lookup_failure(capacity): + server = _make_server(kv_cache_config=KvCacheConfig(tokens_per_block=32)) + if isinstance(capacity, Exception): + server.generator.get_kv_cache_capacity = Mock(side_effect=capacity) + else: + server.generator.get_kv_cache_capacity = Mock(return_value=capacity) + + response = await server.get_server_info() + + content = json.loads(response.body) + assert content["tokens_per_block"] == 32