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>
94 lines
3.1 KiB
Python
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()
|