import threading import unittest from types import SimpleNamespace import torch if not hasattr(torch, "npu"): torch.npu = SimpleNamespace(Event=object) # type: ignore[attr-defined] from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.config_data import ( ChunkedTokenDatabase, KeyMetadata, ReqMeta, ) from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.kv_transfer import ( KVCacheStoreSendingThread, ) class _FakeStore: def __init__(self, exists_result: list[int]): self.exists_result = exists_result self.put_calls: list[tuple[list[str], list[list[int]], list[list[int]]]] = [] def set_device(self): return None def exists(self, keys: list[str]) -> list[int]: # Return exact number of states for requested keys. return self.exists_result[: len(keys)] def put(self, keys, addrs, sizes): self.put_calls.append((list(keys), list(addrs), list(sizes))) class TestKVTransferMissingKeyPut(unittest.TestCase): def test_sending_thread_only_puts_missing_keys(self): store = _FakeStore(exists_result=[1, 0, 1, 0]) token_db = ChunkedTokenDatabase([KeyMetadata("m", 0, 0, 0, 0)], [16], None) token_db.set_group_buffers({0: [1000]}, {0: [16]}, {0: [1]}) thread = KVCacheStoreSendingThread( m_store=store, token_database=token_db, block_size=16, tp_rank=0, dcp_size=1, put_step=1, kv_role="kv_producer", ready_event=threading.Event(), group_uses_align_state=[False], enable_kv_event=False, ) req_meta = ReqMeta( req_id="req-1", token_len_chunk=64, block_ids=[0, 1, 2, 3], block_hashes=[b"h0", b"h1", b"h2", b"h3"], # type: ignore[arg-type] current_event=None, ) thread.add_stored_request("req-1") thread.request_queue.put(req_meta) thread._handle_request(req_meta) self.assertEqual(len(store.put_calls), 1) put_keys, put_addrs, put_sizes = store.put_calls[0] self.assertEqual(len(put_keys), 2) self.assertEqual(put_addrs, [[1001], [1003]]) self.assertEqual(put_sizes, [[16], [16]]) if __name__ == "__main__": unittest.main()