diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index 6d98a974c..7b00f4aa5 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -754,7 +754,11 @@ class Envs: # =================================================================== SGLANG_KV_CANARY_RING_CAPACITY = EnvInt(1024) SGLANG_KV_CANARY_ENABLE_WRITE_INPUT_ASSERT = EnvBool(False) + SGLANG_KV_CANARY_PERTURB_REQ_TO_TOKEN_PROB = EnvFloat(0.0) SGLANG_KV_CANARY_PERTURB_WARMUP_STEPS = EnvInt(50) + SGLANG_KV_CANARY_PERTURB_REAL_KV_USED_PROB = EnvFloat(0.0) + SGLANG_KV_CANARY_PERTURB_REAL_KV_UNUSED_CACHE_PROB = EnvFloat(0.0) + SGLANG_KV_CANARY_PERTURB_REAL_KV_POST_FORWARD_PROB = EnvFloat(0.0) SGLANG_KV_CANARY_PERTURB_TARGET_GROUP = EnvStr(None) SGLANG_KV_CANARY_PERTURB_NEXT_TOKEN_SWAP_PROB = EnvFloat(0.0) SGLANG_KV_CANARY_ENABLE_TOKEN_ORACLE = EnvBool(False) diff --git a/python/sglang/srt/kv_canary/perturb/config.py b/python/sglang/srt/kv_canary/perturb/config.py index 289b1414a..dbf3fec8c 100644 --- a/python/sglang/srt/kv_canary/perturb/config.py +++ b/python/sglang/srt/kv_canary/perturb/config.py @@ -17,14 +17,32 @@ class TargetGroupKind(IntEnum): @dataclass(frozen=True, slots=True, kw_only=True) class PerturbConfig: + req_to_token_prob: float + real_kv_used_prob: float + real_kv_unused_cache_prob: float + real_kv_post_forward_prob: float target_group_kind: TargetGroupKind | None warmup_steps: int @classmethod def from_env(cls) -> "PerturbConfig": + real_kv_used_prob = envs.SGLANG_KV_CANARY_PERTURB_REAL_KV_USED_PROB.get() + real_kv_unused_cache_prob = ( + envs.SGLANG_KV_CANARY_PERTURB_REAL_KV_UNUSED_CACHE_PROB.get() + ) + real_kv_post_forward_prob = ( + envs.SGLANG_KV_CANARY_PERTURB_REAL_KV_POST_FORWARD_PROB.get() + ) return cls( + req_to_token_prob=envs.SGLANG_KV_CANARY_PERTURB_REQ_TO_TOKEN_PROB.get(), + real_kv_used_prob=real_kv_used_prob, + real_kv_unused_cache_prob=real_kv_unused_cache_prob, + real_kv_post_forward_prob=real_kv_post_forward_prob, target_group_kind=_parse_target_group_kind_from_env( raw=envs.SGLANG_KV_CANARY_PERTURB_TARGET_GROUP.get(), + real_kv_used_prob=real_kv_used_prob, + real_kv_unused_cache_prob=real_kv_unused_cache_prob, + real_kv_post_forward_prob=real_kv_post_forward_prob, ), warmup_steps=envs.SGLANG_KV_CANARY_PERTURB_WARMUP_STEPS.get(), ) @@ -33,9 +51,18 @@ class PerturbConfig: def _parse_target_group_kind_from_env( *, raw: str | None, + real_kv_used_prob: float, + real_kv_unused_cache_prob: float, + real_kv_post_forward_prob: float, ) -> TargetGroupKind | None: if raw is not None and raw.strip(): return _parse_target_group_kind(raw) + if ( + real_kv_used_prob > 0.0 + or real_kv_unused_cache_prob > 0.0 + or real_kv_post_forward_prob > 0.0 + ): + return _parse_target_group_kind(raw) return None diff --git a/python/sglang/srt/kv_canary/perturb/manager.py b/python/sglang/srt/kv_canary/perturb/manager.py index 81e71e848..4a21ef26f 100644 --- a/python/sglang/srt/kv_canary/perturb/manager.py +++ b/python/sglang/srt/kv_canary/perturb/manager.py @@ -4,11 +4,18 @@ from collections.abc import Callable from typing import TYPE_CHECKING, Optional from sglang.srt.kv_canary.buffer_group import CanaryBufferGroup +from sglang.srt.kv_canary.perturb import ( + real_kv_post_forward, + real_kv_unused_cache, + real_kv_used, + req_to_token, +) from sglang.srt.kv_canary.perturb.config import PerturbConfig from sglang.srt.kv_canary.perturb.utils import WarmupGate if TYPE_CHECKING: from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache + from sglang.srt.mem_cache.memory_pool import ReqToTokenPool from sglang.srt.model_executor.forward_batch_info import ForwardBatch @@ -17,12 +24,14 @@ class PerturbManager: self, *, config: PerturbConfig, + req_to_token_pool: "ReqToTokenPool", buffer_groups: tuple[CanaryBufferGroup, ...], outer_step_counter_getter: Callable[[], int], swa_window_size: int = 0, sweep_interval: int = 0, ) -> None: self._config = config + self._req_to_token_pool = req_to_token_pool self._buffer_groups = buffer_groups self._outer_step_counter_getter = outer_step_counter_getter self._swa_window_size = swa_window_size @@ -40,4 +49,59 @@ class PerturbManager: *, maybe_inaccurate_forward_batch: Optional["ForwardBatch"], ) -> None: - pass + self.perturb_req_to_token(maybe_inaccurate_forward_batch) + self.perturb_real_kv_used(maybe_inaccurate_forward_batch) + self.perturb_real_kv_unused_cache(maybe_inaccurate_forward_batch) + + def perturb_post_forward( + self, + *, + maybe_inaccurate_forward_batch: Optional["ForwardBatch"], + ) -> None: + self.perturb_real_kv_post_forward(maybe_inaccurate_forward_batch) + + def perturb_req_to_token( + self, maybe_inaccurate_forward_batch: Optional["ForwardBatch"] + ) -> None: + req_to_token.run( + maybe_inaccurate_forward_batch=maybe_inaccurate_forward_batch, + config=self._config, + req_to_token_pool=self._req_to_token_pool, + warmup_gate=self._warmup_gate, + ) + + def perturb_real_kv_used( + self, maybe_inaccurate_forward_batch: Optional["ForwardBatch"] + ) -> None: + real_kv_used.run( + maybe_inaccurate_forward_batch=maybe_inaccurate_forward_batch, + config=self._config, + req_to_token_pool=self._req_to_token_pool, + buffer_groups=self._buffer_groups, + swa_window_size=self._swa_window_size, + warmup_gate=self._warmup_gate, + ) + + def perturb_real_kv_unused_cache( + self, maybe_inaccurate_forward_batch: Optional["ForwardBatch"] + ) -> None: + real_kv_unused_cache.run( + maybe_inaccurate_forward_batch=maybe_inaccurate_forward_batch, + config=self._config, + buffer_groups=self._buffer_groups, + radix_cache=self._radix_cache, + swa_window_size=self._swa_window_size, + sweep_interval=self._sweep_interval, + outer_step_counter=self._outer_step_counter_getter(), + warmup_gate=self._warmup_gate, + ) + + def perturb_real_kv_post_forward( + self, maybe_inaccurate_forward_batch: Optional["ForwardBatch"] + ) -> None: + real_kv_post_forward.run( + maybe_inaccurate_forward_batch=maybe_inaccurate_forward_batch, + config=self._config, + buffer_groups=self._buffer_groups, + warmup_gate=self._warmup_gate, + ) diff --git a/python/sglang/srt/kv_canary/perturb/real_kv_post_forward.py b/python/sglang/srt/kv_canary/perturb/real_kv_post_forward.py new file mode 100644 index 000000000..6db7c3b84 --- /dev/null +++ b/python/sglang/srt/kv_canary/perturb/real_kv_post_forward.py @@ -0,0 +1,78 @@ +"""Pick a random real-KV source byte derived from out_cache_loc and flip it +in-place AFTER the TAIL kernel has captured its canary hash. + +The slot id is taken from maybe_inaccurate_forward_batch.out_cache_loc and used +only as a lookup into the target group's real-KV source buffer; the actual flip +happens inside that buffer. The flip is a PyTorch indexed write on the current +CUDA stream; because TAIL is launched on the same stream, stream ordering +guarantees it happens-after TAIL's canary write. +""" + +from __future__ import annotations + +import logging +from typing import TYPE_CHECKING, Optional + +from sglang.srt.kv_canary.buffer_group import CanaryBufferGroup +from sglang.srt.kv_canary.perturb.config import ( + PerturbConfig, + require_target_group_kind, +) +from sglang.srt.kv_canary.perturb.slot_picker import pick_out_cache_loc_slot +from sglang.srt.kv_canary.perturb.utils import ( + WarmupGate, + flip_random_source_byte_and_log, + pick_target_group, + should_run_perturbation, +) + +if TYPE_CHECKING: + from sglang.srt.model_executor.forward_batch_info import ForwardBatch + +logger = logging.getLogger(__name__) + + +def run( + *, + maybe_inaccurate_forward_batch: Optional["ForwardBatch"], + config: PerturbConfig, + buffer_groups: tuple[CanaryBufferGroup, ...], + warmup_gate: WarmupGate, +) -> None: + if not should_run_perturbation( + perturb_name="real_kv_post_forward", + probability=config.real_kv_post_forward_prob, + warmup_gate=warmup_gate, + maybe_inaccurate_forward_batch=maybe_inaccurate_forward_batch, + ): + return + + slot = pick_out_cache_loc_slot( + maybe_inaccurate_forward_batch=maybe_inaccurate_forward_batch + ) + if slot is None: + logger.info( + "kv_canary perturb real_kv_post_forward: skipped because maybe_inaccurate_forward_batch.out_cache_loc " + "had no valid slot" + ) + return + group = pick_target_group( + buffer_groups=buffer_groups, + target_kind=require_target_group_kind( + target_group_kind=config.target_group_kind, + perturb_name="real_kv_post_forward", + ), + ) + if group is None: + logger.info( + "kv_canary perturb real_kv_post_forward: skipped because no target group matched " + "target_group_kind=%s slot=%d", + config.target_group_kind, + slot, + ) + return + flip_random_source_byte_and_log( + perturb_name="real_kv_post_forward", + group=group, + slot_idx=slot, + ) diff --git a/python/sglang/srt/kv_canary/perturb/real_kv_unused_cache.py b/python/sglang/srt/kv_canary/perturb/real_kv_unused_cache.py new file mode 100644 index 000000000..b18af9a31 --- /dev/null +++ b/python/sglang/srt/kv_canary/perturb/real_kv_unused_cache.py @@ -0,0 +1,162 @@ +"""Flip the first byte of a radix-cached but currently-unused (orphan) slot. + +Detection should come from sweep (per-forward verify won't even look at this +slot). Designed to surface bugs where cached KV is silently corrupted and +sleeps until much later when a prefix happens to match. +""" + +from __future__ import annotations + +import logging +from typing import TYPE_CHECKING, Optional + +import torch + +from sglang.srt.kv_canary.buffer_group import CanaryBufferGroup, PoolKind +from sglang.srt.kv_canary.perturb.config import ( + PerturbConfig, + require_target_group_kind, +) +from sglang.srt.kv_canary.perturb.utils import ( + WarmupGate, + flip_first_byte_in_source, + pick_target_group, + should_run_perturbation, +) +from sglang.srt.kv_canary.radix_cache_walker import walk_radix_cache_for_canary + +if TYPE_CHECKING: + from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache + from sglang.srt.model_executor.forward_batch_info import ForwardBatch + +logger = logging.getLogger(__name__) + + +def run( + *, + maybe_inaccurate_forward_batch: Optional["ForwardBatch"], + config: PerturbConfig, + buffer_groups: tuple[CanaryBufferGroup, ...], + radix_cache: Optional["BasePrefixCache"], + swa_window_size: int, + sweep_interval: int, + outer_step_counter: int, + warmup_gate: WarmupGate, +) -> None: + if sweep_interval <= 0 or outer_step_counter % sweep_interval != 0: + return + + if not should_run_perturbation( + perturb_name="real_kv_unused_cache", + probability=config.real_kv_unused_cache_prob, + warmup_gate=warmup_gate, + maybe_inaccurate_forward_batch=maybe_inaccurate_forward_batch, + require_forward_batch=False, + ): + return + + group = pick_target_group( + buffer_groups=buffer_groups, + target_kind=require_target_group_kind( + target_group_kind=config.target_group_kind, + perturb_name="real_kv_unused_cache", + ), + ) + if group is None: + logger.info( + "kv_canary perturb real_kv_unused_cache: skipped because no target group with " + "real_kv_sources_k matched target_group_kind=%s", + config.target_group_kind, + ) + return + slot = _pick_sweep_slot_for_group( + radix_cache=radix_cache, + group=group, + swa_window_size=swa_window_size, + ) + if slot is None: + logger.info( + "kv_canary perturb real_kv_unused_cache: skipped because no orphan sweep slot " + "was found for group=%s", + group.kind.name, + ) + return + source_pick = int(torch.randint(0, len(group.real_kv_sources_k), (1,)).item()) + source = group.real_kv_sources_k[source_pick] + flip_result = flip_first_byte_in_source( + group=group, + source=source, + slot_idx=slot, + slot_is_physical=True, + ) + if flip_result is None: + logger.info( + "kv_canary perturb real_kv_unused_cache: skipped because slot=%d could not be mapped " + "into group=%s source_idx=%d", + slot, + group.kind.name, + source_pick, + ) + return + row, col, original_byte = flip_result + logger.info( + "kv_canary perturb real_kv_unused_cache: group=%s source_idx=%d slot=%d row=%d col=%d " + "original_byte=0x%02X new_byte=0x%02X", + group.kind.name, + source_pick, + slot, + row, + col, + original_byte, + original_byte ^ 0xFF, + ) + + +def _pick_sweep_slot_for_group( + *, + radix_cache: Optional["BasePrefixCache"], + group: CanaryBufferGroup, + swa_window_size: int, +) -> Optional[int]: + if radix_cache is None: + return None + + walk_result = walk_radix_cache_for_canary( + radix_cache=radix_cache, + unlocked_only=True, + swa_resident_only=group.kind is PoolKind.SWA, + ) + slots = [ + int(raw_slot) + for raw_slot in walk_result.slot_indices.detach().to("cpu").tolist() + if int(raw_slot) >= 0 + ] + if group.kind is PoolKind.SWA: + slots = _translate_full_slots_to_swa_slots( + slots=slots, + full_to_swa_index_mapping=group.swa_index_lut, + ) + if not slots: + return None + + pick = int(torch.randint(0, len(slots), (1,)).item()) + return slots[pick] + + +def _translate_full_slots_to_swa_slots( + *, + slots: list[int], + full_to_swa_index_mapping: Optional[torch.Tensor], +) -> list[int]: + if full_to_swa_index_mapping is None: + return [] + + lut = full_to_swa_index_mapping.detach().to("cpu").to(torch.int64) + translated: list[int] = [] + for slot in slots: + if slot >= int(lut.shape[0]): + continue + physical_slot = int(lut[slot].item()) + if physical_slot >= 0: + translated.append(physical_slot) + return translated diff --git a/python/sglang/srt/kv_canary/perturb/real_kv_used.py b/python/sglang/srt/kv_canary/perturb/real_kv_used.py new file mode 100644 index 000000000..2131147a3 --- /dev/null +++ b/python/sglang/srt/kv_canary/perturb/real_kv_used.py @@ -0,0 +1,134 @@ +"""Flip the first byte of a slot currently being used by an active req. + +Detection should come from per-forward verify (HEAD/TAIL kernel), NOT from +sweep. Designed to surface CUDA-graph-idle-class bugs where production reads +a slot whose KV byte was silently overwritten. +""" + +from __future__ import annotations + +import logging +from typing import TYPE_CHECKING, Optional + +import torch + +from sglang.srt.kv_canary.buffer_group import CanaryBufferGroup, PoolKind +from sglang.srt.kv_canary.perturb.config import ( + PerturbConfig, + require_target_group_kind, +) +from sglang.srt.kv_canary.perturb.slot_picker import ( + ReqToTokenEntry, + collect_active_slots, +) +from sglang.srt.kv_canary.perturb.utils import ( + WarmupGate, + flip_first_byte_in_source, + pick_target_group, + should_run_perturbation, +) + +if TYPE_CHECKING: + from sglang.srt.mem_cache.memory_pool import ReqToTokenPool + from sglang.srt.model_executor.forward_batch_info import ForwardBatch + +logger = logging.getLogger(__name__) + + +def run( + *, + maybe_inaccurate_forward_batch: Optional["ForwardBatch"], + config: PerturbConfig, + req_to_token_pool: "ReqToTokenPool", + buffer_groups: tuple[CanaryBufferGroup, ...], + swa_window_size: int, + warmup_gate: WarmupGate, +) -> None: + if not should_run_perturbation( + perturb_name="real_kv_used", + probability=config.real_kv_used_prob, + warmup_gate=warmup_gate, + maybe_inaccurate_forward_batch=maybe_inaccurate_forward_batch, + ): + return + + group = pick_target_group( + buffer_groups=buffer_groups, + target_kind=require_target_group_kind( + target_group_kind=config.target_group_kind, + perturb_name="real_kv_used", + ), + ) + if group is None: + logger.info( + "kv_canary perturb real_kv_used: skipped because no target group with real_kv_sources_k " + "matched target_group_kind=%s", + config.target_group_kind, + ) + return + + target = _pick_active_slot_for_group( + maybe_inaccurate_forward_batch=maybe_inaccurate_forward_batch, + req_to_token_pool=req_to_token_pool, + group=group, + swa_window_size=swa_window_size, + ) + if target is None: + logger.info( + "kv_canary perturb real_kv_used: skipped because no active slot was found " + "for group=%s", + group.kind.name, + ) + return + + source_pick = int(torch.randint(0, len(group.real_kv_sources_k), (1,)).item()) + source = group.real_kv_sources_k[source_pick] + flip_result = flip_first_byte_in_source( + group=group, source=source, slot_idx=target.value + ) + if flip_result is None: + logger.info( + "kv_canary perturb real_kv_used: skipped because slot=%d could not be mapped into " + "group=%s source_idx=%d", + target.value, + group.kind.name, + source_pick, + ) + return + row, col, original_byte = flip_result + logger.info( + "kv_canary perturb real_kv_used: group=%s source_idx=%d slot=%d row=%d col=%d " + "original_byte=0x%02X new_byte=0x%02X", + group.kind.name, + source_pick, + target.value, + row, + col, + original_byte, + original_byte ^ 0xFF, + ) + + +def _pick_active_slot_for_group( + *, + maybe_inaccurate_forward_batch: "ForwardBatch", + req_to_token_pool: "ReqToTokenPool", + group: CanaryBufferGroup, + swa_window_size: int, +) -> Optional[ReqToTokenEntry]: + candidates = collect_active_slots( + maybe_inaccurate_forward_batch=maybe_inaccurate_forward_batch, + req_to_token_pool=req_to_token_pool, + exclude_out_cache_loc=True, + ) + if group.kind is PoolKind.SWA: + candidates = [ + entry + for entry in candidates + if entry.position >= max(0, entry.seq_len - swa_window_size) + ] + if not candidates: + return None + + pick = int(torch.randint(0, len(candidates), (1,)).item()) + return candidates[pick] diff --git a/python/sglang/srt/kv_canary/perturb/req_to_token.py b/python/sglang/srt/kv_canary/perturb/req_to_token.py new file mode 100644 index 000000000..00b207fbc --- /dev/null +++ b/python/sglang/srt/kv_canary/perturb/req_to_token.py @@ -0,0 +1,75 @@ +"""Flip the req_to_token pointer of a currently-active req. + +The hook picks a random (req_pool_idx, position, value) from active reqs, +filtering out entries whose value is 0 (so slot 0 is excluded), and overwrites +req_to_token[req_pool_idx, position] with another active req's slot id. +KV bytes are not touched. +""" + +from __future__ import annotations + +import logging +from typing import TYPE_CHECKING, Optional + +import torch + +from sglang.srt.kv_canary.perturb.config import PerturbConfig +from sglang.srt.kv_canary.perturb.slot_picker import collect_active_slots +from sglang.srt.kv_canary.perturb.utils import WarmupGate, should_run_perturbation + +if TYPE_CHECKING: + from sglang.srt.mem_cache.memory_pool import ReqToTokenPool + from sglang.srt.model_executor.forward_batch_info import ForwardBatch + +logger = logging.getLogger(__name__) + + +def run( + *, + maybe_inaccurate_forward_batch: Optional["ForwardBatch"], + config: PerturbConfig, + req_to_token_pool: "ReqToTokenPool", + warmup_gate: WarmupGate, +) -> None: + if not should_run_perturbation( + perturb_name="req_to_token", + probability=config.req_to_token_prob, + warmup_gate=warmup_gate, + maybe_inaccurate_forward_batch=maybe_inaccurate_forward_batch, + ): + return + + entries = collect_active_slots( + maybe_inaccurate_forward_batch=maybe_inaccurate_forward_batch, + req_to_token_pool=req_to_token_pool, + exclude_out_cache_loc=True, + ) + entries = [entry for entry in entries if entry.value >= 1] + if not entries: + logger.info( + "kv_canary perturb req_to_token: skipped because no active nonzero slots were found" + ) + return + + pick = int(torch.randint(0, len(entries), (1,)).item()) + target = entries[pick] + replacement_values = [item.value for item in entries if item.value != target.value] + if not replacement_values: + logger.info( + "kv_canary perturb req_to_token: skipped because no replacement slot differs from " + "original_slot=%d", + target.value, + ) + return + replacement_pick = int(torch.randint(0, len(replacement_values), (1,)).item()) + new_value = replacement_values[replacement_pick] + + req_to_token = req_to_token_pool.req_to_token + logger.info( + "kv_canary perturb req_to_token: req_pool_idx=%d position=%d original_slot=%d new_slot=%d", + target.req_pool_idx, + target.position, + target.value, + new_value, + ) + req_to_token[target.req_pool_idx, target.position] = new_value diff --git a/python/sglang/srt/kv_canary/perturb/slot_picker.py b/python/sglang/srt/kv_canary/perturb/slot_picker.py new file mode 100644 index 000000000..e6d9a0ee6 --- /dev/null +++ b/python/sglang/srt/kv_canary/perturb/slot_picker.py @@ -0,0 +1,100 @@ +from __future__ import annotations + +import random +from dataclasses import dataclass +from typing import TYPE_CHECKING, Optional + +import torch + +if TYPE_CHECKING: + from sglang.srt.mem_cache.memory_pool import ReqToTokenPool + from sglang.srt.model_executor.forward_batch_info import ForwardBatch + + +@dataclass(frozen=True, slots=True, kw_only=True) +class ReqToTokenEntry: + req_pool_idx: int + position: int + value: int + seq_len: int = 0 + + +def collect_active_slots( + *, + maybe_inaccurate_forward_batch: "ForwardBatch", + req_to_token_pool: "ReqToTokenPool", + exclude_out_cache_loc: bool = True, +) -> list[ReqToTokenEntry]: + """Collect every (req_pool_idx, position, value) triple for currently-active reqs. + + Excludes slots in ``maybe_inaccurate_forward_batch.out_cache_loc`` when ``exclude_out_cache_loc=True`` + so a slot the current forward is about to write isn't picked (write race). + """ + req_pool_indices = maybe_inaccurate_forward_batch.req_pool_indices + seq_lens = maybe_inaccurate_forward_batch.seq_lens + if req_pool_indices is None or seq_lens is None: + return [] + + req_to_token = req_to_token_pool.req_to_token + if not isinstance(req_to_token, torch.Tensor) or req_to_token.numel() == 0: + return [] + + excluded: set[int] = set() + if exclude_out_cache_loc: + out_cache_loc = maybe_inaccurate_forward_batch.out_cache_loc + if out_cache_loc is not None: + valid_num_tokens = maybe_inaccurate_forward_batch.num_token_non_padded_cpu + if valid_num_tokens is None: + valid_num_tokens = int(out_cache_loc.shape[0]) + excluded = set( + int(x) + for x in out_cache_loc[:valid_num_tokens].detach().to("cpu").tolist() + ) + + req_pool_indices_list = req_pool_indices.detach().to("cpu").tolist() + seq_lens_list = seq_lens.detach().to("cpu").tolist() + rows, cols = int(req_to_token.shape[0]), int(req_to_token.shape[1]) + + candidates: list[ReqToTokenEntry] = [] + for req_pool_idx, seq_len in zip(req_pool_indices_list, seq_lens_list): + req_pool_idx_int = int(req_pool_idx) + seq_len_int = int(seq_len) + if req_pool_idx_int < 0 or req_pool_idx_int >= rows: + continue + upper = min(seq_len_int, cols) + if upper <= 0: + continue + row_values = req_to_token[req_pool_idx_int, :upper].detach().to("cpu").tolist() + candidates.extend( + ReqToTokenEntry( + req_pool_idx=req_pool_idx_int, + position=pos, + value=value, + seq_len=seq_len_int, + ) + for pos, raw_value in enumerate(row_values) + if (value := int(raw_value)) >= 0 and value not in excluded + ) + return candidates + + +def pick_out_cache_loc_slot( + *, maybe_inaccurate_forward_batch: "ForwardBatch" +) -> Optional[int]: + out_cache_loc = maybe_inaccurate_forward_batch.out_cache_loc + if out_cache_loc is None: + return None + total = int(out_cache_loc.shape[0]) + if total <= 0: + return None + valid_num_tokens = maybe_inaccurate_forward_batch.num_token_non_padded_cpu + if valid_num_tokens is None: + valid_num_tokens = total + valid_num_tokens = int(valid_num_tokens) + if valid_num_tokens <= 0: + return None + pick = random.randrange(valid_num_tokens) + slot = int(out_cache_loc[pick].item()) + if slot < 0: + return None + return slot diff --git a/python/sglang/srt/kv_canary/runner/canary_manager.py b/python/sglang/srt/kv_canary/runner/canary_manager.py index 8dfa809a7..7e3cd4fac 100644 --- a/python/sglang/srt/kv_canary/runner/canary_manager.py +++ b/python/sglang/srt/kv_canary/runner/canary_manager.py @@ -101,6 +101,7 @@ class CanaryManager: ) self._perturb_manager = PerturbManager( config=perturb_config, + req_to_token_pool=req_to_token_pool, buffer_groups=self._buffer_groups, outer_step_counter_getter=self._get_outer_step_counter, swa_window_size=self._swa_window_size, @@ -205,6 +206,9 @@ class CanaryManager: ) -> None: for idx in single_forward_indices: self._single_forward_managers[idx].post_ops_outside_graph() + self._perturb_manager.perturb_post_forward( + maybe_inaccurate_forward_batch=maybe_inaccurate_forward_batch + ) self._sweep_orchestrator.maybe_run_sweep() self._outer_step_counter += 1 self._violation_manager.step() diff --git a/python/sglang/test/kv_canary/pd_fixture.py b/python/sglang/test/kv_canary/pd_fixture.py new file mode 100644 index 000000000..ee369aaa8 --- /dev/null +++ b/python/sglang/test/kv_canary/pd_fixture.py @@ -0,0 +1,89 @@ +from __future__ import annotations + +from typing import ClassVar, Literal, Optional + +from sglang.srt.kv_canary.config import CanaryMode +from sglang.test.kv_canary.mode_config import _MODE_CONFIGS, _ModeConfig +from sglang.test.kv_canary.utils import build_canary_server_args, post_parallel_generate +from sglang.test.kv_canary.violation_assert_mixin import CanaryViolationAssertMixin +from sglang.test.server_fixtures.disaggregation_fixture import ( + PDDisaggregationServerBase, +) + +_SHORT_PROMPT_BODY = ("The quick brown fox jumps over the lazy dog. " * 8).strip() + + +class CanaryPDFixture(CanaryViolationAssertMixin, PDDisaggregationServerBase): + capture_per_side_logs = True + + model_mode: ClassVar[Literal["mha", "swa", "dsv4"]] + kv_canary_mode: ClassVar[CanaryMode] = CanaryMode.LOG + extra_server_args: ClassVar[tuple[str, ...]] = ( + "--kv-canary-real-data", + "partial", + "--skip-server-warmup", + ) + + _cfg: ClassVar[Optional[_ModeConfig]] = None + + @classmethod + def setUpClass(cls) -> None: + super().setUpClass() + cls._cfg = _MODE_CONFIGS[cls.model_mode] + cls.model = cls._cfg.model_path + + canary_args = build_canary_server_args( + kv_canary_mode=cls.kv_canary_mode, + mode_cfg=cls._cfg, + extra_server_args=cls.extra_server_args, + ) + cls.extra_prefill_args = list(canary_args) + cls.extra_decode_args = list(canary_args) + if cls.model_mode == "swa": + # SWA mode uses google/gemma-4-E2B-it, whose forward does a + # ``positions += 1`` in-place. canary's WRITE/VERIFY require + # forward_batch.positions to stay 0-indexed, so flip the gemma + # path to out-of-place shift for these tests. + cls.extra_prefill_env = { + **cls.extra_prefill_env, + "SGLANG_GEMMA_OUT_OF_PLACE_POSITION_MUTATION": "1", + } + cls.extra_decode_env = { + **cls.extra_decode_env, + "SGLANG_GEMMA_OUT_OF_PLACE_POSITION_MUTATION": "1", + } + cls.launch_all() + + def send_parallel_short_requests( + self, + n: int, + *, + assert_all_success: bool = True, + max_new_tokens: int = 100, + timeout: float = 60.0, + ) -> list[dict]: + results = post_parallel_generate( + url=self.lb_url + "/generate", + prompts=[_SHORT_PROMPT_BODY] * n, + max_new_tokens=max_new_tokens, + timeout=timeout, + ) + if assert_all_success: + for result in results: + self.assertEqual(result.get("status_code"), 200, result) + return results + + def _captured_log_text( + self, side: Optional[Literal["prefill", "decode"]] = None + ) -> str: + if side == "prefill": + stdout_buf = type(self)._prefill_stdout_buf + stderr_buf = type(self)._prefill_stderr_buf + elif side == "decode": + stdout_buf = type(self)._decode_stdout_buf + stderr_buf = type(self)._decode_stderr_buf + else: + raise ValueError(f"Unsupported side={side!r}") + stdout_text = stdout_buf.getvalue() if stdout_buf is not None else "" + stderr_text = stderr_buf.getvalue() if stderr_buf is not None else "" + return stdout_text + stderr_text diff --git a/python/sglang/test/kv_canary/runner_test_base.py b/python/sglang/test/kv_canary/runner_test_base.py index 077c09a47..3ca41411a 100644 --- a/python/sglang/test/kv_canary/runner_test_base.py +++ b/python/sglang/test/kv_canary/runner_test_base.py @@ -54,6 +54,10 @@ def make_perturb_config() -> PerturbConfig: """Build a PerturbConfig with every probability pinned to 0 so the perturb hooks do nothing during unit tests.""" return PerturbConfig( + req_to_token_prob=0.0, + real_kv_used_prob=0.0, + real_kv_unused_cache_prob=0.0, + real_kv_post_forward_prob=0.0, target_group_kind=None, warmup_steps=0, ) diff --git a/test/registered/kv_canary/test_self_e2e_pd_baseline.py b/test/registered/kv_canary/test_self_e2e_pd_baseline.py new file mode 100644 index 000000000..b13ebf356 --- /dev/null +++ b/test/registered/kv_canary/test_self_e2e_pd_baseline.py @@ -0,0 +1,30 @@ +from __future__ import annotations + +import unittest + +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.kv_canary.pd_fixture import CanaryPDFixture + +register_cuda_ci(est_time=180, stage="extra-a", runner_config="2-gpu-large") + + +class TestPDBaselineMha(CanaryPDFixture): + model_mode = "mha" + + def test_clean_pd_run_produces_no_canary_violation_on_either_side(self) -> None: + self.send_parallel_short_requests(n=4) + self.assert_no_violation(side="prefill") + self.assert_no_violation(side="decode") + + +class TestPDBaselineSwa(CanaryPDFixture): + model_mode = "swa" + + def test_clean_pd_run_produces_no_canary_violation_on_either_side(self) -> None: + self.send_parallel_short_requests(n=4) + self.assert_no_violation(side="prefill") + self.assert_no_violation(side="decode") + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/kv_canary/test_self_e2e_pd_perturb.py b/test/registered/kv_canary/test_self_e2e_pd_perturb.py new file mode 100644 index 000000000..2a7f68a6b --- /dev/null +++ b/test/registered/kv_canary/test_self_e2e_pd_perturb.py @@ -0,0 +1,74 @@ +from __future__ import annotations + +import unittest +from typing import ClassVar + +from sglang.srt.kv_canary.perturb.config import TargetGroupKind +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.kv_canary.pd_fixture import CanaryPDFixture + +register_cuda_ci(est_time=180, stage="extra-a", runner_config="2-gpu-large") + + +class _PDPerturbBase(CanaryPDFixture): + target_group: ClassVar[TargetGroupKind] + + @classmethod + def setUpClass(cls) -> None: + if cls is _PDPerturbBase: + raise unittest.SkipTest( + "abstract base; concrete subclasses set model_mode + target_group" + ) + cls.extra_prefill_env = { + "SGLANG_KV_CANARY_PERTURB_REAL_KV_POST_FORWARD_PROB": "1.0", + "SGLANG_KV_CANARY_PERTURB_TARGET_GROUP": str(cls.target_group), + "SGLANG_KV_CANARY_PERTURB_WARMUP_STEPS": "0", + } + cls.extra_decode_env = { + "SGLANG_KV_CANARY_PERTURB_REAL_KV_POST_FORWARD_PROB": "0", + "SGLANG_KV_CANARY_PERTURB_REAL_KV_USED_PROB": "0", + "SGLANG_KV_CANARY_PERTURB_REAL_KV_UNUSED_CACHE_PROB": "0", + "SGLANG_KV_CANARY_PERTURB_REQ_TO_TOKEN_PROB": "0", + } + super().setUpClass() + + def test_p_side_perturb_surfaces_real_kv_hash_violation_on_decode_side( + self, + ) -> None: + # send_parallel_short_requests defaults to max_new_tokens=100 so D-side runs + # decode forwards that exercise canary verify on the transferred prefix. + self.send_parallel_short_requests(n=4) + # D-side: first decode forward re-verifies the transferred prefix slots, + # so the flip MUST surface as real_kv_hash violation. + self.assert_per_forward_violation_reported( + fail_reason="verify_real_kv_hash", + target_group=self.target_group, + side="decode", + flush_wait_seconds=4.0, + ) + # P-side: flip happens post-TAIL of the prefill forward, and PD prefill + # does not run another forward on P that would verify the perturbed slot, + # so P MUST stay silent (no false-positive violations) for this perturb + # point. If a future canary feature adds post-prefill verify on P, this + # assert will start failing and should be upgraded to assert the + # violation on P-side too. + self.assert_no_violation(side="prefill", wait_seconds=0.5) + + +class TestPDPerturbMhaFull(_PDPerturbBase): + model_mode = "mha" + target_group = TargetGroupKind.FULL + + +class TestPDPerturbSwaFull(_PDPerturbBase): + model_mode = "swa" + target_group = TargetGroupKind.FULL + + +class TestPDPerturbSwaSwa(_PDPerturbBase): + model_mode = "swa" + target_group = TargetGroupKind.SWA + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/kv_canary/test_self_e2e_perturb_raise.py b/test/registered/kv_canary/test_self_e2e_perturb_raise.py new file mode 100644 index 000000000..008057381 --- /dev/null +++ b/test/registered/kv_canary/test_self_e2e_perturb_raise.py @@ -0,0 +1,41 @@ +from __future__ import annotations + +import unittest + +from sglang.srt.kv_canary.config import CanaryMode +from sglang.srt.kv_canary.perturb.config import TargetGroupKind +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.kv_canary.e2e_base import CanaryE2EBase + +register_cuda_ci(est_time=60, stage="extra-a", runner_config="1-gpu-small") + + +class TestPerturbRaiseMha(CanaryE2EBase): + model_mode = "mha" + kv_canary_mode = CanaryMode.RAISE + extra_server_args = ("--kv-canary-real-data", "partial", "--skip-server-warmup") + extra_env = { + "SGLANG_KV_CANARY_PERTURB_REAL_KV_USED_PROB": "0.1", + "SGLANG_KV_CANARY_PERTURB_TARGET_GROUP": "full", + "SGLANG_KV_CANARY_PERTURB_WARMUP_STEPS": "0", + } + + def test_real_kv_used_perturbation_raises_in_raise_mode(self) -> None: + """Verify raise mode surfaces real KV perturbation as a logged violation.""" + try: + self.send_parallel_requests( + n=4, + assert_all_success=False, + timeout=30.0, + ) + except Exception: + pass + self.assert_per_forward_violation_reported( + fail_reason="verify_real_kv_hash", + target_group=TargetGroupKind.FULL, + flush_wait_seconds=3.0, + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/kv_canary/test_self_e2e_perturb_real_kv_unused_cache.py b/test/registered/kv_canary/test_self_e2e_perturb_real_kv_unused_cache.py new file mode 100644 index 000000000..4c2e4d586 --- /dev/null +++ b/test/registered/kv_canary/test_self_e2e_perturb_real_kv_unused_cache.py @@ -0,0 +1,80 @@ +from __future__ import annotations + +import unittest +from typing import ClassVar + +from sglang.srt.kv_canary.config import CanaryMode +from sglang.srt.kv_canary.perturb.config import TargetGroupKind +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.kv_canary.consts import SWA_POOL_SERVER_ARGS +from sglang.test.kv_canary.e2e_base import CanaryE2EBase + +register_cuda_ci(est_time=60, stage="extra-a", runner_config="1-gpu-small") + + +class _PerturbRealKvUnusedCacheBase(CanaryE2EBase): + kv_canary_mode = CanaryMode.LOG + extra_server_args = ( + "--kv-canary-real-data", + "partial", + "--kv-canary-sweep-interval", + "4", + ) + use_unique_prompts = True + + target_group: ClassVar[TargetGroupKind] + + @classmethod + def setUpClass(cls) -> None: + if cls is _PerturbRealKvUnusedCacheBase: + raise unittest.SkipTest( + "abstract base; concrete subclasses set model_mode + target_group" + ) + cls.extra_env = { + "SGLANG_KV_CANARY_PERTURB_REAL_KV_UNUSED_CACHE_PROB": "0.1", + "SGLANG_KV_CANARY_PERTURB_TARGET_GROUP": str(cls.target_group), + "SGLANG_KV_CANARY_PERTURB_WARMUP_STEPS": "0", + } + super().setUpClass() + + def test_real_kv_unused_cache_perturbation_reports_sweep_real_kv_hash_violation( + self, + ) -> None: + """Verify cached unused KV perturbation is caught by sweep verification.""" + # Step 1: first batch builds radix entries that will become orphans once finished. + self.send_parallel_requests(n=8) + # Step 2: second batch drives more forward passes so the sweep cadence fires + # while the orphan slots are still cached. + self.send_parallel_requests(n=8) + self.assert_sweep_violation_reported( + fail_reason="verify_real_kv_hash", + target_group=self.target_group, + flush_wait_seconds=5.0, + ) + + +class TestPerturbRealKvUnusedCacheMhaFull(_PerturbRealKvUnusedCacheBase): + model_mode = "mha" + target_group = TargetGroupKind.FULL + + +class TestPerturbRealKvUnusedCacheSwaFull(_PerturbRealKvUnusedCacheBase): + model_mode = "swa" + target_group = TargetGroupKind.FULL + extra_server_args = ( + *_PerturbRealKvUnusedCacheBase.extra_server_args, + *SWA_POOL_SERVER_ARGS, + ) + + +class TestPerturbRealKvUnusedCacheSwaSwa(_PerturbRealKvUnusedCacheBase): + model_mode = "swa" + target_group = TargetGroupKind.SWA + extra_server_args = ( + *_PerturbRealKvUnusedCacheBase.extra_server_args, + *SWA_POOL_SERVER_ARGS, + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/kv_canary/test_self_e2e_perturb_real_kv_used.py b/test/registered/kv_canary/test_self_e2e_perturb_real_kv_used.py new file mode 100644 index 000000000..4c0274096 --- /dev/null +++ b/test/registered/kv_canary/test_self_e2e_perturb_real_kv_used.py @@ -0,0 +1,68 @@ +from __future__ import annotations + +import unittest +from typing import ClassVar + +from sglang.srt.kv_canary.config import CanaryMode +from sglang.srt.kv_canary.perturb.config import TargetGroupKind +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.kv_canary.consts import SWA_POOL_SERVER_ARGS +from sglang.test.kv_canary.e2e_base import CanaryE2EBase + +register_cuda_ci(est_time=60, stage="extra-a", runner_config="1-gpu-small") + + +class _PerturbRealKvUsedBase(CanaryE2EBase): + kv_canary_mode = CanaryMode.LOG + extra_server_args = ("--kv-canary-real-data", "partial") + + target_group: ClassVar[TargetGroupKind] + + @classmethod + def setUpClass(cls) -> None: + if cls is _PerturbRealKvUsedBase: + raise unittest.SkipTest( + "abstract base; concrete subclasses set model_mode + target_group" + ) + cls.extra_env = { + "SGLANG_KV_CANARY_PERTURB_REAL_KV_USED_PROB": "0.1", + "SGLANG_KV_CANARY_PERTURB_TARGET_GROUP": str(cls.target_group), + "SGLANG_KV_CANARY_PERTURB_WARMUP_STEPS": "0", + } + super().setUpClass() + + def test_real_kv_used_perturbation_reports_real_kv_hash_violation(self) -> None: + """Verify active real KV perturbation reports a real KV hash violation.""" + for _ in range(self.workload_n_batches): + self.send_parallel_requests() + self.assert_per_forward_violation_reported( + fail_reason="verify_real_kv_hash", + target_group=self.target_group, + ) + + +class TestPerturbRealKvUsedMhaFull(_PerturbRealKvUsedBase): + model_mode = "mha" + target_group = TargetGroupKind.FULL + + +class TestPerturbRealKvUsedSwaFull(_PerturbRealKvUsedBase): + model_mode = "swa" + target_group = TargetGroupKind.FULL + extra_server_args = ( + *_PerturbRealKvUsedBase.extra_server_args, + *SWA_POOL_SERVER_ARGS, + ) + + +class TestPerturbRealKvUsedSwaSwa(_PerturbRealKvUsedBase): + model_mode = "swa" + target_group = TargetGroupKind.SWA + extra_server_args = ( + *_PerturbRealKvUsedBase.extra_server_args, + *SWA_POOL_SERVER_ARGS, + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/kv_canary/test_self_e2e_perturb_req_to_token.py b/test/registered/kv_canary/test_self_e2e_perturb_req_to_token.py new file mode 100644 index 000000000..924f38f97 --- /dev/null +++ b/test/registered/kv_canary/test_self_e2e_perturb_req_to_token.py @@ -0,0 +1,49 @@ +from __future__ import annotations + +import unittest + +from sglang.srt.kv_canary.config import CanaryMode +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.kv_canary.consts import SWA_POOL_SERVER_ARGS +from sglang.test.kv_canary.e2e_base import CanaryE2EBase + +register_cuda_ci(est_time=60, stage="extra-a", runner_config="1-gpu-small") + + +class _PerturbReqToTokenBase(CanaryE2EBase): + kv_canary_mode = CanaryMode.LOG + extra_env = { + "SGLANG_KV_CANARY_PERTURB_REQ_TO_TOKEN_PROB": "0.1", + "SGLANG_KV_CANARY_PERTURB_WARMUP_STEPS": "0", + # req_to_token perturbation deliberately corrupts the slot mapping + # by design, which the scheduler's on-idle invariant checker reports + # as a pool memory leak (perturbed slot is freed, original slot + # still looks busy). That's expected for this test; disable strict + # mode so the leak warning doesn't crash the scheduler. + "SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_IDLE": "0", + } + + @classmethod + def setUpClass(cls) -> None: + if cls is _PerturbReqToTokenBase: + raise unittest.SkipTest("abstract base; concrete subclasses set model_mode") + super().setUpClass() + + def test_req_to_token_perturbation_reports_chain_hash_violation(self) -> None: + """Verify req_to_token perturbation reports a chain hash violation.""" + for _ in range(self.workload_n_batches): + self.send_parallel_requests() + self.assert_per_forward_violation_reported(fail_reason="verify_chain_hash") + + +class TestPerturbReqToTokenMha(_PerturbReqToTokenBase): + model_mode = "mha" + + +class TestPerturbReqToTokenSwa(_PerturbReqToTokenBase): + model_mode = "swa" + extra_server_args = SWA_POOL_SERVER_ARGS + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/kv_canary/test_self_unit_perturb.py b/test/registered/kv_canary/test_self_unit_perturb.py index 9ba493fb7..5f3d4ed62 100644 --- a/test/registered/kv_canary/test_self_unit_perturb.py +++ b/test/registered/kv_canary/test_self_unit_perturb.py @@ -2,28 +2,44 @@ from __future__ import annotations import os import unittest -from typing import cast +from typing import TYPE_CHECKING, cast from unittest.mock import patch import torch from sglang.jit_kernel.kv_canary.verify import RealKvSource from sglang.srt.kv_canary.buffer_group import PoolKind +from sglang.srt.kv_canary.perturb import ( + real_kv_post_forward, +) +from sglang.srt.kv_canary.perturb import ( + real_kv_unused_cache as real_kv_unused_cache_module, +) from sglang.srt.kv_canary.perturb.config import ( PerturbConfig, TargetGroupKind, _parse_target_group_kind, ) +from sglang.srt.kv_canary.perturb.manager import PerturbManager +from sglang.srt.kv_canary.perturb.slot_picker import collect_active_slots from sglang.srt.kv_canary.perturb.utils import ( + WarmupGate, flip_first_byte_in_source, pick_target_group, ) from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.kv_canary.fixtures import ( + DEFAULT_DEVICE, make_buffer_group, + make_forward_batch, + make_radix_cache, + make_req_to_token_pool, ) from sglang.test.test_utils import CustomTestCase +if TYPE_CHECKING: + from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache + register_cuda_ci(est_time=10, stage="extra-a", runner_config="1-gpu-small") @@ -56,11 +72,20 @@ class TestParseTargetGroupKind(CustomTestCase): ): _parse_target_group_kind(raw) - def test_from_env_allows_missing_target( + def test_from_env_allows_missing_target_when_real_kv_perturb_is_disabled( self, ) -> None: """Verify normal canary startup does not require a perturb target group.""" - with patch.dict(os.environ, {}, clear=False): + with patch.dict( + os.environ, + { + "SGLANG_KV_CANARY_PERTURB_REQ_TO_TOKEN_PROB": "0", + "SGLANG_KV_CANARY_PERTURB_REAL_KV_USED_PROB": "0", + "SGLANG_KV_CANARY_PERTURB_REAL_KV_UNUSED_CACHE_PROB": "0", + "SGLANG_KV_CANARY_PERTURB_REAL_KV_POST_FORWARD_PROB": "0", + }, + clear=False, + ): os.environ.pop("SGLANG_KV_CANARY_PERTURB_TARGET_GROUP", None) config = PerturbConfig.from_env() @@ -111,7 +136,204 @@ class TestPickTargetGroup(CustomTestCase): self.assertIsNone(group) -class TestPerturbWarmupAndUtils(CustomTestCase): +class TestPerturbManager(CustomTestCase): + def test_perturb_manager_perturb_post_forward_dispatches_real_kv_post_forward( + self, + ) -> None: + """Verify perturb_post_forward() routes only to the post_forward dispatch.""" + device = DEFAULT_DEVICE + manager = PerturbManager( + config=PerturbConfig( + req_to_token_prob=0.0, + real_kv_used_prob=0.0, + real_kv_unused_cache_prob=0.0, + real_kv_post_forward_prob=0.0, + target_group_kind=TargetGroupKind.FULL, + warmup_steps=0, + ), + req_to_token_pool=make_req_to_token_pool(device, max_reqs=4, max_seq_len=8), + buffer_groups=(), + outer_step_counter_getter=lambda: 10, + ) + forward_batch = make_forward_batch(device, bs=1, seq_lens_list=(1,)) + calls: list[str] = [] + + with patch.object( + manager, + "perturb_real_kv_post_forward", + lambda batch: calls.append("real_kv_post_forward"), + ), patch.object( + manager, + "perturb_req_to_token", + lambda batch: calls.append("req_to_token"), + ), patch.object( + manager, + "perturb_real_kv_used", + lambda batch: calls.append("real_kv_used"), + ), patch.object( + manager, + "perturb_real_kv_unused_cache", + lambda batch: calls.append("real_kv_unused_cache"), + ): + manager.perturb_post_forward(maybe_inaccurate_forward_batch=forward_batch) + + self.assertEqual(calls, ["real_kv_post_forward"]) + + +class TestRealKvPostForwardPerturb(CustomTestCase): + def test_real_kv_post_forward_flips_a_byte_in_out_cache_loc_slot(self) -> None: + """Verify post-forward perturbation flips one real-KV byte and leaves canary buffers untouched.""" + device = DEFAULT_DEVICE + group = make_buffer_group(kind=PoolKind.FULL, has_real_kv=True) + source = group.real_kv_sources_k[0] + config = PerturbConfig( + req_to_token_prob=0.0, + real_kv_used_prob=0.0, + real_kv_unused_cache_prob=0.0, + real_kv_post_forward_prob=1.0, + target_group_kind=TargetGroupKind.FULL, + warmup_steps=0, + ) + warmup_gate = WarmupGate(config=config, outer_step_counter_getter=lambda: 10) + + forward_batch = make_forward_batch(device, bs=1, seq_lens_list=(1,)) + forward_batch.out_cache_loc = torch.tensor( + [2], dtype=torch.int32, device=device + ) + forward_batch.num_token_non_padded_cpu = 1 + + head_snapshot = group.k_head.clone() + v_head_snapshot = group.v_head.clone() + k_tail_snapshot = group.k_tail.clone() + v_tail_snapshot = group.v_tail.clone() + source_snapshot = source.tensor.clone() + + with patch.object(torch, "rand", return_value=torch.tensor(0.0)): + real_kv_post_forward.run( + maybe_inaccurate_forward_batch=forward_batch, + config=config, + buffer_groups=(group,), + warmup_gate=warmup_gate, + ) + + diff = source.tensor != source_snapshot + self.assertEqual(int(diff.sum().item()), 1) + self.assertTrue(bool(diff[2, 0].item())) + self.assertEqual(int(source.tensor[2, 0].item()), 0 ^ 0xFF) + self.assertTrue(torch.equal(group.k_head, head_snapshot)) + self.assertTrue(torch.equal(group.v_head, v_head_snapshot)) + self.assertTrue(torch.equal(group.k_tail, k_tail_snapshot)) + self.assertTrue(torch.equal(group.v_tail, v_tail_snapshot)) + + +class TestReqToTokenPerturb(CustomTestCase): + def test_req_to_token_perturb_uses_live_slot_as_replacement(self) -> None: + """Verify req_to_token perturbation replaces a slot with another live slot.""" + device = DEFAULT_DEVICE + pool = make_req_to_token_pool(device, max_reqs=4, max_seq_len=8) + pool.req_to_token[1, :3] = torch.tensor( + [11, 22, 33], dtype=torch.int32, device=device + ) + pool.req_to_token[2, :3] = torch.tensor( + [44, 55, 66], dtype=torch.int32, device=device + ) + manager = PerturbManager( + config=PerturbConfig( + req_to_token_prob=1.0, + real_kv_used_prob=0.0, + real_kv_unused_cache_prob=0.0, + real_kv_post_forward_prob=0.0, + target_group_kind=TargetGroupKind.FULL, + warmup_steps=0, + ), + req_to_token_pool=pool, + buffer_groups=(), + outer_step_counter_getter=lambda: 10, + ) + forward_batch = make_forward_batch(device, bs=2, seq_lens_list=(3, 3)) + forward_batch.out_cache_loc = torch.tensor( + [11], dtype=torch.int32, device=device + ) + + snapshot = pool.req_to_token.clone() + with patch.object(torch, "rand", return_value=torch.tensor(0.0)): + manager.perturb_req_to_token(forward_batch) + + diff = pool.req_to_token != snapshot + self.assertEqual(int(diff.sum().item()), 1) + rows, cols = torch.nonzero(diff, as_tuple=True) + row, col = int(rows[0].item()), int(cols[0].item()) + original = int(snapshot[row, col].item()) + replacement = int(pool.req_to_token[row, col].item()) + live_slots = {11, 22, 33, 44, 55, 66} + self.assertIn(original, live_slots) + self.assertIn(replacement, live_slots) + self.assertNotEqual(replacement, original) + self.assertFalse(bool(diff[1, 0].item())) + + def test_collect_active_slots_ignores_padded_out_cache_loc(self) -> None: + """Verify out_cache_loc padding does not exclude a live slot.""" + device = DEFAULT_DEVICE + pool = make_req_to_token_pool(device, max_reqs=4, max_seq_len=8) + pool.req_to_token[1, :2] = torch.tensor( + [0, 7], dtype=torch.int32, device=device + ) + forward_batch = make_forward_batch(device, bs=1, seq_lens_list=(2,)) + forward_batch.out_cache_loc = torch.tensor( + [7, 0, 0], dtype=torch.int32, device=device + ) + forward_batch.num_token_non_padded_cpu = 1 + + targets = collect_active_slots( + maybe_inaccurate_forward_batch=forward_batch, + req_to_token_pool=pool, + ) + + self.assertEqual([target.value for target in targets], [0]) + + +class TestRealKvUsedPerturb(CustomTestCase): + def test_real_kv_used_flips_first_real_kv_byte_for_active_full_slot( + self, + ) -> None: + """Verify real_kv_used flips only the first real KV byte for an active FULL slot.""" + device = DEFAULT_DEVICE + pool = make_req_to_token_pool(device, max_reqs=4, max_seq_len=8) + pool.req_to_token.fill_(-1) + pool.req_to_token[1, 0] = 2 + group = make_buffer_group(kind=PoolKind.FULL, has_real_kv=True) + source = group.real_kv_sources_k[0] + source.tensor.copy_( + torch.arange(source.tensor.numel(), dtype=torch.uint8).view_as( + source.tensor + ) + ) + manager = PerturbManager( + config=PerturbConfig( + req_to_token_prob=0.0, + real_kv_used_prob=1.0, + real_kv_unused_cache_prob=0.0, + real_kv_post_forward_prob=0.0, + target_group_kind=TargetGroupKind.FULL, + warmup_steps=0, + ), + req_to_token_pool=pool, + buffer_groups=(group,), + outer_step_counter_getter=lambda: 10, + ) + forward_batch = make_forward_batch(device, bs=1, seq_lens_list=(1,)) + forward_batch.out_cache_loc = torch.tensor( + [99], dtype=torch.int32, device=device + ) + + snapshot = source.tensor.clone() + with patch.object(torch, "rand", return_value=torch.tensor(0.0)): + manager.perturb_real_kv_used(forward_batch) + + expected = snapshot.clone() + expected[2, 0] = int(snapshot[2, 0].item()) ^ 0xFF + self.assertTrue(torch.equal(source.tensor, expected)) + def test_flip_first_byte_in_source_maps_swa_logical_slot_through_lut( self, ) -> None: @@ -137,6 +359,161 @@ class TestPerturbWarmupAndUtils(CustomTestCase): expected[1, 16] = int(snapshot[1, 16].item()) ^ 0xFF self.assertTrue(torch.equal(source.tensor, expected)) + def test_warmup_gate_prevents_perturbation_when_probabilities_are_one(self) -> None: + """Verify warmup prevents all perturbations even when every probability is one.""" + device = DEFAULT_DEVICE + pool = make_req_to_token_pool(device, max_reqs=4, max_seq_len=8) + pool.req_to_token.fill_(-1) + pool.req_to_token[1, 0] = 2 + group = make_buffer_group(kind=PoolKind.FULL, has_real_kv=True) + source = group.real_kv_sources_k[0] + source.tensor.copy_( + torch.arange(source.tensor.numel(), dtype=torch.uint8).view_as( + source.tensor + ) + ) + manager = PerturbManager( + config=PerturbConfig( + req_to_token_prob=1.0, + real_kv_used_prob=1.0, + real_kv_unused_cache_prob=1.0, + real_kv_post_forward_prob=0.0, + target_group_kind=TargetGroupKind.FULL, + warmup_steps=20, + ), + req_to_token_pool=pool, + buffer_groups=(group,), + outer_step_counter_getter=lambda: 10, + ) + manager.attach_radix_cache(cast("BasePrefixCache", object())) + forward_batch = make_forward_batch(device, bs=1, seq_lens_list=(1,)) + forward_batch.out_cache_loc = torch.tensor( + [99], dtype=torch.int32, device=device + ) + + pool_snapshot = pool.req_to_token.clone() + source_snapshot = source.tensor.clone() + with patch.object(torch, "rand", return_value=torch.tensor(0.0)), patch.object( + real_kv_unused_cache_module, + "_pick_sweep_slot_for_group", + return_value=3, + ): + manager.perturb(maybe_inaccurate_forward_batch=forward_batch) + + self.assertTrue(torch.equal(pool.req_to_token, pool_snapshot)) + self.assertTrue(torch.equal(source.tensor, source_snapshot)) + + +class TestRealKvUnusedCachePerturb(CustomTestCase): + def test_real_kv_unused_cache_flips_first_real_kv_byte_for_orphan_slot( + self, + ) -> None: + """Verify real_kv_unused_cache flips only the first real KV byte for an orphan slot.""" + device = DEFAULT_DEVICE + pool = make_req_to_token_pool(device, max_reqs=4, max_seq_len=8) + group = make_buffer_group(kind=PoolKind.FULL, has_real_kv=True) + source = group.real_kv_sources_k[0] + source.tensor.copy_( + torch.arange(source.tensor.numel(), dtype=torch.uint8).view_as( + source.tensor + ) + ) + manager = PerturbManager( + config=PerturbConfig( + req_to_token_prob=0.0, + real_kv_used_prob=0.0, + real_kv_unused_cache_prob=1.0, + real_kv_post_forward_prob=0.0, + target_group_kind=TargetGroupKind.FULL, + warmup_steps=0, + ), + req_to_token_pool=pool, + buffer_groups=(group,), + outer_step_counter_getter=lambda: 10, + sweep_interval=1, + ) + manager.attach_radix_cache(make_radix_cache([[], [3]], device=device)) + + snapshot = source.tensor.clone() + with patch.object(torch, "rand", return_value=torch.tensor(0.0)), patch.object( + torch, + "randint", + return_value=torch.tensor(0), + ): + manager.perturb_real_kv_unused_cache(None) + + expected = snapshot.clone() + expected[3, 0] = int(snapshot[3, 0].item()) ^ 0xFF + self.assertTrue(torch.equal(source.tensor, expected)) + + def test_pick_sweep_slot_for_group_skips_locked_radix_nodes(self) -> None: + """Verify unused-cache perturbation chooses only unlocked radix-cache slots.""" + device = DEFAULT_DEVICE + group = make_buffer_group(kind=PoolKind.FULL, has_real_kv=True) + cache = make_radix_cache([[], [1, 2], [3]], device=device) + locked_node = next(iter(cache.root_node.children.values())) + locked_node.lock_ref = 1 + + with patch.object(torch, "randint", return_value=torch.tensor(0)): + slot = real_kv_unused_cache_module._pick_sweep_slot_for_group( + radix_cache=cache, + group=group, + swa_window_size=0, + ) + + self.assertEqual(slot, 3) + + def test_pick_sweep_slot_for_group_translates_swa_slots(self) -> None: + """Verify unused-cache SWA perturbation translates full slots to physical SWA slots.""" + device = DEFAULT_DEVICE + lut = torch.tensor([-1, 2], dtype=torch.int64, device=device) + group = make_buffer_group( + kind=PoolKind.SWA, has_real_kv=True, swa_index_lut=lut + ) + cache = make_radix_cache([[], [1]], device=device) + + with patch.object(torch, "randint", return_value=torch.tensor(0)): + slot = real_kv_unused_cache_module._pick_sweep_slot_for_group( + radix_cache=cache, + group=group, + swa_window_size=4, + ) + + self.assertEqual(slot, 2) + + def test_real_kv_unused_cache_skips_without_radix_cache_when_forward_batch_is_none( + self, + ) -> None: + """Verify unused-cache perturbation accepts no forward batch but skips without radix_cache.""" + device = DEFAULT_DEVICE + group = make_buffer_group(kind=PoolKind.FULL, has_real_kv=True) + source = group.real_kv_sources_k[0] + source.tensor.copy_( + torch.arange(source.tensor.numel(), dtype=torch.uint8).view_as( + source.tensor + ) + ) + manager = PerturbManager( + config=PerturbConfig( + req_to_token_prob=0.0, + real_kv_used_prob=0.0, + real_kv_unused_cache_prob=1.0, + real_kv_post_forward_prob=0.0, + target_group_kind=TargetGroupKind.FULL, + warmup_steps=0, + ), + req_to_token_pool=make_req_to_token_pool(device, max_reqs=4, max_seq_len=8), + buffer_groups=(group,), + outer_step_counter_getter=lambda: 10, + sweep_interval=1, + ) + + snapshot = source.tensor.clone() + with patch.object(torch, "rand", return_value=torch.tensor(0.0)): + manager.perturb_real_kv_unused_cache(None) + + self.assertTrue(torch.equal(source.tensor, snapshot)) + class TestPerturbUtils(CustomTestCase): def test_flip_first_byte_in_physical_swa_slot_does_not_translate_twice( diff --git a/test/registered/mock_model/test_e2e_pd.py b/test/registered/mock_model/test_e2e_pd.py index 40cf6bffb..9594921ed 100644 --- a/test/registered/mock_model/test_e2e_pd.py +++ b/test/registered/mock_model/test_e2e_pd.py @@ -128,5 +128,44 @@ class TestPdTransferCanaryClean(_MockModelPDBase, unittest.TestCase): self.assert_no_canary_violation() +class TestPdTransferChecksumFullRealData(_MockModelPDBase, unittest.TestCase): + """--kv-canary-real-data=all + sweep every step, no perturb, no violation.""" + + extra_prefill_args: ClassVar[List[str]] = mock_model_server_args( + "--skip-server-warmup", + "--kv-canary-real-data", + "all", + "--kv-canary-sweep-interval", + "1", + ) + extra_decode_args: ClassVar[List[str]] = mock_model_server_args( + "--skip-server-warmup", + "--kv-canary-real-data", + "all", + "--kv-canary-sweep-interval", + "1", + "--disaggregation-decode-enable-radix-cache", + ) + + def test_pd_transfer_checksum_full_real_data(self) -> None: + # Step 1: drive traffic through the PD path with full real-KV hashing. + results = _send_parallel_requests( + self.lb_url, + n=_NUM_PROMPTS, + max_new_tokens=_OUTPUT_LEN, + timeout=240.0, + max_workers=_NUM_PROMPTS, + ) + + # Step 2: all requests must succeed. + for result in results: + self.assertEqual(result.get("status_code"), 200, result) + + # Step 3: servers must stay healthy. + self.assertIsNone(self.process_prefill.poll(), "Prefill server died") + self.assertIsNone(self.process_decode.poll(), "Decode server died") + self.assert_no_canary_violation() + + if __name__ == "__main__": unittest.main() diff --git a/test/registered/mock_model/test_self_unit_canary_perturb.py b/test/registered/mock_model/test_self_unit_canary_perturb.py new file mode 100644 index 000000000..b650691e8 --- /dev/null +++ b/test/registered/mock_model/test_self_unit_canary_perturb.py @@ -0,0 +1,47 @@ +import logging +import unittest +from unittest.mock import Mock + +from sglang.srt.kv_canary.perturb import real_kv_used +from sglang.srt.kv_canary.perturb.config import PerturbConfig, TargetGroupKind +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=5, suite="base-a-test-cpu") + + +class TestCanaryPerturb(CustomTestCase): + def test_real_kv_used_logs_when_target_group_has_no_real_kv_sources(self) -> None: + """Verify real KV used perturbation logs when the target group has no sources.""" + config = PerturbConfig( + req_to_token_prob=0.0, + real_kv_used_prob=1.0, + real_kv_unused_cache_prob=0.0, + real_kv_post_forward_prob=0.0, + target_group_kind=TargetGroupKind.FULL, + warmup_steps=0, + ) + warmup_gate = Mock() + warmup_gate.is_in_warmup.return_value = False + + # Empty buffer_groups means pick_target_group returns None, so run() takes + # the early-return branch before any slot is picked. + with self.assertLogs(real_kv_used.logger.name, level=logging.INFO) as logs: + real_kv_used.run( + maybe_inaccurate_forward_batch=Mock(), + config=config, + req_to_token_pool=Mock(), + buffer_groups=(), + swa_window_size=0, + warmup_gate=warmup_gate, + ) + + self.assertIn( + "kv_canary perturb real_kv_used: skipped because no target group with " + "real_kv_sources_k matched target_group_kind=full", + "\n".join(logs.output), + ) + + +if __name__ == "__main__": + unittest.main()