diff --git a/python/sglang/srt/disaggregation/encode_server.py b/python/sglang/srt/disaggregation/encode_server.py index 2ebbbbff8..ef9e6017a 100644 --- a/python/sglang/srt/disaggregation/encode_server.py +++ b/python/sglang/srt/disaggregation/encode_server.py @@ -901,8 +901,9 @@ class MMEncoder: mm_feature, mm_inputs, missing_indices, modality, get_feature_fn ) - # Step 3: Rank 0 prefetches cache-hit embeddings from global cache. - prefetch_status = torch.tensor([1], dtype=torch.int32) + # Step 3: Rank 0 prefetches cache-hit embeddings and builds fallback_mask. + fallback_mask = torch.zeros(num_items, dtype=torch.int32) + cached_slices = [] if self.rank == 0: if hit_indices: @@ -919,32 +920,46 @@ class MMEncoder: await asyncio.sleep(0.005) await asyncio.wait_for(_wait_prefetch(), timeout=60.0) + + # Prefetch IO completed; check which items actually loaded. + cached_slices = self.mm_global_cache.get_embeddings(hit_hashes) + for i, idx in enumerate(hit_indices): + if cached_slices[i] is None: + fallback_mask[idx] = 1 + num_partial_fail = int(fallback_mask.sum().item()) + if num_partial_fail > 0: + logger.warning( + f"Req {req_id}: {num_partial_fail}/{len(hit_indices)} " + f"cache-hit items failed to load (pool full), " + f"falling back to ViT" + ) except (asyncio.TimeoutError, Exception) as e: logger.error( f"Prefetch failed for req {req_id}: {e}. " f"Falling back to ViT for {len(hit_indices)} hit items." ) - prefetch_status[0] = 0 + for idx in hit_indices: + fallback_mask[idx] = 1 - # Step 4: Broadcast prefetch result to all ranks so they stay in sync. + # Step 4: Broadcast fallback_mask to all ranks so they stay in sync. if self.server_args.tp_size > 1: torch.distributed.broadcast( - prefetch_status, + fallback_mask, src=0, group=self.mm_global_cache.prefetch_tp_group, ) - # Step 5: If prefetch failed, all ranks fallback to ViT for the hit mm items. - if prefetch_status.item() == 0 and hit_indices: + # Step 5: All ranks run ViT for items that need fallback recomputation. + fallback_indices = [i for i in range(num_items) if fallback_mask[i].item() == 1] + fallback_slices = None + if fallback_indices: logger.info( - f"Req {req_id}: Prefetch failed, all ranks running ViT fallback " - f"for {len(hit_indices)} mm items." + f"Req {req_id}: All ranks running ViT fallback " + f"for {len(fallback_indices)} items." ) fallback_slices = self._encode_missing( - mm_feature, mm_inputs, hit_indices, modality, get_feature_fn + mm_feature, mm_inputs, fallback_indices, modality, get_feature_fn ) - else: - fallback_slices = None # Step 6: Rank 0 assembles final embedding and prepares for sending. if self.rank == 0: @@ -953,25 +968,37 @@ class MMEncoder: for i, idx in enumerate(missing_indices): final_slices[idx] = new_slices[i] - # Fill in cache-hit embeddings (from prefetch or fallback) - if prefetch_status.item() == 1 and hit_indices: - cached_slices = self.mm_global_cache.get_embeddings( - [str_mm_hashes[i] for i in hit_indices] - ) - for i, idx in enumerate(hit_indices): - final_slices[idx] = cached_slices[i] - elif fallback_slices is not None: + # Fill in successfully loaded cache-hit embeddings + if cached_slices: for i, idx in enumerate(hit_indices): + if cached_slices[i] is not None: + final_slices[idx] = cached_slices[i] + + # Fill in ViT fallback results for failed items + if fallback_slices is not None: + for i, idx in enumerate(fallback_indices): final_slices[idx] = fallback_slices[i] mm_embedding = torch.cat(final_slices, dim=0) + # Release embedding cache references now that torch.cat has + # copied the data into a new tensor. This allows the cache + # entries to be evicted under memory pressure. + if cached_slices: + loaded_hashes = [ + str_mm_hashes[idx] + for idx in hit_indices + if fallback_mask[idx].item() == 0 + ] + if loaded_hashes: + self.mm_global_cache.release_embeddings(loaded_hashes) + # Background insert: store newly computed embeddings into global cache. # Includes both original misses and fallback-recomputed hits. all_new_hashes = [str_mm_hashes[i] for i in missing_indices] all_new_slices = list(new_slices) if fallback_slices is not None: - all_new_hashes += [str_mm_hashes[i] for i in hit_indices] + all_new_hashes += [str_mm_hashes[i] for i in fallback_indices] all_new_slices += list(fallback_slices) if all_new_hashes: diff --git a/python/sglang/srt/mem_cache/storage/mooncake_store/embedding_cache_controller.py b/python/sglang/srt/mem_cache/storage/mooncake_store/embedding_cache_controller.py index 1f834b66b..98dcf37f1 100644 --- a/python/sglang/srt/mem_cache/storage/mooncake_store/embedding_cache_controller.py +++ b/python/sglang/srt/mem_cache/storage/mooncake_store/embedding_cache_controller.py @@ -24,7 +24,8 @@ class ContiguousMemoryAllocator: self.total_size = total_size_bytes # List of (offset, size) for free blocks self.free_blocks = [(0, total_size_bytes)] - self.allocated_map = {} # {handle: (offset, size)} + self.allocated_map = {} # {offset: size_bytes} + self.allocated_size = 0 # Running counter for O(1) get_allocated_size self.lock = threading.Lock() def allocate(self, size_bytes: int) -> Optional[int]: @@ -38,11 +39,18 @@ class ContiguousMemoryAllocator: self.free_blocks[i] = (offset + size_bytes, remaining_size) else: self.free_blocks.pop(i) + self.allocated_map[offset] = size_bytes + self.allocated_size += size_bytes return offset return None def free(self, offset: int, size_bytes: int): with self.lock: + # Remove from allocated map and update counter + if offset in self.allocated_map: + self.allocated_size -= self.allocated_map[offset] + del self.allocated_map[offset] + # Return block and merge adjacent free blocks self.free_blocks.append((offset, size_bytes)) self.free_blocks.sort() @@ -61,6 +69,16 @@ class ContiguousMemoryAllocator: merged.append((curr_offset, curr_size)) self.free_blocks = merged + def get_allocated_size(self) -> int: + """Return total allocated bytes. O(1) operation.""" + with self.lock: + return self.allocated_size + + def get_free_size(self) -> int: + """Return total free bytes.""" + with self.lock: + return sum(block_size for _, block_size in self.free_blocks) + class EmbeddingPrefetchOperation: """Groups all missing images of a request for a single batch GET.""" @@ -98,12 +116,17 @@ class EmbeddingCacheController: hidden_dims: dict = None, tp_group=None, all_rank_get=False, + enable_eviction: bool = True, + max_eviction_batch: int = 100, ): self.tp_world_size = tp_size self.tp_group = tp_group + self.tp_rank = tp_rank self.all_rank_get = all_rank_get self.hidden_dims = hidden_dims or {} self.element_size = torch.float32.itemsize + self.enable_eviction = enable_eviction + self.max_eviction_batch = max_eviction_batch # 1. Mooncake Backend & Pinned Buffer self.mooncake_store = MooncakeEmbeddingStore() @@ -115,10 +138,31 @@ class EmbeddingCacheController: # 2. Variable Size Memory Management self.allocator = ContiguousMemoryAllocator(self.total_pool_size_bytes) - # {hash: (offset, num_tokens, embedding_dim, size_bytes)} + # {hash: (offset, num_tokens, embedding_dim, size_bytes, last_access_time)} self.hash_to_metadata = {} - # 3. Task Tracking + # 3. LRU Tracking + # OrderedDict maintains insertion order, used as LRU cache + # hash -> access_time + self.access_order = {} + self.access_lock = threading.Lock() + + # 4. RDMA / read reference counting + # {hash: ref_count} — entries with ref_count > 0 cannot be evicted. + # Incremented when an RDMA transfer (GET/PUT) is in flight or when + # get_embeddings() returns a view into cpu_pool. Decremented after + # the RDMA completes or the caller releases the view. + self.ref_counts = {} + + # 5. Statistics + self.stats = { + "total_allocated": 0, + "total_evicted": 0, + "eviction_count": 0, + "allocation_failures": 0, + } + + # 6. Task Tracking self.ongoing_prefetch = {} # {req_id: EmbeddingPrefetchOperation} self.prefetch_queue = Queue() self.insert_queue = Queue() @@ -142,6 +186,156 @@ class EmbeddingCacheController: else: self.prefetch_tp_group = None + def _update_access_time(self, image_hash: str): + """Update LRU access time for a hash.""" + with self.access_lock: + # Move to end (most recently used) + if image_hash in self.access_order: + del self.access_order[image_hash] + self.access_order[image_hash] = time.time() + + def _protect_hash(self, image_hash: str): + """Increment ref count to prevent eviction during RDMA or active read. + + NOTE: Caller must hold self.lock. + """ + self.ref_counts[image_hash] = self.ref_counts.get(image_hash, 0) + 1 + + def _release_hash(self, image_hash: str): + """Decrement ref count after RDMA completes or caller releases a view. + + NOTE: Caller must hold self.lock. + """ + if image_hash in self.ref_counts: + self.ref_counts[image_hash] -= 1 + if self.ref_counts[image_hash] <= 0: + del self.ref_counts[image_hash] + + def _select_eviction_candidates(self, required_bytes: int) -> List[str]: + """Select LRU candidates to free up at least required_bytes. + + NOTE: Caller must hold self.lock before calling this method. + """ + candidates = [] + freed_bytes = 0 + + with self.access_lock: + # Sort by access time (oldest first) + # Python dicts are insertion-ordered; the first keys are the oldest. + sorted_hashes = list(self.access_order.items()) + + for image_hash, _ in sorted_hashes: + if image_hash not in self.hash_to_metadata: + with self.access_lock: + self.access_order.pop(image_hash, None) + continue + if self.ref_counts.get(image_hash, 0) > 0: + continue + metadata = self.hash_to_metadata[image_hash] + size_bytes = metadata[3] if len(metadata) > 3 else 0 + candidates.append(image_hash) + freed_bytes += size_bytes + + if freed_bytes >= required_bytes: + break + + if len(candidates) >= self.max_eviction_batch: + break + + return candidates + + def _evict_hashes(self, hashes_to_evict: List[str]) -> int: + """Evict specified hashes and free their memory. Returns freed bytes. + + NOTE: Caller must hold self.lock before calling this method. + """ + total_freed = 0 + + # NOTE: self.lock should be held by the caller (e.g., insert_batch, + # prefetch). Do NOT acquire it here to avoid reentrant deadlock. + for image_hash in hashes_to_evict: + if image_hash not in self.hash_to_metadata: + continue + + # Safety check: skip entries with in-flight RDMA or active reads + if self.ref_counts.get(image_hash, 0) > 0: + logger.warning( + f"[Rank {self.tp_rank}] Skipping eviction of {image_hash}: " + f"ref_count={self.ref_counts[image_hash]} (in-flight RDMA or active read)" + ) + continue + + offset, num_tokens, dim, size_bytes = self.hash_to_metadata[image_hash][:4] + + # Free memory in allocator + self.allocator.free(offset, size_bytes) + + # Remove from metadata and ref counts + del self.hash_to_metadata[image_hash] + self.ref_counts.pop(image_hash, None) + + # Remove from access order + with self.access_lock: + self.access_order.pop(image_hash, None) + + total_freed += size_bytes + self.stats["total_evicted"] += size_bytes + + if total_freed > 0: + self.stats["eviction_count"] += 1 + + if total_freed > 0: + logger.info( + f"[Rank {self.tp_rank}] Evicted {len(hashes_to_evict)} embeddings, " + f"freed {total_freed / 1024**2:.2f} MB" + ) + + return total_freed + + def _allocate_with_eviction(self, size_bytes: int) -> Optional[int]: + """Try to allocate memory, evicting old entries if necessary.""" + # First try direct allocation + offset = self.allocator.allocate(size_bytes) + if offset is not None: + self.stats["total_allocated"] += size_bytes + return offset + + # If failed and eviction is enabled, try eviction + if not self.enable_eviction: + self.stats["allocation_failures"] += 1 + return None + + # Select candidates to evict + candidates = self._select_eviction_candidates(size_bytes) + if not candidates: + n_protected = sum(1 for v in self.ref_counts.values() if v > 0) + logger.warning( + f"[Rank {self.tp_rank}] Cannot allocate {size_bytes / 1024**2:.2f} MB: " + f"pool full ({self.allocator.get_allocated_size() / 1024**2:.1f}/" + f"{self.total_pool_size_bytes / 1024**2:.1f} MB used), " + f"no evictable candidates " + f"({len(self.hash_to_metadata)} entries, {n_protected} protected)" + ) + self.stats["allocation_failures"] += 1 + return None + + # Evict and try again + freed = self._evict_hashes(candidates) + if freed < size_bytes: + logger.warning( + f"[Rank {self.tp_rank}] Could not free enough memory: " + f"needed {size_bytes / 1024**2:.2f} MB, freed {freed / 1024**2:.2f} MB" + ) + + # Try allocation again after eviction + offset = self.allocator.allocate(size_bytes) + if offset is not None: + self.stats["total_allocated"] += size_bytes + else: + self.stats["allocation_failures"] += 1 + + return offset + def prefetch( self, req_id: str, @@ -161,17 +355,25 @@ class EmbeddingCacheController: with self.lock: for h, num_tokens in zip(image_hashes, expected_tokens): if h in self.hash_to_metadata: + # Update access time for LRU + self._update_access_time(h) logger.debug( - f"Req {req_id}: Hash already in local metadata, skipping prefetch." + f"Req {req_id}: Hash already in local metadata, skipping prefetch." ) continue size_bytes = num_tokens * dim * self.element_size - offset = self.allocator.allocate(size_bytes) + offset = self._allocate_with_eviction(size_bytes) if offset is None: + logger.warning( + f"Req {req_id}: Failed to allocate {size_bytes / 1024**2:.2f} MB " + f"for prefetch, skipping this image." + ) continue self.hash_to_metadata[h] = (offset, num_tokens, dim, size_bytes) + self._update_access_time(h) + self._protect_hash(h) keys.append(h) ptrs.append(self.cpu_pool.data_ptr() + offset) sizes.append(size_bytes) @@ -204,6 +406,9 @@ class EmbeddingCacheController: with self.lock: for h, tensor in zip(image_hashes, embedding_tensors): if h in self.hash_to_metadata: + # Update access time for existing entry + self._update_access_time(h) + self._protect_hash(h) # Local cache hit: ensure Mooncake has it offset, num_tokens, dim, size_bytes = self.hash_to_metadata[h][:4] @@ -218,8 +423,12 @@ class EmbeddingCacheController: # Local cache miss: allocate and copy num_tokens, dim = tensor.shape[0], tensor.shape[1] size_bytes = num_tokens * dim * self.element_size - offset = self.allocator.allocate(size_bytes) + offset = self._allocate_with_eviction(size_bytes) if offset is None: + logger.warning( + f"Failed to allocate {size_bytes / 1024**2:.2f} MB for insert, " + f"skipping this embedding." + ) skipped_count += 1 continue @@ -231,6 +440,8 @@ class EmbeddingCacheController: ) target_view.copy_(tensor.cpu()) self.hash_to_metadata[h] = (offset, num_tokens, dim, size_bytes) + self._update_access_time(h) + self._protect_hash(h) keys.append(h) ptrs.append(self.cpu_pool.data_ptr() + offset) @@ -258,6 +469,10 @@ class EmbeddingCacheController: f"Mooncake GET Finished: Req {op.req_id}, Successfully fetched {success_count}/{len(op.keys)} images." ) op.mark_done(all(results)) + # Release ref counts now that RDMA GET is complete + with self.lock: + for h in op.keys: + self._release_hash(h) self.prefetch_queue.task_done() processed_any = True except Empty: @@ -269,6 +484,10 @@ class EmbeddingCacheController: logger.info( f"Mooncake PUT Finished: Successfully stored {len(op.keys)} keys in cluster." ) + # Release ref counts now that RDMA PUT is complete + with self.lock: + for h in op.keys: + self._release_hash(h) self.insert_queue.task_done() processed_any = True except Empty: @@ -306,11 +525,24 @@ class EmbeddingCacheController: return False def get_embeddings(self, image_hashes: List[str]) -> List[torch.Tensor]: - """Final reconstruction for model input.""" + """Final reconstruction for model input. + + Returns views into the pinned cpu_pool. Callers MUST call + release_embeddings() once they no longer need the returned + tensors (e.g. after .to(device) or torch.cat) so that the + entries can be evicted. + """ with self.lock: tensors = [] for h in image_hashes: - offset, num_tokens, dim, size_bytes = self.hash_to_metadata[h] + if h not in self.hash_to_metadata: + logger.warning(f"Hash {h} not found in local cache") + tensors.append(None) + continue + # Update access time for LRU + self._update_access_time(h) + self._protect_hash(h) + offset, num_tokens, dim, size_bytes = self.hash_to_metadata[h][:4] tensors.append( self.cpu_pool[offset : offset + size_bytes] .view(torch.float32) @@ -318,6 +550,29 @@ class EmbeddingCacheController: ) return tensors + def release_embeddings(self, image_hashes: List[str]): + """Release reference counts on embeddings after the caller is done. + + Must be called once for every successful get_embeddings() call, + after the caller no longer needs the returned tensor views + (e.g. after .to(device) or torch.cat has copied the data). + """ + with self.lock: + for h in image_hashes: + self._release_hash(h) + + def get_stats(self) -> dict: + """Return cache statistics.""" + with self.lock: + return { + **self.stats, + "num_cached": len(self.hash_to_metadata), + "num_protected": sum(1 for v in self.ref_counts.values() if v > 0), + "allocated_mb": self.allocator.get_allocated_size() / 1024**2, + "free_mb": self.allocator.get_free_size() / 1024**2, + "total_mb": self.total_pool_size_bytes / 1024**2, + } + async def batch_is_exist(self, image_hashes: List[str]) -> List[bool]: with self.lock: local_results = [h in self.hash_to_metadata for h in image_hashes] diff --git a/test/registered/unit/mem_cache/test_embedding_cache_controller.py b/test/registered/unit/mem_cache/test_embedding_cache_controller.py new file mode 100644 index 000000000..9b65cfe3a --- /dev/null +++ b/test/registered/unit/mem_cache/test_embedding_cache_controller.py @@ -0,0 +1,716 @@ +"""Unit tests for EmbeddingCacheController — LRU eviction and RDMA ref counting.""" + +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=5, suite="base-a-test-cpu") + +import threading +import time +import unittest +from unittest.mock import MagicMock + +import torch + +from sglang.srt.mem_cache.storage.mooncake_store.embedding_cache_controller import ( + ContiguousMemoryAllocator, + EmbeddingCacheController, + EmbeddingInsertOperation, +) + +# --------------------------------------------------------------------------- +# ContiguousMemoryAllocator tests +# --------------------------------------------------------------------------- + + +class TestContiguousMemoryAllocator(unittest.TestCase): + def test_basic_alloc_free(self): + alloc = ContiguousMemoryAllocator(1024) + a = alloc.allocate(256) + self.assertIsNotNone(a) + self.assertEqual(a, 0) + b = alloc.allocate(256) + self.assertEqual(b, 256) + alloc.free(a, 256) + c = alloc.allocate(128) + self.assertEqual(c, 0) # reused freed block + + def test_alloc_fails_when_full(self): + alloc = ContiguousMemoryAllocator(256) + a = alloc.allocate(256) + self.assertIsNotNone(a) + b = alloc.allocate(1) + self.assertIsNone(b) + + def test_free_merges_adjacent(self): + alloc = ContiguousMemoryAllocator(512) + a = alloc.allocate(128) + b = alloc.allocate(128) + c = alloc.allocate(256) + alloc.free(a, 128) + alloc.free(b, 128) + # The two 128-byte blocks should merge into one 256-byte free block + d = alloc.allocate(256) + self.assertIsNotNone(d) + self.assertEqual(d, 0) + + def test_allocated_size_tracking(self): + alloc = ContiguousMemoryAllocator(1024) + self.assertEqual(alloc.get_allocated_size(), 0) + a = alloc.allocate(300) + self.assertEqual(alloc.get_allocated_size(), 300) + b = alloc.allocate(200) + self.assertEqual(alloc.get_allocated_size(), 500) + alloc.free(a, 300) + self.assertEqual(alloc.get_allocated_size(), 200) + alloc.free(b, 200) + self.assertEqual(alloc.get_allocated_size(), 0) + + def test_free_size_tracking(self): + alloc = ContiguousMemoryAllocator(1024) + self.assertEqual(alloc.get_free_size(), 1024) + alloc.allocate(400) + self.assertEqual(alloc.get_free_size(), 624) + + def test_double_free_is_safe(self): + alloc = ContiguousMemoryAllocator(256) + a = alloc.allocate(128) + alloc.free(a, 128) + # Second free of same offset — should not corrupt state + alloc.free(a, 128) + self.assertEqual(alloc.get_allocated_size(), 0) + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + + +def _make_controller( + pool_mb=1.0, enable_eviction=True, hidden_dims=None, max_eviction_batch=10 +): + """Create an EmbeddingCacheController with a mocked MooncakeEmbeddingStore.""" + ctrl = EmbeddingCacheController.__new__(EmbeddingCacheController) + ctrl.tp_world_size = 1 + ctrl.tp_group = None + ctrl.tp_rank = 0 + ctrl.all_rank_get = False + ctrl.hidden_dims = hidden_dims or {"image": 1024} + ctrl.element_size = torch.float32.itemsize + ctrl.enable_eviction = enable_eviction + ctrl.max_eviction_batch = max_eviction_batch + + # Small pool for testing (1 MB by default) + ctrl.total_pool_size_bytes = int(pool_mb * 1024**2) + ctrl.cpu_pool = torch.empty( + ctrl.total_pool_size_bytes, dtype=torch.uint8, pin_memory=False + ) + + # Mock the mooncake store — no real RDMA + ctrl.mooncake_store = MagicMock() + ctrl.mooncake_store.register_buffer = MagicMock() + + ctrl.allocator = ContiguousMemoryAllocator(ctrl.total_pool_size_bytes) + ctrl.hash_to_metadata = {} + ctrl.access_order = {} + ctrl.access_lock = threading.Lock() + ctrl.ref_counts = {} + + ctrl.stats = { + "total_allocated": 0, + "total_evicted": 0, + "eviction_count": 0, + "allocation_failures": 0, + } + + ctrl.ongoing_prefetch = {} + ctrl.prefetch_queue = MagicMock() + ctrl.insert_queue = MagicMock() + + ctrl.lock = threading.Lock() + ctrl.stop_event = threading.Event() + + # Do NOT start the IO thread — tests drive _io_loop logic manually + ctrl.io_thread = MagicMock() + + ctrl.prefetch_tp_group = None + return ctrl + + +def _embedding_bytes(num_tokens, dim): + return num_tokens * dim * torch.float32.itemsize + + +# --------------------------------------------------------------------------- +# LRU eviction tests +# --------------------------------------------------------------------------- + + +class TestLRUEviction(unittest.TestCase): + def test_evict_oldest_first(self): + dim = 64 + size = _embedding_bytes(1, dim) # 256 bytes + # Pool exactly fits 3 entries (768 bytes) + pool_bytes = size * 3 + ctrl = _make_controller(pool_mb=pool_bytes / (1024**2)) + + # Insert 3 entries — fills the pool + for i in range(3): + h = f"hash_{i}" + with ctrl.lock: + offset = ctrl.allocator.allocate(size) + self.assertIsNotNone(offset) + ctrl.hash_to_metadata[h] = (offset, 1, dim, size) + ctrl._update_access_time(h) + + # Pool is full. Inserting a 4th should evict hash_0 (oldest). + h_new = "hash_new" + with ctrl.lock: + offset = ctrl._allocate_with_eviction(size) + + self.assertIsNotNone(offset) + with ctrl.lock: + self.assertNotIn("hash_0", ctrl.hash_to_metadata) + self.assertIn("hash_1", ctrl.hash_to_metadata) + self.assertIn("hash_2", ctrl.hash_to_metadata) + + def test_eviction_disabled(self): + dim = 64 + size = _embedding_bytes(1, dim) # 256 bytes + # Pool exactly fits 1 entry + ctrl = _make_controller(pool_mb=size / (1024**2), enable_eviction=False) + + # Fill the pool + with ctrl.lock: + offset = ctrl.allocator.allocate(size) + self.assertIsNotNone(offset) + + # Try to allocate more — should fail without eviction + with ctrl.lock: + offset2 = ctrl._allocate_with_eviction(size) + self.assertIsNone(offset2) + + def test_access_time_updates_prevent_eviction(self): + ctrl = _make_controller(pool_mb=0.01) + dim = 64 + size = _embedding_bytes(1, dim) + + # Insert 2 entries + hashes = [] + for i in range(2): + h = f"hash_{i}" + hashes.append(h) + with ctrl.lock: + offset = ctrl.allocator.allocate(size) + ctrl.hash_to_metadata[h] = (offset, 1, dim, size) + ctrl._update_access_time(h) + + # Touch hash_0 to make it recently used + time.sleep(0.01) + with ctrl.lock: + ctrl._update_access_time("hash_0") + + # Trigger eviction — hash_1 should be evicted (older) + with ctrl.lock: + candidates = ctrl._select_eviction_candidates(size) + self.assertIn("hash_1", candidates) + self.assertNotIn("hash_0", candidates) + + def test_eviction_stats(self): + ctrl = _make_controller(pool_mb=0.01) + dim = 64 + size = _embedding_bytes(1, dim) + + # Insert and then evict + with ctrl.lock: + offset = ctrl.allocator.allocate(size) + ctrl.hash_to_metadata["h"] = (offset, 1, dim, size) + ctrl._update_access_time("h") + + with ctrl.lock: + freed = ctrl._evict_hashes(["h"]) + self.assertGreater(freed, 0) + self.assertEqual(ctrl.stats["eviction_count"], 1) + self.assertGreater(ctrl.stats["total_evicted"], 0) + + def test_evict_nonexistent_hash(self): + ctrl = _make_controller() + with ctrl.lock: + freed = ctrl._evict_hashes(["nonexistent"]) + self.assertEqual(freed, 0) + + def test_max_eviction_batch(self): + ctrl = _make_controller(pool_mb=1.0, max_eviction_batch=2) + dim = 64 + size = _embedding_bytes(1, dim) + + # Insert many small entries + for i in range(10): + h = f"hash_{i}" + with ctrl.lock: + offset = ctrl.allocator.allocate(size) + if offset is None: + break + ctrl.hash_to_metadata[h] = (offset, 1, dim, size) + ctrl._update_access_time(h) + + # _select_eviction_candidates should return at most 2 + with ctrl.lock: + candidates = ctrl._select_eviction_candidates(size * 100) + self.assertLessEqual(len(candidates), 2) + + +# --------------------------------------------------------------------------- +# Ref counting tests +# --------------------------------------------------------------------------- + + +class TestRefCounting(unittest.TestCase): + def test_protect_prevents_eviction(self): + ctrl = _make_controller() + dim = 64 + size = _embedding_bytes(1, dim) + + with ctrl.lock: + offset = ctrl.allocator.allocate(size) + ctrl.hash_to_metadata["h"] = (offset, 1, dim, size) + ctrl._update_access_time("h") + ctrl._protect_hash("h") + + # Should not be selected for eviction + with ctrl.lock: + candidates = ctrl._select_eviction_candidates(size) + self.assertNotIn("h", candidates) + + # Release and retry + with ctrl.lock: + ctrl._release_hash("h") + candidates = ctrl._select_eviction_candidates(size) + self.assertIn("h", candidates) + + def test_evict_hashes_skips_protected(self): + ctrl = _make_controller() + dim = 64 + size = _embedding_bytes(1, dim) + + with ctrl.lock: + offset = ctrl.allocator.allocate(size) + ctrl.hash_to_metadata["h"] = (offset, 1, dim, size) + ctrl._update_access_time("h") + ctrl._protect_hash("h") + + # Attempt to evict — should be skipped + with ctrl.lock: + freed = ctrl._evict_hashes(["h"]) + self.assertEqual(freed, 0) + with ctrl.lock: + self.assertIn("h", ctrl.hash_to_metadata) + + def test_ref_count_multiple_protects(self): + ctrl = _make_controller() + with ctrl.lock: + ctrl._protect_hash("h") + ctrl._protect_hash("h") + self.assertEqual(ctrl.ref_counts["h"], 2) + + ctrl._release_hash("h") + self.assertEqual(ctrl.ref_counts["h"], 1) + + # Still protected + candidates = ctrl._select_eviction_candidates(1) + self.assertNotIn("h", candidates) + + ctrl._release_hash("h") + self.assertNotIn("h", ctrl.ref_counts) + + def test_release_nonexistent_is_safe(self): + ctrl = _make_controller() + with ctrl.lock: + ctrl._release_hash("nonexistent") # should not raise + + def test_prefetch_sets_ref_count(self): + """prefetch() should set ref_count=1 for each new entry.""" + ctrl = _make_controller() + dim = 64 + ctrl.hidden_dims = {"image": dim} + h = "img_hash_1" + + ctrl.prefetch("req1", [h], [1], modality="image") + + with ctrl.lock: + self.assertEqual(ctrl.ref_counts.get(h), 1) + + # Simulate RDMA completion + with ctrl.lock: + ctrl._release_hash(h) + self.assertNotIn(h, ctrl.ref_counts) + + def test_insert_batch_sets_ref_count(self): + """insert_batch() should set ref_count=1 for each new entry.""" + ctrl = _make_controller() + dim = 64 + h = "img_hash_1" + tensor = torch.randn(1, dim) + + ctrl.insert_batch([h], [tensor]) + + with ctrl.lock: + self.assertEqual(ctrl.ref_counts.get(h), 1) + + # Simulate RDMA completion + with ctrl.lock: + ctrl._release_hash(h) + self.assertNotIn(h, ctrl.ref_counts) + + def test_get_embeddings_sets_ref_count(self): + """get_embeddings() should set ref_count=1 per hash.""" + ctrl = _make_controller() + dim = 64 + size = _embedding_bytes(1, dim) + + with ctrl.lock: + offset = ctrl.allocator.allocate(size) + ctrl.hash_to_metadata["h1"] = (offset, 1, dim, size) + ctrl._update_access_time("h1") + + tensors = ctrl.get_embeddings(["h1"]) + self.assertIsNotNone(tensors[0]) + + with ctrl.lock: + self.assertEqual(ctrl.ref_counts.get("h1"), 1) + + # Protected — eviction should skip + with ctrl.lock: + candidates = ctrl._select_eviction_candidates(size) + self.assertNotIn("h1", candidates) + + # Release + ctrl.release_embeddings(["h1"]) + with ctrl.lock: + self.assertNotIn("h1", ctrl.ref_counts) + + def test_get_embeddings_missing_hash(self): + """Missing hashes should return None and not set ref_count.""" + ctrl = _make_controller() + tensors = ctrl.get_embeddings(["nonexistent"]) + self.assertIsNone(tensors[0]) + with ctrl.lock: + self.assertNotIn("nonexistent", ctrl.ref_counts) + + def test_release_embeddings_missing_hash_is_safe(self): + """Releasing a hash that was never protected should be a no-op.""" + ctrl = _make_controller() + ctrl.release_embeddings(["nonexistent"]) # should not raise + + def test_io_loop_releases_prefetch_ref(self): + """_io_loop should release ref_count after batch_get completes.""" + ctrl = _make_controller() + dim = 64 + ctrl.hidden_dims = {"image": dim} + h = "img_hash_1" + + ctrl.prefetch("req1", [h], [1], modality="image") + + with ctrl.lock: + self.assertEqual(ctrl.ref_counts.get(h), 1) + + # Simulate _io_loop completing the RDMA GET + op = ctrl.ongoing_prefetch.get("req1") + self.assertIsNotNone(op) + ctrl.mooncake_store.batch_get = MagicMock(return_value=[True]) + + # Manually execute what _io_loop does for prefetch + results = ctrl.mooncake_store.batch_get(op.keys, op.ptrs, op.sizes) + op.mark_done(all(results)) + with ctrl.lock: + for k in op.keys: + ctrl._release_hash(k) + + with ctrl.lock: + self.assertNotIn(h, ctrl.ref_counts) + + def test_io_loop_releases_insert_ref(self): + """_io_loop should release ref_count after batch_put completes.""" + ctrl = _make_controller() + dim = 64 + h = "img_hash_1" + tensor = torch.randn(1, dim) + + ctrl.insert_batch([h], [tensor]) + + with ctrl.lock: + self.assertEqual(ctrl.ref_counts.get(h), 1) + + # Get the enqueued insert operation + ctrl.insert_queue.put.assert_called_once() + insert_op = ctrl.insert_queue.put.call_args[0][0] + self.assertIsInstance(insert_op, EmbeddingInsertOperation) + + # Simulate _io_loop completing the RDMA PUT + ctrl.mooncake_store.batch_put = MagicMock() + ctrl.mooncake_store.batch_put(insert_op.keys, insert_op.ptrs, insert_op.sizes) + with ctrl.lock: + for k in insert_op.keys: + ctrl._release_hash(k) + + with ctrl.lock: + self.assertNotIn(h, ctrl.ref_counts) + + +# --------------------------------------------------------------------------- +# Race condition prevention tests +# --------------------------------------------------------------------------- + + +class TestRDMAEvictionRacePrevention(unittest.TestCase): + def test_eviction_during_prefetch_is_blocked(self): + """An entry with in-flight RDMA GET cannot be evicted.""" + ctrl = _make_controller(pool_mb=0.01) + dim = 64 + ctrl.hidden_dims = {"image": dim} + size = _embedding_bytes(1, dim) + + # Prefetch sets ref_count=1 + ctrl.prefetch("req1", ["h1"], [1], modality="image") + + # Now try to evict to make room — should skip h1 + with ctrl.lock: + candidates = ctrl._select_eviction_candidates(size * 10) + self.assertNotIn("h1", candidates) + + # Direct eviction should also be skipped + with ctrl.lock: + freed = ctrl._evict_hashes(["h1"]) + self.assertEqual(freed, 0) + + def test_eviction_during_insert_is_blocked(self): + """An entry with in-flight RDMA PUT cannot be evicted.""" + ctrl = _make_controller(pool_mb=0.01) + dim = 64 + tensor = torch.randn(1, dim) + + ctrl.insert_batch(["h1"], [tensor]) + + with ctrl.lock: + candidates = ctrl._select_eviction_candidates(999999) + self.assertNotIn("h1", candidates) + + with ctrl.lock: + freed = ctrl._evict_hashes(["h1"]) + self.assertEqual(freed, 0) + + def test_eviction_during_get_embeddings_is_blocked(self): + """An entry returned by get_embeddings() cannot be evicted.""" + ctrl = _make_controller(pool_mb=0.01) + dim = 64 + size = _embedding_bytes(1, dim) + + with ctrl.lock: + offset = ctrl.allocator.allocate(size) + ctrl.hash_to_metadata["h1"] = (offset, 1, dim, size) + ctrl._update_access_time("h1") + + tensors = ctrl.get_embeddings(["h1"]) + self.assertIsNotNone(tensors[0]) + + # Try to evict — should be blocked + with ctrl.lock: + candidates = ctrl._select_eviction_candidates(999999) + self.assertNotIn("h1", candidates) + + # Release and verify eviction is now possible + ctrl.release_embeddings(["h1"]) + with ctrl.lock: + candidates = ctrl._select_eviction_candidates(999999) + self.assertIn("h1", candidates) + + def test_concurrent_eviction_while_reading(self): + """Simulate a concurrent eviction attempt while a read holds a ref.""" + ctrl = _make_controller(pool_mb=0.05) + dim = 64 + size = _embedding_bytes(1, dim) + num_entries = 10 + + # Insert entries + for i in range(num_entries): + h = f"hash_{i}" + with ctrl.lock: + offset = ctrl.allocator.allocate(size) + if offset is None: + break + ctrl.hash_to_metadata[h] = (offset, 1, dim, size) + ctrl._update_access_time(h) + + # Simulate get_embeddings holding refs on hash_0..hash_4 + held_hashes = [f"hash_{i}" for i in range(5)] + tensors = ctrl.get_embeddings(held_hashes) + + # Try to evict all — only unprotected entries should be candidates + with ctrl.lock: + candidates = ctrl._select_eviction_candidates(999999) + for h in held_hashes: + self.assertNotIn(h, candidates) + + # Unprotected entries should be candidates + for i in range(5, num_entries): + h = f"hash_{i}" + if h in ctrl.hash_to_metadata: + self.assertIn(h, candidates) + + # Release refs + ctrl.release_embeddings(held_hashes) + with ctrl.lock: + candidates = ctrl._select_eviction_candidates(999999) + for h in held_hashes: + if h in ctrl.hash_to_metadata: + self.assertIn(h, candidates) + + def test_evict_hashes_cleans_up_ref_counts(self): + """After eviction, ref_counts for the evicted hash should be removed.""" + ctrl = _make_controller() + dim = 64 + size = _embedding_bytes(1, dim) + + with ctrl.lock: + offset = ctrl.allocator.allocate(size) + ctrl.hash_to_metadata["h"] = (offset, 1, dim, size) + ctrl._update_access_time("h") + # Stale ref_count (shouldn't happen normally, but test cleanup) + ctrl.ref_counts["h"] = 0 + + with ctrl.lock: + # ref_count is 0, so eviction should proceed + freed = ctrl._evict_hashes(["h"]) + self.assertGreater(freed, 0) + with ctrl.lock: + self.assertNotIn("h", ctrl.hash_to_metadata) + self.assertNotIn("h", ctrl.ref_counts) + + +# --------------------------------------------------------------------------- +# get_embeddings view safety tests +# --------------------------------------------------------------------------- + + +class TestGetEmbeddingsViewSafety(unittest.TestCase): + def test_get_embeddings_returns_view(self): + """get_embeddings returns a view into cpu_pool, not a copy.""" + ctrl = _make_controller() + dim = 64 + size = _embedding_bytes(1, dim) + + with ctrl.lock: + offset = ctrl.allocator.allocate(size) + ctrl.hash_to_metadata["h1"] = (offset, 1, dim, size) + ctrl._update_access_time("h1") + + tensors = ctrl.get_embeddings(["h1"]) + self.assertIsNotNone(tensors[0]) + self.assertEqual(tensors[0].shape, (1, dim)) + + # Verify it's a view into cpu_pool (shares storage) + self.assertTrue( + tensors[0].storage().data_ptr() == ctrl.cpu_pool.storage().data_ptr() + ) + + # Release + ctrl.release_embeddings(["h1"]) + + def test_data_preserved_while_ref_held(self): + """Data should remain intact as long as ref_count > 0.""" + ctrl = _make_controller() + dim = 64 + size = _embedding_bytes(1, dim) + + # Write known data + with ctrl.lock: + offset = ctrl.allocator.allocate(size) + ctrl.hash_to_metadata["h1"] = (offset, 1, dim, size) + ctrl._update_access_time("h1") + view = ( + ctrl.cpu_pool[offset : offset + size].view(torch.float32).view(1, dim) + ) + view.copy_(torch.ones(1, dim)) + + # Read via get_embeddings (holds ref) + tensors = ctrl.get_embeddings(["h1"]) + self.assertTrue(torch.all(tensors[0] == 1.0)) + + # Verify data is still valid + self.assertTrue(torch.all(tensors[0] == 1.0)) + + # Release + ctrl.release_embeddings(["h1"]) + + +# --------------------------------------------------------------------------- +# Stats tests +# --------------------------------------------------------------------------- + + +class TestGetStats(unittest.TestCase): + def test_stats_include_num_protected(self): + ctrl = _make_controller() + dim = 64 + size = _embedding_bytes(1, dim) + + with ctrl.lock: + offset = ctrl.allocator.allocate(size) + ctrl.hash_to_metadata["h1"] = (offset, 1, dim, size) + ctrl._update_access_time("h1") + ctrl._protect_hash("h1") + ctrl._protect_hash("h2") # h2 not in metadata, but has ref + + stats = ctrl.get_stats() + self.assertEqual(stats["num_protected"], 2) + + def test_stats_eviction_tracking(self): + ctrl = _make_controller(pool_mb=0.01) + dim = 64 + size = _embedding_bytes(1, dim) + + with ctrl.lock: + offset = ctrl.allocator.allocate(size) + ctrl.hash_to_metadata["h"] = (offset, 1, dim, size) + ctrl._update_access_time("h") + + with ctrl.lock: + ctrl._evict_hashes(["h"]) + + stats = ctrl.get_stats() + self.assertEqual(stats["eviction_count"], 1) + self.assertGreater(stats["total_evicted"], 0) + self.assertEqual(stats["num_cached"], 0) + + +# --------------------------------------------------------------------------- +# _select_eviction_candidates iterator safety test +# --------------------------------------------------------------------------- + + +class TestEvictionCandidateIteratorSafety(unittest.TestCase): + def test_list_snapshot_prevents_concurrent_mutation(self): + """sorted_hashes should be a list snapshot, not a live dict view.""" + ctrl = _make_controller() + dim = 64 + size = _embedding_bytes(1, dim) + + # Insert entries + for i in range(5): + h = f"hash_{i}" + with ctrl.lock: + offset = ctrl.allocator.allocate(size) + if offset is None: + break + ctrl.hash_to_metadata[h] = (offset, 1, dim, size) + ctrl._update_access_time(h) + + # _select_eviction_candidates should work even if access_order + # is modified during iteration (the snapshot via list() prevents this) + with ctrl.lock: + candidates = ctrl._select_eviction_candidates(size) + # Should return candidates without RuntimeError + self.assertIsInstance(candidates, list) + + +if __name__ == "__main__": + unittest.main()