# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project

import json
import logging
import math
import queue
import sys
import threading
import types
from types import SimpleNamespace
from unittest.mock import MagicMock, patch

import pytest
import torch

from vllm.distributed.kv_transfer.kv_connector.v1.mooncake import (
    rdma_utils,
)
from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store import (
    worker,
)
from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store import (
    worker as mooncake_store_worker,
)
from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store.data import (
    BlobBlockHashes,
    ChunkedTokenDatabase,
    KeyMetadata,
    LoadSpec,
    PoolKey,
    ReqMeta,
)
from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store.metrics import (
    MooncakeStoreConnectorStats,
)
from vllm.v1.core.kv_cache_utils import BlockHash


class _RecordingBlockHashes:
    def __init__(self, values: list[bytes]):
        self.values = values
        self.accessed: list[int] = []

    def __len__(self):
        return len(self.values)

    def __getitem__(self, idx):
        if isinstance(idx, slice):
            return [self[i] for i in range(*idx.indices(len(self)))]
        if idx < 0:
            idx += len(self.values)
        if not 0 <= idx < len(self.values):
            raise IndexError(idx)
        self.accessed.append(idx)
        return self.values[idx]


def _default_send_coord() -> mooncake_store_worker.MooncakeStoreCoordinator:
    from vllm.v1.kv_cache_interface import FullAttentionSpec, KVCacheGroupSpec

    spec = FullAttentionSpec(block_size=16, num_kv_heads=8, head_size=64, dtype=None)
    return mooncake_store_worker.MooncakeStoreCoordinator(
        [KVCacheGroupSpec(["layer0"], spec)],
        scheduler_block_size=16,
        hash_block_size=16,
    )


def _make_store_sending_thread(
    store: MagicMock,
    *,
    coord: mooncake_store_worker.MooncakeStoreCoordinator | None = None,
    token_databases: list[ChunkedTokenDatabase] | None = None,
    block_size: int = 16,
    tp_rank: int = 0,
    put_step: int = 1,
    replicate_config: object | None = None,
    enable_group_semantics: bool = False,
    supports_group_ids: bool = False,
) -> mooncake_store_worker.KVCacheStoreSendingThread:
    if coord is None:
        coord = _default_send_coord()
    if token_databases is None:
        db = ChunkedTokenDatabase(KeyMetadata("test-model", 0, 0, 0, 0), block_size=16)
        db.set_kv_caches_base_addr([0x1000])
        db.set_block_len([256])
        token_databases = [db]
    thread = mooncake_store_worker.KVCacheStoreSendingThread(
        store=store,
        token_databases=token_databases,
        block_size=block_size,
        coord=coord,
        tp_rank=tp_rank,
        group_put_steps=[put_step] * len(token_databases),
        kv_role="kv_producer",
        ready_event=threading.Event(),
        replicate_config=replicate_config,
        enable_group_semantics=enable_group_semantics,
        supports_group_ids=supports_group_ids,
    )
    thread.request_queue.task_done = MagicMock()
    return thread


def _make_store_recving_thread(
    store: MagicMock,
    *,
    tp_rank: int = 0,
    disk_offload_buffer_budget_bytes: int | None = None,
) -> mooncake_store_worker.KVCacheStoreRecvingThread:
    from vllm.v1.kv_cache_interface import FullAttentionSpec, KVCacheGroupSpec

    token_database = ChunkedTokenDatabase(
        KeyMetadata("test-model", 0, 0, 0, 0), block_size=16
    )
    token_database.set_kv_caches_base_addr([0x1000])
    token_database.set_block_len([256])
    spec = FullAttentionSpec(block_size=16, num_kv_heads=8, head_size=64, dtype=None)
    coord = mooncake_store_worker.MooncakeStoreCoordinator(
        [KVCacheGroupSpec(["layer0"], spec)],
        scheduler_block_size=16,
        hash_block_size=16,
    )
    thread = mooncake_store_worker.KVCacheStoreRecvingThread(
        store=store,
        token_databases=[token_database],
        block_size=16,
        tp_rank=tp_rank,
        ready_event=threading.Event(),
        coord=coord,
        disk_offload_buffer_budget_bytes=disk_offload_buffer_budget_bytes,
    )
    thread.request_queue.task_done = MagicMock()
    return thread


def _make_load_req(
    req_id: str,
    block_hashes: list[bytes],
    *,
    token_len: int,
    vllm_cached_tokens: int = 0,
) -> ReqMeta:
    return ReqMeta(
        req_id=req_id,
        token_len_chunk=token_len,
        block_ids=(list(range(len(block_hashes))),),
        block_hashes=block_hashes,
        load_spec=LoadSpec(
            vllm_cached_tokens=vllm_cached_tokens,
            kvpool_cached_tokens=token_len,
            can_load=True,
            token_len=token_len,
        ),
    )


def _make_store_req(req_id: str, block_hashes: list[bytes]) -> ReqMeta:
    return ReqMeta(
        req_id=req_id,
        token_len_chunk=32,
        block_ids=([0, 1],),
        block_hashes=block_hashes,
        can_save=True,
    )


def _make_multi_group_store_req(req_id: str, block_hashes: list[bytes]) -> ReqMeta:
    return ReqMeta(
        req_id=req_id,
        token_len_chunk=32,
        block_ids=([0, 1], [2, 3]),
        block_hashes=block_hashes,
        can_save=True,
    )


_DISK_OFFLOAD_SINGLE_KEY_BYTES = worker._estimate_disk_offload_staging_bytes([256])
_DISK_OFFLOAD_USABLE_BUDGET_RATIO = 0.9
_DISK_OFFLOAD_BUDGET_FOR_THREE_KEYS = 4 * _DISK_OFFLOAD_SINGLE_KEY_BYTES
_DISK_OFFLOAD_BUDGET_FOR_SPLIT = math.ceil(
    2 * _DISK_OFFLOAD_SINGLE_KEY_BYTES / _DISK_OFFLOAD_USABLE_BUDGET_RATIO
)  # Allows two 256-byte chunks but not the third.
_DISK_OFFLOAD_BUDGET_TOO_SMALL = (
    _DISK_OFFLOAD_SINGLE_KEY_BYTES - 1
)  # Smaller than a single 256-byte chunk.


class _FakeKVTransferConfig:
    def __init__(
        self,
        *,
        kv_role: str = "kv_both",
        extra_config: dict[str, object] | None = None,
    ) -> None:
        self.kv_role = kv_role
        self.kv_connector_extra_config = extra_config or {}

    def get_from_extra_config(self, key: str, default: object) -> object:
        return self.kv_connector_extra_config.get(key, default)


class _FakeModelConfig:
    model = "test-model"
    use_mla = False

    def get_num_layers(self, parallel_config) -> int:
        return 1

    def get_total_num_kv_heads(self) -> int:
        return 1


def _make_vllm_config(
    *,
    extra_config: dict[str, object] | None = None,
    rank: int = 0,
    decode_context_parallel_size: int = 1,
) -> SimpleNamespace:
    return SimpleNamespace(
        model_config=_FakeModelConfig(),
        parallel_config=SimpleNamespace(
            pipeline_parallel_size=1,
            rank=rank,
            data_parallel_index=0,
            decode_context_parallel_size=decode_context_parallel_size,
            prefill_context_parallel_size=1,
        ),
        kv_transfer_config=_FakeKVTransferConfig(extra_config=extra_config),
        cache_config=SimpleNamespace(block_size=16, num_gpu_blocks=10),
        kv_events_config=SimpleNamespace(enable_kv_cache_events=False),
        speculative_config=None,
    )


def _make_kv_cache_config(*, block_size: int = 16) -> object:
    """Minimal single-group KVCacheConfig for topology tests."""
    from vllm.v1.kv_cache_interface import (
        FullAttentionSpec,
        KVCacheConfig,
        KVCacheGroupSpec,
    )

    spec = FullAttentionSpec(
        block_size=block_size, num_kv_heads=8, head_size=64, dtype=None
    )
    return KVCacheConfig(
        num_blocks=10,
        kv_cache_tensors=[],
        kv_cache_groups=[KVCacheGroupSpec(["layer0"], spec)],
    )


def _write_mooncake_config(tmp_path, config: dict[str, object]) -> str:
    config_path = tmp_path / "mooncake_config.json"
    config_path.write_text(json.dumps(config), encoding="utf-8")
    return str(config_path)


def _install_fake_mooncake(monkeypatch, store_instance: MagicMock):
    class FakeReplicateConfig:
        def __init__(self) -> None:
            self.preferred_segment = ""

    fake_store_module = types.ModuleType("mooncake.store")
    fake_store_module.MooncakeDistributedStore = lambda: store_instance  # type: ignore[attr-defined]
    fake_store_module.ReplicateConfig = FakeReplicateConfig  # type: ignore[attr-defined]
    fake_mooncake_module = types.ModuleType("mooncake")
    fake_mooncake_module.store = fake_store_module  # type: ignore[attr-defined]
    monkeypatch.setitem(sys.modules, "mooncake", fake_mooncake_module)
    monkeypatch.setitem(sys.modules, "mooncake.store", fake_store_module)
    return FakeReplicateConfig


def _patch_worker_runtime(
    monkeypatch,
    *,
    local_ip: str = "10.0.0.7",
    tp_rank: int = 0,
    tp_size: int = 1,
    dcp_size: int = 1,
) -> None:
    single_rank_group = SimpleNamespace(world_size=1, rank_in_group=0)
    # DCP groups are contiguous splits of the TP group (see
    # parallel_state.py), so dcp_rank == tp_rank % dcp_size.
    dcp_group = SimpleNamespace(world_size=dcp_size, rank_in_group=tp_rank % dcp_size)
    monkeypatch.setattr(worker, "get_tensor_model_parallel_rank", lambda: tp_rank)
    monkeypatch.setattr(worker, "get_tensor_model_parallel_world_size", lambda: tp_size)
    monkeypatch.setattr(worker, "get_pcp_group", lambda: single_rank_group)
    monkeypatch.setattr(worker, "get_dcp_group", lambda: dcp_group)
    monkeypatch.setattr(worker, "get_ip", lambda: local_ip)
    monkeypatch.setattr(worker, "LookupKeyServer", MagicMock())


def test_pool_key_to_string_without_prefix_is_unchanged():
    """Default (empty) cache_prefix keeps keys byte-identical to the
    historical unprefixed format so existing deployments keep their hits."""
    key = PoolKey(KeyMetadata("test-model", 0, 0, 0, 0), "deadbeef")
    assert (
        key.to_string() == "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:0@deadbeef"
    )


def test_pool_key_cache_prefix_namespaces_and_disambiguates():
    """A non-empty cache_prefix is prepended, and two instances with
    different prefixes never collide on identical block hashes."""
    md_a = KeyMetadata("test-model", 0, 0, 0, 0, cache_prefix="depA")
    md_b = KeyMetadata("test-model", 0, 0, 0, 0, cache_prefix="depB")

    key_a = PoolKey(md_a, "deadbeef")
    key_b = PoolKey(md_b, "deadbeef")

    assert key_a.to_string() == (
        "depA@test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:0@deadbeef"
    )
    assert key_a.to_string() != key_b.to_string()
    assert hash(key_a) != hash(key_b)


def test_default_local_buffer_size_matches_pr40900():
    """PR-40900 shipped a 4 GiB default for local_buffer_size; the dual-mode
    patch preserves it (and the JSON key) so unchanged PR-40900 configs work."""
    assert worker.DEFAULT_LOCAL_BUFFER_SIZE == 4 * 1024**3


def test_get_requester_local_hostname_prefers_override(monkeypatch):
    monkeypatch.setenv("MOONCAKE_REQUESTER_LOCAL_HOSTNAME", "worker-a:50053")

    assert rdma_utils.get_requester_local_hostname("10.0.0.7") == "worker-a:50053"


def test_get_configured_preferred_segment_returns_explicit_override():
    assert (
        rdma_utils.get_configured_preferred_segment(
            {"preferred_segment": "10.0.0.7:50053"}
        )
        == "10.0.0.7:50053"
    )


def test_get_configured_preferred_segment_prefers_explicit_over_env(monkeypatch):
    monkeypatch.setenv("MOONCAKE_PREFERRED_SEGMENT", "10.0.0.8:50053")

    assert (
        rdma_utils.get_configured_preferred_segment(
            {"preferred_segment": "10.0.0.7:50053"}
        )
        == "10.0.0.7:50053"
    )


def test_get_configured_preferred_segment_returns_env_override(monkeypatch):
    monkeypatch.setenv("MOONCAKE_PREFERRED_SEGMENT", "10.0.0.8:50053")

    assert rdma_utils.get_configured_preferred_segment({}) == "10.0.0.8:50053"


def test_get_configured_preferred_segment_rejects_empty_override():
    with pytest.raises(ValueError, match="preferred_segment"):
        rdma_utils.get_configured_preferred_segment({"preferred_segment": "  "})


class _ReplicaDesc:
    def __init__(self, tier: str):
        self.tier = tier

    def is_memory_replica(self) -> bool:
        return self.tier == "memory"

    def is_disk_replica(self) -> bool:
        return self.tier == "disk"

    def is_local_disk_replica(self) -> bool:
        return self.tier == "disk"


def test_store_sending_thread_skips_request_during_cpu_pressure():
    store = MagicMock()
    store.batch_is_exist.side_effect = lambda keys: [0] * len(keys)
    store.batch_put_from_multi_buffers.side_effect = [
        [-200, -200],
        [256, 256],
        [256, 256],
    ]
    thread = _make_store_sending_thread(store)

    thread.add_stored_request("req-a")
    thread._handle_request(_make_store_req("req-a", [b"a0", b"a1"]))

    assert thread._store_pressure_active is True
    assert "req-a" in thread._skip_store_requests
    assert store.batch_put_from_multi_buffers.call_count == 1

    thread.add_stored_request("req-a")
    thread._handle_request(_make_store_req("req-a", [b"a2", b"a3"]))

    assert store.batch_put_from_multi_buffers.call_count == 1

    thread.add_stored_request("req-b")
    thread._handle_request(_make_store_req("req-b", [b"b0", b"b1"]))

    assert thread._store_pressure_active is False
    assert "req-a" not in thread._skip_store_requests
    assert store.batch_put_from_multi_buffers.call_count == 2

    thread.add_stored_request("req-a")
    thread._handle_request(_make_store_req("req-a", [b"a4", b"a5"]))

    assert store.batch_put_from_multi_buffers.call_count == 3


