diff --git a/python/sglang/srt/managers/cache_controller.py b/python/sglang/srt/managers/cache_controller.py index 3d4905e50..a5959e9ee 100644 --- a/python/sglang/srt/managers/cache_controller.py +++ b/python/sglang/srt/managers/cache_controller.py @@ -22,12 +22,10 @@ from typing import TYPE_CHECKING, List, NamedTuple, Optional import torch +from sglang.srt.mem_cache.hicache_storage import HiCacheStorageConfig + if TYPE_CHECKING: from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator - from sglang.srt.mem_cache.hicache_storage import ( - HiCacheStorageConfig, - HiCacheStorageExtraInfo, - ) from sglang.srt.mem_cache.memory_pool_host import HostKVCache from sglang.srt.distributed import ( diff --git a/test/registered/unit/managers/test_hicache_controller_cp.py b/test/registered/unit/managers/test_hicache_controller_cp.py index e122fa051..cbaa1c381 100644 --- a/test/registered/unit/managers/test_hicache_controller_cp.py +++ b/test/registered/unit/managers/test_hicache_controller_cp.py @@ -216,6 +216,36 @@ class TestHiCacheControllerCPWrite(CustomTestCase): self.assertEqual(result.required_host_slots, 4) + def test_generate_storage_config_constructs_config_at_runtime(self): + controller = HiCacheController.__new__(HiCacheController) + controller.mem_pool_device = FakeDevicePool() + controller.mem_pool_host = FakeHostPool(torch.tensor([], dtype=torch.int64)) + controller.pp_rank = 1 + controller.pp_size = 2 + controller.enable_storage_metrics = True + + with patch( + "sglang.srt.managers.cache_controller.is_dp_attention_enabled", + return_value=False, + ), patch( + "sglang.srt.managers.cache_controller.get_tensor_model_parallel_rank", + return_value=3, + ), patch( + "sglang.srt.managers.cache_controller.get_tensor_model_parallel_world_size", + return_value=4, + ): + config = controller._generate_storage_config( + model_name="test-model", + storage_backend_extra_config={"tp_lcm_size": 8}, + ) + + self.assertEqual(config.tp_rank, 3) + self.assertEqual(config.tp_size, 4) + self.assertEqual(config.pp_rank, 1) + self.assertEqual(config.pp_size, 2) + self.assertEqual(config.model_name, "test-model") + self.assertEqual(config.tp_lcm_size, 8) + class TestHiCacheControllerCPLoad(TestHiCacheControllerCPWrite): def test_cp_load_allocates_full_logical_locs_and_transfers_owned_physical_locs(self):