diff --git a/tensorrt_llm/llmapi/disagg_utils.py b/tensorrt_llm/llmapi/disagg_utils.py index 8404cbaf7ad3..e3d7a5d8dfaa 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 @@ -43,6 +43,21 @@ class ConditionalDisaggConfig(): max_local_prefill_length: int = 0 +@dataclass +class MinimalInstances: + 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 # 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_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 class DisaggServerConfig(): server_configs: List[CtxGenServerConfig] diff --git a/tensorrt_llm/serve/cluster_storage.py b/tensorrt_llm/serve/cluster_storage.py new file mode 100644 index 000000000000..462c247cb232 --- /dev/null +++ b/tensorrt_llm/serve/cluster_storage.py @@ -0,0 +1,379 @@ +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): + ... + + # 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, + overwrite_if_exists=False, + ttl: int = -1) -> bool: + ... + + # refresh the key’s ttl + 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]: + ... + + +def create_cluster_storage(cluster_uri, cluster_name, **kwargs): + if cluster_uri.startswith("http"): + return HttpClusterStorageServer(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) + 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 + self._check_expired_interval = 1 # in seconds + 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_prefix: str) -> None: + async with self._watch_lock: + 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): + 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(self._check_expired_interval) + try: + before_len = len(self._storage) + current_time = key_time() + async with self._lock: + 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: {kv_to_delete.keys()}" + ) + 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=int(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/tensorrt_llm/serve/disagg_auto_scaling.py b/tensorrt_llm/serve/disagg_auto_scaling.py new file mode 100644 index 000000000000..51c688b2b03e --- /dev/null +++ b/tensorrt_llm/serve/disagg_auto_scaling.py @@ -0,0 +1,301 @@ +import asyncio +import json +import os +import random +import time +from dataclasses import asdict, dataclass +from typing import Any, Dict, 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 + + +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 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 + 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 = {} # worker_id -> WorkerInfo + self._current_gen_workers = {} # worker_id -> WorkerInfo + self._watch_handle = None + + 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) -> None: + await self._cluster_storage.stop() + + async def cluster_info(self) -> Dict[str, Any]: + 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) -> int: + return len(self._current_ctx_workers) + + @property + def current_gen_worker_num(self) -> int: + return len(self._current_gen_workers) + + @property + 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): + 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) -> 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: + try: + worker_info = self._parse_worker_info(event) + worker_events.append((worker_info, event.event_type)) + except Exception as e: + logger.error( + 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.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}" + ) + + 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) -> 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 + 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}" + ) + # 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}") + 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 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): + 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}" + + 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 + + @property + def worker_info(self) -> WorkerInfo: + return WorkerInfo(worker_id=self._worker_id, + role=self._role, + host=self._host, + port=self._port) + + @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) -> bool: + 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_sec) + 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(retry_interval) + return await self.register_worker(validator, retry_interval) + else: + logger.info( + f"Worker {self.worker_info.worker_id} registration successful") + self._last_heartbeat = key_time() + 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_sec} is not positive or less than inactive timeout {self._config.inactive_timeout_sec}, heartbeat is disabled" + ) + return True + + async def deregister_worker(self) -> bool: + self._stop = True + 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: + 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_sec - ( + 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_sec) + 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/tests/integration/test_lists/test-db/l0_h100.yml b/tests/integration/test_lists/test-db/l0_h100.yml index 8b4d3261be8b..78e4efd97b73 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_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 - accuracy/test_llm_api_pytorch.py::TestGemma3_1BInstruct::test_auto_dtype_vswa_reuse diff --git a/tests/unittest/disaggregated/test_cluster_storage.py b/tests/unittest/disaggregated/test_cluster_storage.py new file mode 100644 index 000000000000..3bce3b4b5485 --- /dev/null +++ b/tests/unittest/disaggregated/test_cluster_storage.py @@ -0,0 +1,224 @@ +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) + +_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_asyncio.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.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.mark.asyncio(loop_scope="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.mark.asyncio(loop_scope="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.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): + assert await storage_client.set(key, + value, + overwrite_if_exists=True) + + 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.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") + 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.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") + 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.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") + 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): + # Disable this test until Etcd functionality is ready. + __test__ = False + + @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() diff --git a/tests/unittest/disaggregated/test_disagg_cluster_manager_worker.py b/tests/unittest/disaggregated/test_disagg_cluster_manager_worker.py new file mode 100644 index 000000000000..979106c35f3a --- /dev/null +++ b/tests/unittest/disaggregated/test_disagg_cluster_manager_worker.py @@ -0,0 +1,264 @@ +import asyncio +import subprocess +import tempfile +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.cluster_storage import (WatchEventType, + create_cluster_storage, + create_cluster_storage_client) +from tensorrt_llm.serve.disagg_auto_scaling import (DisaggClusterManager, + DisaggClusterWorker) + +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_sec=INACTIVE_TIMEOUT, + heartbeat_interval_sec=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_asyncio.fixture(scope="function") +async def storage_client(storage_server): + _, cluster_uri = storage_server + return create_cluster_storage_client(cluster_uri, "test") + + +@pytest_asyncio.fixture(scope="function") +async def cluster_manager(config, storage_server): + storage, cluster_uri = storage_server + manager = DisaggClusterManager(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 = 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 = DisaggClusterManager(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() + + +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_watch_workers(cluster_manager, storage_client, config): + try: + 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() + + +@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() + 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_heartbeat(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 = 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 + ) + 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_sec + 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()