def test_store_sending_thread_records_mooncake_metrics():
    store = MagicMock()
    store.batch_is_exist.return_value = [0, 0]
    store.batch_put_from_multi_buffers.return_value = [256, 256]
    thread = _make_store_sending_thread(store)
    stats = MooncakeStoreConnectorStats()
    thread._record_operation_cb = stats.record_operation

    thread.add_stored_request("req-a")
    thread._handle_request(_make_store_req("req-a", [b"a0", b"a1"]))

    assert len(stats.data["save_exists"]) == 1
    assert stats.data["save_exists"][0]["num_keys"] == 2
    assert len(stats.data["save_put"]) == 1
    assert stats.data["save_put"][0]["num_bytes"] == 512
    assert stats.data["save_put"][0]["status"] == "ok"


def test_process_tokens_uses_mask_num_as_start_chunk():
    db = ChunkedTokenDatabase(
        KeyMetadata("test-model", 0, 0, 0, 0),
        block_size=32,
        hash_block_size=8,
    )
    block_hashes = _RecordingBlockHashes([bytes([i]) for i in range(16)])

    results = list(
        db.process_tokens(
            token_len=96,
            block_hashes=block_hashes,
            mask_num=64,
        )
    )

    assert block_hashes.accessed == [11]
    assert results == [(64, 96, bytes([11]))]


def test_process_tokens_applies_chunk_mask_before_hash_access():
    db = ChunkedTokenDatabase(
        KeyMetadata("test-model", 0, 0, 0, 0),
        block_size=32,
        hash_block_size=8,
    )
    block_hashes = _RecordingBlockHashes([bytes([i]) for i in range(16)])

    results = list(
        db.process_tokens(
            token_len=128,
            block_hashes=block_hashes,
            mask_num=64,
            chunk_mask=[False, True],
        )
    )

    assert block_hashes.accessed == [15]
    assert results == [(96, 128, bytes([15]))]


def test_process_tokens_applies_stride_before_hash_access():
    db = ChunkedTokenDatabase(
        KeyMetadata("test-model", 0, 0, 0, 0),
        block_size=32,
        hash_block_size=8,
    )
    block_hashes = _RecordingBlockHashes([bytes([i]) for i in range(16)])

    results = list(
        db.process_tokens(
            token_len=128,
            block_hashes=block_hashes,
            mask_num=64,
            put_step=2,
            put_step_rank=1,
        )
    )

    assert block_hashes.accessed == [15]
    assert results == [(96, 128, bytes([15]))]


def test_store_sending_thread_delta_saves_only_new_full_attention_chunks():
    store = MagicMock()
    store.batch_is_exist.side_effect = lambda keys: [0] * len(keys)
    store.batch_put_from_multi_buffers.return_value = [256, 256]
    thread = _make_store_sending_thread(store)

    thread.add_stored_request("req-a")
    thread._saved_offset["req-a"] = 32
    thread._handle_request(
        ReqMeta(
            req_id="req-a",
            token_len_chunk=64,
            block_ids=([0, 1, 2, 3],),
            block_hashes=[b"a0", b"a1", b"a2", b"a3"],
            can_save=True,
        )
    )

    keys = store.batch_is_exist.call_args.args[0]
    assert keys == [
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:0@6132",
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:0@6133",
    ]
    assert store.batch_put_from_multi_buffers.call_args.args[0] == keys


def test_store_sending_thread_delta_strides_with_local_phase():
    store = MagicMock()
    store.batch_is_exist.side_effect = lambda keys: [0] * len(keys)
    store.batch_put_from_multi_buffers.return_value = [256]
    thread = _make_store_sending_thread(store, tp_rank=0, put_step=2)

    thread.add_stored_request("req-a")
    thread._saved_offset["req-a"] = 16
    thread._handle_request(
        ReqMeta(
            req_id="req-a",
            token_len_chunk=64,
            block_ids=([0, 1, 2, 3],),
            block_hashes=[b"a0", b"a1", b"a2", b"a3"],
            can_save=True,
        )
    )

    keys = store.batch_is_exist.call_args.args[0]
    assert keys == [
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:0@6132",
    ]
    assert store.batch_put_from_multi_buffers.call_args.args[0] == keys


def test_tp_sharded_group_saves_every_block_on_every_rank():
    """Sharded ranks must write every block because peers hold different bytes."""
    store = MagicMock()
    store.batch_is_exist.side_effect = lambda keys: [0] * len(keys)
    store.batch_put_from_multi_buffers.side_effect = lambda keys, *a: [256] * len(keys)
    thread = _make_store_sending_thread(store, tp_rank=0, put_step=2)
    thread.group_put_steps = [1]

    thread.add_stored_request("req-a")
    thread._handle_request(
        ReqMeta(
            req_id="req-a",
            token_len_chunk=64,
            block_ids=([0, 1, 2, 3],),
            block_hashes=[b"a0", b"a1", b"a2", b"a3"],
            can_save=True,
        )
    )

    keys = store.batch_is_exist.call_args.args[0]
    assert len(keys) == 4


def test_store_sending_thread_retries_skipped_range_after_pressure():
    store = MagicMock()
    store.batch_is_exist.side_effect = lambda keys: [0] * len(keys)
    store.batch_put_from_multi_buffers.side_effect = lambda keys, *a: [256] * len(keys)
    thread = _make_store_sending_thread(store, tp_rank=0, put_step=2)

    # Under pressure the request is skipped without persisting anything, so the
    # saved offset stays at 0.
    thread._store_pressure_active = True
    thread._skip_store_requests.add("req-a")

    thread.add_stored_request("req-a")
    thread._handle_request(
        ReqMeta(
            req_id="req-a",
            token_len_chunk=16,
            block_ids=([0],),
            block_hashes=[b"a0"],
            can_save=True,
        )
    )

    store.batch_is_exist.assert_not_called()
    assert thread._saved_offset.get("req-a", 0) == 0

    thread._store_pressure_active = False
    thread._skip_store_requests.clear()

    # The next batch resumes from offset 0, re-covering the chunk skipped under
    # pressure (chunk 0) rather than losing it.
    thread.add_stored_request("req-a")
    thread._handle_request(
        ReqMeta(
            req_id="req-a",
            token_len_chunk=64,
            block_ids=([0, 1, 2, 3],),
            block_hashes=[b"a0", b"a1", b"a2", b"a3"],
            can_save=True,
        )
    )

    keys = store.batch_is_exist.call_args.args[0]
    assert keys == [
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:0@6130",
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:0@6132",
    ]
    assert store.batch_put_from_multi_buffers.call_args.args[0] == keys


def _make_partial_tail_send_thread(
    store,
    *,
    replicate_config=None,
    enable_group_semantics=False,
    supports_group_ids=False,
):
    coord = SimpleNamespace(
        enable_partial_hash_hits=True,
        hash_block_size=4,
        lcm_block_size=16,
    )
    db = ChunkedTokenDatabase(
        KeyMetadata("test-model", 0, 0, 0, 0),
        block_size=4,
        hash_block_size=4,
    )
    db.set_kv_caches_base_addr([0x1000])
    db.set_block_len([256])
    return _make_store_sending_thread(
        store,
        coord=coord,
        token_databases=[db],
        replicate_config=replicate_config,
        enable_group_semantics=enable_group_semantics,
        supports_group_ids=supports_group_ids,
    )


def _make_partial_tail_req(block_ids: list[int]) -> ReqMeta:
    return ReqMeta(
        req_id="req-a",
        token_len_chunk=0,
        block_ids=(block_ids,),
        block_hashes=[b"a0", b"a1", b"a2"],
        can_save=True,
        partial_tail_offloads=[(1, 7, 12)],
    )


def test_partial_tail_offload_skips_null_source_blocks():
    store = MagicMock()
    store.batch_is_exist.side_effect = lambda keys: [0] * len(keys)
    store.batch_put_from_multi_buffers.return_value = [256, 256]
    thread = _make_partial_tail_send_thread(store)

    assert thread._maybe_offload_partial_tail(_make_partial_tail_req([0, 2, 3]))

    keys, addrs, _sizes, _replicate_config = (
        store.batch_put_from_multi_buffers.call_args.args
    )
    assert keys == [
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:0@6131",
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:0@6132",
    ]
    assert addrs == [[0x1000 + 2 * 256], [0x1000 + 3 * 256]]


def test_partial_tail_offload_replaces_stale_group_ids_after_filtering():
    store = MagicMock()
    store.batch_is_exist.return_value = [1, 0, 0]
    store.batch_put_from_multi_buffers.return_value = [256, 256]
    replicate_config = SimpleNamespace(group_ids=["stale"])
    thread = _make_partial_tail_send_thread(
        store,
        replicate_config=replicate_config,
        enable_group_semantics=True,
        supports_group_ids=True,
    )

    assert thread._maybe_offload_partial_tail(_make_partial_tail_req([1, 2, 3]))

    keys, _addrs, _sizes, config = store.batch_put_from_multi_buffers.call_args.args
    assert keys == [
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:0@6131",
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:0@6132",
    ]
    assert config is replicate_config
    assert config.group_ids == [
        "vllm-mooncake-store:test-model@6131",
        "vllm-mooncake-store:test-model@6132",
    ]


def test_partial_tail_offload_honors_active_pressure_gate():
    store = MagicMock()
    thread = _make_partial_tail_send_thread(store)
    thread._store_pressure_active = True
    thread._skip_store_requests.add("req-a")
    thread.add_stored_request("req-a")

    thread._handle_request(_make_partial_tail_req([1, 2, 3]))

    store.batch_is_exist.assert_not_called()
    store.batch_put_from_multi_buffers.assert_not_called()
    assert thread.stored_requests["req-a"] == 0


def test_partial_tail_put_failure_activates_pressure_gate():
    store = MagicMock()
    store.batch_is_exist.side_effect = lambda keys: [0] * len(keys)
    store.batch_put_from_multi_buffers.return_value = [256, -200, 256]
    thread = _make_partial_tail_send_thread(store)
    thread.add_stored_request("req-a")

    thread._handle_request(_make_partial_tail_req([1, 2, 3]))

    assert thread._store_pressure_active is True
    assert thread._skip_store_requests == {"req-a"}
    assert thread._saved_offset.get("req-a", 0) == 0
    assert thread.stored_requests["req-a"] == 0

    thread.add_stored_request("req-a")
    thread._handle_request(_make_partial_tail_req([1, 2, 3]))
    assert store.batch_put_from_multi_buffers.call_count == 1


def test_store_sending_thread_delta_start_rank_saves_second_local_chunk():
    store = MagicMock()
    store.batch_is_exist.side_effect = lambda keys: [0] * len(keys)
    store.batch_put_from_multi_buffers.return_value = [256, 256]
    thread = _make_store_sending_thread(store, tp_rank=1, put_step=2)

    thread.add_stored_request("req-a")
    thread._saved_offset["req-a"] = 16
    thread._handle_request(
        ReqMeta(
            req_id="req-a",
            token_len_chunk=64,
            block_ids=([0, 1, 2, 3],),
            block_hashes=[b"a0", b"a1", b"a2", b"a3"],
            can_save=True,
        )
    )

    keys = store.batch_is_exist.call_args.args[0]
    assert keys == [
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:0@6131",
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:0@6133",
    ]
    assert store.batch_put_from_multi_buffers.call_args.args[0] == keys


def test_store_sending_thread_delta_saves_only_new_masked_chunks():
    store = MagicMock()
    store.batch_is_exist.side_effect = lambda keys: [0] * len(keys)
    store.batch_put_from_multi_buffers.side_effect = (
        lambda keys, addrs, sizes, replicate_config: [256] * len(keys)
    )
    coord = SimpleNamespace(
        lcm_block_size=16,
        store_mask=lambda token_len, start_token, num_prompt_tokens=None: (
            None,
            [True, False],
        ),
    )

    db_full = ChunkedTokenDatabase(
        KeyMetadata("test-model", 0, 0, 0, 0, group_id=0),
        block_size=16,
    )
    db_full.set_kv_caches_base_addr([0x1000])
    db_full.set_block_len([256])
    db_masked = ChunkedTokenDatabase(
        KeyMetadata("test-model", 0, 0, 0, 0, group_id=1),
        block_size=16,
    )
    db_masked.set_kv_caches_base_addr([0x2000])
    db_masked.set_block_len([256])

    thread = _make_store_sending_thread(
        store,
        coord=coord,
        token_databases=[db_full, db_masked],
    )

    thread.add_stored_request("req-a")
    thread._saved_offset["req-a"] = 32
    thread._handle_request(
        ReqMeta(
            req_id="req-a",
            token_len_chunk=64,
            block_ids=([0, 1, 2, 3], [0, 1, 2, 3]),
            block_hashes=[b"a0", b"a1", b"a2", b"a3"],
            can_save=True,
        )
    )

    keys = store.batch_is_exist.call_args.args[0]
    full_hashes = [k.rsplit("@", 1)[-1] for k in keys if "@group:0" in k]
    masked_hashes = [k.rsplit("@", 1)[-1] for k in keys if "@group:1" in k]

    assert full_hashes == [b"a2".hex(), b"a3".hex()]
    assert masked_hashes == [b"a2".hex()]


