Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
73 changes: 68 additions & 5 deletions tensorrt_llm/serve/cluster_storage.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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"])
Expand All @@ -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,
Expand All @@ -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)
Expand All @@ -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,
Expand All @@ -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

Expand Down Expand Up @@ -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
Expand Down
33 changes: 22 additions & 11 deletions tensorrt_llm/serve/disagg_auto_scaling.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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))
Comment thread
reasonsolo marked this conversation as resolved.
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
Comment thread
coderabbitai[bot] marked this conversation as resolved.
continue
if refreshed:
self._stamp_registration_expiry(attempt_start)
Expand Down
28 changes: 28 additions & 0 deletions tensorrt_llm/serve/openai_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Comment thread
coderabbitai[bot] marked this conversation as resolved.
content["startup_metrics"] = getattr(self.generator, "startup_metrics",
{})
return JSONResponse(content=content)
Expand Down
147 changes: 145 additions & 2 deletions tests/unittest/disaggregated/test_cluster_storage.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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
Expand All @@ -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"))
Expand Down Expand Up @@ -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)
Comment thread
coderabbitai[bot] marked this conversation as resolved.

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
Comment thread
reasonsolo marked this conversation as resolved.
await asyncio.gather(stall_then_read(), stall_again())
assert await storage.get(key) == "worker", (
Comment thread
reasonsolo marked this conversation as resolved.
"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()
Loading
Loading