From 5ba27bafb695fcd4d8b63f7dfbf7098a462ace67 Mon Sep 17 00:00:00 2001 From: Lizhi Zhou <1432185+reasonsolo@users.noreply.github.com> Date: Thu, 18 Sep 2025 19:52:41 -0700 Subject: [PATCH 1/3] tested Signed-off-by: Lizhi Zhou <1432185+reasonsolo@users.noreply.github.com> --- tensorrt_llm/llmapi/disagg_utils.py | 4 +- tensorrt_llm/serve/auto_scaling.py | 283 +++++++++++++ tensorrt_llm/serve/cluster_storage.py | 376 ++++++++++++++++++ tests/unittest/serve/__init__.py | 0 .../serve/test_cluster_manager_worker.py | 211 ++++++++++ tests/unittest/serve/test_cluster_storage.py | 212 ++++++++++ 6 files changed, 1084 insertions(+), 2 deletions(-) create mode 100644 tensorrt_llm/serve/auto_scaling.py create mode 100644 tensorrt_llm/serve/cluster_storage.py create mode 100644 tests/unittest/serve/__init__.py create mode 100644 tests/unittest/serve/test_cluster_manager_worker.py create mode 100644 tests/unittest/serve/test_cluster_storage.py diff --git a/tensorrt_llm/llmapi/disagg_utils.py b/tensorrt_llm/llmapi/disagg_utils.py index 8404cbaf7ad3..cef72c687791 100644 --- a/tensorrt_llm/llmapi/disagg_utils.py +++ b/tensorrt_llm/llmapi/disagg_utils.py @@ -1,6 +1,6 @@ import logging from dataclasses import dataclass, field -from enum import Enum +from enum import IntEnum from typing import Any, List, Literal, Optional, Tuple import yaml @@ -16,7 +16,7 @@ ] -class ServerRole(Enum): +class ServerRole(IntEnum): CONTEXT = 0 GENERATION = 1 MM_ENCODER = 2 diff --git a/tensorrt_llm/serve/auto_scaling.py b/tensorrt_llm/serve/auto_scaling.py new file mode 100644 index 000000000000..9d098a04dec2 --- /dev/null +++ b/tensorrt_llm/serve/auto_scaling.py @@ -0,0 +1,283 @@ +import asyncio +import json +import os +import random +import time +from dataclasses import asdict, dataclass +from typing import List, Tuple + +from tensorrt_llm.llmapi.disagg_utils import DisaggClusterConfig, ServerRole +from tensorrt_llm.logger import logger + +from .cluster_storage import (ClusterStorage, StorageItem, WatchEvent, + WatchEventType, key_time) + + +@dataclass +class WorkerInfo: + worker_id: str + host: str = "" + port: int = 0 + role: ServerRole = ServerRole.CONTEXT + status: str = "" + + +def get_worker_key_prefix(cluster_name: str): + return f"/trtllm-disagg/{cluster_name}/workers" + + +def get_worker_key(name: str, role: ServerRole, worker_id: str = "") -> str: + return f"{get_worker_key_prefix(name)}/{worker_id}" + + +class ClusterManager: + + def __init__(self, config: DisaggClusterConfig, storage: ClusterStorage): + self._config = config + self._cluster_storage = storage + self._lock = asyncio.Lock() + self._minimal_ctx_worker_num = config.minimal_instances.context_servers + self._minimal_gen_worker_num = config.minimal_instances.generation_servers + self._current_ctx_workers = {} + self._current_gen_workers = {} + self._watch_handle = None + + async def start(self): + await self._cluster_storage.start() + + async def stop(self): + await self._cluster_storage.stop() + + async def cluster_info(self) -> dict: + async with self._lock: + return { + "current_workers": { + "context_servers": [ + asdict(worker) + for worker in self._current_ctx_workers.values() + ], + "generation_servers": [ + asdict(worker) + for worker in self._current_gen_workers.values() + ] + }, + "minimal_instances": { + "context_servers": self._minimal_ctx_worker_num, + "generation_servers": self._minimal_gen_worker_num + }, + } + + @property + def current_ctx_worker_num(self): + return len(self._current_ctx_workers) + + @property + def current_gen_worker_num(self): + return len(self._current_gen_workers) + + @property + def worker_key_prefix(self): + return get_worker_key_prefix(self._config.cluster_name) + + async def watch_workers(self, get_existing_first: bool = True): + workers = [] + if get_existing_first: + # There is a tiny gap between getting existing workers and watching the key, + # which may cause we missing some workers registered in between. + resp = await self._cluster_storage.get_prefix( + self.worker_key_prefix, keys_only=False) + for worker_id, data in resp.items(): + event = WatchEvent(storage_item=StorageItem(key=worker_id, + value=data), + event_type=WatchEventType.SET) + workers.append(self._parse_worker_info(event)) + self._watch_handle = await self._cluster_storage.watch( + self.worker_key_prefix) + return workers + + async def unwatch_workers(self): + await self._cluster_storage.unwatch([self.worker_key_prefix]) + self._watch_handle = None + + async def get_worker_events( + self) -> List[Tuple[WorkerInfo, WatchEventType]]: + events = await self._watch_handle.drain() + worker_events = [] + for event in events: + try: + print( + f"Processing event: {event.event_type} for key: {event.storage_item.key} value {event.storage_item.value}" + ) + worker_info = self._parse_worker_info(event) + worker_events.append((worker_info, event.event_type)) + except Exception as e: + logger.error( + f"Error parsing worker info: {event.storage_item.value}, error: {e}" + ) + continue + return worker_events + + def _log_cluster_status(self, worker_info: WorkerInfo, change_event: str): + logger.error( + f"Worker {worker_info.worker_id} becomes {change_event}, current context worker: {self.current_ctx_worker_num}/{self._minimal_ctx_worker_num}, current generation worker: {self.current_gen_worker_num}/{self._minimal_gen_worker_num}" + ) + + def _get_workers(self, role: ServerRole) -> dict[str, WorkerInfo]: + if role == ServerRole.CONTEXT: + return self._current_ctx_workers + elif role == ServerRole.GENERATION: + return self._current_gen_workers + else: + raise ValueError(f"Invalid worker role: {role}") + + def _get_workers_by_id(self, worker_id: str) -> dict[str, WorkerInfo]: + if worker_id in self._current_ctx_workers: + return self._current_ctx_workers + elif worker_id in self._current_gen_workers: + return self._current_gen_workers + else: + raise ValueError(f"Worker {worker_id} is unknown") + + def _parse_worker_info(self, event: WatchEvent) -> Tuple[WorkerInfo, bool]: + # parse the worker info from the event, if it's a delete event, pop the corresponding worker from the current workers + # if it's a set event, parse the worker info from the value and add it to the current workers + # return the worker info and whether to notify the event + if event.event_type == WatchEventType.DELETE: + workers = self._get_workers_by_id(event.storage_item.key) + if workers is None: + logger.warning( + f"Failed to parse delete event: Worker {event.storage_item.key} is unknown, " + ) + worker_info = WorkerInfo(worker_id=event.storage_item.key) + else: + worker_info = workers.pop(event.storage_item.key) + elif event.event_type == WatchEventType.SET: + try: + worker_info = WorkerInfo(**json.loads(event.storage_item.value)) + worker_info.role = ServerRole(worker_info.role) + workers = self._get_workers(worker_info.role) + workers[event.storage_item.key] = worker_info + + except Exception as e: + logger.error( + f"Failed to parse set event: {event.storage_item.key}: {event.storage_item.value}, error: {e}" + ) + worker_info = WorkerInfo(worker_id=event.storage_item.key) + else: + raise ValueError(f"Invalid event type: {event.event_type}") + self._log_cluster_status( + worker_info, "active/updated" + if event.event_type == WatchEventType.SET else "inactive") + return worker_info + + async def is_ready(self) -> bool: + return self.current_ctx_worker_num >= self._minimal_ctx_worker_num and self.current_gen_worker_num >= self._minimal_gen_worker_num + + async def is_ready_with_router(self, router_ctx_worker_num: int, + router_gen_worker_num: int) -> bool: + return router_ctx_worker_num >= self._minimal_ctx_worker_num and router_gen_worker_num >= self._minimal_gen_worker_num + + +class ClusterWorker: + + def __init__(self, role: ServerRole, host: str, port: int, + config: DisaggClusterConfig, storage: ClusterStorage): + self._role = role + self._host = host + self._port = port + self._config = config + self._cluster_storage = storage + self._stop = False + self._heartbeat_task = None + self._last_heartbeat = 0 + self._worker_id = f"{role.name}-{host}:{port}-{int(time.time()*1000)}-{os.getpid()}-{random.randint(0, 1000):03}" + + @property + def worker_id(self) -> str: + return self._worker_id + + @property + def worker_info(self) -> WorkerInfo: + return WorkerInfo(worker_id=self._worker_id, + role=self._role, + host=self._host, + port=self._port, + status="") + + @property + def worker_key(self) -> str: + return get_worker_key(self._config.cluster_name, self._role, + self._worker_id) + + async def register_worker(self, validator=None, retry_interval=5): + self._stop = False + await self._cluster_storage.start() + if validator and not validator(): + logger.warning( + f"Worker {self.worker_info.worker_id} is not valid, skipping registration" + ) + return False + worker_info = self.worker_info + logger.debug( + f"Worker {self.worker_info.worker_id} registering, {asdict(worker_info)}" + ) + success = await self._cluster_storage.set( + self.worker_key, + json.dumps(asdict(worker_info)), + ttl=self._config.inactive_timeout) + if not success: + if retry_interval > 0: + logger.warning( + f"Worker {self.worker_info.worker_id} registration failed, retry in {retry_interval} seconds" + ) + await asyncio.sleep(max(10, retry_interval)) + return await self.register_worker(validator, retry_interval + 1) + else: + logger.info( + f"Worker {self.worker_info.worker_id} registration successful") + self._last_heartbeat = key_time() + if self._config.heartbeat_interval > 0 and self._config.heartbeat_interval < self._config.inactive_timeout: + if not self._heartbeat_task: + self._heartbeat_task = asyncio.create_task( + self._heartbeat(validator)) + else: + logger.warning( + f"Heartbeat interval {self._config.heartbeat_interval} is not positive or less than inactive timeout {self._config.inactive_timeout}, heartbeat is disabled" + ) + return True + + async def deregister_worker(self): + self._stop = True + self._heartbeat_task.cancel() + self._heartbeat_task = None + await self._cluster_storage.stop() + success = await self._cluster_storage.delete(self.worker_key) + if not success: + logger.warning( + f"Worker {self.worker_info.worker_id} deregistration failed") + return success + + async def _heartbeat(self, validator=None): + logger.info(f"Worker {self.worker_info.worker_id} heartbeat started") + while not self._stop: + remaining_time = self._config.heartbeat_interval - ( + key_time() - self._last_heartbeat) + if remaining_time > 0: + await asyncio.sleep(remaining_time) + self._last_heartbeat = key_time() + if validator and not validator(): + logger.warning( + f"Worker {self.worker_info.worker_id} is not valid, skipping heartbeat {key_time()}" + ) + continue + expire_res = await self._cluster_storage.expire( + self.worker_key, self._config.inactive_timeout) + if not expire_res: + logger.warning( + f"Worker {self.worker_info.worker_id} heartbeat failed, re-registering {key_time()}" + ) + await self.register_worker(validator) + else: + logger.debug( + f"Worker {self.worker_info.worker_id} heartbeat successful {key_time()}" + ) diff --git a/tensorrt_llm/serve/cluster_storage.py b/tensorrt_llm/serve/cluster_storage.py new file mode 100644 index 000000000000..6f14e96c2a9d --- /dev/null +++ b/tensorrt_llm/serve/cluster_storage.py @@ -0,0 +1,376 @@ +import abc +import asyncio +import logging +import time +from dataclasses import dataclass +from enum import IntEnum +from functools import wraps +from typing import Dict, List, Optional + +import aiohttp +from fastapi import FastAPI +from fastapi.responses import JSONResponse +from pydantic import BaseModel + +logger = logging.getLogger('uvicorn.error') + + +class StorageItem(BaseModel): + key: str + value: Optional[str] = "" + expire_time: Optional[int] = -1 + ttl: Optional[int] = -1 + overwrite_if_exists: Optional[bool] = False + + +class WatchEventType(IntEnum): + SET = 0 + DELETE = 1 + + +@dataclass +class WatchEvent: + storage_item: StorageItem + event_type: WatchEventType + + +class WatchEventQueue: + + def __init__(self, key_prefixes: List[str], + events: asyncio.Queue[WatchEvent]): + self.key_prefixes = key_prefixes + self.events = events + + async def drain(self): + events = [] + event = await self.events.get() + logger.debug(f"Draining watch event: {self.events.qsize()}") + events.append(event) + while not self.events.empty(): + event = self.events.get_nowait() + events.append(event) + self.events.task_done() + logger.debug(f"after draining watch event: {self.events.qsize()}") + return events + + +class ClusterStorage(abc.ABC): + + def __init__(self, cluster_uri: str, cluster_name: str): + ... + + # start the storage, if it's already started, do nothing + async def start(self): + ... + + # stop the storage, if it's already stopped, do nothing + async def stop(self): + ... + + async def set(self, + key: str, + value: str, + overwrite_if_exists=False, + ttl: int = -1) -> bool: + ... + + # refresh the key’s ttl + async def expire(self, key: str, ttl: int) -> bool: + ... + + async def get(self, key: str) -> str: + ... + + async def delete(self, key: str) -> bool: + ... + + async def watch(self, key_prefix: str) -> WatchEventQueue: + ... + + async def unwatch(self, key_prefix: str) -> None: + ... + + async def get_prefix(self, + key_prefix: str, + keys_only: bool = False) -> Dict[str, str]: + ... + + +def create_cluster_storage(cluster_uri, cluster_name, **kwargs): + if cluster_uri.startswith("http"): + return HttpClusterStorageServer(cluster_uri, cluster_name, **kwargs) + elif cluster_uri.startswith("etcd"): + from tensorrt_llm.serve.cluster_storage_etcd import Etcd3ClusterStorage + return Etcd3ClusterStorage(cluster_uri, cluster_name, **kwargs) + raise ValueError(f"Invalid cluster storage URI: {cluster_uri}") + + +def create_cluster_storage_client(cluster_uri, cluster_name): + if cluster_uri.startswith("http"): + return HttpClusterStorageClient(cluster_uri, cluster_name) + elif cluster_uri.startswith("etcd"): + from tensorrt_llm.serve.cluster_storage_etcd import Etcd3ClusterStorage + return Etcd3ClusterStorage(cluster_uri, cluster_name) + raise ValueError(f"Invalid cluster storage URI: {cluster_uri}") + + +# All Http endpoints return {"result": } and status code 400 +# if result is False or None, 200 otherwise +def jsonify(f): + + @wraps(f) + async def wrapper(*args, **kwargs): + result = await f(*args, **kwargs) + return JSONResponse({"result": result}, + status_code=200 if result else 400) + + return wrapper + + +def key_time(): + return time.monotonic() + + +class HttpClusterStorageServer(ClusterStorage): + + def __init__(self, cluster_uri, cluster_name, server: FastAPI = None): + self._storage = {} + self._lock = asyncio.Lock() + self._watch_handles = {} + self._watch_lock = asyncio.Lock() + self._check_expired_task = None + if server: + self.add_routes(server) + + 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"]) + server.add_api_route("/delete", + jsonify(self.delete), + methods=["DELETE"]) + server.add_api_route("/expire", jsonify(self.expire), methods=["GET"]) + server.add_api_route("/get_prefix", + jsonify(self.get_prefix), + methods=["GET"]) + + async def start(self): + if self._check_expired_task: + return + self._check_expired_task = asyncio.create_task(self._check_expired()) + + async def stop(self): + if self._check_expired_task: + self._check_expired_task.cancel() + self._check_expired_task = None + + async def set(self, + key: str, + value: str, + overwrite_if_exists: bool = False, + ttl: int = -1) -> bool: + storage_item = StorageItem(key=key, + value=value, + overwrite_if_exists=overwrite_if_exists, + ttl=ttl) + return await self._set(storage_item) + + async def _set(self, storage_item: StorageItem) -> bool: + async with self._lock: + 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 + self._storage[storage_item.key] = storage_item + await self._notify_watch_event(storage_item.key, storage_item, + WatchEventType.SET) + return True + + 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(): + return item.value + else: + await self._notify_watch_event(key, item, + WatchEventType.DELETE) + self._storage.pop(key) + return None + + 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) + return True + return False + + async def delete(self, key: str) -> bool: + async with self._lock: + if key in self._storage: + storage_item = self._storage[key] + await self._notify_watch_event(key, storage_item, + WatchEventType.DELETE) + self._storage.pop(key) + return True + return False + + async def get_prefix(self, + key_prefix: str, + keys_only: bool = False) -> List[str]: + async with self._lock: + return { + k: "" if keys_only else v.value + for k, v in self._storage.items() if k.startswith(key_prefix) + } + + async def watch(self, key_prefix: str) -> WatchEventQueue: + async with self._watch_lock: + if key_prefix in self._watch_handles: + logger.debug( + f"Watch handle for key prefix {key_prefix} already exists, skip" + ) + else: + self._watch_handles[key_prefix] = WatchEventQueue( + key_prefixes=[key_prefix], events=asyncio.Queue()) + return self._watch_handles[key_prefix] + + async def unwatch(self, key_prefixes: List[str]) -> None: + async with self._watch_lock: + for key_prefix in key_prefixes: + if key_prefix in self._watch_handles: + self._watch_handles.pop(key_prefix) + else: + raise ValueError( + f"Key prefix {key_prefix} not in watch list") + + async def _notify_watch_event(self, key, storage_item: StorageItem, + event_type: WatchEventType): + loop = asyncio.get_event_loop() + async with self._watch_lock: + for watch_key, handle in self._watch_handles.items(): + if key.startswith(watch_key): + # update queue immediately and wake up the event loop + handle.events.put_nowait( + WatchEvent(storage_item, event_type)) + logger.info( + f"Notifying watch event for watch key {watch_key} with type {event_type}" + ) + loop._write_to_self() + logger.info( + f"Notified watch event for key {key} with type {event_type}") + + async def _check_expired(self): + while True: + await asyncio.sleep(1) + try: + before_len = len(self._storage) + async with self._lock: + key_to_delete = [ + key for key in self._storage.keys() + if self._storage[key].expire_time > 0 + and self._storage[key].expire_time < key_time() + ] + for key in key_to_delete: + await self._notify_watch_event(key, self._storage[key], + WatchEventType.DELETE) + self._storage.pop(key) + logger.debug( + f"Checked expired, {before_len} -> {len(self._storage)}, keys to delete: {key_to_delete}" + ) + except Exception as e: + logger.error(f"Error checking expired: {e}") + + +class HttpClusterStorageClient(ClusterStorage): + + def __init__(self, cluster_uri, cluster_name): + self._session = aiohttp.ClientSession(timeout=aiohttp.ClientTimeout( + total=5)) + self._cluster_uri = cluster_uri if cluster_uri.startswith( + "http") else f"http://{cluster_uri}" + self._cluster_name = cluster_name + + def __del__(self): + if asyncio.get_event_loop(): + asyncio.run_coroutine_threadsafe(self._session.close(), + asyncio.get_event_loop()) + + def _url_for(self, endpoint: str) -> str: + return f"{self._cluster_uri}/{endpoint}" + + async def _post_json(self, + endpoint: str, + data: StorageItem, + headers: dict = {}, + ignore_result: bool = False) -> bool: + headers["Content-Type"] = "application/json" + try: + async with self._session.post(self._url_for(endpoint), + headers=headers, + json=data.model_dump()) as resp: + if resp.status == 200: + json = await resp.json() + return json.get("result") if not ignore_result else True + return None + except (aiohttp.ClientError, OSError) as e: + logger.warning(f"Failed to post {endpoint}, error: {e}") + return False + + async def _get(self, + endpoint: str, + ignore_result: bool = False, + **kwargs) -> bool: + try: + async with self._session.get(self._url_for(endpoint), + params=kwargs) as resp: + if resp.status == 200: + json = await resp.json() + return json.get("result") if not ignore_result else True + return None if not ignore_result else False + except (aiohttp.ClientError, OSError) as e: + logger.warning(f"Failed to get {endpoint}, error: {e}") + return False + + async def set(self, + key: str, + value: str, + overwrite_if_exists: bool = False, + ttl: int = -1) -> bool: + storage_item = StorageItem(key=key, + value=value, + overwrite_if_exists=overwrite_if_exists, + ttl=ttl) + return await self._post_json("set", storage_item, ignore_result=True) + + async def expire(self, key: str, ttl: int) -> bool: + return await self._get("expire", + key=key, + ttl=str(ttl), + ignore_result=True) + + async def get(self, key: str) -> str: + return await self._get("get", key=key) + + async def get_prefix(self, + key_prefix: str, + keys_only: bool = False) -> Dict[str, str]: + return await self._get("get_prefix", + key_prefix=key_prefix, + keys_only=keys_only) + + async def delete(self, key: str) -> bool: + try: + async with self._session.delete(self._url_for("delete"), + params={"key": key}) as resp: + return resp.status == 200 + except (aiohttp.ClientError, OSError) as e: + logger.warning(f"Failed to delete key {key}, error: {e}") + return False + + async def watch(self, key_prefix: str) -> WatchEventQueue: + raise NotImplementedError( + "Watch functionality not implemented for HTTP client") + + async def unwatch(self, key_prefix: str) -> None: + raise NotImplementedError( + "Unwatch functionality not implemented for HTTP client") diff --git a/tests/unittest/serve/__init__.py b/tests/unittest/serve/__init__.py new file mode 100644 index 000000000000..e69de29bb2d1 diff --git a/tests/unittest/serve/test_cluster_manager_worker.py b/tests/unittest/serve/test_cluster_manager_worker.py new file mode 100644 index 000000000000..6d0f2acf4b0a --- /dev/null +++ b/tests/unittest/serve/test_cluster_manager_worker.py @@ -0,0 +1,211 @@ +import asyncio +import subprocess +import sys +import tempfile +import time + +import pytest + +from tensorrt_llm.llmapi.disagg_utils import (DisaggClusterConfig, + MinimalInstances, ServerRole) +from tensorrt_llm.serve.auto_scaling import ClusterManager, ClusterWorker +from tensorrt_llm.serve.cluster_storage import (WatchEventType, + create_cluster_storage, + create_cluster_storage_client) + +from .test_cluster_storage import (http_server_storage, pytest_async_fixture, + pytest_async_module) + + +def exception_handler(exc_t, *args, **kwargs): + print(f"Exception handler: {exc_t} {args} {kwargs}") + raise exc_t(*args, **kwargs) + + +sys.excepthook = exception_handler + +INACTIVE_TIMEOUT = 4 +HEARTBEAT_INTERVAL = 2 + + +def get_uri(storage_type): + if storage_type == "http": + return f"http://localhost:18000" + elif storage_type == "etcd": + return f"etcd://localhost:2379" + else: + raise ValueError(f"Invalid storage type: {storage_type}") + + +@pytest.fixture(scope="module") +def config(request): + cluster_uri = get_uri(request.param) + return DisaggClusterConfig(cluster_uri=cluster_uri, + cluster_name="test", + minimal_instances=MinimalInstances( + context_servers=1, generation_servers=1), + inactive_timeout=INACTIVE_TIMEOUT, + heartbeat_interval=HEARTBEAT_INTERVAL) + + +@pytest.fixture(scope="module") +def storage_server(config): + if config.cluster_uri.startswith("http"): + port = 18000 + server, cluster_storage = http_server_storage(port) + with server.run_in_thread(): + yield cluster_storage, config.cluster_uri + elif config.cluster_uri.startswith("etcd"): + with tempfile.TemporaryDirectory() as temp_dir: + etcd = subprocess.Popen( + ["etcd", "--data-dir", temp_dir, "--log-level", "debug"]) + time.sleep(2) # wait for etcd to start + yield create_cluster_storage( + config.cluster_uri, config.cluster_name), config.cluster_uri + etcd.kill() + etcd.wait() + else: + raise ValueError(f"Invalid cluster storage URI: {config.cluster_uri}") + + +@pytest_async_fixture(scope="module") +async def cluster_manager(config, storage_server): + storage, cluster_uri = storage_server + manager = ClusterManager(config, storage) + await manager.start() + yield manager + await manager.stop() + + +@pytest.mark.parametrize("config", ["http", "etcd"], indirect=True) +@pytest.mark.threadleak(enabled=False) +@pytest_async_module +async def test_init_workers_first(config, storage_server): + # init workers before initializing the manager, so the manager should be able to + # get the pre-registered workers + server, storage_uri = storage_server + storage_client = create_cluster_storage_client(storage_uri, "test") + ctx_worker = ClusterWorker(ServerRole.CONTEXT, "127.0.0.1", 8001, config, + storage_client) + gen_worker = ClusterWorker(ServerRole.GENERATION, "127.0.0.1", 8002, config, + storage_client) + await ctx_worker.register_worker() + await gen_worker.register_worker() + + cluster_manager = ClusterManager(config, server) + await cluster_manager.start() + existing_workers = await cluster_manager.watch_workers( + get_existing_first=True) + assert set([worker.worker_id for worker in existing_workers]) == { + ctx_worker.worker_id, + gen_worker.worker_id, + } + + assert await cluster_manager.is_ready() == True + await cluster_manager.stop() + + +@pytest.mark.parametrize("config", ["http", "etcd"], indirect=True) +@pytest.mark.threadleak(enabled=False) +@pytest.mark.timeout(20) +@pytest_async_module +async def test_cluster_manager(cluster_manager, storage_client, config): + cluster_manager.current_ctx_worker_num == 0 + cluster_manager.current_gen_worker_num == 0 + await cluster_manager.watch_workers() + with pytest.raises(asyncio.TimeoutError): + await asyncio.wait_for(cluster_manager.get_worker_events(), timeout=1) + assert await cluster_manager.is_ready() == False + + ctx_worker = ClusterWorker(ServerRole.CONTEXT, "127.0.0.1", 8001, config, + storage_client) + await cluster_manager.watch_workers() + await ctx_worker.register_worker() + worker_events = await cluster_manager.get_worker_events() + assert worker_events == [(ctx_worker.worker_info, WatchEventType.SET)] + assert cluster_manager.current_ctx_worker_num == 1 + assert cluster_manager.current_gen_worker_num == 0 + assert await cluster_manager.is_ready() == False + + gen_worker = ClusterWorker(ServerRole.GENERATION, "127.0.0.1", 8002, config, + storage_client) + await gen_worker.register_worker() + worker_events = await cluster_manager.get_worker_events() + assert worker_events == [(gen_worker.worker_info, WatchEventType.SET)] + assert cluster_manager.current_ctx_worker_num == 1 + assert cluster_manager.current_gen_worker_num == 1 + assert await cluster_manager.is_ready() == True + + await ctx_worker.deregister_worker() + worker_events = await cluster_manager.get_worker_events() + assert worker_events == [(ctx_worker.worker_info, WatchEventType.DELETE)] + assert cluster_manager.current_ctx_worker_num == 0 + assert cluster_manager.current_gen_worker_num == 1 + assert await cluster_manager.is_ready() == False + + await gen_worker.deregister_worker() + worker_events = await cluster_manager.get_worker_events() + assert worker_events == [(gen_worker.worker_info, WatchEventType.DELETE)] + assert cluster_manager.current_ctx_worker_num == 0 + assert cluster_manager.current_gen_worker_num == 0 + assert await cluster_manager.is_ready() == False + + +# @pytest.mark.timeout(20) +@pytest.mark.parametrize("config", ["http", "etcd"], indirect=True) +@pytest.mark.threadleak(enabled=False) +@pytest_async_module +async def test_cluster_worker(cluster_manager, storage_client, config): + + async def wait_for_worker_events(expected_new_event_num, + expected_dead_event_num): + new_worker_ids = [] + dead_workers_ids = [] + while len(new_worker_ids) < expected_new_event_num or len( + dead_workers_ids) < expected_dead_event_num: + try: + worker_events = await asyncio.wait_for( + cluster_manager.get_worker_events(), timeout=2) + new_workers = [ + worker_info.worker_id + for worker_info, event_type in worker_events + if event_type == WatchEventType.SET + ] + dead_workers = [ + worker_info.worker_id + for worker_info, event_type in worker_events + if event_type == WatchEventType.DELETE + ] + print(f"Worker events: {worker_events} {time.time()}") + new_worker_ids += new_workers + dead_workers_ids += dead_workers + except asyncio.TimeoutError: + pass + return new_worker_ids, dead_workers_ids + + await cluster_manager.start() + await cluster_manager.watch_workers() + ctx_worker = ClusterWorker(ServerRole.CONTEXT, "127.0.0.1", 8001, config, + storage_client) + gen_worker = ClusterWorker(ServerRole.GENERATION, "127.0.0.1", 8002, config, + storage_client) + + keep_heartbeat = True + assert await ctx_worker.register_worker(validator=lambda: keep_heartbeat) + assert await gen_worker.register_worker(validator=lambda: keep_heartbeat) + worker_ids = set([ctx_worker.worker_id, gen_worker.worker_id]) + new_worker_ids, dead_workers_ids = await wait_for_worker_events(2, 0) + assert set(new_worker_ids) == worker_ids + assert len(dead_workers_ids) == 0 + assert await cluster_manager.is_ready() == True + + await asyncio.sleep(config.inactive_timeout + 1) + assert await cluster_manager.is_ready() == True + + # stop heartbeat, then we should see two workers deleted + keep_heartbeat = False + new_worker_ids, dead_workers_ids = await wait_for_worker_events(0, 2) + assert len(new_worker_ids) == 0 + assert len(dead_workers_ids) == 2 + assert set(dead_workers_ids) == worker_ids + assert await cluster_manager.is_ready() == False diff --git a/tests/unittest/serve/test_cluster_storage.py b/tests/unittest/serve/test_cluster_storage.py new file mode 100644 index 000000000000..e62f85cb680e --- /dev/null +++ b/tests/unittest/serve/test_cluster_storage.py @@ -0,0 +1,212 @@ +import asyncio +import contextlib +import subprocess +import tempfile +import threading +import time + +import pytest +import pytest_asyncio +import uvicorn +from fastapi import FastAPI + +from tensorrt_llm.serve.cluster_storage import (HttpClusterStorageServer, + StorageItem, WatchEvent, + WatchEventType, + create_cluster_storage, + create_cluster_storage_client) + +pytest_async_module = pytest.mark.asyncio(loop_scope="module") +pytest_async_fixture = pytest_asyncio.fixture +pytest_ignore_tleak = pytest.mark.threadleak(enabled=False) + +_counter = 0 + + +# generate unique keys so that tests can run without affecting each other +def gen_key(prefix): + global _counter + _counter += 1 + return f"{prefix}_{_counter}" + + +class Server(uvicorn.Server): + + @contextlib.contextmanager + def run_in_thread(self): + thread = threading.Thread(target=self.run) + thread.start() + try: + while not self.started: + time.sleep(0.01) + yield + finally: + self.should_exit = True + thread.join() + + +timeout = pytest.mark.timeout + + +@pytest_async_fixture(scope="function") +async def storage_client(storage_server): + _, cluster_uri = storage_server + return create_cluster_storage_client(cluster_uri, "test") + + +# storage server client is the server itself in HTTP tests +@pytest.fixture +def storage_server_client(storage_server): + _, cluster_uri = storage_server + yield create_cluster_storage(cluster_uri, "test") + + +@pytest.mark.usefixtures("storage_client", "storage_server_client") +class TestClusterStorage: + __test__ = False + + @timeout(5) + @pytest_async_module + async def test_set(self, storage_server, storage_client): + assert await storage_client.set("test_key", + "test_value", + overwrite_if_exists=True) + assert await storage_client.get("test_key") == "test_value" + + @timeout(5) + @pytest_async_module + async def test_get(self, storage_server, storage_client): + assert await storage_client.set("test_key", + "test_value", + overwrite_if_exists=True) + assert await storage_client.get("test_key") == "test_value" + + @timeout(5) + @pytest_async_module + async def test_expire(self, storage_server, storage_client): + assert await storage_client.set("test_key", + "test_value", + overwrite_if_exists=True, + ttl=2) + assert await storage_client.get("test_key") == "test_value" + time.sleep(1) + assert await storage_client.get("test_key") == "test_value" + time.sleep(2) + assert await storage_client.get("test_key") is None + + @timeout(5) + @pytest_async_module + async def test_get_prefix(self, storage_server, storage_client): + keys = [gen_key("test_key_unique") for _ in range(3)] + for key in keys: + assert await storage_client.set(key, + "test_value1", + overwrite_if_exists=True) + + answer_keys = await storage_client.get_prefix("test_key_unique") + assert set(keys) == set(answer_keys) + answer_keys = await storage_client.get_prefix(keys[0]) + assert answer_keys == [keys[0]] + answer_keys = await storage_client.get_prefix(keys[1]) + assert answer_keys == [keys[1]] + + @pytest_ignore_tleak + @pytest_async_module + @timeout(5) + async def test_watch(self, storage_server_client, storage_client): + item1 = StorageItem(key=gen_key("test_key"), value="test_value1") + event_queue = await storage_server_client.watch("test_key") + await storage_server_client.set(key=item1.key, value=item1.value) + await asyncio.sleep(1) + watch_events = await event_queue.drain() + assert watch_events == [ + WatchEvent(storage_item=item1, event_type=WatchEventType.SET) + ] + assert await storage_server_client.get(item1.key) == item1.value + + @pytest_ignore_tleak + @pytest_async_module + @timeout(10) + async def test_watch_multiple(self, storage_server_client): + item1 = StorageItem(key=gen_key("test_key"), value="test_value1") + item2 = StorageItem(key=gen_key("test_key"), value="test_value2") + event_queue = await storage_server_client.watch("test_key") + await storage_server_client.set(key=item1.key, value=item1.value) + await storage_server_client.set(key=item2.key, value=item2.value) + await asyncio.sleep(1) + watch_events = await event_queue.drain() + assert len(watch_events) == 2 + keys = set([event.storage_item.key for event in watch_events]) + assert keys == {item1.key, item2.key} + assert set([event.event_type + for event in watch_events]) == {WatchEventType.SET} + + @pytest_ignore_tleak + @pytest_async_module + @timeout(10) + async def test_watch_set_and_delete(self, storage_server_client): + item1 = StorageItem(key=gen_key("test_key"), value="test_value1") + item2 = StorageItem(key=gen_key("test_key"), value="test_value2") + item3 = StorageItem(key=gen_key("test_key"), value="test_value3") + event_queue = await storage_server_client.watch("test_key") + await storage_server_client.set(key=item1.key, value=item1.value) + await storage_server_client.set(key=item2.key, value=item2.value) + await asyncio.sleep(1) + watch_events = await event_queue.drain() + assert len(watch_events) == 2 + assert set([event.storage_item.key + for event in watch_events]) == {item1.key, item2.key} + assert set([event.event_type + for event in watch_events]) == {WatchEventType.SET} + + event_queue = await storage_server_client.watch("test_key") + await storage_server_client.delete(item1.key) + await storage_server_client.set(key=item3.key, value=item3.value) + await asyncio.sleep(1) + watch_events = await event_queue.drain() + assert len(watch_events) == 2 + assert set([event.storage_item.key + for event in watch_events]) == {item1.key, item3.key} + assert set([event.event_type for event in watch_events + ]) == {WatchEventType.DELETE, WatchEventType.SET} + + +def http_server_storage(port): + cluster_storage = HttpClusterStorageServer("", "") + + @contextlib.asynccontextmanager + async def lifespan(app: FastAPI): + await cluster_storage.start() + yield + await cluster_storage.stop() + + app = FastAPI(lifespan=lifespan) + cluster_storage.add_routes(app) + server = Server( + uvicorn.Config(app=app, host="localhost", port=port, log_level="info")) + return server, cluster_storage + + +class TestHttpClusterStorage(TestClusterStorage): + __test__ = True + + @pytest.fixture(scope="class") + def storage_server(self): + port = 18000 + server, cluster_storage = http_server_storage(port) + with server.run_in_thread(): + yield cluster_storage, f"http://localhost:{port}" + + +class TestEtcdClusterStorage(TestClusterStorage): + __test__ = True + + @pytest.fixture(scope="class") + def storage_server(self): + with tempfile.TemporaryDirectory() as temp_dir: + self.etcd = subprocess.Popen( + ["etcd", "--data-dir", temp_dir, "--log-level", "debug"]) + time.sleep(2) # wait for etcd to start + yield self.etcd, "etcd://localhost:2379" + self.etcd.kill() + self.etcd.wait() From da42ff28780d41381bdc70dd4d29b9e6e796edfe Mon Sep 17 00:00:00 2001 From: Lizhi Zhou <1432185+reasonsolo@users.noreply.github.com> Date: Fri, 26 Sep 2025 00:06:00 -0700 Subject: [PATCH 2/3] add tests defines Signed-off-by: Lizhi Zhou <1432185+reasonsolo@users.noreply.github.com> --- tensorrt_llm/llmapi/disagg_utils.py | 15 ++ tensorrt_llm/serve/auto_scaling.py | 28 ++- tensorrt_llm/serve/cluster_storage.py | 8 +- .../test_lists/test-db/l0_h100.yml | 2 + .../{serve => disaggregated}/__init__.py | 0 .../test_cluster_manager_worker.py | 227 ++++++++++++++++++ .../test_cluster_storage.py | 24 +- .../serve/test_cluster_manager_worker.py | 211 ---------------- 8 files changed, 274 insertions(+), 241 deletions(-) rename tests/unittest/{serve => disaggregated}/__init__.py (100%) create mode 100644 tests/unittest/disaggregated/test_cluster_manager_worker.py rename tests/unittest/{serve => disaggregated}/test_cluster_storage.py (91%) delete mode 100644 tests/unittest/serve/test_cluster_manager_worker.py diff --git a/tensorrt_llm/llmapi/disagg_utils.py b/tensorrt_llm/llmapi/disagg_utils.py index cef72c687791..5122343e2c67 100644 --- a/tensorrt_llm/llmapi/disagg_utils.py +++ b/tensorrt_llm/llmapi/disagg_utils.py @@ -43,6 +43,21 @@ class ConditionalDisaggConfig(): max_local_prefill_length: int = 0 +@dataclass +class MinimalInstances: + context_servers: int = 1 + generation_servers: int = 1 + + +@dataclass +class DisaggClusterConfig: + cluster_uri: str + cluster_name: str = "" + minimal_instances: Optional[MinimalInstances] = None + heartbeat_interval: int = 5 + inactive_timeout: int = 10 + + @dataclass class DisaggServerConfig(): server_configs: List[CtxGenServerConfig] diff --git a/tensorrt_llm/serve/auto_scaling.py b/tensorrt_llm/serve/auto_scaling.py index 9d098a04dec2..ea3942801289 100644 --- a/tensorrt_llm/serve/auto_scaling.py +++ b/tensorrt_llm/serve/auto_scaling.py @@ -19,7 +19,6 @@ class WorkerInfo: host: str = "" port: int = 0 role: ServerRole = ServerRole.CONTEXT - status: str = "" def get_worker_key_prefix(cluster_name: str): @@ -105,20 +104,17 @@ async def get_worker_events( worker_events = [] for event in events: try: - print( - f"Processing event: {event.event_type} for key: {event.storage_item.key} value {event.storage_item.value}" - ) worker_info = self._parse_worker_info(event) worker_events.append((worker_info, event.event_type)) except Exception as e: logger.error( - f"Error parsing worker info: {event.storage_item.value}, error: {e}" + f"Failed to parse worker info: {event.storage_item.value}, error: {e}" ) continue return worker_events def _log_cluster_status(self, worker_info: WorkerInfo, change_event: str): - logger.error( + logger.info( f"Worker {worker_info.worker_id} becomes {change_event}, current context worker: {self.current_ctx_worker_num}/{self._minimal_ctx_worker_num}, current generation worker: {self.current_gen_worker_num}/{self._minimal_gen_worker_num}" ) @@ -138,7 +134,7 @@ def _get_workers_by_id(self, worker_id: str) -> dict[str, WorkerInfo]: else: raise ValueError(f"Worker {worker_id} is unknown") - def _parse_worker_info(self, event: WatchEvent) -> Tuple[WorkerInfo, bool]: + def _parse_worker_info(self, event: WatchEvent) -> WorkerInfo: # parse the worker info from the event, if it's a delete event, pop the corresponding worker from the current workers # if it's a set event, parse the worker info from the value and add it to the current workers # return the worker info and whether to notify the event @@ -162,6 +158,7 @@ def _parse_worker_info(self, event: WatchEvent) -> Tuple[WorkerInfo, bool]: logger.error( f"Failed to parse set event: {event.storage_item.key}: {event.storage_item.value}, error: {e}" ) + # Generate a dummy worker info with id only, router should be able to ignore it worker_info = WorkerInfo(worker_id=event.storage_item.key) else: raise ValueError(f"Invalid event type: {event.event_type}") @@ -192,6 +189,11 @@ def __init__(self, role: ServerRole, host: str, port: int, self._last_heartbeat = 0 self._worker_id = f"{role.name}-{host}:{port}-{int(time.time()*1000)}-{os.getpid()}-{random.randint(0, 1000):03}" + def __del__(self): + if asyncio.get_event_loop(): + asyncio.run_coroutine_threadsafe(self.deregister_worker(), + asyncio.get_event_loop()) + @property def worker_id(self) -> str: return self._worker_id @@ -201,8 +203,7 @@ def worker_info(self) -> WorkerInfo: return WorkerInfo(worker_id=self._worker_id, role=self._role, host=self._host, - port=self._port, - status="") + port=self._port) @property def worker_key(self) -> str: @@ -230,8 +231,8 @@ async def register_worker(self, validator=None, retry_interval=5): logger.warning( f"Worker {self.worker_info.worker_id} registration failed, retry in {retry_interval} seconds" ) - await asyncio.sleep(max(10, retry_interval)) - return await self.register_worker(validator, retry_interval + 1) + await asyncio.sleep(retry_interval) + return await self.register_worker(validator, retry_interval) else: logger.info( f"Worker {self.worker_info.worker_id} registration successful") @@ -248,8 +249,9 @@ async def register_worker(self, validator=None, retry_interval=5): async def deregister_worker(self): self._stop = True - self._heartbeat_task.cancel() - self._heartbeat_task = None + if self._heartbeat_task: + self._heartbeat_task.cancel() + self._heartbeat_task = None await self._cluster_storage.stop() success = await self._cluster_storage.delete(self.worker_key) if not success: diff --git a/tensorrt_llm/serve/cluster_storage.py b/tensorrt_llm/serve/cluster_storage.py index 6f14e96c2a9d..15e60a5f527b 100644 --- a/tensorrt_llm/serve/cluster_storage.py +++ b/tensorrt_llm/serve/cluster_storage.py @@ -99,18 +99,12 @@ async def get_prefix(self, def create_cluster_storage(cluster_uri, cluster_name, **kwargs): if cluster_uri.startswith("http"): return HttpClusterStorageServer(cluster_uri, cluster_name, **kwargs) - elif cluster_uri.startswith("etcd"): - from tensorrt_llm.serve.cluster_storage_etcd import Etcd3ClusterStorage - return Etcd3ClusterStorage(cluster_uri, cluster_name, **kwargs) raise ValueError(f"Invalid cluster storage URI: {cluster_uri}") def create_cluster_storage_client(cluster_uri, cluster_name): if cluster_uri.startswith("http"): return HttpClusterStorageClient(cluster_uri, cluster_name) - elif cluster_uri.startswith("etcd"): - from tensorrt_llm.serve.cluster_storage_etcd import Etcd3ClusterStorage - return Etcd3ClusterStorage(cluster_uri, cluster_name) raise ValueError(f"Invalid cluster storage URI: {cluster_uri}") @@ -356,7 +350,7 @@ async def get_prefix(self, keys_only: bool = False) -> Dict[str, str]: return await self._get("get_prefix", key_prefix=key_prefix, - keys_only=keys_only) + keys_only=int(keys_only)) async def delete(self, key: str) -> bool: try: diff --git a/tests/integration/test_lists/test-db/l0_h100.yml b/tests/integration/test_lists/test-db/l0_h100.yml index 8b4d3261be8b..8aecb26c439a 100644 --- a/tests/integration/test_lists/test-db/l0_h100.yml +++ b/tests/integration/test_lists/test-db/l0_h100.yml @@ -35,6 +35,8 @@ l0_h100: - unittest/disaggregated/test_disagg_utils.py - unittest/disaggregated/test_router.py - unittest/disaggregated/test_remoteDictionary.py + - unittest/disaggregated/test_cluster_manager_worker.py + - unittest/disaggregated/test_cluster_storage.py - accuracy/test_llm_api_pytorch.py::TestGemma3_1BInstruct::test_auto_dtype - accuracy/test_llm_api_pytorch.py::TestGemma3_1BInstruct::test_auto_dtype_vswa_without_reuse - accuracy/test_llm_api_pytorch.py::TestGemma3_1BInstruct::test_auto_dtype_vswa_reuse diff --git a/tests/unittest/serve/__init__.py b/tests/unittest/disaggregated/__init__.py similarity index 100% rename from tests/unittest/serve/__init__.py rename to tests/unittest/disaggregated/__init__.py diff --git a/tests/unittest/disaggregated/test_cluster_manager_worker.py b/tests/unittest/disaggregated/test_cluster_manager_worker.py new file mode 100644 index 000000000000..bd1700daf94a --- /dev/null +++ b/tests/unittest/disaggregated/test_cluster_manager_worker.py @@ -0,0 +1,227 @@ +import asyncio +import subprocess +import tempfile +import time + +import pytest + +from tensorrt_llm.llmapi.disagg_utils import (DisaggClusterConfig, + MinimalInstances, ServerRole) +from tensorrt_llm.serve.auto_scaling import ClusterManager, ClusterWorker +from tensorrt_llm.serve.cluster_storage import (WatchEventType, + create_cluster_storage, + create_cluster_storage_client) + +from .test_cluster_storage import http_server_storage, pytest_async_fixture + +INACTIVE_TIMEOUT = 4 +HEARTBEAT_INTERVAL = 2 + +storage_types = ["http"] + + +def get_uri(storage_type): + if storage_type == "http": + return f"http://localhost:18000" + elif storage_type == "etcd": + return f"etcd://localhost:2379" + else: + raise ValueError(f"Invalid storage type: {storage_type}") + + +@pytest.fixture(scope="module") +def config(request): + cluster_uri = get_uri(request.param) + return DisaggClusterConfig(cluster_uri=cluster_uri, + cluster_name="test", + minimal_instances=MinimalInstances( + context_servers=1, generation_servers=1), + inactive_timeout=INACTIVE_TIMEOUT, + heartbeat_interval=HEARTBEAT_INTERVAL) + + +@pytest.fixture(scope="module") +def storage_server(config): + if config.cluster_uri.startswith("http"): + port = 18000 + server, cluster_storage = http_server_storage(port) + with server.run_in_thread(): + yield cluster_storage, config.cluster_uri + elif config.cluster_uri.startswith("etcd"): + with tempfile.TemporaryDirectory() as temp_dir: + etcd = subprocess.Popen( + ["etcd", "--data-dir", temp_dir, "--log-level", "debug"]) + time.sleep(2) # wait for etcd to start + yield create_cluster_storage( + config.cluster_uri, config.cluster_name), config.cluster_uri + etcd.kill() + etcd.wait() + else: + raise ValueError(f"Invalid cluster storage URI: {config.cluster_uri}") + + +@pytest_async_fixture(scope="module") +async def storage_client(storage_server): + _, cluster_uri = storage_server + return create_cluster_storage_client(cluster_uri, "test") + + +@pytest_async_fixture(scope="module") +async def cluster_manager(config, storage_server): + storage, cluster_uri = storage_server + manager = ClusterManager(config, storage) + await manager.start() + yield manager + await manager.stop() + + +@pytest.mark.parametrize("config", storage_types, indirect=True) +@pytest.mark.threadleak(enabled=False) +@pytest.mark.asyncio(scope="module") +async def test_init_workers_first(config, storage_server): + try: + # init workers before initializing the manager, so the manager should be able to + # get the pre-registered workers + server, storage_uri = storage_server + storage_client = create_cluster_storage_client(storage_uri, "test") + ctx_worker = ClusterWorker(ServerRole.CONTEXT, "127.0.0.1", 8001, + config, storage_client) + gen_worker = ClusterWorker(ServerRole.GENERATION, "127.0.0.1", 8002, + config, storage_client) + await ctx_worker.register_worker() + await gen_worker.register_worker() + + cluster_manager = ClusterManager(config, server) + await cluster_manager.start() + existing_workers = await cluster_manager.watch_workers( + get_existing_first=True) + assert set([worker.worker_id for worker in existing_workers]) == { + ctx_worker.worker_id, + gen_worker.worker_id, + } + + assert await cluster_manager.is_ready() == True + finally: + await ctx_worker.deregister_worker() + await gen_worker.deregister_worker() + + +@pytest.mark.parametrize("config", storage_types, indirect=True) +@pytest.mark.threadleak(enabled=False) +@pytest.mark.timeout(20) +@pytest.mark.asyncio(scope="module") +async def test_cluster_manager(cluster_manager, storage_client, config): + try: + cluster_manager.current_ctx_worker_num == 0 + cluster_manager.current_gen_worker_num == 0 + await cluster_manager.watch_workers() + try: + await asyncio.wait_for(cluster_manager.get_worker_events(), + timeout=1) + except asyncio.TimeoutError: + pass + assert await cluster_manager.is_ready() == False + + ctx_worker = ClusterWorker(ServerRole.CONTEXT, "127.0.0.1", 8001, + config, storage_client) + await cluster_manager.watch_workers() + await ctx_worker.register_worker() + worker_events = await cluster_manager.get_worker_events() + assert worker_events == [(ctx_worker.worker_info, WatchEventType.SET)] + assert cluster_manager.current_ctx_worker_num == 1 + assert cluster_manager.current_gen_worker_num == 0 + assert await cluster_manager.is_ready() == False + + gen_worker = ClusterWorker(ServerRole.GENERATION, "127.0.0.1", 8002, + config, storage_client) + await gen_worker.register_worker() + worker_events = await cluster_manager.get_worker_events() + assert worker_events == [(gen_worker.worker_info, WatchEventType.SET)] + assert cluster_manager.current_ctx_worker_num == 1 + assert cluster_manager.current_gen_worker_num == 1 + assert await cluster_manager.is_ready() == True + + await ctx_worker.deregister_worker() + worker_events = await cluster_manager.get_worker_events() + assert worker_events == [(ctx_worker.worker_info, WatchEventType.DELETE) + ] + assert cluster_manager.current_ctx_worker_num == 0 + assert cluster_manager.current_gen_worker_num == 1 + assert await cluster_manager.is_ready() == False + + await gen_worker.deregister_worker() + worker_events = await cluster_manager.get_worker_events() + assert worker_events == [(gen_worker.worker_info, WatchEventType.DELETE) + ] + assert cluster_manager.current_ctx_worker_num == 0 + assert cluster_manager.current_gen_worker_num == 0 + assert await cluster_manager.is_ready() == False + finally: + await ctx_worker.deregister_worker() + await gen_worker.deregister_worker() + + +@pytest.mark.timeout(20) +@pytest.mark.parametrize("config", storage_types, indirect=True) +@pytest.mark.threadleak(enabled=False) +@pytest.mark.asyncio(scope="module") +async def test_cluster_worker(cluster_manager, storage_client, config): + + async def wait_for_worker_events(expected_new_event_num, + expected_dead_event_num): + new_worker_ids = [] + dead_workers_ids = [] + while len(new_worker_ids) < expected_new_event_num or len( + dead_workers_ids) < expected_dead_event_num: + try: + worker_events = await asyncio.wait_for( + cluster_manager.get_worker_events(), timeout=2) + new_workers = [ + worker_info.worker_id + for worker_info, event_type in worker_events + if event_type == WatchEventType.SET + ] + dead_workers = [ + worker_info.worker_id + for worker_info, event_type in worker_events + if event_type == WatchEventType.DELETE + ] + print(f"Worker events: {worker_events} {time.time()}") + new_worker_ids += new_workers + dead_workers_ids += dead_workers + except asyncio.TimeoutError: + pass + return new_worker_ids, dead_workers_ids + + try: + await cluster_manager.start() + await cluster_manager.watch_workers() + ctx_worker = ClusterWorker(ServerRole.CONTEXT, "127.0.0.1", 8001, + config, storage_client) + gen_worker = ClusterWorker(ServerRole.GENERATION, "127.0.0.1", 8002, + config, storage_client) + + keep_heartbeat = True + assert await ctx_worker.register_worker(validator=lambda: keep_heartbeat + ) + assert await gen_worker.register_worker(validator=lambda: keep_heartbeat + ) + worker_ids = set([ctx_worker.worker_id, gen_worker.worker_id]) + new_worker_ids, dead_workers_ids = await wait_for_worker_events(2, 0) + assert set(new_worker_ids) == worker_ids + assert len(dead_workers_ids) == 0 + assert await cluster_manager.is_ready() == True + + await asyncio.sleep(config.inactive_timeout + 1) + assert await cluster_manager.is_ready() == True + + # stop heartbeat, then we should see two workers deleted + keep_heartbeat = False + new_worker_ids, dead_workers_ids = await wait_for_worker_events(0, 2) + assert len(new_worker_ids) == 0 + assert len(dead_workers_ids) == 2 + assert set(dead_workers_ids) == worker_ids + assert await cluster_manager.is_ready() == False + finally: + await ctx_worker.deregister_worker() + await gen_worker.deregister_worker() diff --git a/tests/unittest/serve/test_cluster_storage.py b/tests/unittest/disaggregated/test_cluster_storage.py similarity index 91% rename from tests/unittest/serve/test_cluster_storage.py rename to tests/unittest/disaggregated/test_cluster_storage.py index e62f85cb680e..d2fe1facf738 100644 --- a/tests/unittest/serve/test_cluster_storage.py +++ b/tests/unittest/disaggregated/test_cluster_storage.py @@ -96,19 +96,22 @@ async def test_expire(self, storage_server, storage_client): @timeout(5) @pytest_async_module - async def test_get_prefix(self, storage_server, storage_client): + async def test_get_keys(self, storage_server, storage_client): keys = [gen_key("test_key_unique") for _ in range(3)] - for key in keys: + values = [f"test_value{i}" for i in range(3)] + for key, value in zip(keys, values): assert await storage_client.set(key, - "test_value1", + value, overwrite_if_exists=True) - answer_keys = await storage_client.get_prefix("test_key_unique") - assert set(keys) == set(answer_keys) - answer_keys = await storage_client.get_prefix(keys[0]) - assert answer_keys == [keys[0]] - answer_keys = await storage_client.get_prefix(keys[1]) - assert answer_keys == [keys[1]] + answer_keys = await storage_client.get_prefix("test_key_unique", + keys_only=False) + assert set(keys) == set(answer_keys.keys()) + assert set(values) == set(answer_keys.values()) + answer_keys = await storage_client.get_prefix(keys[0], keys_only=True) + assert answer_keys == {keys[0]: ""} + answer_keys = await storage_client.get_prefix(keys[1], keys_only=True) + assert answer_keys == {keys[1]: ""} @pytest_ignore_tleak @pytest_async_module @@ -199,7 +202,8 @@ def storage_server(self): class TestEtcdClusterStorage(TestClusterStorage): - __test__ = True + # Disable this test until Etcd functionality is ready. + __test__ = False @pytest.fixture(scope="class") def storage_server(self): diff --git a/tests/unittest/serve/test_cluster_manager_worker.py b/tests/unittest/serve/test_cluster_manager_worker.py deleted file mode 100644 index 6d0f2acf4b0a..000000000000 --- a/tests/unittest/serve/test_cluster_manager_worker.py +++ /dev/null @@ -1,211 +0,0 @@ -import asyncio -import subprocess -import sys -import tempfile -import time - -import pytest - -from tensorrt_llm.llmapi.disagg_utils import (DisaggClusterConfig, - MinimalInstances, ServerRole) -from tensorrt_llm.serve.auto_scaling import ClusterManager, ClusterWorker -from tensorrt_llm.serve.cluster_storage import (WatchEventType, - create_cluster_storage, - create_cluster_storage_client) - -from .test_cluster_storage import (http_server_storage, pytest_async_fixture, - pytest_async_module) - - -def exception_handler(exc_t, *args, **kwargs): - print(f"Exception handler: {exc_t} {args} {kwargs}") - raise exc_t(*args, **kwargs) - - -sys.excepthook = exception_handler - -INACTIVE_TIMEOUT = 4 -HEARTBEAT_INTERVAL = 2 - - -def get_uri(storage_type): - if storage_type == "http": - return f"http://localhost:18000" - elif storage_type == "etcd": - return f"etcd://localhost:2379" - else: - raise ValueError(f"Invalid storage type: {storage_type}") - - -@pytest.fixture(scope="module") -def config(request): - cluster_uri = get_uri(request.param) - return DisaggClusterConfig(cluster_uri=cluster_uri, - cluster_name="test", - minimal_instances=MinimalInstances( - context_servers=1, generation_servers=1), - inactive_timeout=INACTIVE_TIMEOUT, - heartbeat_interval=HEARTBEAT_INTERVAL) - - -@pytest.fixture(scope="module") -def storage_server(config): - if config.cluster_uri.startswith("http"): - port = 18000 - server, cluster_storage = http_server_storage(port) - with server.run_in_thread(): - yield cluster_storage, config.cluster_uri - elif config.cluster_uri.startswith("etcd"): - with tempfile.TemporaryDirectory() as temp_dir: - etcd = subprocess.Popen( - ["etcd", "--data-dir", temp_dir, "--log-level", "debug"]) - time.sleep(2) # wait for etcd to start - yield create_cluster_storage( - config.cluster_uri, config.cluster_name), config.cluster_uri - etcd.kill() - etcd.wait() - else: - raise ValueError(f"Invalid cluster storage URI: {config.cluster_uri}") - - -@pytest_async_fixture(scope="module") -async def cluster_manager(config, storage_server): - storage, cluster_uri = storage_server - manager = ClusterManager(config, storage) - await manager.start() - yield manager - await manager.stop() - - -@pytest.mark.parametrize("config", ["http", "etcd"], indirect=True) -@pytest.mark.threadleak(enabled=False) -@pytest_async_module -async def test_init_workers_first(config, storage_server): - # init workers before initializing the manager, so the manager should be able to - # get the pre-registered workers - server, storage_uri = storage_server - storage_client = create_cluster_storage_client(storage_uri, "test") - ctx_worker = ClusterWorker(ServerRole.CONTEXT, "127.0.0.1", 8001, config, - storage_client) - gen_worker = ClusterWorker(ServerRole.GENERATION, "127.0.0.1", 8002, config, - storage_client) - await ctx_worker.register_worker() - await gen_worker.register_worker() - - cluster_manager = ClusterManager(config, server) - await cluster_manager.start() - existing_workers = await cluster_manager.watch_workers( - get_existing_first=True) - assert set([worker.worker_id for worker in existing_workers]) == { - ctx_worker.worker_id, - gen_worker.worker_id, - } - - assert await cluster_manager.is_ready() == True - await cluster_manager.stop() - - -@pytest.mark.parametrize("config", ["http", "etcd"], indirect=True) -@pytest.mark.threadleak(enabled=False) -@pytest.mark.timeout(20) -@pytest_async_module -async def test_cluster_manager(cluster_manager, storage_client, config): - cluster_manager.current_ctx_worker_num == 0 - cluster_manager.current_gen_worker_num == 0 - await cluster_manager.watch_workers() - with pytest.raises(asyncio.TimeoutError): - await asyncio.wait_for(cluster_manager.get_worker_events(), timeout=1) - assert await cluster_manager.is_ready() == False - - ctx_worker = ClusterWorker(ServerRole.CONTEXT, "127.0.0.1", 8001, config, - storage_client) - await cluster_manager.watch_workers() - await ctx_worker.register_worker() - worker_events = await cluster_manager.get_worker_events() - assert worker_events == [(ctx_worker.worker_info, WatchEventType.SET)] - assert cluster_manager.current_ctx_worker_num == 1 - assert cluster_manager.current_gen_worker_num == 0 - assert await cluster_manager.is_ready() == False - - gen_worker = ClusterWorker(ServerRole.GENERATION, "127.0.0.1", 8002, config, - storage_client) - await gen_worker.register_worker() - worker_events = await cluster_manager.get_worker_events() - assert worker_events == [(gen_worker.worker_info, WatchEventType.SET)] - assert cluster_manager.current_ctx_worker_num == 1 - assert cluster_manager.current_gen_worker_num == 1 - assert await cluster_manager.is_ready() == True - - await ctx_worker.deregister_worker() - worker_events = await cluster_manager.get_worker_events() - assert worker_events == [(ctx_worker.worker_info, WatchEventType.DELETE)] - assert cluster_manager.current_ctx_worker_num == 0 - assert cluster_manager.current_gen_worker_num == 1 - assert await cluster_manager.is_ready() == False - - await gen_worker.deregister_worker() - worker_events = await cluster_manager.get_worker_events() - assert worker_events == [(gen_worker.worker_info, WatchEventType.DELETE)] - assert cluster_manager.current_ctx_worker_num == 0 - assert cluster_manager.current_gen_worker_num == 0 - assert await cluster_manager.is_ready() == False - - -# @pytest.mark.timeout(20) -@pytest.mark.parametrize("config", ["http", "etcd"], indirect=True) -@pytest.mark.threadleak(enabled=False) -@pytest_async_module -async def test_cluster_worker(cluster_manager, storage_client, config): - - async def wait_for_worker_events(expected_new_event_num, - expected_dead_event_num): - new_worker_ids = [] - dead_workers_ids = [] - while len(new_worker_ids) < expected_new_event_num or len( - dead_workers_ids) < expected_dead_event_num: - try: - worker_events = await asyncio.wait_for( - cluster_manager.get_worker_events(), timeout=2) - new_workers = [ - worker_info.worker_id - for worker_info, event_type in worker_events - if event_type == WatchEventType.SET - ] - dead_workers = [ - worker_info.worker_id - for worker_info, event_type in worker_events - if event_type == WatchEventType.DELETE - ] - print(f"Worker events: {worker_events} {time.time()}") - new_worker_ids += new_workers - dead_workers_ids += dead_workers - except asyncio.TimeoutError: - pass - return new_worker_ids, dead_workers_ids - - await cluster_manager.start() - await cluster_manager.watch_workers() - ctx_worker = ClusterWorker(ServerRole.CONTEXT, "127.0.0.1", 8001, config, - storage_client) - gen_worker = ClusterWorker(ServerRole.GENERATION, "127.0.0.1", 8002, config, - storage_client) - - keep_heartbeat = True - assert await ctx_worker.register_worker(validator=lambda: keep_heartbeat) - assert await gen_worker.register_worker(validator=lambda: keep_heartbeat) - worker_ids = set([ctx_worker.worker_id, gen_worker.worker_id]) - new_worker_ids, dead_workers_ids = await wait_for_worker_events(2, 0) - assert set(new_worker_ids) == worker_ids - assert len(dead_workers_ids) == 0 - assert await cluster_manager.is_ready() == True - - await asyncio.sleep(config.inactive_timeout + 1) - assert await cluster_manager.is_ready() == True - - # stop heartbeat, then we should see two workers deleted - keep_heartbeat = False - new_worker_ids, dead_workers_ids = await wait_for_worker_events(0, 2) - assert len(new_worker_ids) == 0 - assert len(dead_workers_ids) == 2 - assert set(dead_workers_ids) == worker_ids - assert await cluster_manager.is_ready() == False From 62c250db1605fed2ab2e2120e9c23348147263e8 Mon Sep 17 00:00:00 2001 From: Lizhi Zhou <1432185+reasonsolo@users.noreply.github.com> Date: Tue, 7 Oct 2025 19:40:59 -0700 Subject: [PATCH 3/3] fix by review comments Signed-off-by: Lizhi Zhou <1432185+reasonsolo@users.noreply.github.com> --- tensorrt_llm/llmapi/disagg_utils.py | 12 +- tensorrt_llm/serve/cluster_storage.py | 45 +++--- ...auto_scaling.py => disagg_auto_scaling.py} | 56 +++++--- .../test_lists/test-db/l0_h100.yml | 2 +- tests/unittest/disaggregated/__init__.py | 0 .../disaggregated/test_cluster_storage.py | 40 +++--- ... => test_disagg_cluster_manager_worker.py} | 129 +++++++++++------- 7 files changed, 177 insertions(+), 107 deletions(-) rename tensorrt_llm/serve/{auto_scaling.py => disagg_auto_scaling.py} (85%) delete mode 100644 tests/unittest/disaggregated/__init__.py rename tests/unittest/disaggregated/{test_cluster_manager_worker.py => test_disagg_cluster_manager_worker.py} (62%) diff --git a/tensorrt_llm/llmapi/disagg_utils.py b/tensorrt_llm/llmapi/disagg_utils.py index 5122343e2c67..e3d7a5d8dfaa 100644 --- a/tensorrt_llm/llmapi/disagg_utils.py +++ b/tensorrt_llm/llmapi/disagg_utils.py @@ -45,17 +45,17 @@ class ConditionalDisaggConfig(): @dataclass class MinimalInstances: - context_servers: int = 1 - generation_servers: int = 1 + context_servers: int = 1 # the minimal number of context servers + generation_servers: int = 1 # the minimal number of generation servers @dataclass class DisaggClusterConfig: - cluster_uri: str - cluster_name: str = "" + cluster_uri: str # the uri of the cluster storage + cluster_name: str = "" # the name of the cluster, used like a namespace minimal_instances: Optional[MinimalInstances] = None - heartbeat_interval: int = 5 - inactive_timeout: int = 10 + heartbeat_interval_sec: int = 5 # the worker will send heartbeat to the cluster storage every heartbeat_interval_sec seconds + inactive_timeout_sec: int = 10 # the worker will be considered inactive if it doesn't send heartbeat for inactive_timeout_sec seconds @dataclass diff --git a/tensorrt_llm/serve/cluster_storage.py b/tensorrt_llm/serve/cluster_storage.py index 15e60a5f527b..462c247cb232 100644 --- a/tensorrt_llm/serve/cluster_storage.py +++ b/tensorrt_llm/serve/cluster_storage.py @@ -67,6 +67,7 @@ async def start(self): async def stop(self): ... + # set the key with the value, if the key already exists and overwrite_if_exists is False, return False async def set(self, key: str, value: str, @@ -78,18 +79,24 @@ async def set(self, async def expire(self, key: str, ttl: int) -> bool: ... + # get the value of the key, return None if the key does not exist or is expired async def get(self, key: str) -> str: ... + # delete the key, return True if the key is deleted, False otherwise async def delete(self, key: str) -> bool: ... + # watch the key prefix, return the watch event queue async def watch(self, key_prefix: str) -> WatchEventQueue: ... + # unwatch the key prefix, if the key prefix is not in the watch list, raise an error async def unwatch(self, key_prefix: str) -> None: ... + # get the value of the key prefix, return the dict of key and value + # if keys_only is True, the value will be empty string async def get_prefix(self, key_prefix: str, keys_only: bool = False) -> Dict[str, str]: @@ -133,6 +140,7 @@ def __init__(self, cluster_uri, cluster_name, server: FastAPI = None): self._watch_handles = {} self._watch_lock = asyncio.Lock() self._check_expired_task = None + self._check_expired_interval = 1 # in seconds if server: self.add_routes(server) @@ -228,14 +236,14 @@ async def watch(self, key_prefix: str) -> WatchEventQueue: key_prefixes=[key_prefix], events=asyncio.Queue()) return self._watch_handles[key_prefix] - async def unwatch(self, key_prefixes: List[str]) -> None: + async def unwatch(self, key_prefix: str) -> None: async with self._watch_lock: - for key_prefix in key_prefixes: - if key_prefix in self._watch_handles: - self._watch_handles.pop(key_prefix) - else: - raise ValueError( - f"Key prefix {key_prefix} not in watch list") + if key_prefix in self._watch_handles: + self._watch_handles.pop(key_prefix) + else: + raise ValueError( + f"Key prefix {key_prefix} not in watch list, {self._watch_handles.keys()}" + ) async def _notify_watch_event(self, key, storage_item: StorageItem, event_type: WatchEventType): @@ -255,21 +263,22 @@ async def _notify_watch_event(self, key, storage_item: StorageItem, async def _check_expired(self): while True: - await asyncio.sleep(1) + await asyncio.sleep(self._check_expired_interval) try: before_len = len(self._storage) + current_time = key_time() async with self._lock: - key_to_delete = [ - key for key in self._storage.keys() - if self._storage[key].expire_time > 0 - and self._storage[key].expire_time < key_time() - ] - for key in key_to_delete: - await self._notify_watch_event(key, self._storage[key], - WatchEventType.DELETE) - self._storage.pop(key) + kv_to_delete = { + k: v + for k, v in self._storage.items() + if v.expire_time > 0 and v.expire_time < current_time + } + for k in kv_to_delete.keys(): + self._storage.pop(k) + for k, v in kv_to_delete.items(): + await self._notify_watch_event(k, v, WatchEventType.DELETE) logger.debug( - f"Checked expired, {before_len} -> {len(self._storage)}, keys to delete: {key_to_delete}" + f"Checked expired, {before_len} -> {len(self._storage)}, keys to delete: {kv_to_delete.keys()}" ) except Exception as e: logger.error(f"Error checking expired: {e}") diff --git a/tensorrt_llm/serve/auto_scaling.py b/tensorrt_llm/serve/disagg_auto_scaling.py similarity index 85% rename from tensorrt_llm/serve/auto_scaling.py rename to tensorrt_llm/serve/disagg_auto_scaling.py index ea3942801289..51c688b2b03e 100644 --- a/tensorrt_llm/serve/auto_scaling.py +++ b/tensorrt_llm/serve/disagg_auto_scaling.py @@ -4,7 +4,7 @@ import random import time from dataclasses import asdict, dataclass -from typing import List, Tuple +from typing import Any, Dict, List, Tuple from tensorrt_llm.llmapi.disagg_utils import DisaggClusterConfig, ServerRole from tensorrt_llm.logger import logger @@ -29,7 +29,11 @@ def get_worker_key(name: str, role: ServerRole, worker_id: str = "") -> str: return f"{get_worker_key_prefix(name)}/{worker_id}" -class ClusterManager: +class DisaggClusterManager: + """ + The cluster manager is responsible for managing the workers in the cluster. + It will watch the workers and notify the router when the workers are changed. + """ def __init__(self, config: DisaggClusterConfig, storage: ClusterStorage): self._config = config @@ -37,17 +41,22 @@ def __init__(self, config: DisaggClusterConfig, storage: ClusterStorage): self._lock = asyncio.Lock() self._minimal_ctx_worker_num = config.minimal_instances.context_servers self._minimal_gen_worker_num = config.minimal_instances.generation_servers - self._current_ctx_workers = {} - self._current_gen_workers = {} + self._current_ctx_workers = {} # worker_id -> WorkerInfo + self._current_gen_workers = {} # worker_id -> WorkerInfo self._watch_handle = None - async def start(self): + def __del__(self): + if asyncio.get_event_loop(): + asyncio.run_coroutine_threadsafe(self.unwatch_workers(), + asyncio.get_event_loop()) + + async def start(self) -> None: await self._cluster_storage.start() - async def stop(self): + async def stop(self) -> None: await self._cluster_storage.stop() - async def cluster_info(self) -> dict: + async def cluster_info(self) -> Dict[str, Any]: async with self._lock: return { "current_workers": { @@ -67,15 +76,15 @@ async def cluster_info(self) -> dict: } @property - def current_ctx_worker_num(self): + def current_ctx_worker_num(self) -> int: return len(self._current_ctx_workers) @property - def current_gen_worker_num(self): + def current_gen_worker_num(self) -> int: return len(self._current_gen_workers) @property - def worker_key_prefix(self): + def worker_key_prefix(self) -> str: return get_worker_key_prefix(self._config.cluster_name) async def watch_workers(self, get_existing_first: bool = True): @@ -94,12 +103,14 @@ async def watch_workers(self, get_existing_first: bool = True): self.worker_key_prefix) return workers - async def unwatch_workers(self): - await self._cluster_storage.unwatch([self.worker_key_prefix]) + async def unwatch_workers(self) -> None: + await self._cluster_storage.unwatch(self.worker_key_prefix) self._watch_handle = None async def get_worker_events( self) -> List[Tuple[WorkerInfo, WatchEventType]]: + if self._watch_handle is None: + raise ValueError("Watch handle is not initialized") events = await self._watch_handle.drain() worker_events = [] for event in events: @@ -175,7 +186,12 @@ async def is_ready_with_router(self, router_ctx_worker_num: int, return router_ctx_worker_num >= self._minimal_ctx_worker_num and router_gen_worker_num >= self._minimal_gen_worker_num -class ClusterWorker: +class DisaggClusterWorker: + """ + The cluster worker is responsible for registering and deregistering the worker to the cluster storage. + It will send heartbeat to the cluster storage every heartbeat_interval_sec seconds. + If the worker heartbeat fails, it will re-register itself. + """ def __init__(self, role: ServerRole, host: str, port: int, config: DisaggClusterConfig, storage: ClusterStorage): @@ -210,7 +226,7 @@ def worker_key(self) -> str: return get_worker_key(self._config.cluster_name, self._role, self._worker_id) - async def register_worker(self, validator=None, retry_interval=5): + async def register_worker(self, validator=None, retry_interval=5) -> bool: self._stop = False await self._cluster_storage.start() if validator and not validator(): @@ -225,7 +241,7 @@ async def register_worker(self, validator=None, retry_interval=5): success = await self._cluster_storage.set( self.worker_key, json.dumps(asdict(worker_info)), - ttl=self._config.inactive_timeout) + ttl=self._config.inactive_timeout_sec) if not success: if retry_interval > 0: logger.warning( @@ -237,17 +253,17 @@ async def register_worker(self, validator=None, retry_interval=5): logger.info( f"Worker {self.worker_info.worker_id} registration successful") self._last_heartbeat = key_time() - if self._config.heartbeat_interval > 0 and self._config.heartbeat_interval < self._config.inactive_timeout: + if self._config.heartbeat_interval_sec > 0 and self._config.heartbeat_interval_sec < self._config.inactive_timeout_sec: if not self._heartbeat_task: self._heartbeat_task = asyncio.create_task( self._heartbeat(validator)) else: logger.warning( - f"Heartbeat interval {self._config.heartbeat_interval} is not positive or less than inactive timeout {self._config.inactive_timeout}, heartbeat is disabled" + f"Heartbeat interval {self._config.heartbeat_interval_sec} is not positive or less than inactive timeout {self._config.inactive_timeout_sec}, heartbeat is disabled" ) return True - async def deregister_worker(self): + async def deregister_worker(self) -> bool: self._stop = True if self._heartbeat_task: self._heartbeat_task.cancel() @@ -262,7 +278,7 @@ async def deregister_worker(self): async def _heartbeat(self, validator=None): logger.info(f"Worker {self.worker_info.worker_id} heartbeat started") while not self._stop: - remaining_time = self._config.heartbeat_interval - ( + remaining_time = self._config.heartbeat_interval_sec - ( key_time() - self._last_heartbeat) if remaining_time > 0: await asyncio.sleep(remaining_time) @@ -273,7 +289,7 @@ async def _heartbeat(self, validator=None): ) continue expire_res = await self._cluster_storage.expire( - self.worker_key, self._config.inactive_timeout) + self.worker_key, self._config.inactive_timeout_sec) if not expire_res: logger.warning( f"Worker {self.worker_info.worker_id} heartbeat failed, re-registering {key_time()}" diff --git a/tests/integration/test_lists/test-db/l0_h100.yml b/tests/integration/test_lists/test-db/l0_h100.yml index 8aecb26c439a..78e4efd97b73 100644 --- a/tests/integration/test_lists/test-db/l0_h100.yml +++ b/tests/integration/test_lists/test-db/l0_h100.yml @@ -35,7 +35,7 @@ l0_h100: - unittest/disaggregated/test_disagg_utils.py - unittest/disaggregated/test_router.py - unittest/disaggregated/test_remoteDictionary.py - - unittest/disaggregated/test_cluster_manager_worker.py + - unittest/disaggregated/test_disagg_cluster_manager_worker.py - unittest/disaggregated/test_cluster_storage.py - accuracy/test_llm_api_pytorch.py::TestGemma3_1BInstruct::test_auto_dtype - accuracy/test_llm_api_pytorch.py::TestGemma3_1BInstruct::test_auto_dtype_vswa_without_reuse diff --git a/tests/unittest/disaggregated/__init__.py b/tests/unittest/disaggregated/__init__.py deleted file mode 100644 index e69de29bb2d1..000000000000 diff --git a/tests/unittest/disaggregated/test_cluster_storage.py b/tests/unittest/disaggregated/test_cluster_storage.py index d2fe1facf738..3bce3b4b5485 100644 --- a/tests/unittest/disaggregated/test_cluster_storage.py +++ b/tests/unittest/disaggregated/test_cluster_storage.py @@ -16,10 +16,6 @@ create_cluster_storage, create_cluster_storage_client) -pytest_async_module = pytest.mark.asyncio(loop_scope="module") -pytest_async_fixture = pytest_asyncio.fixture -pytest_ignore_tleak = pytest.mark.threadleak(enabled=False) - _counter = 0 @@ -48,7 +44,7 @@ def run_in_thread(self): timeout = pytest.mark.timeout -@pytest_async_fixture(scope="function") +@pytest_asyncio.fixture(scope="function") async def storage_client(storage_server): _, cluster_uri = storage_server return create_cluster_storage_client(cluster_uri, "test") @@ -66,15 +62,18 @@ class TestClusterStorage: __test__ = False @timeout(5) - @pytest_async_module + @pytest.mark.asyncio(loop_scope="module") async def test_set(self, storage_server, storage_client): assert await storage_client.set("test_key", "test_value", overwrite_if_exists=True) assert await storage_client.get("test_key") == "test_value" + assert not await storage_client.set( + "test_key", "test_value", overwrite_if_exists=False) + assert await storage_client.get("test_key") == "test_value" @timeout(5) - @pytest_async_module + @pytest.mark.asyncio(loop_scope="module") async def test_get(self, storage_server, storage_client): assert await storage_client.set("test_key", "test_value", @@ -82,7 +81,7 @@ async def test_get(self, storage_server, storage_client): assert await storage_client.get("test_key") == "test_value" @timeout(5) - @pytest_async_module + @pytest.mark.asyncio(loop_scope="module") async def test_expire(self, storage_server, storage_client): assert await storage_client.set("test_key", "test_value", @@ -95,8 +94,8 @@ async def test_expire(self, storage_server, storage_client): assert await storage_client.get("test_key") is None @timeout(5) - @pytest_async_module - async def test_get_keys(self, storage_server, storage_client): + @pytest.mark.asyncio(loop_scope="module") + async def test_get_prefix(self, storage_server, storage_client): keys = [gen_key("test_key_unique") for _ in range(3)] values = [f"test_value{i}" for i in range(3)] for key, value in zip(keys, values): @@ -113,8 +112,8 @@ async def test_get_keys(self, storage_server, storage_client): answer_keys = await storage_client.get_prefix(keys[1], keys_only=True) assert answer_keys == {keys[1]: ""} - @pytest_ignore_tleak - @pytest_async_module + @pytest.mark.threadleak(enabled=False) + @pytest.mark.asyncio(loop_scope="module") @timeout(5) async def test_watch(self, storage_server_client, storage_client): item1 = StorageItem(key=gen_key("test_key"), value="test_value1") @@ -127,8 +126,17 @@ async def test_watch(self, storage_server_client, storage_client): ] assert await storage_server_client.get(item1.key) == item1.value - @pytest_ignore_tleak - @pytest_async_module + @pytest.mark.threadleak(enabled=False) + @pytest.mark.asyncio(loop_scope="module") + @timeout(10) + async def test_unwatch(self, storage_server_client, storage_client): + assert await storage_server_client.watch("test_key") + await storage_server_client.unwatch("test_key") + with pytest.raises(ValueError): + await storage_server_client.unwatch("test_key") + + @pytest.mark.threadleak(enabled=False) + @pytest.mark.asyncio(loop_scope="module") @timeout(10) async def test_watch_multiple(self, storage_server_client): item1 = StorageItem(key=gen_key("test_key"), value="test_value1") @@ -144,8 +152,8 @@ async def test_watch_multiple(self, storage_server_client): assert set([event.event_type for event in watch_events]) == {WatchEventType.SET} - @pytest_ignore_tleak - @pytest_async_module + @pytest.mark.threadleak(enabled=False) + @pytest.mark.asyncio(loop_scope="module") @timeout(10) async def test_watch_set_and_delete(self, storage_server_client): item1 = StorageItem(key=gen_key("test_key"), value="test_value1") diff --git a/tests/unittest/disaggregated/test_cluster_manager_worker.py b/tests/unittest/disaggregated/test_disagg_cluster_manager_worker.py similarity index 62% rename from tests/unittest/disaggregated/test_cluster_manager_worker.py rename to tests/unittest/disaggregated/test_disagg_cluster_manager_worker.py index bd1700daf94a..979106c35f3a 100644 --- a/tests/unittest/disaggregated/test_cluster_manager_worker.py +++ b/tests/unittest/disaggregated/test_disagg_cluster_manager_worker.py @@ -4,15 +4,16 @@ import time import pytest +import pytest_asyncio +from test_cluster_storage import http_server_storage from tensorrt_llm.llmapi.disagg_utils import (DisaggClusterConfig, MinimalInstances, ServerRole) -from tensorrt_llm.serve.auto_scaling import ClusterManager, ClusterWorker from tensorrt_llm.serve.cluster_storage import (WatchEventType, create_cluster_storage, create_cluster_storage_client) - -from .test_cluster_storage import http_server_storage, pytest_async_fixture +from tensorrt_llm.serve.disagg_auto_scaling import (DisaggClusterManager, + DisaggClusterWorker) INACTIVE_TIMEOUT = 4 HEARTBEAT_INTERVAL = 2 @@ -36,8 +37,8 @@ def config(request): cluster_name="test", minimal_instances=MinimalInstances( context_servers=1, generation_servers=1), - inactive_timeout=INACTIVE_TIMEOUT, - heartbeat_interval=HEARTBEAT_INTERVAL) + inactive_timeout_sec=INACTIVE_TIMEOUT, + heartbeat_interval_sec=HEARTBEAT_INTERVAL) @pytest.fixture(scope="module") @@ -60,16 +61,16 @@ def storage_server(config): raise ValueError(f"Invalid cluster storage URI: {config.cluster_uri}") -@pytest_async_fixture(scope="module") +@pytest_asyncio.fixture(scope="function") async def storage_client(storage_server): _, cluster_uri = storage_server return create_cluster_storage_client(cluster_uri, "test") -@pytest_async_fixture(scope="module") +@pytest_asyncio.fixture(scope="function") async def cluster_manager(config, storage_server): storage, cluster_uri = storage_server - manager = ClusterManager(config, storage) + manager = DisaggClusterManager(config, storage) await manager.start() yield manager await manager.stop() @@ -84,14 +85,14 @@ async def test_init_workers_first(config, storage_server): # get the pre-registered workers server, storage_uri = storage_server storage_client = create_cluster_storage_client(storage_uri, "test") - ctx_worker = ClusterWorker(ServerRole.CONTEXT, "127.0.0.1", 8001, - config, storage_client) - gen_worker = ClusterWorker(ServerRole.GENERATION, "127.0.0.1", 8002, - config, storage_client) + ctx_worker = DisaggClusterWorker(ServerRole.CONTEXT, "127.0.0.1", 8001, + config, storage_client) + gen_worker = DisaggClusterWorker(ServerRole.GENERATION, "127.0.0.1", + 8002, config, storage_client) await ctx_worker.register_worker() await gen_worker.register_worker() - cluster_manager = ClusterManager(config, server) + cluster_manager = DisaggClusterManager(config, server) await cluster_manager.start() existing_workers = await cluster_manager.watch_workers( get_existing_first=True) @@ -106,40 +107,75 @@ async def test_init_workers_first(config, storage_server): await gen_worker.deregister_worker() +async def register_worker_and_watch(cluster_manager, storage_client, config): + assert cluster_manager.current_ctx_worker_num == 0 + assert cluster_manager.current_gen_worker_num == 0 + await cluster_manager.watch_workers() + try: + await asyncio.wait_for(cluster_manager.get_worker_events(), timeout=1) + except asyncio.TimeoutError: + pass + assert await cluster_manager.is_ready() == False + + ctx_worker = DisaggClusterWorker(ServerRole.CONTEXT, "127.0.0.1", 8001, + config, storage_client) + await cluster_manager.watch_workers() + await ctx_worker.register_worker() + worker_events = await cluster_manager.get_worker_events() + assert worker_events == [(ctx_worker.worker_info, WatchEventType.SET)] + assert cluster_manager.current_ctx_worker_num == 1 + assert cluster_manager.current_gen_worker_num == 0 + assert await cluster_manager.is_ready() == False + + gen_worker = DisaggClusterWorker(ServerRole.GENERATION, "127.0.0.1", 8002, + config, storage_client) + await gen_worker.register_worker() + worker_events = await cluster_manager.get_worker_events() + assert worker_events == [(gen_worker.worker_info, WatchEventType.SET)] + assert cluster_manager.current_ctx_worker_num == 1 + assert cluster_manager.current_gen_worker_num == 1 + assert await cluster_manager.is_ready() == True + return ctx_worker, gen_worker + + @pytest.mark.parametrize("config", storage_types, indirect=True) @pytest.mark.threadleak(enabled=False) @pytest.mark.timeout(20) @pytest.mark.asyncio(scope="module") -async def test_cluster_manager(cluster_manager, storage_client, config): +async def test_watch_workers(cluster_manager, storage_client, config): try: - cluster_manager.current_ctx_worker_num == 0 - cluster_manager.current_gen_worker_num == 0 - await cluster_manager.watch_workers() - try: - await asyncio.wait_for(cluster_manager.get_worker_events(), - timeout=1) - except asyncio.TimeoutError: - pass - assert await cluster_manager.is_ready() == False + ctx_worker, gen_worker = await register_worker_and_watch( + cluster_manager, storage_client, config) + finally: + await ctx_worker.deregister_worker() + await gen_worker.deregister_worker() - ctx_worker = ClusterWorker(ServerRole.CONTEXT, "127.0.0.1", 8001, - config, storage_client) - await cluster_manager.watch_workers() - await ctx_worker.register_worker() - worker_events = await cluster_manager.get_worker_events() - assert worker_events == [(ctx_worker.worker_info, WatchEventType.SET)] - assert cluster_manager.current_ctx_worker_num == 1 - assert cluster_manager.current_gen_worker_num == 0 - assert await cluster_manager.is_ready() == False - gen_worker = ClusterWorker(ServerRole.GENERATION, "127.0.0.1", 8002, - config, storage_client) - await gen_worker.register_worker() - worker_events = await cluster_manager.get_worker_events() - assert worker_events == [(gen_worker.worker_info, WatchEventType.SET)] - assert cluster_manager.current_ctx_worker_num == 1 - assert cluster_manager.current_gen_worker_num == 1 - assert await cluster_manager.is_ready() == True +@pytest.mark.parametrize("config", storage_types, indirect=True) +@pytest.mark.threadleak(enabled=False) +@pytest.mark.timeout(20) +@pytest.mark.asyncio(scope="module") +async def test_unwatch_workers(cluster_manager, storage_client, config): + try: + ctx_worker, gen_worker = await register_worker_and_watch( + cluster_manager, storage_client, config) + await cluster_manager.unwatch_workers() + with pytest.raises(ValueError): + await cluster_manager.get_worker_events() + finally: + await ctx_worker.deregister_worker() + await gen_worker.deregister_worker() + + +@pytest.mark.parametrize("config", storage_types, indirect=True) +@pytest.mark.threadleak(enabled=False) +@pytest.mark.timeout(20) +@pytest.mark.asyncio(scope="module") +async def test_watch_register_then_deregister(cluster_manager, storage_client, + config): + try: + ctx_worker, gen_worker = await register_worker_and_watch( + cluster_manager, storage_client, config) await ctx_worker.deregister_worker() worker_events = await cluster_manager.get_worker_events() @@ -165,7 +201,8 @@ async def test_cluster_manager(cluster_manager, storage_client, config): @pytest.mark.parametrize("config", storage_types, indirect=True) @pytest.mark.threadleak(enabled=False) @pytest.mark.asyncio(scope="module") -async def test_cluster_worker(cluster_manager, storage_client, config): +async def test_cluster_worker_heartbeat(cluster_manager, storage_client, + config): async def wait_for_worker_events(expected_new_event_num, expected_dead_event_num): @@ -196,10 +233,10 @@ async def wait_for_worker_events(expected_new_event_num, try: await cluster_manager.start() await cluster_manager.watch_workers() - ctx_worker = ClusterWorker(ServerRole.CONTEXT, "127.0.0.1", 8001, - config, storage_client) - gen_worker = ClusterWorker(ServerRole.GENERATION, "127.0.0.1", 8002, - config, storage_client) + ctx_worker = DisaggClusterWorker(ServerRole.CONTEXT, "127.0.0.1", 8001, + config, storage_client) + gen_worker = DisaggClusterWorker(ServerRole.GENERATION, "127.0.0.1", + 8002, config, storage_client) keep_heartbeat = True assert await ctx_worker.register_worker(validator=lambda: keep_heartbeat @@ -212,7 +249,7 @@ async def wait_for_worker_events(expected_new_event_num, assert len(dead_workers_ids) == 0 assert await cluster_manager.is_ready() == True - await asyncio.sleep(config.inactive_timeout + 1) + await asyncio.sleep(config.inactive_timeout_sec + 1) assert await cluster_manager.is_ready() == True # stop heartbeat, then we should see two workers deleted