def test_store_sending_thread_prepares_missing_chunks_once_per_group():
    store = MagicMock()
    store.batch_is_exist.return_value = [0, 1, 0, 1, 0, 0]
    store.batch_put_from_multi_buffers.return_value = [256, 256, 512, 512]
    coord = SimpleNamespace(
        lcm_block_size=16,
        store_mask=lambda token_len, start_token, num_prompt_tokens=None: (
            None,
            None,
        ),
    )

    db0 = ChunkedTokenDatabase(
        KeyMetadata("test-model", 0, 0, 0, 0, group_id=0),
        block_size=16,
    )
    db0.set_kv_caches_base_addr([0x1000])
    db0.set_block_len([256])
    db0.prepare_values = MagicMock(wraps=db0.prepare_values)
    db0.prepare_value = MagicMock(side_effect=AssertionError("scalar path called"))

    db1 = ChunkedTokenDatabase(
        KeyMetadata("test-model", 0, 0, 0, 0, group_id=1),
        block_size=16,
    )
    db1.set_kv_caches_base_addr([0x2000])
    db1.set_block_len([512])
    db1.prepare_values = MagicMock(wraps=db1.prepare_values)
    db1.prepare_value = MagicMock(side_effect=AssertionError("scalar path called"))

    thread = _make_store_sending_thread(
        store,
        coord=coord,
        token_databases=[db0, db1],
    )
    thread.add_stored_request("req-a")
    thread._handle_request(
        ReqMeta(
            req_id="req-a",
            token_len_chunk=48,
            block_ids=([0, 1, 2], [2, 1, 0]),
            block_hashes=[b"a0", b"a1", b"a2"],
            can_save=True,
        )
    )

    db0.prepare_value.assert_not_called()
    db1.prepare_value.assert_not_called()
    db0.prepare_values.assert_called_once_with([(0, 16), (32, 48)], [0, 1, 2])
    db1.prepare_values.assert_called_once_with([(16, 32), (32, 48)], [2, 1, 0])

    keys, addrs, sizes, _ = store.batch_put_from_multi_buffers.call_args.args
    assert [key.rsplit("@", 1)[-1] for key in keys] == [
        "6130",
        "6132",
        "6131",
        "6132",
    ]
    assert addrs == [[0x1000], [0x1200], [0x2200], [0x2000]]
    assert sizes == [[256], [256], [512], [512]]


def test_store_sending_thread_only_skips_on_no_available_handle():
    store = MagicMock()
    store.batch_is_exist.side_effect = lambda keys: [0] * len(keys)
    store.batch_put_from_multi_buffers.side_effect = [
        [-500, -500],
        [256, 256],
    ]
    thread = _make_store_sending_thread(store)

    thread.add_stored_request("req-a")
    thread._handle_request(_make_store_req("req-a", [b"a0", b"a1"]))

    assert thread._store_pressure_active is False
    assert "req-a" not in thread._skip_store_requests
    assert store.batch_put_from_multi_buffers.call_count == 1

    thread.add_stored_request("req-a")
    thread._handle_request(_make_store_req("req-a", [b"a2", b"a3"]))

    assert store.batch_put_from_multi_buffers.call_count == 2


def test_store_sending_thread_releases_pin_on_batch_is_exist_failure():
    # `batch_is_exist` raising must still decrement `stored_requests` so the
    # scheduler can drop `delay_free_blocks` and release the pinned GPU blocks.
    store = MagicMock()
    store.batch_is_exist.side_effect = RuntimeError("mooncake down")
    thread = _make_store_sending_thread(store)

    thread.add_stored_request("req-a")
    with pytest.raises(RuntimeError):
        thread._handle_request(_make_store_req("req-a", [b"a0", b"a1"]))

    assert thread.stored_requests["req-a"] == 0
    store.batch_put_from_multi_buffers.assert_not_called()


def test_store_sending_thread_releases_pin_on_batch_put_failure():
    # `batch_put_from_multi_buffers` raising is logged (not re-raised), and the
    # pin must still be released through the finally block.
    store = MagicMock()
    store.batch_is_exist.return_value = [0, 0]
    store.batch_put_from_multi_buffers.side_effect = RuntimeError("rdma error")
    thread = _make_store_sending_thread(store)

    thread.add_stored_request("req-a")
    thread._handle_request(_make_store_req("req-a", [b"a0", b"a1"]))

    assert thread.stored_requests["req-a"] == 0


def test_store_recving_thread_reports_failed_block_ids():
    store = MagicMock()
    store.batch_get_into_multi_buffers.return_value = [256, -5, -7]
    thread = _make_store_recving_thread(store)

    thread._handle_request(
        _make_load_req(
            "req-a",
            [b"a0", b"a1", b"a2"],
            token_len=48,
        )
    )

    assert thread.get_and_clear_finished_requests() == {"req-a"}
    assert thread.get_and_clear_block_ids_with_load_errors() == {1, 2}
    assert thread.get_and_clear_block_ids_with_load_errors() == set()


def test_store_recving_thread_reports_failed_block_ids_after_rotation():
    store = MagicMock()
    store.batch_get_into_multi_buffers.return_value = [256, -5, 256]
    thread = _make_store_recving_thread(store, tp_rank=1)

    thread._handle_request(
        _make_load_req(
            "req-a",
            [b"a0", b"a1", b"a2"],
            token_len=48,
        )
    )

    assert thread.get_and_clear_block_ids_with_load_errors() == {2}


def test_store_recving_thread_reports_all_attempted_blocks_on_exception():
    store = MagicMock()
    store.batch_get_into_multi_buffers.side_effect = RuntimeError("boom")
    thread = _make_store_recving_thread(store)

    thread._handle_request(
        _make_load_req(
            "req-a",
            [b"a0", b"a1", b"a2"],
            token_len=48,
        )
    )

    assert thread.get_and_clear_finished_requests() == {"req-a"}
    assert thread.get_and_clear_block_ids_with_load_errors() == {0, 1, 2}


def test_store_worker_get_block_ids_with_load_errors_delegates_to_recv_thread():
    recv_thread = MagicMock()
    recv_thread.get_and_clear_block_ids_with_load_errors.return_value = {3, 4}
    w = _make_bare_worker()
    w.kv_recv_threads = [recv_thread]

    assert w.get_block_ids_with_load_errors() == {3, 4}
    recv_thread.get_and_clear_block_ids_with_load_errors.assert_called_once_with()


def test_store_sending_thread_passes_replicate_config_when_preferred_segment_set():
    store = MagicMock()
    store.batch_is_exist.side_effect = lambda keys: [0] * len(keys)
    store.batch_put_from_multi_buffers.return_value = [256, 256]
    replicate_config = SimpleNamespace(preferred_segment="10.0.0.7:50053")
    thread = _make_store_sending_thread(store, replicate_config=replicate_config)

    thread.add_stored_request("req-a")
    thread._handle_request(_make_store_req("req-a", [b"a0", b"a1"]))

    assert store.batch_put_from_multi_buffers.call_count == 1
    call_args = store.batch_put_from_multi_buffers.call_args.args
    assert len(call_args) == 4
    assert call_args[3] is replicate_config


def test_store_sending_thread_passes_default_replicate_config_when_no_preferred_segment():  # noqa: E501
    """Without a preferred_segment the SendingThread still forwards a
    (default-constructed) ReplicateConfig so the C++ side always sees a
    well-defined config object."""
    store = MagicMock()
    store.batch_is_exist.side_effect = lambda keys: [0] * len(keys)
    store.batch_put_from_multi_buffers.return_value = [256, 256]
    replicate_config = SimpleNamespace()
    thread = _make_store_sending_thread(store, replicate_config=replicate_config)

    thread.add_stored_request("req-a")
    thread._handle_request(_make_store_req("req-a", [b"a0", b"a1"]))

    assert store.batch_put_from_multi_buffers.call_count == 1
    call_args = store.batch_put_from_multi_buffers.call_args.args
    assert len(call_args) == 4
    assert call_args[3] is replicate_config


def test_group_id_detection_uses_class_attribute_without_instantiating():
    class ReplicateConfigWithClassGroupIdsAndRequiredInit:
        group_ids = None

        def __init__(self, required):
            self.group_ids = required

    assert worker._replicate_config_supports_group_ids(
        ReplicateConfigWithClassGroupIdsAndRequiredInit,
        object(),
    )


def test_group_id_detection_uses_instance_attribute_when_class_lacks_attribute():
    class ReplicateConfigWithInstanceGroupIds:
        def __init__(self) -> None:
            self.group_ids = None

    replicate_config = ReplicateConfigWithInstanceGroupIds()

    assert worker._replicate_config_supports_group_ids(
        ReplicateConfigWithInstanceGroupIds,
        replicate_config,
    )


def test_group_id_detection_rejects_old_replicate_config():
    class ReplicateConfigWithoutGroupIds:
        __slots__ = ()

    assert not worker._replicate_config_supports_group_ids(
        ReplicateConfigWithoutGroupIds,
        ReplicateConfigWithoutGroupIds(),
    )


def test_store_sending_thread_sets_group_ids_when_enabled():
    store = MagicMock()
    store.batch_is_exist.side_effect = lambda keys: [0] * len(keys)
    store.batch_put_from_multi_buffers.return_value = [256, 256]
    replicate_config = SimpleNamespace(group_ids=None)
    thread = _make_store_sending_thread(
        store,
        replicate_config=replicate_config,
        enable_group_semantics=True,
        supports_group_ids=True,
    )

    thread.add_stored_request("req-a")
    thread._handle_request(_make_store_req("req-a", [b"a0", b"a1"]))

    assert store.batch_put_from_multi_buffers.call_count == 1
    keys, _addrs, _sizes, config = store.batch_put_from_multi_buffers.call_args.args
    assert config is replicate_config
    assert config.group_ids == [
        "vllm-mooncake-store:test-model@6130",
        "vllm-mooncake-store:test-model@6131",
    ]
    assert len(config.group_ids) == len(keys)


def test_store_sending_thread_leaves_group_ids_unchanged_when_flag_disabled():
    store = MagicMock()
    store.batch_is_exist.side_effect = lambda keys: [0] * len(keys)
    store.batch_put_from_multi_buffers.return_value = [256, 256]
    replicate_config = SimpleNamespace(group_ids=["existing"])
    thread = _make_store_sending_thread(
        store,
        replicate_config=replicate_config,
        enable_group_semantics=False,
        supports_group_ids=True,
    )

    thread.add_stored_request("req-a")
    thread._handle_request(_make_store_req("req-a", [b"a0", b"a1"]))

    assert store.batch_put_from_multi_buffers.call_count == 1
    assert store.batch_put_from_multi_buffers.call_args.args[3] is replicate_config
    assert replicate_config.group_ids == ["existing"]


def test_store_sending_thread_leaves_group_ids_unchanged_when_unsupported():
    store = MagicMock()
    store.batch_is_exist.side_effect = lambda keys: [0] * len(keys)
    store.batch_put_from_multi_buffers.return_value = [256, 256]
    replicate_config = SimpleNamespace(group_ids=["existing"])
    thread = _make_store_sending_thread(
        store,
        replicate_config=replicate_config,
        enable_group_semantics=True,
        supports_group_ids=False,
    )

    thread.add_stored_request("req-a")
    thread._handle_request(_make_store_req("req-a", [b"a0", b"a1"]))

    assert store.batch_put_from_multi_buffers.call_count == 1
    assert store.batch_put_from_multi_buffers.call_args.args[3] is replicate_config
    assert replicate_config.group_ids == ["existing"]


def test_mooncake_group_id_uses_logical_chunk_and_excludes_physical_ranks():
    group_id = worker._make_mooncake_group_id(
        KeyMetadata(
            "test-model",
            tp_rank=3,
            pcp_rank=5,
            dcp_rank=7,
            pp_rank=11,
            group_id=13,
        ),
        "abcdef",
    )

    assert group_id == "vllm-mooncake-store:test-model@abcdef"
    assert "tp_rank" not in group_id
    assert "pp_rank" not in group_id
    assert "pcp" not in group_id
    assert "dcp" not in group_id
    assert "group:" not in group_id


def test_mooncake_group_id_uses_cache_prefix_namespace():
    metadata_a = KeyMetadata("test-model", 0, 0, 0, 0, cache_prefix="deployment-a")
    metadata_b = KeyMetadata("test-model", 0, 0, 0, 0, cache_prefix="deployment-b")

    group_id_a = worker._make_mooncake_group_id(metadata_a, "abcdef")
    group_id_b = worker._make_mooncake_group_id(metadata_b, "abcdef")

    assert group_id_a == "vllm-mooncake-store:deployment-a@test-model@abcdef"
    assert group_id_b == "vllm-mooncake-store:deployment-b@test-model@abcdef"
    assert group_id_a != group_id_b


def test_store_sending_thread_group_id_excludes_physical_sharding():
    store = MagicMock()
    store.batch_is_exist.return_value = [0, 0]
    store.batch_put_from_multi_buffers.return_value = [256, 256]
    replicate_config = SimpleNamespace(group_ids=None)
    db = ChunkedTokenDatabase(
        KeyMetadata(
            "test-model",
            tp_rank=2,
            pcp_rank=0,
            dcp_rank=0,
            pp_rank=0,
            group_id=0,
        ),
        block_size=16,
    )
    db.set_kv_caches_base_addr([0x1000])
    db.set_block_len([256])
    thread = _make_store_sending_thread(
        store,
        token_databases=[db],
        tp_rank=2,
        replicate_config=replicate_config,
        enable_group_semantics=True,
        supports_group_ids=True,
    )

    thread.add_stored_request("req-a")
    thread._handle_request(_make_store_req("req-a", [b"a0", b"a1"]))

    keys, _addrs, _sizes, config = store.batch_put_from_multi_buffers.call_args.args
    assert keys == [
        "test-model@tp_rank:2@pcp0@dcp0@pp_rank:0@group:0@6130",
        "test-model@tp_rank:2@pcp0@dcp0@pp_rank:0@group:0@6131",
    ]
    assert config.group_ids == [
        "vllm-mooncake-store:test-model@6130",
        "vllm-mooncake-store:test-model@6131",
    ]


def test_store_sending_thread_multiple_segments_share_logical_group_id():
    store = MagicMock()
    store.batch_is_exist.return_value = [0, 0]
    store.batch_put_from_multi_buffers.return_value = [512, 512]
    replicate_config = SimpleNamespace(group_ids=None)
    db = ChunkedTokenDatabase(KeyMetadata("test-model", 0, 0, 0, 0), block_size=16)
    db.set_kv_caches_base_addr([0x1000, 0x2000])
    db.set_block_len([256, 256])
    thread = _make_store_sending_thread(
        store,
        token_databases=[db],
        replicate_config=replicate_config,
        enable_group_semantics=True,
        supports_group_ids=True,
    )

    thread.add_stored_request("req-a")
    thread._handle_request(_make_store_req("req-a", [b"a0", b"a1"]))

    keys, addrs, sizes, config = store.batch_put_from_multi_buffers.call_args.args
    assert keys == [
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:0@6130",
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:0@6131",
    ]
    assert addrs == [[0x1000, 0x2000], [0x1100, 0x2100]]
    assert sizes == [[256, 256], [256, 256]]
    assert config.group_ids == [
        "vllm-mooncake-store:test-model@6130",
        "vllm-mooncake-store:test-model@6131",
    ]


