From f444b5897bf0328e60e4f9a1056300804deb9b7b Mon Sep 17 00:00:00 2001 From: Jackey Hua <107608053+zhendonghua@users.noreply.github.com> Date: Tue, 23 Jun 2026 17:04:45 -0700 Subject: [PATCH] [Spec][1/N] Decoupled speculative decoding: IPC protocol + cross-process request id + server flags (#27634) --- python/sglang/srt/server_args.py | 55 ++- .../srt/speculative/decoupled_spec_io.py | 384 ++++++++++++++++++ .../unit/server_args/test_server_args.py | 94 +++++ .../unit/spec/test_decoupled_spec_io.py | 272 +++++++++++++ 4 files changed, 804 insertions(+), 1 deletion(-) create mode 100644 python/sglang/srt/speculative/decoupled_spec_io.py create mode 100644 test/registered/unit/spec/test_decoupled_spec_io.py diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 125940710..82f8a3d74 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -59,6 +59,7 @@ from sglang.srt.model_executor.cuda_graph_config import ( ) from sglang.srt.parser.reasoning_parser import ReasoningParser from sglang.srt.platforms import current_platform +from sglang.srt.speculative.decoupled_spec_io import DecoupledSpecIpcConfig from sglang.srt.utils.common import ( LORA_TARGET_ALL_MODULES, SUPPORTED_LORA_TARGET_MODULES, @@ -1558,7 +1559,36 @@ class ServerArgs: "Path to a JSON config file for adaptive speculative decoding tuning knobs.", ] = None - # ------------------------------------------------------------------------- + # Decoupled speculative decoding: draft and verify run as + # separate engines, currently connected by a ZMQ IPC mesh. + decoupled_spec_bind_endpoint: A[ + Optional[str], + "ZMQ endpoint this engine binds for its inbound channel in decoupled " + "speculative decoding (verifier: result PULL; drafter: control PULL).", + ] = None + decoupled_spec_connect_endpoints: A[ + Optional[List[str]], + Arg( + help="Peer inbound (bind) endpoints to connect to, ordered by peer " + "rank, for decoupled speculative decoding.", + type_parser=json_list_type, + ), + ] = None + decoupled_spec_rank: A[ + Optional[int], + "This engine's rank within its own role space (verifier-rank or " + "drafter-rank) for decoupled speculative decoding.", + ] = None + decoupled_spec_role: A[ + Literal["null", "verifier", "drafter"], + "Role in decoupled speculative decoding: 'null' disables it, 'verifier' " + "runs the target/verify half, 'drafter' runs the draft half.", + ] = "null" + spec_trace_dir: A[ + Optional[str], + "Directory to write decoupled speculative decoding trace files.", + ] = None + # Speculative decoding (ngram) # ------------------------------------------------------------------------- speculative_ngram_min_bfs_breadth: A[ @@ -7297,6 +7327,9 @@ class PortArgs: # The ipc filename for MultiTokenizerRouter to receive inputs from TokenizerWorker processes (zmq) tokenizer_worker_ipc_name: Optional[str] + # The ipc endpoints between verifier scheduler and drafter scheduler + decoupled_spec_ipc_config: Optional[DecoupledSpecIpcConfig] + # zmq address for load snapshot PUSH/PULL (dp-attention TCP mode only; # empty when IPC mode derives the address from instance_id). load_collector_ipc_name: str = "" @@ -7325,6 +7358,24 @@ class PortArgs: instance_id = uuid.uuid4().hex[:12] + decoupled_spec_ipc_config = None + if server_args.decoupled_spec_role != "null": + if ( + server_args.decoupled_spec_bind_endpoint is None + or server_args.decoupled_spec_connect_endpoints is None + or server_args.decoupled_spec_rank is None + ): + raise ValueError( + "--decoupled-spec-bind-endpoint, " + "--decoupled-spec-connect-endpoints, and " + "--decoupled-spec-rank are required for decoupled speculative decoding." + ) + decoupled_spec_ipc_config = DecoupledSpecIpcConfig( + bind_endpoint=server_args.decoupled_spec_bind_endpoint, + connect_endpoints=tuple(server_args.decoupled_spec_connect_endpoints), + rank=int(server_args.decoupled_spec_rank), + ) + if not server_args.enable_dp_attention: # Normal case, use IPC within a single node return PortArgs( @@ -7335,6 +7386,7 @@ class PortArgs: rpc_ipc_name=f"ipc://{tempfile.NamedTemporaryFile(delete=False).name}", metrics_ipc_name=f"ipc://{tempfile.NamedTemporaryFile(delete=False).name}", tokenizer_worker_ipc_name=tokenizer_worker_ipc_name, + decoupled_spec_ipc_config=decoupled_spec_ipc_config, instance_id=instance_id, ) else: @@ -7405,6 +7457,7 @@ class PortArgs: rpc_ipc_name=NetworkAddress(dist_init_host, rpc_port).to_tcp(), metrics_ipc_name=NetworkAddress(dist_init_host, metrics_port).to_tcp(), tokenizer_worker_ipc_name=tokenizer_worker_ipc_name, + decoupled_spec_ipc_config=decoupled_spec_ipc_config, load_collector_ipc_name=NetworkAddress( dist_init_host, load_collector_port ).to_tcp(), diff --git a/python/sglang/srt/speculative/decoupled_spec_io.py b/python/sglang/srt/speculative/decoupled_spec_io.py new file mode 100644 index 000000000..65744bbb5 --- /dev/null +++ b/python/sglang/srt/speculative/decoupled_spec_io.py @@ -0,0 +1,384 @@ +from __future__ import annotations + +from dataclasses import dataclass, field +from enum import Enum +from typing import Callable, Optional + + +class DraftMeshMessageType(str, Enum): + CONTROL_BATCH = "control_batch" + TAIL_STREAM_OUTPUT_BATCH = "tail_stream_output_batch" + + +@dataclass(frozen=True) +class DraftReqKey: + """Request identity on the drafter side. + + The original request_id is only unique within the verifier that owns it. + src_verifier_rank keeps the drafter-side request table unambiguous when + multiple verifier ranks send work to the same drafter rank. + """ + + src_verifier_rank: int + request_id: str + + +def build_draft_scheduler_rid(draft_key: DraftReqKey) -> str: + return f"draft:{int(draft_key.src_verifier_rank)}:{draft_key.request_id}" + + +def parse_draft_scheduler_rid(rid: str) -> DraftReqKey: + if rid.startswith("draft:"): + encoded = rid[len("draft:") :] + rank_text, sep, request_id = encoded.partition(":") + if sep and request_id: + return DraftReqKey( + src_verifier_rank=int(rank_text), + request_id=request_id, + ) + + raise ValueError(f"Invalid decoupled draft scheduler rid: {rid}") + + +@dataclass +class DraftSync: + """Open or re-open a drafter request from a verifier-owned prefix. + + The verifier is the source of truth for committed tokens. DraftSync gives + the drafter the prompt and already committed output prefix that it must + align to before it can emit draft tail tokens. + """ + + request_id: str + src_verifier_rank: int + dst_drafter_rank: int + prompt_token_ids: list[int] = field(default_factory=list) + committed_outputs: list[int] = field(default_factory=list) + + @property + def draft_key(self) -> DraftReqKey: + return DraftReqKey( + src_verifier_rank=int(self.src_verifier_rank), + request_id=self.request_id, + ) + + +@dataclass +class VerifyCommit: + """ + Sent from verifier to drafter to commit a portion of the draft outputs. + + committed_tokens is the verifier-committed contiguous output segment: + output_ids[ + pre_verify_committed_len: + pre_verify_committed_len + len(committed_tokens) + ]. + Drafter must align its reqs to these committed tokens, + and sometimes needs to truncate tokens / reprefill. + """ + + request_id: str + src_verifier_rank: int + dst_drafter_rank: int + pre_verify_committed_len: int + committed_tokens: list[int] + + @property + def draft_key(self) -> DraftReqKey: + return DraftReqKey( + src_verifier_rank=int(self.src_verifier_rank), + request_id=self.request_id, + ) + + def validate_committed_tokens(self) -> None: + if not self.committed_tokens: + raise ValueError( + "VerifyCommit committed_tokens must be non-empty: " + f"request_id={self.request_id} " + f"pre_verify_committed_len={self.pre_verify_committed_len}" + ) + if int(self.pre_verify_committed_len) < 0: + raise ValueError( + "VerifyCommit pre_verify_committed_len must be non-negative: " + f"request_id={self.request_id} " + f"pre_verify_committed_len={self.pre_verify_committed_len}" + ) + + +@dataclass +class DraftClose: + request_id: str + src_verifier_rank: int + dst_drafter_rank: int + reason: str + + @property + def draft_key(self) -> DraftReqKey: + return DraftReqKey( + src_verifier_rank=int(self.src_verifier_rank), + request_id=self.request_id, + ) + + +@dataclass +class DraftTailStreamOutput: + """ + Drafter sends one output token to the verifier-side DraftTailBuffer. + + base_committed_len records the verifier prefix length that the drafter used + as the base when this token was emitted. The verifier compares it with its + stale-base boundary before accepting the token as tail data or as + pending-prefix confirmation. + + new_token_pos is the 0-based output token position for new_token. Normal + decode streams send the latest generated token. + """ + + src_drafter_rank: int + dst_verifier_rank: int + request_id: str + base_committed_len: int + new_token_pos: int + new_token: int + + +@dataclass +class DraftTailStreamOutputBatch: + outputs: list[DraftTailStreamOutput] = field(default_factory=list) + + +@dataclass +class DraftControlBatch: + dst_drafter_rank: int + sync_messages: list[DraftSync] = field(default_factory=list) + verify_commit_messages: list[VerifyCommit] = field(default_factory=list) + close_messages: list[DraftClose] = field(default_factory=list) + + +@dataclass +class VerifierCommitSegment: + """Contiguous VerifyCommit messages coalesced for one drafter request. + + When receiving contiguous VerifyCommit messages for the same draft req, + the transport thread(TokenSync thread at drafter side) coalesces them into a single VerifierCommitSegment. + + VerifierCommitSegment represents a contiguous verifier-committed token segment for drafter, + and drafter scheduler should align with these segments before emitting tail tokens + """ + + draft_key: DraftReqKey + dst_drafter_rank: int + pre_verify_committed_len: int + committed_tokens: list[int] = field(default_factory=list) + + @property + def end_committed_len(self) -> int: + return int(self.pre_verify_committed_len) + len(self.committed_tokens) + + def append_message(self, message: VerifyCommit) -> None: + """ + It runs on TokenSyncThread under _pending_lock. That loop only + catches zmq.error.ContextTerminated, so a raise here escapes _run and + silently kills the drafter control thread. It then stops applying + ALL requests' controls while the verifier keeps pushing. + + TODO: 1. peer-data violations (non-contiguous / invalid len) + should quarantine just that request (drop + add to close_keys), not + crash the thread. 2. phase 5.c will handle the drafter failure by + degrading the verifier into normal autoregressive decoding. + """ + if message.draft_key != self.draft_key: + raise RuntimeError( + "Verifier commit segment received a commit for a different " + f"request: segment_key={self.draft_key} message_key={message.draft_key}" + ) + if int(message.dst_drafter_rank) != int(self.dst_drafter_rank): + raise RuntimeError( + "Verifier commit segment received a commit for a different " + "drafter rank: " + f"request_id={message.request_id} " + f"segment_drafter_rank={self.dst_drafter_rank} " + f"message_drafter_rank={message.dst_drafter_rank}" + ) + message.validate_committed_tokens() + pre_verify_committed_len = int(message.pre_verify_committed_len) + if pre_verify_committed_len != self.end_committed_len: + raise RuntimeError( + "Verifier commit segment requires contiguous VerifyCommit " + "messages: " + f"request_id={message.request_id} " + f"expected_pre_verify_committed_len={self.end_committed_len} " + f"actual_pre_verify_committed_len={pre_verify_committed_len}" + ) + + token_ids = [int(token_id) for token_id in message.committed_tokens] + self.committed_tokens.extend(token_ids) + + def extract_prefix(self, num_tokens: int) -> VerifierCommitSegment: + num_tokens = int(num_tokens) + if num_tokens <= 0: + raise ValueError( + "Verifier commit segment prefix length must be positive: " + f"request_id={self.draft_key.request_id} num_tokens={num_tokens}" + ) + if num_tokens > len(self.committed_tokens): + raise ValueError( + "Verifier commit segment prefix length exceeds segment length: " + f"request_id={self.draft_key.request_id} " + f"num_tokens={num_tokens} " + f"segment_len={len(self.committed_tokens)}" + ) + + prefix_tokens = [ + int(token_id) for token_id in self.committed_tokens[:num_tokens] + ] + remaining_tokens = [ + int(token_id) for token_id in self.committed_tokens[num_tokens:] + ] + prefix_segment = VerifierCommitSegment( + draft_key=self.draft_key, + dst_drafter_rank=int(self.dst_drafter_rank), + pre_verify_committed_len=int(self.pre_verify_committed_len), + committed_tokens=prefix_tokens, + ) + self.pre_verify_committed_len = int(self.pre_verify_committed_len) + num_tokens + self.committed_tokens = remaining_tokens + return prefix_segment + + +@dataclass +class DraftControlInbox: + """Drafter-side inbox for verifier control messages. + + The TokenSync thread temporarily stores incoming control messages here. + The drafter scheduler extracts and consumes them each time it finishes a decoding step. + """ + + sync_messages: list[DraftSync] = field(default_factory=list) + verifier_commit_segments: dict[DraftReqKey, VerifierCommitSegment] = field( + default_factory=dict + ) + close_keys: set[DraftReqKey] = field(default_factory=set) + + def is_empty(self) -> bool: + return ( + not self.sync_messages + and not self.verifier_commit_segments + and not self.close_keys + ) + + def pending_control_count(self) -> int: + return ( + len(self.sync_messages) + + len(self.verifier_commit_segments) + + len(self.close_keys) + ) + + def add_control_batch_locked(self, batch: DraftControlBatch) -> None: + for message in batch.close_messages: + self.add_close_key_locked(message.draft_key) + for message in batch.sync_messages: + if message.draft_key not in self.close_keys: + self.sync_messages.append(message) + for message in batch.verify_commit_messages: + self.add_verify_commit_locked(message) + + def add_close_key_locked(self, draft_key: DraftReqKey) -> None: + self.close_keys.add(draft_key) + self.verifier_commit_segments.pop(draft_key, None) + self.sync_messages = [ + message for message in self.sync_messages if message.draft_key != draft_key + ] + + def add_verify_commit_locked(self, message: VerifyCommit) -> None: + if message.draft_key in self.close_keys: + return + segment = self.verifier_commit_segments.get(message.draft_key) + if segment is None: + segment = VerifierCommitSegment( + draft_key=message.draft_key, + dst_drafter_rank=int(message.dst_drafter_rank), + pre_verify_committed_len=int(message.pre_verify_committed_len), + ) + segment.append_message(message) + self.verifier_commit_segments[message.draft_key] = segment + return + segment.append_message(message) + + def extract_ready_controls_locked( + self, + consumable_commit_len: Callable[[VerifierCommitSegment], int], + ) -> ReadyDraftControls: + ready_controls = ReadyDraftControls() + + if self.close_keys: + ready_controls.close_keys = self.close_keys + self.close_keys = set() + + if self.sync_messages: + ready_controls.sync_messages = self.sync_messages + self.sync_messages = [] + + for draft_key, segment in list(self.verifier_commit_segments.items()): + consumable_len = consumable_commit_len(segment) + if consumable_len <= 0: + continue + + ready_controls.ready_commit_segments.append( + segment.extract_prefix(consumable_len) + ) + if not segment.committed_tokens: + self.verifier_commit_segments.pop(draft_key, None) + + return ready_controls + + +@dataclass +class ReadyDraftControls: + sync_messages: list[DraftSync] = field(default_factory=list) + close_keys: set[DraftReqKey] = field(default_factory=set) + ready_commit_segments: list[VerifierCommitSegment] = field(default_factory=list) + + def is_empty(self) -> bool: + return ( + not self.sync_messages + and not self.close_keys + and not self.ready_commit_segments + ) + + def extracted_control_count(self) -> int: + return ( + len(self.sync_messages) + + len(self.close_keys) + + len(self.ready_commit_segments) + ) + + +@dataclass +class DraftMeshMessage: + message_type: DraftMeshMessageType + control_batch: Optional[DraftControlBatch] = None + tail_stream_output_batch: Optional[DraftTailStreamOutputBatch] = None + + @staticmethod + def from_control_batch(message: DraftControlBatch) -> DraftMeshMessage: + return DraftMeshMessage( + message_type=DraftMeshMessageType.CONTROL_BATCH, + control_batch=message, + ) + + @staticmethod + def from_tail_stream_output_batch( + message: DraftTailStreamOutputBatch, + ) -> DraftMeshMessage: + return DraftMeshMessage( + message_type=DraftMeshMessageType.TAIL_STREAM_OUTPUT_BATCH, + tail_stream_output_batch=message, + ) + + +@dataclass(frozen=True) +class DecoupledSpecIpcConfig: + bind_endpoint: str + connect_endpoints: tuple[str, ...] + rank: int diff --git a/test/registered/unit/server_args/test_server_args.py b/test/registered/unit/server_args/test_server_args.py index 5ed5b9687..4af1f3c51 100644 --- a/test/registered/unit/server_args/test_server_args.py +++ b/test/registered/unit/server_args/test_server_args.py @@ -518,6 +518,52 @@ class TestPortArgs(unittest.TestCase): self.assertTrue(port_args.detokenizer_ipc_name.startswith("ipc://")) self.assertIsInstance(port_args.nccl_port, int) + @patch("sglang.srt.server_args.tempfile.NamedTemporaryFile") + def test_init_new_builds_decoupled_spec_ipc_config(self, mock_temp_file): + mock_temp_file.return_value.name = "temp_file" + + server_args = ServerArgs(model_path="dummy") + server_args.nccl_port = None + server_args.enable_dp_attention = False + server_args.decoupled_spec_role = "verifier" + server_args.decoupled_spec_bind_endpoint = "ipc:///tmp/v" + server_args.decoupled_spec_connect_endpoints = ["ipc:///tmp/d"] + server_args.decoupled_spec_rank = 0 + + port_args = PortArgs.init_new(server_args) + + self.assertIsNotNone(port_args.decoupled_spec_ipc_config) + self.assertEqual(port_args.decoupled_spec_ipc_config.rank, 0) + self.assertEqual( + port_args.decoupled_spec_ipc_config.bind_endpoint, "ipc:///tmp/v" + ) + self.assertEqual( + port_args.decoupled_spec_ipc_config.connect_endpoints, ("ipc:///tmp/d",) + ) + + @patch("sglang.srt.server_args.tempfile.NamedTemporaryFile") + def test_init_new_no_decoupled_config_when_role_null(self, mock_temp_file): + mock_temp_file.return_value.name = "temp_file" + + server_args = ServerArgs(model_path="dummy") + server_args.nccl_port = None + server_args.enable_dp_attention = False + # decoupled_spec_role defaults to "null" + + port_args = PortArgs.init_new(server_args) + + self.assertIsNone(port_args.decoupled_spec_ipc_config) + + def test_init_new_decoupled_role_requires_endpoints(self): + server_args = ServerArgs(model_path="dummy") + server_args.nccl_port = None + server_args.enable_dp_attention = False + server_args.decoupled_spec_role = "drafter" + # endpoints intentionally left as their None defaults + + with self.assertRaises(ValueError): + PortArgs.init_new(server_args) + def test_init_new_with_single_node_dp_attention(self): server_args = ServerArgs(model_path="dummy") @@ -919,6 +965,54 @@ class TestNgramExternalSamArgs(CustomTestCase): self.assertIn("external-corpus-max-tokens", str(context.exception)) +class TestDecoupledSpecArgs(CustomTestCase): + """Decoupled speculative-decoding CLI flags. + + These flags are auto-derived from the ``A[...]`` field metadata on + ``ServerArgs``; a bare annotation is silently skipped by + ``add_cli_args_from_dataclass``. This guards against the regression where + the flags went missing (e.g. after rebasing onto the auto-gen + ``add_cli_args``), which the direct-attribute ``PortArgs`` tests cannot + catch because they never exercise the CLI. + """ + + def test_decoupled_spec_cli_flags_round_trip(self): + server_args = prepare_server_args( + [ + "--model-path", + "dummy", + "--decoupled-spec-role", + "verifier", + "--decoupled-spec-bind-endpoint", + "ipc:///tmp/v", + "--decoupled-spec-connect-endpoints", + '["ipc:///tmp/d"]', + "--decoupled-spec-rank", + "0", + "--spec-trace-dir", + "/tmp/tr", + ] + ) + self.assertEqual(server_args.decoupled_spec_role, "verifier") + self.assertEqual(server_args.decoupled_spec_bind_endpoint, "ipc:///tmp/v") + self.assertEqual(server_args.decoupled_spec_connect_endpoints, ["ipc:///tmp/d"]) + self.assertEqual(server_args.decoupled_spec_rank, 0) + self.assertEqual(server_args.spec_trace_dir, "/tmp/tr") + + def test_decoupled_spec_role_defaults_to_null(self): + server_args = prepare_server_args(["--model-path", "dummy"]) + self.assertEqual(server_args.decoupled_spec_role, "null") + self.assertIsNone(server_args.decoupled_spec_bind_endpoint) + self.assertIsNone(server_args.decoupled_spec_connect_endpoints) + self.assertIsNone(server_args.decoupled_spec_rank) + + def test_decoupled_spec_role_rejects_invalid_choice(self): + with self.assertRaises(SystemExit): + prepare_server_args( + ["--model-path", "dummy", "--decoupled-spec-role", "bogus"] + ) + + class TestAdaptiveSpecArgs(CustomTestCase): def test_adaptive_defaults_to_config_step_when_spec_params_omitted(self): with tempfile.NamedTemporaryFile("w", suffix=".json") as f: diff --git a/test/registered/unit/spec/test_decoupled_spec_io.py b/test/registered/unit/spec/test_decoupled_spec_io.py new file mode 100644 index 000000000..8a7c81839 --- /dev/null +++ b/test/registered/unit/spec/test_decoupled_spec_io.py @@ -0,0 +1,272 @@ +"""Unit tests for srt/speculative/decoupled_spec_io. + +decoupled_spec_io is the schema-only IPC layer for decoupled speculative +decoding: protocol message dataclasses, the cross-process request id codec, and +the drafter-side reconciliation helpers. These tests drive the real logic (id +round-trip + parse errors, commit validation, segment coalescing / contiguity / +prefix extraction, and inbox routing) on CPU; there is no GPU or transport here. +""" + +import unittest + +from sglang.srt.speculative.decoupled_spec_io import ( + DraftClose, + DraftControlBatch, + DraftMeshMessage, + DraftMeshMessageType, + DraftReqKey, + DraftSync, + VerifierCommitSegment, + VerifyCommit, + build_draft_scheduler_rid, + parse_draft_scheduler_rid, +) +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=6, suite="base-a-test-cpu") + + +def _commit(rid, *, pre, tokens, src_verifier_rank=0, drafter_rank=0) -> VerifyCommit: + return VerifyCommit( + request_id=rid, + src_verifier_rank=src_verifier_rank, + dst_drafter_rank=drafter_rank, + pre_verify_committed_len=pre, + committed_tokens=list(tokens), + ) + + +def _segment( + rid, *, pre=0, drafter_rank=0, src_verifier_rank=0 +) -> VerifierCommitSegment: + return VerifierCommitSegment( + draft_key=DraftReqKey(src_verifier_rank=src_verifier_rank, request_id=rid), + dst_drafter_rank=drafter_rank, + pre_verify_committed_len=pre, + ) + + +class TestDraftSchedulerRid(CustomTestCase): + def test_round_trip(self): + key = DraftReqKey(src_verifier_rank=3, request_id="req-1") + rid = build_draft_scheduler_rid(key) + self.assertEqual(rid, "draft:3:req-1") + self.assertEqual(parse_draft_scheduler_rid(rid), key) + + def test_request_id_containing_colon_round_trips(self): + # request_id may contain ':' — parse splits on the first ':' only. + key = DraftReqKey(src_verifier_rank=1, request_id="a:b:c") + self.assertEqual(parse_draft_scheduler_rid(build_draft_scheduler_rid(key)), key) + + def test_parse_invalid_rid_raises(self): + for bad in ["no-prefix", "draft:", "draft:0:", "draft:notint:r"]: + with self.subTest(rid=bad): + with self.assertRaises(ValueError): + parse_draft_scheduler_rid(bad) + + +class TestVerifyCommitValidation(CustomTestCase): + def test_empty_tokens_raises(self): + with self.assertRaises(ValueError): + _commit("r", pre=0, tokens=[]).validate_committed_tokens() + + def test_negative_pre_len_raises(self): + with self.assertRaises(ValueError): + _commit("r", pre=-1, tokens=[1]).validate_committed_tokens() + + def test_valid_commit_passes(self): + # Should not raise. + _commit("r", pre=0, tokens=[1, 2]).validate_committed_tokens() + + +class TestVerifierCommitSegment(CustomTestCase): + def test_append_coalesces_contiguous_commits(self): + seg = _segment("r", pre=0) + seg.append_message(_commit("r", pre=0, tokens=[10, 11])) + self.assertEqual(seg.committed_tokens, [10, 11]) + self.assertEqual(seg.end_committed_len, 2) + seg.append_message(_commit("r", pre=2, tokens=[12])) + self.assertEqual(seg.committed_tokens, [10, 11, 12]) + self.assertEqual(seg.end_committed_len, 3) + + def test_append_wrong_request_raises(self): + seg = _segment("r", pre=0) + with self.assertRaises(RuntimeError): + seg.append_message(_commit("other", pre=0, tokens=[1])) + + def test_append_wrong_drafter_rank_raises(self): + seg = _segment("r", pre=0, drafter_rank=0) + with self.assertRaises(RuntimeError): + seg.append_message(_commit("r", pre=0, tokens=[1], drafter_rank=7)) + + def test_append_non_contiguous_raises(self): + seg = _segment("r", pre=0) + seg.append_message(_commit("r", pre=0, tokens=[10])) # end -> 1 + with self.assertRaises(RuntimeError): + seg.append_message(_commit("r", pre=5, tokens=[11])) # gap + + def test_append_runs_message_validation(self): + seg = _segment("r", pre=0) + with self.assertRaises(ValueError): + seg.append_message(_commit("r", pre=0, tokens=[])) # empty -> validate + + def test_extract_prefix_splits_segment(self): + seg = _segment("r", pre=0) + seg.append_message(_commit("r", pre=0, tokens=[10, 11, 12, 13])) + prefix = seg.extract_prefix(2) + self.assertEqual(prefix.committed_tokens, [10, 11]) + self.assertEqual(prefix.pre_verify_committed_len, 0) + # Remainder stays in the original segment, with pre advanced by 2. + self.assertEqual(seg.committed_tokens, [12, 13]) + self.assertEqual(seg.pre_verify_committed_len, 2) + self.assertEqual(seg.end_committed_len, 4) + + def test_extract_prefix_bounds(self): + seg = _segment("r", pre=0) + seg.append_message(_commit("r", pre=0, tokens=[10, 11])) + with self.assertRaises(ValueError): + seg.extract_prefix(0) + with self.assertRaises(ValueError): + seg.extract_prefix(3) # exceeds segment length + + +class TestDraftControlInbox(CustomTestCase): + def _inbox(self): + from sglang.srt.speculative.decoupled_spec_io import DraftControlInbox + + return DraftControlInbox() + + def _sync(self, rid, drafter_rank=0): + return DraftSync( + request_id=rid, src_verifier_rank=0, dst_drafter_rank=drafter_rank + ) + + def _close(self, rid, drafter_rank=0): + return DraftClose( + request_id=rid, + src_verifier_rank=0, + dst_drafter_rank=drafter_rank, + reason="x", + ) + + def test_add_control_batch_routes_each_message_type(self): + inbox = self._inbox() + inbox.add_control_batch_locked( + DraftControlBatch( + dst_drafter_rank=0, + sync_messages=[self._sync("s")], + verify_commit_messages=[_commit("c", pre=0, tokens=[1])], + close_messages=[self._close("x")], + ) + ) + self.assertEqual([m.request_id for m in inbox.sync_messages], ["s"]) + self.assertIn(DraftReqKey(0, "c"), inbox.verifier_commit_segments) + self.assertIn(DraftReqKey(0, "x"), inbox.close_keys) + + def test_close_drops_pending_segment_and_sync(self): + inbox = self._inbox() + inbox.add_control_batch_locked( + DraftControlBatch( + dst_drafter_rank=0, + sync_messages=[self._sync("r")], + verify_commit_messages=[_commit("r", pre=0, tokens=[1])], + ) + ) + inbox.add_close_key_locked(DraftReqKey(0, "r")) + self.assertEqual(inbox.sync_messages, []) + self.assertNotIn(DraftReqKey(0, "r"), inbox.verifier_commit_segments) + self.assertIn(DraftReqKey(0, "r"), inbox.close_keys) + + def test_verify_commit_for_closed_key_is_ignored(self): + inbox = self._inbox() + inbox.add_close_key_locked(DraftReqKey(0, "r")) + inbox.add_verify_commit_locked(_commit("r", pre=0, tokens=[1])) + self.assertNotIn(DraftReqKey(0, "r"), inbox.verifier_commit_segments) + + def test_extract_ready_controls_full_consume(self): + inbox = self._inbox() + inbox.add_control_batch_locked( + DraftControlBatch( + dst_drafter_rank=0, + sync_messages=[self._sync("s")], + verify_commit_messages=[_commit("c", pre=0, tokens=[1, 2])], + close_messages=[self._close("x")], + ) + ) + ready = inbox.extract_ready_controls_locked( + lambda seg: len(seg.committed_tokens) + ) + self.assertEqual([m.request_id for m in ready.sync_messages], ["s"]) + self.assertEqual({k.request_id for k in ready.close_keys}, {"x"}) + self.assertEqual(len(ready.ready_commit_segments), 1) + self.assertEqual(ready.ready_commit_segments[0].committed_tokens, [1, 2]) + # Fully consumed -> the segment is gone; inbox drained. + self.assertTrue(inbox.is_empty()) + + def test_extract_ready_controls_zero_consumable_keeps_segment(self): + inbox = self._inbox() + inbox.add_verify_commit_locked(_commit("c", pre=0, tokens=[1, 2])) + ready = inbox.extract_ready_controls_locked(lambda seg: 0) + self.assertEqual(ready.ready_commit_segments, []) + # Segment is left buffered for a later step. + self.assertIn(DraftReqKey(0, "c"), inbox.verifier_commit_segments) + + def test_extract_ready_controls_partial_consume_buffers_remainder(self): + inbox = self._inbox() + inbox.add_verify_commit_locked(_commit("c", pre=0, tokens=[1, 2, 3])) + ready = inbox.extract_ready_controls_locked(lambda seg: 1) + self.assertEqual(ready.ready_commit_segments[0].committed_tokens, [1]) + # Remainder [2, 3] stays buffered with pre advanced to 1. + seg = inbox.verifier_commit_segments[DraftReqKey(0, "c")] + self.assertEqual(seg.committed_tokens, [2, 3]) + self.assertEqual(seg.pre_verify_committed_len, 1) + + def test_close_in_same_batch_drops_sync_and_commit(self): + # add_control_batch applies close first, so a same-key sync/commit in the + # same batch is dropped/ignored: close wins within one batch. + inbox = self._inbox() + inbox.add_control_batch_locked( + DraftControlBatch( + dst_drafter_rank=0, + sync_messages=[self._sync("r")], + verify_commit_messages=[_commit("r", pre=0, tokens=[1])], + close_messages=[self._close("r")], + ) + ) + self.assertEqual(inbox.sync_messages, []) + self.assertNotIn(DraftReqKey(0, "r"), inbox.verifier_commit_segments) + self.assertEqual({k.request_id for k in inbox.close_keys}, {"r"}) + + def test_two_commits_same_key_coalesce_in_inbox(self): + # A second commit for an existing key appends to the buffered segment. + inbox = self._inbox() + inbox.add_verify_commit_locked(_commit("r", pre=0, tokens=[1, 2])) + inbox.add_verify_commit_locked(_commit("r", pre=2, tokens=[3])) + seg = inbox.verifier_commit_segments[DraftReqKey(0, "r")] + self.assertEqual(seg.committed_tokens, [1, 2, 3]) + self.assertEqual(seg.end_committed_len, 3) + + +class TestDraftMeshMessageEnvelope(CustomTestCase): + def test_from_control_batch_sets_discriminant_and_slot(self): + batch = DraftControlBatch(dst_drafter_rank=0) + msg = DraftMeshMessage.from_control_batch(batch) + self.assertEqual(msg.message_type, DraftMeshMessageType.CONTROL_BATCH) + self.assertIs(msg.control_batch, batch) + self.assertIsNone(msg.tail_stream_output_batch) + + def test_from_tail_stream_output_batch_sets_discriminant_and_slot(self): + from sglang.srt.speculative.decoupled_spec_io import DraftTailStreamOutputBatch + + batch = DraftTailStreamOutputBatch() + msg = DraftMeshMessage.from_tail_stream_output_batch(batch) + self.assertEqual( + msg.message_type, DraftMeshMessageType.TAIL_STREAM_OUTPUT_BATCH + ) + self.assertIs(msg.tail_stream_output_batch, batch) + self.assertIsNone(msg.control_batch) + + +if __name__ == "__main__": + unittest.main(verbosity=3)