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>
This commit is contained in:
@@ -0,0 +1,93 @@
|
||||
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()
|
||||
Reference in New Issue
Block a user