def test_store_sending_thread_group_ids_share_across_kv_cache_groups():
    from vllm.v1.kv_cache_interface import FullAttentionSpec, KVCacheGroupSpec

    store = MagicMock()
    store.batch_is_exist.return_value = [0, 0, 0, 0]
    store.batch_put_from_multi_buffers.return_value = [256, 256, 256, 256]
    replicate_config = SimpleNamespace(group_ids=None)
    spec = FullAttentionSpec(block_size=16, num_kv_heads=8, head_size=64, dtype=None)
    coord = mooncake_store_worker.MooncakeStoreCoordinator(
        [
            KVCacheGroupSpec(["layer0"], spec),
            KVCacheGroupSpec(["layer1"], spec),
        ],
        scheduler_block_size=16,
        hash_block_size=16,
    )
    token_databases = []
    for group_id, base_addr in enumerate([0x1000, 0x3000]):
        db = ChunkedTokenDatabase(
            KeyMetadata("test-model", 0, 0, 0, 0, group_id=group_id),
            block_size=16,
        )
        db.set_kv_caches_base_addr([base_addr])
        db.set_block_len([256])
        token_databases.append(db)
    thread = _make_store_sending_thread(
        store,
        coord=coord,
        token_databases=token_databases,
        replicate_config=replicate_config,
        enable_group_semantics=True,
        supports_group_ids=True,
    )

    thread.add_stored_request("req-a")
    thread._handle_request(_make_multi_group_store_req("req-a", [b"a0", b"a1"]))

    keys, addrs, sizes, config = store.batch_put_from_multi_buffers.call_args.args
    assert keys == [
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:0@6130",
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:0@6131",
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:1@6130",
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:1@6131",
    ]
    assert addrs == [[0x1000], [0x1100], [0x3200], [0x3300]]
    assert sizes == [[256], [256], [256], [256]]
    # Different vLLM KV cache groups for the same prefix chunk share the
    # same Mooncake lifecycle group id.
    assert config.group_ids == [
        "vllm-mooncake-store:test-model@6130",
        "vllm-mooncake-store:test-model@6131",
        "vllm-mooncake-store:test-model@6130",
        "vllm-mooncake-store:test-model@6131",
    ]
    assert config.group_ids[0] == config.group_ids[2]
    assert config.group_ids[1] == config.group_ids[3]


def test_store_sending_thread_group_ids_follow_missing_key_filter():
    store = MagicMock()
    store.batch_is_exist.return_value = [1, 0]
    store.batch_put_from_multi_buffers.return_value = [256]
    replicate_config = SimpleNamespace(group_ids=None)
    thread = _make_store_sending_thread(
        store,
        replicate_config=replicate_config,
        enable_group_semantics=True,
        supports_group_ids=True,
    )

    thread.add_stored_request("req-a")
    thread._handle_request(_make_store_req("req-a", [b"a0", b"a1"]))

    keys, _addrs, _sizes, config = store.batch_put_from_multi_buffers.call_args.args
    assert keys == ["test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:0@6131"]
    assert config.group_ids == ["vllm-mooncake-store:test-model@6131"]


def test_estimate_disk_offload_staging_bytes_sums_multi_segment_sizes():
    assert worker._estimate_disk_offload_staging_bytes([256, 512]) == 12288


def test_recv_thread_uses_single_batch_when_no_disk_offload_budget(monkeypatch):
    monkeypatch.delenv("VLLM_MOONCAKE_STORE_TIER_LOG", raising=False)
    store = MagicMock()
    store.batch_get_into_multi_buffers.return_value = [256, 256, 256]
    thread = _make_store_recving_thread(store, disk_offload_buffer_budget_bytes=None)

    req = _make_load_req(
        "req-a",
        [b"a0", b"a1", b"a2"],
        token_len=48,
    )

    thread._handle_request(req)

    assert store.batch_get_into_multi_buffers.call_count == 1
    keys, addrs, sizes = store.batch_get_into_multi_buffers.call_args.args
    assert keys == [
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:0@6130",
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:0@6131",
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:0@6132",
    ]
    assert sizes == [[256], [256], [256]]
    store.batch_get_replica_desc.assert_not_called()


def test_recv_thread_logs_tier_summary_when_enabled(monkeypatch, caplog_vllm):
    monkeypatch.setenv("VLLM_MOONCAKE_STORE_TIER_LOG", "1")
    caplog_vllm.set_level(logging.INFO, logger=worker.logger.name)

    store = MagicMock()
    store.batch_get_into_multi_buffers.return_value = [256, 256, -10]
    thread = _make_store_recving_thread(store, disk_offload_buffer_budget_bytes=None)

    req = _make_load_req(
        "req-a",
        [b"a0", b"a1", b"a2"],
        token_len=48,
    )
    expected_keys = [
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:0@6130",
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:0@6131",
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:0@6132",
    ]
    store.batch_get_replica_desc.return_value = {
        expected_keys[0]: [_ReplicaDesc("memory")],
        expected_keys[1]: [_ReplicaDesc("disk")],
        expected_keys[2]: [],
    }

    thread._handle_request(req)

    assert store.batch_get_replica_desc.call_args.args == (expected_keys,)
    assert store.method_calls[0][0] == "batch_get_replica_desc"
    assert store.method_calls[1][0] == "batch_get_into_multi_buffers"

    messages = [record.getMessage() for record in caplog_vllm.records]
    assert any(
        "Mooncake load tier summary" in message
        and "req_id=req-a" in message
        and "batch_keys=3" in message
        and "memory_keys=1" in message
        and "disk_keys=1" in message
        and "unknown_keys=1" in message
        and "success_keys=2" in message
        and "failed_keys=1" in message
        and "bytes_by_tier={'memory': 256, 'disk': 256, 'unknown': 0}" in message
        for message in messages
    )


def test_recv_thread_records_partial_failure_metrics(monkeypatch):
    monkeypatch.delenv("VLLM_MOONCAKE_STORE_TIER_LOG", raising=False)
    store = MagicMock()
    store.batch_get_into_multi_buffers.return_value = [256, -10]
    thread = _make_store_recving_thread(store, disk_offload_buffer_budget_bytes=None)
    stats = MooncakeStoreConnectorStats()
    thread._record_operation_cb = stats.record_operation

    req = _make_load_req(
        "req-a",
        [b"a0", b"a1"],
        token_len=32,
    )

    thread._handle_request(req)

    assert len(stats.data["load_get"]) == 1
    assert stats.data["load_get"][0]["num_keys"] == 2
    assert stats.data["load_get"][0]["num_bytes"] == 512
    assert stats.data["load_get"][0]["status"] == "partial_failure"
    assert stats.data["load_get"][0]["num_failed_keys"] == 1


def test_recv_thread_uses_ratio_scaled_budget_for_first_pass_split():
    store = MagicMock()
    store.batch_get_into_multi_buffers.side_effect = [
        [256],
        [256],
    ]
    thread = _make_store_recving_thread(
        store,
        disk_offload_buffer_budget_bytes=2 * _DISK_OFFLOAD_SINGLE_KEY_BYTES,
    )

    req = _make_load_req(
        "req-a",
        [b"a0", b"a1"],
        token_len=32,
    )

    thread._handle_request(req)

    assert store.batch_get_into_multi_buffers.call_count == 2
    first_keys = store.batch_get_into_multi_buffers.call_args_list[0].args[0]
    second_keys = store.batch_get_into_multi_buffers.call_args_list[1].args[0]
    assert first_keys == [
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:0@6130",
    ]
    assert second_keys == [
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:0@6131",
    ]


def test_recv_thread_splits_disk_offload_loads_by_budget():
    store = MagicMock()
    store.batch_get_into_multi_buffers.side_effect = [
        [256, 256],
        [256],
    ]
    thread = _make_store_recving_thread(
        store,
        disk_offload_buffer_budget_bytes=_DISK_OFFLOAD_BUDGET_FOR_SPLIT,
    )

    req = _make_load_req(
        "req-a",
        [b"a0", b"a1", b"a2"],
        token_len=48,
    )

    thread._handle_request(req)

    assert store.batch_get_into_multi_buffers.call_count == 2

    first_keys = store.batch_get_into_multi_buffers.call_args_list[0].args[0]
    second_keys = store.batch_get_into_multi_buffers.call_args_list[1].args[0]
    first_addrs = store.batch_get_into_multi_buffers.call_args_list[0].args[1]
    second_addrs = store.batch_get_into_multi_buffers.call_args_list[1].args[1]
    first_sizes = store.batch_get_into_multi_buffers.call_args_list[0].args[2]
    second_sizes = store.batch_get_into_multi_buffers.call_args_list[1].args[2]
    assert first_keys == [
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:0@6130",
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:0@6131",
    ]
    assert second_keys == [
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:0@6132",
    ]
    base_addr = thread.token_databases[0].kv_caches_base_addr[0]
    block_len = thread.token_databases[0].block_len[0]
    assert first_addrs == [[base_addr], [base_addr + block_len]]
    assert second_addrs == [[base_addr + 2 * block_len]]
    expected_size = block_len
    assert first_sizes == [[expected_size], [expected_size]]
    assert second_sizes == [[expected_size]]


def test_recv_thread_stops_after_first_failing_disk_offload_sub_batch():
    store = MagicMock()
    store.batch_get_into_multi_buffers.return_value = [-10, -10]
    thread = _make_store_recving_thread(
        store,
        disk_offload_buffer_budget_bytes=_DISK_OFFLOAD_BUDGET_FOR_SPLIT,
    )

    req = _make_load_req(
        "req-a",
        [b"a0", b"a1", b"a2"],
        token_len=48,
    )

    thread._handle_request(req)

    assert store.batch_get_into_multi_buffers.call_count == 1
    assert thread.get_and_clear_block_ids_with_load_errors() == {0, 1}


def test_recv_thread_skips_split_when_budget_holds_all_keys():
    """PR-36 removed the count-based split trigger; with budget for 3 keys,
    all three should be requested in a single call."""
    store = MagicMock()
    store.batch_get_into_multi_buffers.return_value = [256, 256, 256]
    thread = _make_store_recving_thread(
        store,
        disk_offload_buffer_budget_bytes=_DISK_OFFLOAD_BUDGET_FOR_THREE_KEYS,
    )

    req = _make_load_req(
        "req-a",
        [b"a0", b"a1", b"a2"],
        token_len=48,
    )

    thread._handle_request(req)

    assert store.batch_get_into_multi_buffers.call_count == 1
    assert store.batch_get_into_multi_buffers.call_args_list[0].args[0] == [
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:0@6130",
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:0@6131",
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:0@6132",
    ]


def test_recv_thread_reports_unsplittable_key_larger_than_budget():
    # Non-zero tp_rank exercises the rotation: the oversized key may not sit
    # at the original request's first block. Every block must still be marked
    # invalid, since none are loaded.
    store = MagicMock()
    thread = _make_store_recving_thread(
        store,
        tp_rank=2,
        disk_offload_buffer_budget_bytes=_DISK_OFFLOAD_BUDGET_TOO_SMALL,
    )

    req = _make_load_req(
        "req-a",
        [b"a0", b"a1", b"a2"],
        token_len=48,
    )

    thread._handle_request(req)

    assert store.batch_get_into_multi_buffers.call_count == 0
    assert thread.get_and_clear_block_ids_with_load_errors() == {0, 1, 2}


def test_requester_worker_init_uses_positional_setup(tmp_path, monkeypatch):
    store = MagicMock()
    store.setup.return_value = 0
    _install_fake_mooncake(monkeypatch, store)
    _patch_worker_runtime(monkeypatch)
    monkeypatch.setenv(
        "MOONCAKE_CONFIG_PATH",
        _write_mooncake_config(
            tmp_path,
            {
                "metadata_server": "http://metadata/endpoint",
                "global_segment_size": "4gb",
                "local_buffer_size": "64mb",
                "protocol": "rdma",
                "device_name": "mlx5_0",
                "master_server_address": "10.0.0.7:50051",
                "enable_offload": True,
            },
        ),
    )
    w = worker.MooncakeStoreWorker(_make_vllm_config(), _make_kv_cache_config())

    assert not hasattr(w, "_isolate_offload_resources")
    assert store.setup.call_args.args == (
        "10.0.0.7",
        "http://metadata/endpoint",
        4 * 1024 * 1024 * 1024,  # global_segment_size: "4gb" honored
        64 * 1024 * 1024,
        "rdma",
        "mlx5_0",
        "10.0.0.7:50051",
    )


def test_requester_worker_init_prefers_local_hostname_override(
    tmp_path,
    monkeypatch,
):
    store = MagicMock()
    store.setup.return_value = 0
    _install_fake_mooncake(monkeypatch, store)
    _patch_worker_runtime(monkeypatch)
    monkeypatch.setenv("MOONCAKE_REQUESTER_LOCAL_HOSTNAME", "worker-a:50053")
    monkeypatch.setenv(
        "MOONCAKE_CONFIG_PATH",
        _write_mooncake_config(
            tmp_path,
            {
                "metadata_server": "http://metadata/endpoint",
                "local_buffer_size": "64mb",
                "protocol": "tcp",
                "device_name": "",
                "master_server_address": "10.0.0.7:50051",
            },
        ),
    )
    worker.MooncakeStoreWorker(_make_vllm_config(rank=1), _make_kv_cache_config())

    assert store.setup.call_args.args[0] == "worker-a:50053"


def test_requester_worker_init_skips_disk_budget_when_offload_disabled(
    tmp_path,
    monkeypatch,
):
    """enable_offload=False zeroes out the disk budget so we don't generate
    redundant owner GET-RPCs."""
    store = MagicMock()
    store.setup.return_value = 0
    _install_fake_mooncake(monkeypatch, store)
    _patch_worker_runtime(monkeypatch)
    monkeypatch.setenv(
        "MOONCAKE_CONFIG_PATH",
        _write_mooncake_config(
            tmp_path,
            {
                "metadata_server": "http://metadata/endpoint",
                "protocol": "tcp",
                "device_name": "",
                "master_server_address": "10.0.0.7:50051",
                "enable_offload": False,
            },
        ),
    )
    w = worker.MooncakeStoreWorker(_make_vllm_config(), _make_kv_cache_config())

    assert w.disk_offload_buffer_budget_bytes is None


