Files
sglang/test/registered/unit/disaggregation/test_common_conn_runtime_fingerprint.py
laoyao0822 f50e2b1e00 Prevent stale CP shared-KV contracts from corrupting prefill
CP shared-KV now uses CP-local current rows consistently across MLA/index current reuse, passes fp8 current-index K through the tai-kernel uint8 ABI, and clears the transient EAGLE CP-local hidden marker after draft capture. The disaggregation bootstrap also fingerprints the runtime source contract so prefill/decode mismatches fail fast instead of silently exchanging incompatible KV metadata.

Constraint: CP shared-KV batch paths flatten current K/V rows in CP-rank-local valid order, not global request order.

Constraint: tai-kernel current-index prepare validates current_index_k as uint8 bytes for fp8 payloads.

Rejected: Keep using global extend offsets for bs>1 current-index reuse | corrupts request-local bases once current_index_kv is CP-local.

Rejected: Infer CP-local EAGLE hidden semantics from tensor shape | static padding and bs>1 can make shape-based inference unsafe.

Confidence: medium

Scope-risk: moderate

Directive: Do not reintroduce forward_batch.out_cache_loc slicing in CP shared-KV current reuse without verifying CP-local owner-lane layout.

Tested: Remote container py_compile for touched runtime/test files.

Tested: Remote PYTHONPATH=python pytest -q test/registered/unit/layers/test_nsa_cp_utils.py test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py test/registered/unit/disaggregation/test_common_conn_runtime_fingerprint.py (198 passed, 2 subtests passed).

Not-tested: Full remote ETE traffic after this commit; accept length and garbage-output recovery still require a fresh prefill/decode run.

Co-authored-by: OmX <omx@oh-my-codex.dev>
2026-06-04 20:22:29 +08:00

94 lines
3.1 KiB
Python

import types
import unittest
from unittest.mock import Mock, patch
from sglang.srt.disaggregation.common import conn as common_conn
class TestDisaggregationRuntimeFingerprint(unittest.TestCase):
def _manager(self):
mgr = common_conn.CommonKVManager.__new__(common_conn.CommonKVManager)
mgr.prefill_info_table = {}
mgr.kv_args = types.SimpleNamespace(page_size=64, engine_rank=0)
mgr.server_args = types.SimpleNamespace(kv_cache_dtype="fp8_e4m3")
mgr._resolve_rank_mapping = Mock()
return mgr
def _response(self, payload, status_code=200):
response = Mock()
response.status_code = status_code
response.json.return_value = payload
response.text = ""
return response
def _payload(self, **overrides):
payload = {
"attn_tp_size": 8,
"attn_cp_size": 8,
"dp_size": 1,
"pp_size": 1,
"page_size": 64,
"kv_cache_dtype": "fp8_e4m3",
"follow_bootstrap_room": True,
}
payload.update(overrides)
return payload
def test_try_ensure_parallel_info_rejects_runtime_source_fingerprint_mismatch(self):
mgr = self._manager()
with patch.object(
common_conn.requests,
"get",
return_value=self._response(
self._payload(
runtime_source_fingerprint="prefill-source-hash",
runtime_source_root="/sgl-workspace/sglang-tai",
)
),
), patch.object(
common_conn,
"get_runtime_source_fingerprint",
return_value="decode-source-hash",
), patch.object(
common_conn,
"get_runtime_source_root",
return_value="/sgl-workspace/sglang",
):
with self.assertRaisesRegex(
RuntimeError,
"Runtime source fingerprint mismatch",
):
mgr.try_ensure_parallel_info("g0034:18992")
mgr._resolve_rank_mapping.assert_not_called()
self.assertEqual(mgr.prefill_info_table, {})
def test_try_ensure_parallel_info_accepts_matching_runtime_source_fingerprint(self):
mgr = self._manager()
with patch.object(
common_conn.requests,
"get",
return_value=self._response(
self._payload(
runtime_source_fingerprint="same-source-hash",
runtime_source_root="/sgl-workspace/sglang-tai",
)
),
), patch.object(
common_conn,
"get_runtime_source_fingerprint",
return_value="same-source-hash",
), patch.object(
common_conn,
"get_runtime_source_root",
return_value="/sgl-workspace/sglang-tai",
):
self.assertTrue(mgr.try_ensure_parallel_info("g0034:18992"))
mgr._resolve_rank_mapping.assert_called_once()
self.assertIn("g0034:18992", mgr.prefill_info_table)
if __name__ == "__main__":
unittest.main()