def test_requester_worker_init_builds_replicate_config_for_preferred_segment(
    tmp_path,
    monkeypatch,
):
    store = MagicMock()
    store.setup.return_value = 0
    fake_replicate_config_cls = _install_fake_mooncake(monkeypatch, store)
    _patch_worker_runtime(monkeypatch)
    monkeypatch.setenv(
        "MOONCAKE_CONFIG_PATH",
        _write_mooncake_config(
            tmp_path,
            {
                "metadata_server": "http://metadata/endpoint",
                "protocol": "tcp",
                "device_name": "",
                "master_server_address": "10.0.0.7:50051",
            },
        ),
    )
    w = worker.MooncakeStoreWorker(
        _make_vllm_config(
            extra_config={
                "preferred_segment": "10.0.0.7:50053",
            }
        ),
        _make_kv_cache_config(),
    )

    assert isinstance(w.store_replicate_config, fake_replicate_config_cls)
    assert w.store_replicate_config.preferred_segment == "10.0.0.7:50053"


@pytest.mark.parametrize("dcp_size", [1, 4])
def test_worker_put_striding_covers_every_rank_get_namespace(
    tmp_path, monkeypatch, dcp_size
):
    """Every key a rank GETs must have been PUT by some rank.

    When num_kv_head < tp_size, ranks holding the same KV heads stripe
    their PUTs across one shared key namespace. That dedup is only valid
    when those ranks really share a namespace: with DCP > 1 each rank GETs
    every key from its own ``@dcpN`` namespace, so striding must be
    disabled.
    """
    tp_size = 4
    store = MagicMock()
    store.setup.return_value = 0
    _install_fake_mooncake(monkeypatch, store)
    monkeypatch.setenv(
        "MOONCAKE_CONFIG_PATH",
        _write_mooncake_config(
            tmp_path,
            {
                "metadata_server": "http://metadata/endpoint",
                "protocol": "tcp",
                "device_name": "",
                "master_server_address": "10.0.0.7:50051",
            },
        ),
    )

    # _FakeModelConfig has num_kv_head=1 < tp_size, which enables striding.
    block_hashes = [f"hash-{i}".encode() for i in range(4)]
    put_keys: set[str] = set()
    get_keys_per_rank: dict[int, set[str]] = {}
    for tp_rank in range(tp_size):
        _patch_worker_runtime(
            monkeypatch, tp_rank=tp_rank, tp_size=tp_size, dcp_size=dcp_size
        )
        w = worker.MooncakeStoreWorker(
            _make_vllm_config(rank=tp_rank, decode_context_parallel_size=dcp_size),
            _make_kv_cache_config(),
        )
        db = w.token_dbs[0]
        token_len = len(block_hashes) * db.block_size
        keys = [
            PoolKey(db.metadata, block_hash.hex()).to_string()
            for _, _, block_hash in db.process_tokens(token_len, block_hashes)
        ]
        assert len(keys) == len(block_hashes)
        # PUT side: mirrors KVCacheStoreSendingThread's striding slice.
        put_step = w._group_tp_replication_factors[0]
        put_keys.update(keys[w.tp_rank % put_step :: put_step])
        # GET side: KVCacheStoreRecvingThread fetches every key.
        get_keys_per_rank[tp_rank] = set(keys)

    for tp_rank, rank_keys in get_keys_per_rank.items():
        missing = rank_keys - put_keys
        assert not missing, (
            f"tp_rank={tp_rank} would GET {len(missing)}/{len(rank_keys)} keys "
            f"that no rank PUT (Mooncake OBJECT_NOT_FOUND): {sorted(missing)}"
        )


def test_requester_worker_group_semantics_falls_back_without_group_ids(
    tmp_path,
    monkeypatch,
):
    store = MagicMock()
    store.setup.return_value = 0
    _install_fake_mooncake(monkeypatch, store)
    _patch_worker_runtime(monkeypatch)
    warning = MagicMock()
    monkeypatch.setattr(worker.logger, "warning", warning)
    monkeypatch.setenv(
        "MOONCAKE_CONFIG_PATH",
        _write_mooncake_config(
            tmp_path,
            {
                "metadata_server": "http://metadata/endpoint",
                "protocol": "tcp",
                "device_name": "",
                "master_server_address": "10.0.0.7:50051",
            },
        ),
    )

    w = worker.MooncakeStoreWorker(
        _make_vllm_config(
            extra_config={
                "enable_group_semantics": True,
            }
        ),
        _make_kv_cache_config(),
    )

    assert w.enable_group_semantics is True
    assert w._supports_group_ids is False
    assert any(
        "does not support ReplicateConfig.group_ids" in call.args[0]
        for call in warning.call_args_list
    )


def test_requester_worker_group_semantics_string_false_does_not_enable(
    tmp_path,
    monkeypatch,
):
    store = MagicMock()
    store.setup.return_value = 0
    _install_fake_mooncake(monkeypatch, store)
    _patch_worker_runtime(monkeypatch)
    warning = MagicMock()
    monkeypatch.setattr(worker.logger, "warning", warning)
    monkeypatch.setenv(
        "MOONCAKE_CONFIG_PATH",
        _write_mooncake_config(
            tmp_path,
            {
                "metadata_server": "http://metadata/endpoint",
                "protocol": "tcp",
                "device_name": "",
                "master_server_address": "10.0.0.7:50051",
            },
        ),
    )

    w = worker.MooncakeStoreWorker(
        _make_vllm_config(
            extra_config={
                "enable_group_semantics": "false",
            }
        ),
        _make_kv_cache_config(),
    )

    assert w.enable_group_semantics is False
    assert not any(
        "does not support ReplicateConfig.group_ids" in call.args[0]
        for call in warning.call_args_list
    )


def test_requester_worker_group_semantics_string_true_enables(
    tmp_path,
    monkeypatch,
):
    store = MagicMock()
    store.setup.return_value = 0

    class FakeReplicateConfig:
        def __init__(self) -> None:
            self.group_ids = None

    fake_store_module = types.ModuleType("mooncake.store")
    fake_store_module.MooncakeDistributedStore = lambda: store  # type: ignore[attr-defined]
    fake_store_module.ReplicateConfig = FakeReplicateConfig  # type: ignore[attr-defined]
    fake_mooncake_module = types.ModuleType("mooncake")
    fake_mooncake_module.store = fake_store_module  # type: ignore[attr-defined]
    monkeypatch.setitem(sys.modules, "mooncake", fake_mooncake_module)
    monkeypatch.setitem(sys.modules, "mooncake.store", fake_store_module)
    _patch_worker_runtime(monkeypatch)
    monkeypatch.setenv(
        "MOONCAKE_CONFIG_PATH",
        _write_mooncake_config(
            tmp_path,
            {
                "metadata_server": "http://metadata/endpoint",
                "protocol": "tcp",
                "device_name": "",
                "master_server_address": "10.0.0.7:50051",
            },
        ),
    )

    w = worker.MooncakeStoreWorker(
        _make_vllm_config(
            extra_config={
                "enable_group_semantics": "true",
            }
        ),
        _make_kv_cache_config(),
    )

    assert w.enable_group_semantics is True
    assert w._supports_group_ids is True


# ---------------------------------------------------------------------------
# Helpers for register_kv_caches tests
# ---------------------------------------------------------------------------


def test_store_sending_thread_clamps_token_len_to_lcm():
    """Partial chunks past the last lcm boundary aren't stored — cache hits
    are always lcm-aligned (mirrors HybridKVCacheCoordinator)."""
    store = MagicMock()
    store.batch_is_exist.return_value = [0, 0]
    store.batch_put_from_multi_buffers.return_value = [256, 256]
    # Default coord: single full-attn block_size=16, lcm=16.
    # token_len_chunk=33 clamps to 32 → 2 chunks (not 3 with a partial 1-token chunk).
    thread = _make_store_sending_thread(store)

    thread.add_stored_request("r0")
    thread._handle_request(
        ReqMeta(
            req_id="r0",
            token_len_chunk=33,
            block_ids=([0, 1, 2],),
            block_hashes=[b"a0", b"a1", b"a2"],
            can_save=True,
        )
    )

    keys = store.batch_put_from_multi_buffers.call_args.args[0]
    assert len(keys) == 2


def test_store_sending_thread_skips_when_token_len_below_lcm():
    """Requests shorter than lcm_block_size cannot produce any aligned chunk,
    so neither the existence check nor the put should be issued."""
    from vllm.v1.kv_cache_interface import FullAttentionSpec, KVCacheGroupSpec

    store = MagicMock()
    # lcm=64 via single full-attn block_size=64.
    spec = FullAttentionSpec(block_size=64, num_kv_heads=8, head_size=64, dtype=None)
    coord = mooncake_store_worker.MooncakeStoreCoordinator(
        [KVCacheGroupSpec(["L"], spec)],
        scheduler_block_size=64,
        hash_block_size=64,
    )
    db = ChunkedTokenDatabase(
        KeyMetadata("test-model", 0, 0, 0, 0, group_id=0),
        block_size=64,
        hash_block_size=64,
    )
    db.set_kv_caches_base_addr([0x1000])
    db.set_block_len([1024])
    thread = _make_store_sending_thread(
        store, coord=coord, token_databases=[db], block_size=64
    )

    thread.add_stored_request("r0")
    thread._handle_request(
        ReqMeta(
            req_id="r0",
            token_len_chunk=32,
            block_ids=([0, 1],),
            block_hashes=[b"a0", b"a1"],
            can_save=True,
        )
    )

    store.batch_is_exist.assert_not_called()
    store.batch_put_from_multi_buffers.assert_not_called()
    assert thread.stored_requests["r0"] == 0


def test_store_sending_thread_only_stores_swa_blocks_in_window():
    """For SWA groups, only blocks within ``sliding_window`` of an
    lcm-aligned boundary are stored — a block outside every such window
    can never participate in any future hit (mirrors recv-side masking).

    Setup: full-attn (block_size=32) + SWA (block_size=8, sliding_window=8)
    → lcm=32. With token_len=64 there are two LCM boundaries (32, 64);
    the SWA group should store exactly one block per boundary: the block
    ending at 32 and the block ending at 64.
    """
    from vllm.v1.kv_cache_interface import (
        FullAttentionSpec,
        KVCacheGroupSpec,
        SlidingWindowSpec,
    )

    store = MagicMock()
    store.batch_is_exist.side_effect = lambda keys: [0] * len(keys)
    store.batch_put_from_multi_buffers.side_effect = (
        lambda keys, addrs, sizes, *_args: [256] * len(keys)
    )

    full_spec = FullAttentionSpec(
        block_size=32, num_kv_heads=8, head_size=64, dtype=None
    )
    swa_spec = SlidingWindowSpec(
        block_size=8,
        num_kv_heads=8,
        head_size=64,
        dtype=None,
        sliding_window=8,
    )
    coord = mooncake_store_worker.MooncakeStoreCoordinator(
        [KVCacheGroupSpec(["L0"], full_spec), KVCacheGroupSpec(["L1"], swa_spec)],
        scheduler_block_size=32,
        hash_block_size=8,
    )

    db_full = ChunkedTokenDatabase(
        KeyMetadata("test-model", 0, 0, 0, 0, group_id=0),
        block_size=32,
        hash_block_size=8,
    )
    db_full.set_kv_caches_base_addr([0x1000])
    db_full.set_block_len([512])
    db_swa = ChunkedTokenDatabase(
        KeyMetadata("test-model", 0, 0, 0, 0, group_id=1),
        block_size=8,
        hash_block_size=8,
    )
    db_swa.set_kv_caches_base_addr([0x2000])
    db_swa.set_block_len([128])

    thread = _make_store_sending_thread(
        store,
        coord=coord,
        token_databases=[db_full, db_swa],
        block_size=32,
    )

    hs = [bytes([i + 1]) * 4 for i in range(8)]
    thread.add_stored_request("r0")
    thread._handle_request(
        ReqMeta(
            req_id="r0",
            token_len_chunk=64,
            block_ids=([0, 1], list(range(8))),
            block_hashes=hs,
            can_save=True,
        )
    )

    keys = store.batch_put_from_multi_buffers.call_args.args[0]
    full_keys = [k for k in keys if "@group:0" in k]
    swa_keys = [k for k in keys if "@group:1" in k]
    # Full-attn: 2 blocks (chunks ending at 32 and 64).
    assert len(full_keys) == 2
    # SWA: only the two blocks ending at lcm boundaries (32 and 64), i.e.
    # blocks covering tokens [24,32) and [56,64) — hashes hs[3] and hs[7].
    assert len(swa_keys) == 2
    swa_hashes = {k.rsplit("@", 1)[-1] for k in swa_keys}
    assert swa_hashes == {hs[3].hex(), hs[7].hex()}


def test_store_sending_thread_delta_saves_only_new_swa_boundary_chunks():
    from vllm.v1.kv_cache_interface import (
        FullAttentionSpec,
        KVCacheGroupSpec,
        SlidingWindowSpec,
    )

    store = MagicMock()
    store.batch_is_exist.side_effect = lambda keys: [0] * len(keys)
    store.batch_put_from_multi_buffers.side_effect = (
        lambda keys, addrs, sizes, replicate_config: [256] * len(keys)
    )

    full_spec = FullAttentionSpec(
        block_size=32, num_kv_heads=8, head_size=64, dtype=None
    )
    swa_spec = SlidingWindowSpec(
        block_size=8,
        num_kv_heads=8,
        head_size=64,
        dtype=None,
        sliding_window=8,
    )
    coord = mooncake_store_worker.MooncakeStoreCoordinator(
        [KVCacheGroupSpec(["L0"], full_spec), KVCacheGroupSpec(["L1"], swa_spec)],
        scheduler_block_size=32,
        hash_block_size=8,
    )

    db_full = ChunkedTokenDatabase(
        KeyMetadata("test-model", 0, 0, 0, 0, group_id=0),
        block_size=32,
        hash_block_size=8,
    )
    db_full.set_kv_caches_base_addr([0x1000])
    db_full.set_block_len([512])
    db_swa = ChunkedTokenDatabase(
        KeyMetadata("test-model", 0, 0, 0, 0, group_id=1),
        block_size=8,
        hash_block_size=8,
    )
    db_swa.set_kv_caches_base_addr([0x2000])
    db_swa.set_block_len([128])

    thread = _make_store_sending_thread(
        store,
        coord=coord,
        token_databases=[db_full, db_swa],
        block_size=32,
    )

    hs = [bytes([i + 1]) * 4 for i in range(8)]
    thread.add_stored_request("r0")
    thread._saved_offset["r0"] = 32
    thread._handle_request(
        ReqMeta(
            req_id="r0",
            token_len_chunk=64,
            block_ids=([0, 1], list(range(8))),
            block_hashes=hs,
            can_save=True,
        )
    )

    keys = store.batch_put_from_multi_buffers.call_args.args[0]
    full_hashes = [k.rsplit("@", 1)[-1] for k in keys if "@group:0" in k]
    swa_hashes = [k.rsplit("@", 1)[-1] for k in keys if "@group:1" in k]
    assert full_hashes == [hs[7].hex()]
    assert swa_hashes == [hs[7].hex()]


def test_store_sending_thread_kv_events_use_group_chunk_metadata():
    from vllm.v1.core.kv_cache_utils import BlockHash, maybe_convert_block_hash
    from vllm.v1.kv_cache_interface import (
        FullAttentionSpec,
        KVCacheGroupSpec,
        SlidingWindowSpec,
    )

    store = MagicMock()
    store.batch_is_exist.side_effect = lambda keys: [0] * len(keys)
    store.batch_put_from_multi_buffers.return_value = [256, 256]

    full_spec = FullAttentionSpec(
        block_size=32, num_kv_heads=8, head_size=64, dtype=None
    )
    swa_spec = SlidingWindowSpec(
        block_size=8,
        num_kv_heads=8,
        head_size=64,
        dtype=None,
        sliding_window=8,
    )
    coord = mooncake_store_worker.MooncakeStoreCoordinator(
        [KVCacheGroupSpec(["L0"], full_spec), KVCacheGroupSpec(["L1"], swa_spec)],
        scheduler_block_size=32,
        hash_block_size=8,
    )

    db_full = ChunkedTokenDatabase(
        KeyMetadata("test-model", 0, 0, 0, 0, group_id=0),
        block_size=32,
        hash_block_size=8,
    )
    db_full.set_kv_caches_base_addr([0x1000])
    db_full.set_block_len([512])
    db_swa = ChunkedTokenDatabase(
        KeyMetadata("test-model", 0, 0, 0, 0, group_id=1),
        block_size=8,
        hash_block_size=8,
    )
    db_swa.set_kv_caches_base_addr([0x2000])
    db_swa.set_block_len([128])

    thread = _make_store_sending_thread(
        store,
        coord=coord,
        token_databases=[db_full, db_swa],
        block_size=32,
    )
    thread.enable_kv_event = True

    hs = [bytes([i + 1]) * 4 for i in range(4)]
    thread.add_stored_request("r0")
    thread._handle_request(
        ReqMeta(
            req_id="r0",
            token_len_chunk=32,
            block_ids=([0], list(range(4))),
            block_hashes=hs,
            can_save=True,
            token_ids=list(range(32)),
        )
    )

    full_event, swa_event = thread.get_kv_events()
    assert full_event.group_idx == 0
    assert full_event.block_size == 32
    assert full_event.token_ids == list(range(32))
    # block_size=32 over hash_block_size=8 (scale 4): the chunk is keyed by its
    # last sub-hash, not the concatenation of all four.
    assert full_event.block_hashes == [maybe_convert_block_hash(BlockHash(hs[3]))]

    assert swa_event.group_idx == 1
    assert swa_event.block_size == 8
    assert swa_event.token_ids == list(range(24, 32))
    assert swa_event.block_hashes == [maybe_convert_block_hash(BlockHash(hs[3]))]


def _auto_set_ready_event(*args, **kwargs):
    """Side effect for mocked thread constructors that auto-sets ready_event."""
    for arg in args:
        if isinstance(arg, threading.Event):
            arg.set()
    for val in kwargs.values():
        if isinstance(val, threading.Event):
            val.set()
    return MagicMock()


def _register_with_mocked_threads(
    worker: mooncake_store_worker.MooncakeStoreWorker,
    kv_caches: dict[str, torch.Tensor],
) -> None:
    """Call register_kv_caches with the I/O transfer threads mocked out."""
    prefix = "vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store.worker."
    with (
        patch(prefix + "KVCacheStoreSendingThread", side_effect=_auto_set_ready_event),
        patch(prefix + "KVCacheStoreRecvingThread", side_effect=_auto_set_ready_event),
    ):
        worker.register_kv_caches(kv_caches)


def _refresh_group_tp_replication_factors(
    worker: mooncake_store_worker.MooncakeStoreWorker,
) -> None:
    worker._group_tp_replication_factors = (
        worker._compute_group_tp_replication_factors()
    )
    worker._init_lookup_key_prefixes()


def _make_bare_worker(
    *,
    num_gpu_blocks: int = 10,
    block_size: int = 16,
    kv_role: str = "kv_both",
) -> mooncake_store_worker.MooncakeStoreWorker:
    """Construct a MooncakeStoreWorker via __new__, bypassing __init__.

    Sets only the attributes that register_kv_caches() reads so we can
    test the stride-based layout detection without a real
    MooncakeDistributedStore.
    """
    worker = object.__new__(mooncake_store_worker.MooncakeStoreWorker)
    worker.cache_config = MagicMock()
    worker.cache_config.num_gpu_blocks = num_gpu_blocks
    worker.store = MagicMock()
    worker.store.register_buffer.return_value = 0
    worker.kv_role = kv_role
    worker._capacity_only = False
    worker.block_size = block_size
    worker.tp_rank = 0
    worker.enable_kv_events = False
    worker.kv_send_thread = None
    worker.kv_recv_threads = []
    worker.num_recv_threads = 1
    worker.recv_request_queue = queue.Queue()
    worker.tp_size = 1
    worker.num_kv_head = 1
    worker.pp_size = 1
    # Minimal single-full-attention-group config so the coordinator-based
    # lookup path works (the connector no longer carries a legacy single-group
    # path; everything flows through the coordinator).
    from vllm.v1.kv_cache_interface import (
        FullAttentionSpec,
        KVCacheGroupSpec,
    )

    worker.disk_offload_buffer_budget_bytes = None
    worker.store_replicate_config = SimpleNamespace()
    worker.enable_group_semantics = False
    worker._supports_group_ids = False
    worker._kv_connector_stats_lock = threading.Lock()
    worker.kv_connector_stats = MooncakeStoreConnectorStats()

    spec = FullAttentionSpec(
        block_size=block_size, num_kv_heads=8, head_size=64, dtype=None
    )
    group = KVCacheGroupSpec(["layer0", "__cross_layer__"], spec)
    worker._kv_cache_groups = [group]
    worker.pcp_size = 1
    worker.dcp_size = 1
    worker.hash_block_size = block_size
    # Pre-build a single-group token_dbs so lookup-only tests don't have to
    # call register_kv_caches.
    worker.token_dbs = [
        ChunkedTokenDatabase(
            KeyMetadata("test-model", 0, 0, 0, 0, group_id=0),
            block_size=block_size,
            hash_block_size=block_size,
        )
    ]
    worker.coord = mooncake_store_worker.MooncakeStoreCoordinator(
        worker._kv_cache_groups,
        scheduler_block_size=block_size,
        hash_block_size=block_size,
    )
    _refresh_group_tp_replication_factors(worker)
    return worker


def test_lookup_key_prefixes_cover_dcp_rank_namespaces():
    worker = _make_bare_worker()
    worker.tp_size = 4
    worker.num_kv_head = 1
    worker.dcp_size = 4
    _refresh_group_tp_replication_factors(worker)

    assert worker._lookup_key_prefixes[0] == (
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:0",
        "test-model@tp_rank:1@pcp0@dcp1@pp_rank:0@group:0",
        "test-model@tp_rank:2@pcp0@dcp2@pp_rank:0@group:0",
        "test-model@tp_rank:3@pcp0@dcp3@pp_rank:0@group:0",
    )


def test_lookup_key_prefixes_cover_pcp_rank_namespaces():
    worker = _make_bare_worker()
    worker.tp_size = 4
    worker.num_kv_head = 1
    worker.pcp_size = 2
    worker.dcp_size = 1
    _refresh_group_tp_replication_factors(worker)

    assert worker._lookup_key_prefixes[0] == (
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:0",
        "test-model@tp_rank:0@pcp1@dcp0@pp_rank:0@group:0",
    )


def test_lookup_key_prefixes_expand_tp_sharded_groups_per_rank():
    """Replicated attention needs one namespace; sharded Mamba needs every rank."""
    from vllm.v1.kv_cache_interface import (
        FullAttentionSpec,
        KVCacheGroupSpec,
        MambaSpec,
    )

    worker = _make_bare_worker(block_size=16)
    worker.tp_size = 2
    worker.num_kv_head = 1
    fa = FullAttentionSpec(block_size=16, num_kv_heads=8, head_size=64, dtype=None)
    mamba = MambaSpec(
        block_size=16,
        shapes=((1, 1),),
        dtypes=(torch.float32,),
        mamba_cache_mode="align",
    )
    worker._kv_cache_groups = [
        KVCacheGroupSpec(["l0"], fa),
        KVCacheGroupSpec(["l1"], mamba),
    ]
    worker.token_dbs = [
        ChunkedTokenDatabase(
            KeyMetadata("test-model", 0, 0, 0, 0, group_id=0), block_size=16
        ),
        ChunkedTokenDatabase(
            KeyMetadata("test-model", 1, 0, 0, 0, group_id=1), block_size=16
        ),
    ]
    _refresh_group_tp_replication_factors(worker)

    assert worker._lookup_key_prefixes[0] == (
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:0",
    )
    assert worker._lookup_key_prefixes[1] == (
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:1",
        "test-model@tp_rank:1@pcp0@dcp0@pp_rank:0@group:1",
    )


def test_group_tp_replication_factors_mixed_mla_gqa_mamba():
    from vllm.v1.kv_cache_interface import (
        FullAttentionSpec,
        KVCacheGroupSpec,
        MambaSpec,
        MLAAttentionSpec,
    )

    worker = _make_bare_worker(block_size=16)
    worker.tp_size = 4
    worker.num_kv_head = 2
    mla = MLAAttentionSpec(block_size=16, num_kv_heads=1, head_size=64, dtype=None)
    gqa = FullAttentionSpec(block_size=16, num_kv_heads=8, head_size=64, dtype=None)
    mamba = MambaSpec(
        block_size=16,
        shapes=((1, 1),),
        dtypes=(torch.float32,),
        mamba_cache_mode="align",
    )
    worker._kv_cache_groups = [
        KVCacheGroupSpec(["l0"], mla),
        KVCacheGroupSpec(["l1"], gqa),
        KVCacheGroupSpec(["l2"], mamba),
    ]
    worker.token_dbs = [
        ChunkedTokenDatabase(
            KeyMetadata("test-model", 0, 0, 0, 0, group_id=g_idx), block_size=16
        )
        for g_idx in range(3)
    ]

    _refresh_group_tp_replication_factors(worker)
    assert worker._group_tp_replication_factors == (4, 2, 1)
    assert worker._lookup_key_prefixes[0] == (
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:0",
    )
    assert worker._lookup_key_prefixes[1] == (
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:1",
        "test-model@tp_rank:1@pcp0@dcp0@pp_rank:0@group:1",
    )
    assert worker._lookup_key_prefixes[2] == (
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:2",
        "test-model@tp_rank:1@pcp0@dcp0@pp_rank:0@group:2",
        "test-model@tp_rank:2@pcp0@dcp0@pp_rank:0@group:2",
        "test-model@tp_rank:3@pcp0@dcp0@pp_rank:0@group:2",
    )


@pytest.mark.parametrize("spec_order", [("mla", "gqa"), ("gqa", "mla")])
def test_uniform_group_uses_common_inner_replication_factor(spec_order):
    from vllm.v1.kv_cache_interface import (
        FullAttentionSpec,
        KVCacheGroupSpec,
        MLAAttentionSpec,
        UniformTypeKVCacheSpecs,
    )

    worker = _make_bare_worker(block_size=16)
    worker.tp_size = 4
    worker.num_kv_head = 2
    specs_by_name = {
        "mla": MLAAttentionSpec(
            block_size=16, num_kv_heads=1, head_size=64, dtype=None
        ),
        "gqa": FullAttentionSpec(
            block_size=16, num_kv_heads=1, head_size=64, dtype=None
        ),
    }
    inner_specs = {name: specs_by_name[name] for name in spec_order}
    uniform_spec = UniformTypeKVCacheSpecs(
        block_size=16,
        kv_cache_specs=inner_specs,
    )
    worker._kv_cache_groups = [
        KVCacheGroupSpec(list(inner_specs), uniform_spec),
    ]

    _refresh_group_tp_replication_factors(worker)

    assert worker._group_tp_replication_factors == (2,)
    assert worker._lookup_key_prefixes[0] == (
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:0",
        "test-model@tp_rank:1@pcp0@dcp0@pp_rank:0@group:0",
    )


def test_lookup_rejects_boundary_missing_one_mamba_shard():
    from vllm.v1.kv_cache_interface import (
        FullAttentionSpec,
        KVCacheGroupSpec,
        MambaSpec,
    )

    worker = _make_bare_worker(block_size=16)
    worker.tp_size = 2
    worker.num_kv_head = 1
    fa = FullAttentionSpec(block_size=16, num_kv_heads=8, head_size=64, dtype=None)
    mamba = MambaSpec(
        block_size=16,
        shapes=((1, 1),),
        dtypes=(torch.float32,),
        mamba_cache_mode="align",
    )
    worker._kv_cache_groups = [
        KVCacheGroupSpec(["l0"], fa),
        KVCacheGroupSpec(["l1"], mamba),
    ]
    worker.token_dbs = [
        ChunkedTokenDatabase(
            KeyMetadata("test-model", 0, 0, 0, 0, group_id=0), block_size=16
        ),
        ChunkedTokenDatabase(
            KeyMetadata("test-model", 1, 0, 0, 0, group_id=1), block_size=16
        ),
    ]
    worker.coord = mooncake_store_worker.MooncakeStoreCoordinator(
        worker._kv_cache_groups,
        scheduler_block_size=16,
        hash_block_size=16,
    )
    _refresh_group_tp_replication_factors(worker)

    # 33 tokens for two 16-token blocks: the hit stops below the request end,
    # so the full-hit re-derivation stays out of the shard accounting.
    worker.store.batch_is_exist.side_effect = lambda keys: [1] * len(keys)
    assert worker.lookup(33, [b"h0", b"h1"]) == 32

    worker.store.batch_is_exist.side_effect = lambda keys: [
        0 if "tp_rank:1" in k and "group:1" in k else 1 for k in keys
    ]
    assert worker.lookup(33, [b"h0", b"h1"]) == 0


def test_lookup_requires_all_dcp_rank_namespaces():
    worker = _make_bare_worker(block_size=16)
    worker.tp_size = 4
    worker.num_kv_head = 1
    worker.dcp_size = 4
    _refresh_group_tp_replication_factors(worker)
    worker.store.batch_is_exist.return_value = [1, 1, 0, 1]

    assert worker.lookup(16, [b"a0"]) == 0
    assert worker.store.batch_is_exist.call_args.args[0] == [
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:0@6130",
        "test-model@tp_rank:1@pcp0@dcp1@pp_rank:0@group:0@6130",
        "test-model@tp_rank:2@pcp0@dcp2@pp_rank:0@group:0@6130",
        "test-model@tp_rank:3@pcp0@dcp3@pp_rank:0@group:0@6130",
    ]


def test_lookup_partial_prefix_returns_first_hit_length():
    worker = _make_bare_worker()
    worker.store.batch_is_exist.return_value = [1, 1, 0]
    assert worker.lookup(48, [b"a0", b"a1", b"a2"]) == 32


def test_lookup_partial_tail_uses_hash_alignment():
    """A stored sub-block tail can serve a request extending past it."""
    from vllm.v1.kv_cache_interface import (
        FullAttentionSpec,
        KVCacheGroupSpec,
        MambaSpec,
    )

    worker = _make_bare_worker(block_size=16)
    full = FullAttentionSpec(block_size=16, num_kv_heads=8, head_size=64, dtype=None)
    mamba = MambaSpec(
        block_size=16,
        shapes=((1, 1),),
        dtypes=(torch.float32,),
        mamba_cache_mode="align",
    )
    worker._kv_cache_groups = [
        KVCacheGroupSpec(["full"], full),
        KVCacheGroupSpec(["mamba"], mamba),
    ]
    worker.hash_block_size = 4
    worker.token_dbs = [
        ChunkedTokenDatabase(
            KeyMetadata("test-model", 0, 0, 0, 0, group_id=group_id),
            block_size=16,
            hash_block_size=4,
        )
        for group_id in range(2)
    ]
    worker.coord = mooncake_store_worker.MooncakeStoreCoordinator(
        worker._kv_cache_groups,
        scheduler_block_size=16,
        hash_block_size=4,
    )
    _refresh_group_tp_replication_factors(worker)
    worker.store.batch_is_exist.return_value = [0, 0, 1, 0, 0, 1]

    assert worker.lookup(13, [b"h0", b"h1", b"h2"]) == 12


def test_lookup_full_hit_reuses_existing_boundary():
    """A full hit is re-derived below the request end without another RPC."""
    worker = _make_bare_worker(block_size=16)
    worker.store.batch_is_exist.return_value = [1, 1]

    assert worker.lookup(32, [b"h0", b"h1"]) == 16
    assert worker.store.batch_is_exist.call_count == 1


def test_lookup_full_hit_with_eagle_pops_once_not_twice():
    """Eagle already leaves the last block for the drafter, so a
    full-prompt re-derivation must never fire for eagle-governed hits:
    firing would anchor the search one block lower and pop a second
    block, regressing the hit by an extra producer boundary."""
    worker = _make_bare_worker(block_size=16)
    worker.coord = mooncake_store_worker.MooncakeStoreCoordinator(
        worker._kv_cache_groups,
        scheduler_block_size=16,
        hash_block_size=16,
        use_eagle=True,
    )
    worker.store.batch_is_exist.return_value = [1, 1, 1, 1]

    # 64-token exact-multiple prompt, all 4 blocks stored: one eagle pop
    # gives 48; a spurious re-derivation (anchored at 48) would pop again
    # and return 32.
    assert worker.lookup(64, [b"h0", b"h1", b"h2", b"h3"]) == 48
    assert worker.store.batch_is_exist.call_count == 1


def test_lookup_full_hit_swa_degrades_when_no_stored_boundary_is_usable():
    """The motivating livelock: the producer of a 64-token prompt stored
    only its SWA tail window (blocks 2-3). The old arithmetic clamp turned
    the full hit into 48, whose SWA window needs the never-written block 1,
    so every load failed and the recompute re-entered the same lookup. The
    re-derivation must report that no stored boundary below the request end
    is usable."""
    from vllm.v1.kv_cache_interface import KVCacheGroupSpec, SlidingWindowSpec

    worker = _make_bare_worker(block_size=16)
    swa = SlidingWindowSpec(
        block_size=16, num_kv_heads=8, head_size=64, dtype=None, sliding_window=32
    )
    worker._kv_cache_groups = [KVCacheGroupSpec(["layer0"], swa)]
    worker.coord = mooncake_store_worker.MooncakeStoreCoordinator(
        worker._kv_cache_groups,
        scheduler_block_size=worker.hash_block_size,
        hash_block_size=worker.hash_block_size,
    )
    worker.store.batch_is_exist.return_value = [0, 0, 1, 1]

    assert worker.lookup(64, [b"h0", b"h1", b"h2", b"h3"]) == 0
    assert worker.store.batch_is_exist.call_count == 1


def test_lookup_swa_single_group_returns_full_when_tail_window_present():
    """Single-SWA, sliding_window=32 (= 2 blocks): producer stored only the
    tail. Coordinator-driven lookup returns full prefix even though the
    pre-window blocks are absent."""
    from vllm.v1.kv_cache_interface import KVCacheGroupSpec, SlidingWindowSpec

    worker = _make_bare_worker(block_size=16)
    swa = SlidingWindowSpec(
        block_size=16, num_kv_heads=8, head_size=64, dtype=None, sliding_window=32
    )
    worker._kv_cache_groups = [KVCacheGroupSpec(["layer0"], swa)]
    worker.coord = mooncake_store_worker.MooncakeStoreCoordinator(
        worker._kv_cache_groups,
        scheduler_block_size=worker.hash_block_size,
        hash_block_size=worker.hash_block_size,
    )
    worker.store.batch_is_exist.return_value = [0, 0, 1, 1]
    assert worker.lookup(65, [b"h0", b"h1", b"h2", b"h3"]) == 64


def test_lookup_checks_all_potential_swa_hit_boundaries():
    """Lookup should skip SWA chunks that can never validate a hit, but still
    check earlier aligned boundaries when sparse retention stores only the
    current request's replay boundary.
    """
    from vllm.v1.kv_cache_interface import (
        FullAttentionSpec,
        KVCacheGroupSpec,
        SlidingWindowSpec,
    )

    worker = _make_bare_worker(block_size=8)
    full = FullAttentionSpec(block_size=32, num_kv_heads=8, head_size=64, dtype=None)
    swa = SlidingWindowSpec(
        block_size=8, num_kv_heads=8, head_size=64, dtype=None, sliding_window=8
    )
    worker._kv_cache_groups = [
        KVCacheGroupSpec(["full"], full),
        KVCacheGroupSpec(["swa"], swa),
    ]
    worker.token_dbs = [
        ChunkedTokenDatabase(
            KeyMetadata("test-model", 0, 0, 0, 0, group_id=0),
            block_size=32,
            hash_block_size=8,
        ),
        ChunkedTokenDatabase(
            KeyMetadata("test-model", 0, 0, 0, 0, group_id=1),
            block_size=8,
            hash_block_size=8,
        ),
    ]
    worker.coord = mooncake_store_worker.MooncakeStoreCoordinator(
        worker._kv_cache_groups,
        scheduler_block_size=32,
        hash_block_size=8,
        retention_interval=0,
    )
    _refresh_group_tp_replication_factors(worker)
    # Candidate order: 3 full-attention chunks, then SWA chunks 3, 7, 11.
    # Only the first full chunk and the SWA chunk ending at token 32 exist, so
    # lookup should recover a 32-token external prefix hit. A sparse
    # prompt-specific store mask for num_prompt_tokens=96 would only check SWA
    # chunk 7 and miss this earlier reusable prefix.
    worker.store.batch_is_exist.return_value = [1, 0, 0, 1, 0, 0]

    result = worker.lookup(
        96,
        [f"h{i}".encode() for i in range(12)],
    )

    assert result == 32
    keys = worker.store.batch_is_exist.call_args.args[0]
    assert len(keys) == 6
    swa_keys = [key for key in keys if "@group:1@" in key]
    assert swa_keys == [
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:1@6833",
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:1@6837",
        "test-model@tp_rank:0@pcp0@dcp0@pp_rank:0@group:1@683131",
    ]


def test_lookup_applies_swa_mask_before_accessing_hashes():
    """Lookup must apply the sparse SWA mask before touching group hashes, so
    false-mask chunks pay neither the hash access nor the key-construction cost.
    """
    from vllm.v1.kv_cache_interface import (
        FullAttentionSpec,
        KVCacheGroupSpec,
        SlidingWindowSpec,
    )

    worker = _make_bare_worker(block_size=8)
    full = FullAttentionSpec(block_size=32, num_kv_heads=8, head_size=64, dtype=None)
    swa = SlidingWindowSpec(
        block_size=8, num_kv_heads=8, head_size=64, dtype=None, sliding_window=8
    )
    worker._kv_cache_groups = [
        KVCacheGroupSpec(["full"], full),
        KVCacheGroupSpec(["swa"], swa),
    ]
    worker.token_dbs = [
        ChunkedTokenDatabase(
            KeyMetadata("test-model", 0, 0, 0, 0, group_id=0),
            block_size=32,
            hash_block_size=8,
        ),
        ChunkedTokenDatabase(
            KeyMetadata("test-model", 0, 0, 0, 0, group_id=1),
            block_size=8,
            hash_block_size=8,
        ),
    ]
    worker.coord = mooncake_store_worker.MooncakeStoreCoordinator(
        worker._kv_cache_groups,
        scheduler_block_size=32,
        hash_block_size=8,
        retention_interval=0,
    )
    _refresh_group_tp_replication_factors(worker)

    block_hashes = _RecordingBlockHashes([f"h{i}".encode() for i in range(12)])
    accessed_before_rpc: list[int] = []

    def exists(keys):
        accessed_before_rpc.extend(block_hashes.accessed)
        return [0] * len(keys)

    worker.store.batch_is_exist.side_effect = exists

    worker.lookup(96, block_hashes)

    # Full-attention chunks (compact hashes at 3, 7, 11) plus only the reachable
    # SWA boundary tails (chunks 3, 7, 11) are hash-accessed. Every masked SWA
    # chunk in between is skipped before both the hash access and key build.
    assert accessed_before_rpc == [3, 7, 11, 3, 7, 11]


# ---------------------------------------------------------------------------
# register_kv_caches tests
# ---------------------------------------------------------------------------


def test_register_kv_caches_blocks_first_single_segment():
    """Blocks-first layout (FlashInfer/MLA): one segment per layer."""
    num_blocks = 10
    page_size_elements = 64
    worker = _make_bare_worker(num_gpu_blocks=num_blocks)

    # Shape: (num_blocks, page_size_elements) — blocks outermost, no outer_dims
    tensor = torch.zeros(num_blocks, page_size_elements, dtype=torch.float16)
    _register_with_mocked_threads(worker, {"layer0": tensor})

    db = worker.token_dbs[0]
    assert db.kv_caches_base_addr == [tensor.untyped_storage().data_ptr()]
    assert db.block_len == [tensor.untyped_storage().nbytes() // num_blocks]
    worker.store.register_buffer.assert_called_once_with(
        tensor.untyped_storage().data_ptr(),
        tensor.untyped_storage().nbytes(),
    )


def test_register_kv_caches_kv_first_two_segments():
    """K/V-first layout (FlashAttn): two segments (K, V) per layer."""
    num_blocks = 10
    block_size_tokens = 16
    num_kv_heads = 4
    head_size = 8
    worker = _make_bare_worker(num_gpu_blocks=num_blocks)

    # Shape: (2, num_blocks, block_size, num_kv_heads, head_size) — K/V outermost
    tensor = torch.zeros(
        2,
        num_blocks,
        block_size_tokens,
        num_kv_heads,
        head_size,
        dtype=torch.float16,
    )
    _register_with_mocked_threads(worker, {"layer0": tensor})

    db = worker.token_dbs[0]
    seg_stride = tensor.stride(0) * tensor.element_size()
    base = tensor.untyped_storage().data_ptr()
    assert db.kv_caches_base_addr == [base, base + seg_stride]
    assert db.block_len == [seg_stride // num_blocks] * 2


def test_register_kv_caches_cross_layer_single_segment():
    """Cross-layer tensor: single segment with block_len = page_size * num_layers."""
    num_blocks = 10
    num_layers = 4
    per_layer_page_elements = 64  # elements per layer per block

    worker = _make_bare_worker(num_gpu_blocks=num_blocks)

    # Cross-layer blocks-first tensor: all layers packed into a single
    # contiguous block.  Shape (num_blocks, num_layers * per_layer_page)
    # mimics the physical layout after stride reordering.
    total_page_elements = num_layers * per_layer_page_elements
    tensor = torch.zeros(num_blocks, total_page_elements, dtype=torch.float16)

    with (
        patch(
            "vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store."
            "worker.KVCacheStoreSendingThread",
            side_effect=_auto_set_ready_event,
        ),
        patch(
            "vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store."
            "worker.KVCacheStoreRecvingThread",
            side_effect=_auto_set_ready_event,
        ),
    ):
        # Use the cross-layer wrapper key, same as register_cross_layers_kv_caches
        worker.register_kv_caches({"__cross_layer__": tensor})

    db = worker.token_dbs[0]
    assert len(db.kv_caches_base_addr) == 1
    assert db.kv_caches_base_addr[0] == tensor.untyped_storage().data_ptr()

    expected_block_len = tensor.untyped_storage().nbytes() // num_blocks
    # block_len should be per_layer_page_size * num_layers
    assert (
        expected_block_len
        == num_layers * per_layer_page_elements * tensor.element_size()
    )
    assert len(db.block_len) == 1
    assert db.block_len[0] == expected_block_len

    # Also verify via register_cross_layers_kv_caches wrapper
    worker2 = _make_bare_worker(num_gpu_blocks=num_blocks)
    with (
        patch(
            "vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store."
            "worker.KVCacheStoreSendingThread",
            side_effect=_auto_set_ready_event,
        ),
        patch(
            "vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store."
            "worker.KVCacheStoreRecvingThread",
            side_effect=_auto_set_ready_event,
        ),
    ):
        worker2.register_cross_layers_kv_caches(tensor)

    db2 = worker2.token_dbs[0]
    assert db2.kv_caches_base_addr == db.kv_caches_base_addr
    assert db2.block_len == db.block_len


# ---------------------------------------------------------------------------
# Dual-mode (embedded / standalone-store) config validation tests
# ---------------------------------------------------------------------------


def _make_config(**overrides):
    """Build a MooncakeStoreConfig with sensible defaults for validation tests.

    Required dataclass fields are populated; callers override only the field
    under test.
    """
    base = dict(
        metadata_server="http://metadata/endpoint",
        master_server_address="10.0.0.7:50051",
        protocol="rdma",
        device_name="mlx5_0",
    )
    base.update(overrides)
    return worker.MooncakeStoreConfig(**base)


def test_config_defaults_to_embedded():
    """A JSON without explicit mode parses as embedded with 4 GiB segment."""
    cfg = _make_config()
    assert cfg.mode == "embedded"
    assert cfg.global_segment_size == worker.DEFAULT_GLOBAL_SEGMENT_SIZE
    assert cfg.local_buffer_size == worker.DEFAULT_LOCAL_BUFFER_SIZE
    assert cfg.enable_offload is False
    assert cfg.tenant_id == worker.DEFAULT_TENANT_ID


def test_config_pr40900_unchanged(tmp_path):
    """A literal PR-40900 config (no mode, no enable_offload, no preferred_segment)
    parses without raising and resolves to embedded mode."""
    config_path = _write_mooncake_config(
        tmp_path,
        {
            "metadata_server": "http://metadata/endpoint",
            "global_segment_size": "4GB",
            "local_buffer_size": "4GB",
            "protocol": "rdma",
            "device_name": "mlx5_0",
            "master_server_address": "10.0.0.7:50051",
        },
    )
    cfg = worker.MooncakeStoreConfig.from_file(config_path)
    assert cfg.mode == "embedded"
    assert cfg.global_segment_size == 4 * 1024**3
    assert cfg.local_buffer_size == 4 * 1024**3
    assert cfg.enable_offload is False
    assert cfg.tenant_id == worker.DEFAULT_TENANT_ID


def test_config_from_file_normalizes_tenant_id(tmp_path):
    config_path = _write_mooncake_config(
        tmp_path,
        {
            "metadata_server": "http://metadata/endpoint",
            "master_server_address": "10.0.0.7:50051",
            "tenant_id": "  tenant-a  ",
        },
    )

    cfg = worker.MooncakeStoreConfig.from_file(config_path)

    assert cfg.tenant_id == "tenant-a"


def test_config_from_file_normalizes_empty_tenant_id_to_default(tmp_path):
    config_path = _write_mooncake_config(
        tmp_path,
        {
            "metadata_server": "http://metadata/endpoint",
            "master_server_address": "10.0.0.7:50051",
            "tenant_id": "   ",
        },
    )

    cfg = worker.MooncakeStoreConfig.from_file(config_path)

    assert cfg.tenant_id == worker.DEFAULT_TENANT_ID


def test_config_from_file_rejects_non_string_tenant_id(tmp_path):
    config_path = _write_mooncake_config(
        tmp_path,
        {
            "metadata_server": "http://metadata/endpoint",
            "master_server_address": "10.0.0.7:50051",
            "tenant_id": False,
        },
    )

    with pytest.raises(TypeError, match="tenant_id must be a string or null"):
        worker.MooncakeStoreConfig.from_file(config_path)


def test_config_embedded_rejects_zero_segment():
    with pytest.raises(
        ValueError, match=r"embedded mode requires global_segment_size > 0"
    ):
        _make_config(mode="embedded", global_segment_size=0)


def test_config_standalone_store_rejects_nonzero_segment():
    with pytest.raises(
        ValueError,
        match=r"standalone-store mode requires global_segment_size == 0",
    ):
        _make_config(mode="standalone-store", global_segment_size=4 * 1024**3)


def test_config_standalone_store_accepts_zero_segment():
    cfg = _make_config(mode="standalone-store", global_segment_size=0)
    assert cfg.mode == "standalone-store"
    assert cfg.global_segment_size == 0


def test_config_unknown_mode():
    with pytest.raises(ValueError, match=r"unknown Mooncake mode"):
        _make_config(mode="something-else")


def test_config_zero_local_buffer():
    with pytest.raises(ValueError, match=r"local_buffer_size must be > 0"):
        _make_config(local_buffer_size=0)


# ---------------------------------------------------------------------------
# End-to-end topology tests
# Covers the two supported recipes:
#   (A) standalone-store mode + disk offload (mode="standalone-store",
#       segment=0, enable_offload=true, preferred_segment set)
#   (B) embedded mode + CPU only      (mode default, segment>0,
#       enable_offload=false, no preferred_segment)
# ---------------------------------------------------------------------------


def test_topology_standalone_store_with_disk_offload(tmp_path, monkeypatch):
    """standalone-store + disk: global_segment_size=0, enable_offload=True,
    preferred_segment set. Assert setup() positional args, ReplicateConfig
    wiring, and that the disk-offload buffer budget is allocated."""
    store = MagicMock()
    store.setup.return_value = 0
    fake_replicate_config_cls = _install_fake_mooncake(monkeypatch, store)
    _patch_worker_runtime(monkeypatch)
    monkeypatch.setenv(
        "MOONCAKE_CONFIG_PATH",
        _write_mooncake_config(
            tmp_path,
            {
                "mode": "standalone-store",
                "metadata_server": "http://metadata/endpoint",
                "global_segment_size": 0,
                "local_buffer_size": "1GB",
                "protocol": "rdma",
                "device_name": "mlx5_0",
                "master_server_address": "10.0.0.7:50051",
                "enable_offload": True,
            },
        ),
    )

    w = worker.MooncakeStoreWorker(
        _make_vllm_config(extra_config={"preferred_segment": "10.0.0.7:50053"}),
        _make_kv_cache_config(),
    )

    # setup() receives global_segment_size=0 and the configured local buffer.
    assert store.setup.call_args.args == (
        "10.0.0.7",
        "http://metadata/endpoint",
        0,
        1024 * 1024 * 1024,
        "rdma",
        "mlx5_0",
        "10.0.0.7:50051",
    )
    assert store.setup.call_args.kwargs == {}
    # ReplicateConfig is built and carries the preferred_segment.
    assert isinstance(w.store_replicate_config, fake_replicate_config_cls)
    assert w.store_replicate_config.preferred_segment == "10.0.0.7:50053"
    # Disk-offload staging budget is allocated (enable_offload=True).
    assert w.disk_offload_buffer_budget_bytes is not None
    assert w.disk_offload_buffer_budget_bytes > 0


def test_topology_embedded_cpu_only(tmp_path, monkeypatch):
    """embedded + CPU-only: no mode key (defaults to embedded),
    global_segment_size>0, enable_offload absent, no preferred_segment.
    This is the PR-40900 baseline recipe."""
    store = MagicMock()
    store.setup.return_value = 0
    fake_replicate_config_cls = _install_fake_mooncake(monkeypatch, store)
    _patch_worker_runtime(monkeypatch)
    monkeypatch.setenv(
        "MOONCAKE_CONFIG_PATH",
        _write_mooncake_config(
            tmp_path,
            {
                "metadata_server": "http://metadata/endpoint",
                "global_segment_size": "4GB",
                "local_buffer_size": "4GB",
                "protocol": "rdma",
                "device_name": "mlx5_0",
                "master_server_address": "10.0.0.7:50051",
            },
        ),
    )

    w = worker.MooncakeStoreWorker(_make_vllm_config(), _make_kv_cache_config())

    # setup() receives global_segment_size=4 GiB (rank contributes a segment).
    assert store.setup.call_args.args == (
        "10.0.0.7",
        "http://metadata/endpoint",
        4 * 1024 * 1024 * 1024,
        4 * 1024 * 1024 * 1024,
        "rdma",
        "mlx5_0",
        "10.0.0.7:50051",
    )
    assert store.setup.call_args.kwargs == {}
    # No preferred_segment — ReplicateConfig is default-constructed (so the
    # preferred_segment field keeps its default value).
    assert w.preferred_segment is None
    assert isinstance(w.store_replicate_config, fake_replicate_config_cls)
    assert w.store_replicate_config.preferred_segment == ""
    # No disk budget — enable_offload was absent (defaults to False).
    assert w.disk_offload_buffer_budget_bytes is None


def test_topology_forwards_non_default_tenant_id(tmp_path, monkeypatch):
    store = MagicMock()
    store.setup.return_value = 0
    _install_fake_mooncake(monkeypatch, store)
    _patch_worker_runtime(monkeypatch)
    monkeypatch.setenv(
        "MOONCAKE_CONFIG_PATH",
        _write_mooncake_config(
            tmp_path,
            {
                "metadata_server": "http://metadata/endpoint",
                "global_segment_size": "4GB",
                "local_buffer_size": "4GB",
                "protocol": "rdma",
                "device_name": "mlx5_0",
                "master_server_address": "10.0.0.7:50051",
                "tenant_id": "  tenant-a  ",
            },
        ),
    )

    worker.MooncakeStoreWorker(_make_vllm_config(rank=1), _make_kv_cache_config())

    assert store.setup.call_args.kwargs == {"tenant_id": "tenant-a"}


def test_non_default_tenant_preserves_setup_type_error(tmp_path, monkeypatch):
    store = MagicMock()
    setup_error = TypeError(
        "setup(): incompatible function arguments; "
        "supported signature includes tenant_id"
    )
    store.setup.side_effect = setup_error
    _install_fake_mooncake(monkeypatch, store)
    _patch_worker_runtime(monkeypatch)
    monkeypatch.setenv(
        "MOONCAKE_CONFIG_PATH",
        _write_mooncake_config(
            tmp_path,
            {
                "metadata_server": "http://metadata/endpoint",
                "global_segment_size": "4GB",
                "local_buffer_size": "4GB",
                "protocol": "rdma",
                "device_name": "mlx5_0",
                "master_server_address": "10.0.0.7:50051",
                "tenant_id": "tenant-a",
            },
        ),
    )

    with pytest.raises(TypeError) as exc_info:
        worker.MooncakeStoreWorker(_make_vllm_config(rank=1), _make_kv_cache_config())

    assert exc_info.value is setup_error


# ---------------------------------------------------------------------------
# Stats/metrics tests (PR-35 port)
# ---------------------------------------------------------------------------


def test_mooncake_store_stats_aggregate_reduce():
    stats = MooncakeStoreConnectorStats()
    stats.record_operation("save_put", 0.01, 2, num_bytes=128)
    other = MooncakeStoreConnectorStats()
    other.record_operation(
        "save_put",
        0.03,
        1,
        num_bytes=64,
        status="error",
        num_failed_keys=1,
    )

    reduced = stats.aggregate(other).reduce()

    assert reduced["save_put_count"] == 2
    assert reduced["save_put_total_keys"] == 3
    assert reduced["save_put_total_bytes"] == 192
    assert reduced["save_put_failed_keys"] == 1
    assert reduced["save_put_error_count"] == 1


def test_worker_get_kv_connector_stats_resets_after_read():
    worker = _make_bare_worker()
    worker._record_kv_connector_operation(
        "save_put",
        0.01,
        2,
        num_bytes=128,
    )

    stats = worker.get_kv_connector_stats()

    assert isinstance(stats, MooncakeStoreConnectorStats)
    assert stats.data["save_put"][0]["num_bytes"] == 128
    assert worker.get_kv_connector_stats() is None


def test_lookup_records_mooncake_metrics():
    worker = _make_bare_worker()
    worker.store.batch_is_exist.return_value = [1, 1]

    result = worker.lookup(33, [b"a0", b"a1"])
    stats = worker.get_kv_connector_stats()

    assert result == 32
    assert isinstance(stats, MooncakeStoreConnectorStats)
    assert len(stats.data["lookup_exists"]) == 1
    assert stats.data["lookup_exists"][0]["num_keys"] == 2


def test_store_worker_close_releases_store():
    worker = _make_bare_worker()
    store = worker.store

    worker.close()

    store.close.assert_called_once_with()
    assert worker.store is None


def test_store_worker_close_is_idempotent():
    worker = _make_bare_worker()
    store = worker.store

    worker.close()
    worker.close()

    # Second call short-circuits because store was already released.
    store.close.assert_called_once_with()


def test_store_worker_close_swallows_store_errors():
    worker = _make_bare_worker()
    worker.store.close.side_effect = RuntimeError("boom")

    # A failure tearing down the store must not propagate out of close().
    worker.close()

    assert worker.store is None


def test_blob_block_hashes_wire_roundtrip():
    """The lookup wire format sends a ``hash_len`` frame plus the raw hashes
    concatenated back-to-back; the server rebuilds them through a zero-copy
    ``BlobBlockHashes`` view over the frame buffer."""
    hashes = [BlockHash(bytes([i]) * 16) for i in range(5)]
    hash_len = len(hashes[0])

    # Client side (LookupKeyClient._lookup): flat payload frame.
    blob = b"".join(hashes)

    # Server side (LookupKeyServer): view over the frame buffer (a memoryview),
    # never materializing the full hash list upfront.
    view = BlobBlockHashes(memoryview(blob), hash_len)

    assert len(view) == 5
    assert list(view) == hashes  # default Sequence iter terminates via IndexError
    assert [bytes(h) for h in view] == hashes
    assert bytes(view[-1]) == hashes[-1]
    assert [bytes(h) for h in view[1:3]] == hashes[1:3]
    with pytest.raises(IndexError):
        _ = view[5]


def test_blob_block_hashes_empty():
    """Empty lookups send hash_len=0 and an empty payload."""
    view = BlobBlockHashes(memoryview(b""), 0)
    assert len(view) == 0
    assert list(view) == []
