diff --git a/.github/workflows/test.yaml b/.github/workflows/test.yaml index 724f9c107..5585a61a1 100644 --- a/.github/workflows/test.yaml +++ b/.github/workflows/test.yaml @@ -63,7 +63,8 @@ jobs: run: | source sf/bin/activate export PYTHONPATH=$PWD - bash scripts/apply_sglang_spec_capture_patch.sh + bash scripts/apply_sglang_spec_capture_patch.sh \ + --python "$PWD/sf/bin/python" --apply python -c 'import importlib.util, shutil, torch; assert torch.cuda.is_available(), "live capture gate requires CUDA"; assert importlib.util.find_spec("mooncake.store"), "live capture gate requires mooncake.store"; assert shutil.which("mooncake_master") or __import__("os").environ.get("MOONCAKE_MASTER_SERVER_ADDR"), "live capture gate requires mooncake_master"' python -m unittest tests.test_runtime.test_server_capture_gate -v diff --git a/examples/configs/README.md b/examples/configs/README.md index 025b3943b..6075ecbe5 100644 --- a/examples/configs/README.md +++ b/examples/configs/README.md @@ -38,7 +38,8 @@ training, `*-offline.yaml` consumes precomputed features, and recipe is disaggregated even when its historical filename only says `online`. VLM training is not supported, so the catalog contains text-only recipes. -The `qwen3-8b-dflash-1server-dp7-disaggregated.yaml`, +The `qwen3-8b-dflash-windowed-fanout.yaml`, +`qwen3-8b-dflash-1server-dp7-disaggregated.yaml`, `qwen3-8b-domino-1server-dp7-disaggregated.yaml`, `qwen3-8b-domino-multiserver-disaggregated.yaml`, `qwen3.6-27b-dflash-1server-dp2-disaggregated.yaml`, and @@ -114,6 +115,7 @@ assume the command runs from the repository root. | Colocated offline | `qwen3-8b-eagle3-offline.yaml` | | External-service online | `qwen3-8b-eagle3-disaggregated.yaml` | | Managed-local disaggregated online | `qwen3-8b-domino-multiserver-disaggregated.yaml` | +| Independent heterogeneous consumers | `qwen3-8b-dflash-windowed-fanout.yaml` | | Disaggregated offline | `qwen3-8b-eagle3-offline-disaggregated.yaml` | The online/offline mode is derived from the selected `data` source, not from @@ -310,7 +312,7 @@ Managed-local fields: | Field | Default | What to write | | --- | --- | --- | -| `deployment.disaggregated.managed_local.trainer_cuda_visible_devices` | required | One device token per `nproc_per_node`; trainer and capture devices must not overlap. | +| `deployment.disaggregated.managed_local.trainer_cuda_visible_devices` | required | One device token per trainer rank, or one per independent `windowed_fanout` consumer; trainer and capture devices must not overlap. | | `deployment.disaggregated.managed_local.mooncake` | default object | Owned loopback Mooncake configuration described by the nested fields below. | | `deployment.disaggregated.managed_local.capture_servers` | required | One or more owned patched SGLang server definitions. | | `deployment.disaggregated.managed_local.shutdown_grace_s` | `30` | Positive graceful process-group shutdown window. | @@ -336,6 +338,51 @@ endpoints, `store_root`, or `producer_segment_size`. It does not support resume, an existing torchrun, or `--node-rank`. All owned ports and GPU assignments must be disjoint. +### Independent windowed fanout + +`deployment.disaggregated.windowed_fanout` runs one producer and multiple +single-process consumers from the same YAML. Each child still enters through +`specforge train`; `launch_plan` selects a consumer with `--consumer-id` and +projects its loss, block size, anchors, optimizer schedule, seed, checkpoint, +and output directory onto the canonical trainer configuration. + +This topology is intentionally different from data-parallel training. Consumers +advance independently and may have different step costs. A shared SQLite +registry records each cursor and bounded capture interest, while Mooncake owns +the tensor payloads. The producer captures a sample only when a consumer's +window requests it, reuses a live compatible capture across consumers, and +reclaims entries outside every legal window. + +Required constraints are: + +- online DFlash with Mooncake and exactly one prompt epoch; +- a positive fixed `data.max_prompts`, divisible by batch size times + accumulation; +- `deployment.trainer` fixed at 1x1 because each consumer is its own process; +- exactly one capture server, shared by every consumer; +- one managed-local trainer device per consumer, or one explicit + `cuda_visible_device` on every external consumer; +- `max_live_refs`, `max_live_bytes`, and `max_outstanding_per_consumer` large + enough for the configured batch and capture reservation. + +The main bounds and timing controls are: + +| Field | Default | Meaning | +| --- | --- | --- | +| `window_lookbehind` / `window_lookahead` | `2` / `40` | Shared default legal window around each consumer cursor. | +| `max_prefetch_per_consumer` | `8` | Maximum speculative requests inside the lookahead window. | +| `max_outstanding_per_consumer` | `8` | Hard per-consumer acquired-ref bound; must cover one accumulated optimizer step. | +| `max_live_refs` / `max_live_bytes` | required byte bound | Global capture capacity shared by all consumers. | +| `capture_reservation_bytes` | 128 MiB | Bytes reserved transactionally before a capture starts. | +| `capture_max_sample_bytes` | 128 MiB | Server-side rejection bound for one captured sample; it cannot exceed the transactional reservation. | +| `capture_batch_size` / `capture_batch_wait_s` | `8` / `0.002` | Producer request batching. | +| `consumer_prefetch_batches` | `1` | Loader-side prefetch workers for each consumer. | + +Each consumer may override `window_lookbehind`, `window_lookahead`, and +`max_prefetch` in addition to its required training fields. External deployments +may resume consumers independently with `resume_from`; managed-local runs are +fresh attempts and reject resume checkpoints. + ### `runtime`: streaming backpressure This section affects disaggregated streaming producers and is normally omitted diff --git a/examples/configs/qwen3-8b-dflash-windowed-fanout.yaml b/examples/configs/qwen3-8b-dflash-windowed-fanout.yaml new file mode 100644 index 000000000..9fd52ceb5 --- /dev/null +++ b/examples/configs/qwen3-8b-dflash-windowed-fanout.yaml @@ -0,0 +1,98 @@ +model: + target_model_path: Qwen/Qwen3-8B + draft_model_config: configs/qwen3-8b-dflash.json + target_backend: sglang + trust_remote_code: true + embedding_key: model.embed_tokens.weight + torch_dtype: bfloat16 + mask_token_id: 151669 + sglang_attention_backend: flashinfer + sglang_mem_fraction_static: 0.5 + +data: + train_data_path: ./cache/dataset/perfectblend_qwen3-8b_regen.jsonl + max_prompts: 48 + max_length: 3072 + chat_template: qwen + build_dataset_num_proc: 32 + cache_dir: ./cache + +training: + strategy: dflash + num_epochs: 1 + max_steps: 24 + total_steps: 24 + batch_size: 2 + accumulation_steps: 1 + attention_backend: flex_attention + max_grad_norm: 1.0 + save_interval: 0 + log_interval: 1 + dist_timeout: 30 + +tracking: + report_to: none + +run_id: qwen3-8b-dflash-windowed-fanout +output_dir: ./outputs/qwen3-8b-dflash-windowed-fanout + +deployment: + mode: disaggregated + trainer: + nnodes: 1 + nproc_per_node: 1 + disaggregated: + control_dir: ./outputs/qwen3-8b-dflash-windowed-fanout/control + backend: mooncake + client_buffer_size: 1073741824 + windowed_fanout: + window_lookbehind: 2 + window_lookahead: 16 + max_prefetch_per_consumer: 8 + max_outstanding_per_consumer: 8 + max_live_refs: 48 + max_live_bytes: 25769803776 + capture_reservation_bytes: 536870912 + capture_max_sample_bytes: 536870912 + capture_batch_size: 8 + consumer_prefetch_batches: 1 + consumers: + - consumer_id: dflash-b4 + seed: 42 + loss_type: dflash + loss_decay_gamma: 7.0 + dpace_alpha: 0.5 + draft_block_size: 4 + num_anchors: 64 + learning_rate: 0.0006 + warmup_ratio: 0.04 + - consumer_id: dflash-b8 + seed: 43 + loss_type: dflash + loss_decay_gamma: 7.0 + dpace_alpha: 0.5 + draft_block_size: 8 + num_anchors: 128 + learning_rate: 0.0006 + warmup_ratio: 0.04 + - consumer_id: dflash-b16 + seed: 44 + loss_type: dflash + loss_decay_gamma: 7.0 + dpace_alpha: 0.5 + draft_block_size: 16 + num_anchors: 256 + learning_rate: 0.0006 + warmup_ratio: 0.04 + managed_local: + trainer_cuda_visible_devices: ["1", "2", "3"] + shutdown_grace_s: 120 + mooncake: + protocol: tcp + global_segment_size_bytes: 34359738368 + local_buffer_size_bytes: 1073741824 + capture_servers: + - port: 30000 + cuda_visible_devices: ["0"] + tp_size: 1 + mem_fraction_static: 0.5 diff --git a/examples/disagg/README.md b/examples/disagg/README.md index 6be6e0841..9b52f8943 100644 --- a/examples/disagg/README.md +++ b/examples/disagg/README.md @@ -109,12 +109,20 @@ trainer in one YAML: specforge train -c \ examples/configs/qwen3-8b-dflash-1server-dp7-disaggregated.yaml +specforge train -c \ + examples/configs/qwen3-8b-dflash-windowed-fanout.yaml + specforge train -c \ examples/configs/qwen3-8b-domino-1server-dp7-disaggregated.yaml ``` -These recipes preserve the old DFlash and Domino one-server + DP7 -self-contained topologies. The genuine two-server Domino recipe is: +The first recipe preserves the DFlash one-server + DP7 topology. The windowed +recipe instead launches one producer on GPU 0 and three independent DFlash +consumers on GPUs 1-3. They share compatible captures but keep separate cursors, +windows, trainer state, and hyperparameters. Its block/anchor pairs are 4/64, +8/128, and 16/256 so the example does not depend on optional compact-loss or +custom-kernel optimizations. The Domino recipe preserves its +one-server + DP7 topology. The genuine two-server Domino recipe is: ```bash specforge train -c \ diff --git a/patches/sglang/v0.5.14/spec-capture-base-to-current.patch b/patches/sglang/v0.5.14/spec-capture-base-to-current.patch new file mode 100644 index 000000000..e7114f51d --- /dev/null +++ b/patches/sglang/v0.5.14/spec-capture-base-to-current.patch @@ -0,0 +1,712 @@ +diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py +index a93a466..e4e84cd 100755 +--- a/python/sglang/srt/managers/schedule_batch.py ++++ b/python/sglang/srt/managers/schedule_batch.py +@@ -772,10 +772,10 @@ class Req(ReqDllmMixin): + } + self.sampling_params = sampling_params + self.custom_logit_processor = custom_logit_processor +- # Spec-training capture: piggyback the return_hidden_states path (fires +- # CaptureHiddenMode.FULL); the sink consumes the slices, not the response. ++ # Keep the client response flag separate from internal spec capture. ++ # ScheduleBatch enables CaptureHiddenMode.FULL for either use case. + self.spec_capture = spec_capture +- self.return_hidden_states = return_hidden_states or spec_capture is not None ++ self.return_hidden_states = return_hidden_states + self.spec_capture_aux: List[torch.Tensor] = [] + self.spec_capture_last_hidden: List[torch.Tensor] = [] + self.spec_capture_result = None # per-request sink result -> output field +@@ -1872,7 +1872,10 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): + has_grammar=any(req.grammar for req in reqs), + device=req_to_token_pool.device, + spec_algorithm=spec_algorithm, +- return_hidden_states=any(req.return_hidden_states for req in reqs), ++ return_hidden_states=any( ++ req.return_hidden_states or req.spec_capture is not None ++ for req in reqs ++ ), + is_prefill_only=all(req.is_prefill_only for req in reqs), + chunked_req=chunked_req, + chunked_req_next_prompt_token=_compute_chunked_req_next_prompt_token( +diff --git a/python/sglang/srt/managers/scheduler_components/batch_result_processor.py b/python/sglang/srt/managers/scheduler_components/batch_result_processor.py +index 001eff8..afe32e9 100644 +--- a/python/sglang/srt/managers/scheduler_components/batch_result_processor.py ++++ b/python/sglang/srt/managers/scheduler_components/batch_result_processor.py +@@ -484,15 +484,19 @@ class SchedulerBatchResultProcessor: + """Accumulate captured rows as CPU tensors for the Mooncake sink. + + Same offset arithmetic as ``_append_prefill_hidden_states`` but keeps +- tensor slices (aux concat in ``hidden_states``, post-norm last in +- ``last_hidden_states``) rather than the JSON-able response payload. ++ only the artifacts requested by this capture as CPU tensor slices. + """ + start = hidden_state_offset + end = start + len(req.origin_input_ids) +- req.spec_capture_aux.append( +- logits_output.hidden_states[start:end].cpu().clone() +- ) +- if logits_output.last_hidden_states is not None: ++ requested_artifacts = req.spec_capture.get("features") or {} ++ if "aux" in requested_artifacts: ++ req.spec_capture_aux.append( ++ logits_output.hidden_states[start:end].cpu().clone() ++ ) ++ if ( ++ "last_hidden" in requested_artifacts ++ and logits_output.last_hidden_states is not None ++ ): + req.spec_capture_last_hidden.append( + logits_output.last_hidden_states[start:end].cpu().clone() + ) +diff --git a/python/sglang/srt/models/qwen2.py b/python/sglang/srt/models/qwen2.py +index 744e9b1..a52317c 100644 +--- a/python/sglang/srt/models/qwen2.py ++++ b/python/sglang/srt/models/qwen2.py +@@ -386,6 +386,12 @@ class Qwen2Model(nn.Module): + } + ) + else: ++ # Capture after the final transformer layer but before final norm. ++ # The ordinary pre-layer hook cannot observe end_layer itself. ++ if self.end_layer in self.layers_to_capture: ++ aux_hidden_states.append( ++ hidden_states + residual if residual is not None else hidden_states ++ ) + if hidden_states.shape[0] != 0: + if residual is None: + hidden_states = self.norm(hidden_states) +diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py +index d515b3c..03632e7 100644 +--- a/python/sglang/srt/server_args.py ++++ b/python/sglang/srt/server_args.py +@@ -2146,6 +2146,25 @@ class ServerArgs: + "match the draft strategy being trained; they wire capture onto " + "different submodules (VL targets only populate the dflash path).", + ] = "eagle3" ++ spec_capture_store_id: A[ ++ Optional[str], ++ "Server-owned Mooncake namespace for spec capture. Required with " ++ "--enable-spec-capture; requests cannot override it.", ++ ] = None ++ spec_capture_max_sample_bytes: A[ ++ int, ++ "Maximum total bytes a single spec-capture request may hard-pin.", ++ ] = 1 << 30 ++ spec_capture_inventory_db: A[ ++ Optional[str], ++ "Server-side durable capture transaction inventory. Required with " ++ "--enable-spec-capture so response-loss retries can reclaim old keys.", ++ ] = None ++ spec_capture_lifecycle_db: A[ ++ Optional[str], ++ "Shared SpecForge owner lifecycle DB. Every planned Mooncake key is " ++ "recorded here before its first hard-pinned write.", ++ ] = None + enable_return_routed_experts: A[ + bool, + "Enable returning routed experts of each layer with responses.", +diff --git a/python/sglang/srt/spec_capture_sink.py b/python/sglang/srt/spec_capture_sink.py +index 038a084..8281292 100644 +--- a/python/sglang/srt/spec_capture_sink.py ++++ b/python/sglang/srt/spec_capture_sink.py +@@ -6,9 +6,9 @@ + # http://www.apache.org/licenses/LICENSE-2.0 + """Server-side spec-training capture sink (SpecForge DataFlow transport). + +-Under ``--enable-spec-capture``, a request's ``spec_capture`` dict tells this +-sink to write the prefill's captured tensors straight into a Mooncake store +-(one hard-pinned object per tensor at ``{store_id}/{sample_id}/g{gen}/{name}``, ++Under ``--enable-spec-capture``, an authenticated request's ``spec_capture`` ++dict tells this sink to write captured tensors into the server-owned Mooncake ++namespace (one object per tensor at ``{store_id}/{sample_id}/g{gen}/{name}``, + raw bytes — shape/dtype travel on the returned spec). Feature tensors never + touch the response path; ``meta_info["spec_capture"]`` returns only keys + + shapes/dtypes. Strategy naming is the client's (the ``features`` mapping); the +@@ -17,7 +17,8 @@ one-liners, every capture decision lives here. + + Request schema:: + +- {"store_id", "sample_id", "gen", "replace", # key namespace / retry policy ++ {"auth_token", "store_id", "sample_id", "gen", "replace", ++ # capability / namespace / retry + "features": {"aux": , "last_hidden": }, # artifact -> feature + "passthrough": [{"name", "data", "shape", "dtype"}]} # client tensors verbatim + +@@ -30,9 +31,15 @@ Mooncake connection uses the standard ``MOONCAKE_*`` env vars (see + + from __future__ import annotations + ++import hmac ++import json + import logging ++import math + import os ++import re ++import sqlite3 + import threading ++import time + from typing import Any, Dict, List, Optional + + import torch +@@ -57,18 +64,94 @@ _STR_DTYPE = {v: k for k, v in _DTYPE_STR.items()} + + _ARTIFACT_AUX = "aux" + _ARTIFACT_LAST_HIDDEN = "last_hidden" ++_ALLOWED_FEATURE_NAMES = { ++ "attention_mask", ++ "hidden_state", ++ "hidden_states", ++ "input_ids", ++ "loss_mask", ++ "target", ++} ++_SAFE_SAMPLE_ID = re.compile(r"^[A-Za-z0-9_.:-]{1,256}$") ++_MAX_PASSTHROUGH_FEATURES = 8 + + + class SpecCaptureSink: + """Writes captured per-request tensors into Mooncake in SpecForge layout.""" + +- def __init__(self, aux_layer_ids: Optional[List[int]] = None) -> None: ++ def __init__( ++ self, ++ *, ++ store_id: str, ++ auth_token: str, ++ max_sample_bytes: int, ++ inventory_db_path: str, ++ lifecycle_db_path: str, ++ aux_layer_ids: Optional[List[int]] = None, ++ ) -> None: ++ if not store_id or "/" in store_id or store_id in {".", ".."}: ++ raise ValueError("spec-capture store id must be one safe path segment") ++ if not auth_token: ++ raise ValueError("SGLANG_SPEC_CAPTURE_TOKEN must be non-empty") ++ if max_sample_bytes <= 0: ++ raise ValueError("spec-capture max sample bytes must be positive") ++ if not inventory_db_path: ++ raise ValueError("spec-capture inventory DB path must be non-empty") ++ if not lifecycle_db_path: ++ raise ValueError("spec-capture lifecycle DB path must be non-empty") ++ self.store_id = store_id ++ self.auth_token = auth_token ++ self.max_sample_bytes = int(max_sample_bytes) + self.aux_layer_ids = list(aux_layer_ids) if aux_layer_ids else None ++ self.inventory_db_path = os.path.abspath(inventory_db_path) ++ os.makedirs(os.path.dirname(self.inventory_db_path), exist_ok=True) ++ self._inventory = sqlite3.connect( ++ self.inventory_db_path, check_same_thread=False, timeout=30.0 ++ ) ++ self._inventory.execute("PRAGMA journal_mode=WAL") ++ self._inventory.execute("PRAGMA synchronous=FULL") ++ self._inventory.execute("PRAGMA busy_timeout=30000") ++ self._inventory.execute( ++ "CREATE TABLE IF NOT EXISTS captures (" ++ "sample_id TEXT PRIMARY KEY, generation INTEGER NOT NULL, " ++ "keys_json TEXT NOT NULL, result_json TEXT, state TEXT NOT NULL, " ++ "prior_generation INTEGER, prior_keys_json TEXT)" ++ ) ++ capture_columns = { ++ row[1] for row in self._inventory.execute("PRAGMA table_info(captures)") ++ } ++ if "prior_generation" not in capture_columns: ++ self._inventory.execute( ++ "ALTER TABLE captures ADD COLUMN prior_generation INTEGER" ++ ) ++ if "prior_keys_json" not in capture_columns: ++ self._inventory.execute( ++ "ALTER TABLE captures ADD COLUMN prior_keys_json TEXT" ++ ) ++ self._inventory.commit() ++ self.lifecycle_db_path = os.path.abspath(lifecycle_db_path) ++ os.makedirs(os.path.dirname(self.lifecycle_db_path), exist_ok=True) ++ self._lifecycle = sqlite3.connect( ++ self.lifecycle_db_path, check_same_thread=False, timeout=30.0 ++ ) ++ self._lifecycle.execute("PRAGMA journal_mode=WAL") ++ self._lifecycle.execute("PRAGMA synchronous=FULL") ++ self._lifecycle.execute("PRAGMA busy_timeout=30000") ++ self._lifecycle.execute( ++ "CREATE TABLE IF NOT EXISTS mooncake_objects (" ++ "store_id TEXT NOT NULL, sample_id TEXT NOT NULL, " ++ "generation INTEGER NOT NULL, feature_names_json TEXT NOT NULL, " ++ "estimated_bytes INTEGER NOT NULL, state TEXT NOT NULL, " ++ "reason TEXT, updated_at REAL NOT NULL, " ++ "PRIMARY KEY (store_id, sample_id, generation))" ++ ) ++ self._lifecycle.commit() + self._store = None + self._put_config = None + self._lock = threading.Lock() + # Retried HTTP requests reuse deterministic keys. Striped locks keep +- # replacement atomic per key without retaining one lock per sample. ++ # replacement atomic without retaining one lock per sample or key. ++ self._sample_locks = [threading.RLock() for _ in range(256)] + self._write_locks = [threading.Lock() for _ in range(256)] + + # -- connection --------------------------------------------------------- +@@ -127,7 +210,7 @@ class SpecCaptureSink: + if replace: + # Do not probe with is_exist(): Mooncake existence checks can + # acquire a read lease that prevents the following removal. +- self._remove_quiet(key) ++ self._remove_exact(key, force=True) + try: + store.register_buffer(t.data_ptr(), nbytes) + except Exception: +@@ -142,11 +225,281 @@ class SpecCaptureSink: + if rc is not None and int(rc) < 0: + raise RuntimeError(f"spec-capture put_from failed (status {rc}) for {key}") + +- def _remove_quiet(self, key: str) -> None: +- try: +- self._connect().remove(key) +- except Exception: +- pass ++ def _remove_exact(self, key: str, *, force: bool) -> None: ++ store = self._connect() ++ rc = store.remove(key, force) ++ if rc is not None and int(rc) not in (0, -704): ++ raise RuntimeError( ++ f"spec-capture remove failed (status {rc}) for {key}" ++ ) ++ ++ def _inventory_row(self, sample_id: str): ++ return self._inventory.execute( ++ "SELECT generation, keys_json, result_json, state, " ++ "prior_generation, prior_keys_json FROM captures " ++ "WHERE sample_id=?", ++ (sample_id,), ++ ).fetchone() ++ ++ def _record_lifecycle( ++ self, ++ sample_id: str, ++ gen: int, ++ feature_names: List[str], ++ estimated_bytes: int, ++ state: str, ++ ) -> None: ++ names_json = json.dumps(sorted(feature_names), separators=(",", ":")) ++ row = self._lifecycle.execute( ++ "SELECT feature_names_json, estimated_bytes, state FROM " ++ "mooncake_objects WHERE store_id=? AND sample_id=? AND generation=?", ++ (self.store_id, sample_id, gen), ++ ).fetchone() ++ if row is not None: ++ if row[0] != names_json or int(row[1]) != int(estimated_bytes): ++ raise RuntimeError( ++ f"spec-capture lifecycle identity changed for {sample_id} " ++ f"generation {gen}" ++ ) ++ if row[2] in {"tombstoned", "cleaned"}: ++ raise RuntimeError( ++ f"spec-capture lifecycle refuses {sample_id} generation {gen} " ++ f"from state {row[2]!r}" ++ ) ++ if state == "planned" or row[2] == "resident": ++ return ++ self._lifecycle.execute( ++ "UPDATE mooncake_objects SET state='resident', updated_at=? WHERE " ++ "store_id=? AND sample_id=? AND generation=? AND state='planned'", ++ (time.time(), self.store_id, sample_id, gen), ++ ) ++ else: ++ self._lifecycle.execute( ++ "INSERT INTO mooncake_objects " ++ "(store_id, sample_id, generation, feature_names_json, " ++ "estimated_bytes, state, reason, updated_at) " ++ "VALUES (?, ?, ?, ?, ?, ?, NULL, ?)", ++ ( ++ self.store_id, ++ sample_id, ++ gen, ++ names_json, ++ int(estimated_bytes), ++ state, ++ time.time(), ++ ), ++ ) ++ self._lifecycle.commit() ++ ++ def _mark_lifecycle_cleaned(self, sample_id: str, gen: int) -> None: ++ self._lifecycle.execute( ++ "UPDATE mooncake_objects SET state='cleaned', updated_at=? WHERE " ++ "store_id=? AND sample_id=? AND generation=?", ++ (time.time(), self.store_id, sample_id, gen), ++ ) ++ self._lifecycle.commit() ++ ++ def _reject_stale_generation(self, sample_id: str, gen: int) -> None: ++ row = self._inventory_row(sample_id) ++ if row is not None and int(row[0]) > gen: ++ raise ValueError( ++ f"stale spec-capture generation {gen}; latest is {row[0]}" ++ ) ++ ++ def _finish_replacement( ++ self, ++ sample_id: str, ++ gen: int, ++ prior_gen: int, ++ prior_keys: List[str], ++ ) -> None: ++ for key in prior_keys: ++ self._remove_exact(key, force=False) ++ self._mark_lifecycle_cleaned(sample_id, prior_gen) ++ cursor = self._inventory.execute( ++ "UPDATE captures SET state='writing', prior_generation=NULL, " ++ "prior_keys_json=NULL WHERE sample_id=? AND generation=? " ++ "AND state='replacing'", ++ (sample_id, gen), ++ ) ++ if cursor.rowcount != 1: ++ raise RuntimeError( ++ f"spec-capture replacement journal changed for {sample_id} " ++ f"generation {gen}" ++ ) ++ self._inventory.commit() ++ ++ def _prepare_write(self, sample_id: str, gen: int, keys: List[str]): ++ row = self._inventory_row(sample_id) ++ if row is None: ++ self._inventory.execute( ++ "INSERT INTO captures " ++ "(sample_id, generation, keys_json, result_json, state, " ++ "prior_generation, prior_keys_json) " ++ "VALUES (?, ?, ?, NULL, 'writing', NULL, NULL)", ++ (sample_id, gen, json.dumps(keys, separators=(",", ":"))), ++ ) ++ self._inventory.commit() ++ return None ++ ++ current_gen, keys_json, result_json, state, prior_gen, prior_keys_json = row ++ current_gen = int(current_gen) ++ if current_gen > gen: ++ raise ValueError( ++ f"stale spec-capture generation {gen}; latest is {current_gen}" ++ ) ++ if current_gen == gen and state == "committed": ++ return json.loads(result_json) ++ if state == "replacing": ++ if prior_gen is None or prior_keys_json is None: ++ raise RuntimeError( ++ f"incomplete replacement journal for {sample_id} generation " ++ f"{current_gen}" ++ ) ++ self._finish_replacement( ++ sample_id, ++ current_gen, ++ int(prior_gen), ++ list(json.loads(prior_keys_json)), ++ ) ++ row = self._inventory_row(sample_id) ++ current_gen, keys_json, result_json, state, _, _ = row ++ current_gen = int(current_gen) ++ if current_gen == gen: ++ if state != "writing": ++ raise RuntimeError( ++ f"invalid spec-capture inventory state {state!r} for " ++ f"{sample_id}" ++ ) ++ for key in json.loads(keys_json): ++ self._remove_exact(key, force=True) ++ return None ++ if state not in {"writing", "committed"}: ++ raise RuntimeError( ++ f"invalid spec-capture inventory state {state!r} for {sample_id}" ++ ) ++ ++ prior_keys = list(json.loads(keys_json)) ++ self._inventory.execute( ++ "INSERT OR REPLACE INTO captures " ++ "(sample_id, generation, keys_json, result_json, state, " ++ "prior_generation, prior_keys_json) " ++ "VALUES (?, ?, ?, NULL, 'replacing', ?, ?)", ++ ( ++ sample_id, ++ gen, ++ json.dumps(keys, separators=(",", ":")), ++ current_gen, ++ json.dumps(prior_keys, separators=(",", ":")), ++ ), ++ ) ++ self._inventory.commit() ++ self._finish_replacement(sample_id, gen, current_gen, prior_keys) ++ return None ++ ++ def _commit_write(self, sample_id: str, gen: int, result: Dict[str, Any]): ++ cursor = self._inventory.execute( ++ "UPDATE captures SET result_json=?, state='committed' " ++ "WHERE sample_id=? AND generation=? AND state='writing'", ++ (json.dumps(result, separators=(",", ":")), sample_id, gen), ++ ) ++ if cursor.rowcount != 1: ++ raise RuntimeError( ++ f"spec-capture write transaction changed for {sample_id} " ++ f"generation {gen}" ++ ) ++ self._inventory.commit() ++ ++ @staticmethod ++ def _validate_feature_name(name: Any) -> str: ++ name = str(name) ++ if name not in _ALLOWED_FEATURE_NAMES: ++ raise ValueError(f"spec-capture feature name {name!r} is not allowed") ++ return name ++ ++ def _validate_request( ++ self, ++ spec: Dict[str, Any], ++ aux: Optional[torch.Tensor], ++ last_hidden: Optional[torch.Tensor], ++ ) -> tuple[ ++ str, str, int, bool, Dict[str, str], List[Dict[str, Any]], int ++ ]: ++ supplied = str(spec.get("auth_token") or "") ++ if not hmac.compare_digest(supplied, self.auth_token): ++ raise PermissionError("invalid spec-capture capability") ++ store_id = str(spec.get("store_id") or "") ++ if store_id != self.store_id: ++ raise ValueError( ++ "spec-capture store_id must match the server-owned namespace" ++ ) ++ sample_id = str(spec.get("sample_id") or "") ++ if not _SAFE_SAMPLE_ID.fullmatch(sample_id) or not sample_id.startswith( ++ self.store_id + ":" ++ ): ++ raise ValueError( ++ "spec-capture sample_id must be safe and prefixed by the " ++ "server-owned store id" ++ ) ++ gen = int(spec.get("gen", 1)) ++ if not 1 <= gen <= 2**31 - 1: ++ raise ValueError("spec-capture generation is out of range") ++ replace = spec.get("replace", False) ++ if type(replace) is not bool: ++ raise ValueError("spec-capture replace must be a boolean") ++ ++ features = dict(spec.get("features") or {}) ++ unknown_artifacts = set(features) - {_ARTIFACT_AUX, _ARTIFACT_LAST_HIDDEN} ++ if unknown_artifacts: ++ raise ValueError( ++ f"unknown spec-capture artifacts: {sorted(unknown_artifacts)}" ++ ) ++ features = { ++ artifact: self._validate_feature_name(name) ++ for artifact, name in features.items() ++ } ++ passthrough = list(spec.get("passthrough") or []) ++ if len(passthrough) > _MAX_PASSTHROUGH_FEATURES: ++ raise ValueError("too many spec-capture passthrough features") ++ ++ names = list(features.values()) ++ total_bytes = 0 ++ selected_artifacts = ( ++ (_ARTIFACT_AUX, aux), ++ (_ARTIFACT_LAST_HIDDEN, last_hidden), ++ ) ++ for artifact, tensor in selected_artifacts: ++ if artifact in features and tensor is not None: ++ total_bytes += tensor.element_size() * tensor.numel() ++ for item in passthrough: ++ name = self._validate_feature_name(item.get("name")) ++ names.append(name) ++ dtype = _STR_DTYPE.get(str(item.get("dtype", "int64"))) ++ if dtype is None: ++ raise ValueError( ++ f"spec-capture passthrough {name!r} has unsupported dtype " ++ f"{item.get('dtype')!r}" ++ ) ++ shape = [int(d) for d in item.get("shape") or []] ++ if not shape or len(shape) > 4 or any(d <= 0 for d in shape): ++ raise ValueError( ++ f"spec-capture passthrough {name!r} has invalid shape {shape}" ++ ) ++ numel = math.prod(shape) ++ if numel != len(item.get("data") or []): ++ raise ValueError( ++ f"spec-capture passthrough {name!r} shape requires {numel} " ++ f"values, got {len(item.get('data') or [])}" ++ ) ++ total_bytes += numel * torch.empty((), dtype=dtype).element_size() ++ if len(names) != len(set(names)): ++ raise ValueError("spec-capture feature names must be unique") ++ if total_bytes > self.max_sample_bytes: ++ raise ValueError( ++ f"spec-capture sample is {total_bytes} bytes, above the " ++ f"{self.max_sample_bytes}-byte limit" ++ ) ++ return store_id, sample_id, gen, replace, features, passthrough, total_bytes + + # -- the one entry point -------------------------------------------------- + def put_sample( +@@ -162,65 +515,98 @@ class SpecCaptureSink: + stored with a leading batch dim of 1. On any failure the keys already + written are best-effort removed (no partial sample is consumable). + """ +- store_id = str(spec["store_id"]) +- sample_id = str(spec["sample_id"]) +- gen = int(spec.get("gen", 1)) +- replace = bool(spec.get("replace", False)) +- features: Dict[str, str] = dict(spec.get("features") or {}) +- ++ ( ++ store_id, ++ sample_id, ++ gen, ++ replace, ++ features, ++ passthrough, ++ total_bytes, ++ ) = self._validate_request(spec, aux, last_hidden) ++ ++ artifact_names = list(features.values()) + [ ++ str(item["name"]) for item in passthrough ++ ] ++ planned_keys = [ ++ self._tkey(store_id, sample_id, gen, name) ++ for name in artifact_names ++ ] + written: List[str] = [] + result_feats: Dict[str, Dict[str, Any]] = {} + + def _write(name: str, t: torch.Tensor) -> None: + key = self._tkey(store_id, sample_id, gen, name) +- self._put_tensor(key, t, replace=replace) ++ # Include the attempted key in rollback: a negative put status may ++ # still follow a partial remote write. + written.append(key) ++ self._put_tensor(key, t, replace=replace) + result_feats[name] = { + "shape": list(t.shape), + "dtype": _DTYPE_STR.get(t.dtype, str(t.dtype).replace("torch.", "")), + } + +- try: +- aux_name = features.get(_ARTIFACT_AUX) +- if aux_name is not None: +- if aux is None: +- raise RuntimeError( ++ sample_lock = self._sample_locks[ ++ hash((store_id, sample_id)) % len(self._sample_locks) ++ ] ++ with sample_lock: ++ self._reject_stale_generation(sample_id, gen) ++ self._record_lifecycle( ++ sample_id, gen, artifact_names, total_bytes, "planned" ++ ) ++ cached = self._prepare_write(sample_id, gen, planned_keys) ++ if cached is not None: ++ self._record_lifecycle( ++ sample_id, gen, artifact_names, total_bytes, "resident" ++ ) ++ return cached ++ try: ++ aux_name = features.get(_ARTIFACT_AUX) ++ if aux_name is not None: ++ if aux is None: ++ raise RuntimeError( + "spec_capture requested 'aux' but no aux hidden states were " + "captured — launch the server with --enable-spec-capture " + "(and optionally --spec-capture-aux-layer-ids)" + ) +- _write(aux_name, aux.unsqueeze(0)) +- lh_name = features.get(_ARTIFACT_LAST_HIDDEN) +- if lh_name is not None: +- if last_hidden is None: +- raise RuntimeError( ++ _write(aux_name, aux.unsqueeze(0)) ++ lh_name = features.get(_ARTIFACT_LAST_HIDDEN) ++ if lh_name is not None: ++ if last_hidden is None: ++ raise RuntimeError( + "spec_capture requested 'last_hidden' but the logits " + "processor did not return it (is aux capture enabled?)" + ) +- _write(lh_name, last_hidden.unsqueeze(0)) +- for item in spec.get("passthrough") or []: +- dtype = _STR_DTYPE.get(str(item.get("dtype", "int64"))) +- if dtype is None: +- raise RuntimeError( ++ _write(lh_name, last_hidden.unsqueeze(0)) ++ for item in passthrough: ++ dtype = _STR_DTYPE.get(str(item.get("dtype", "int64"))) ++ if dtype is None: ++ raise RuntimeError( + f"spec_capture passthrough {item.get('name')!r}: " + f"unsupported dtype {item.get('dtype')!r}" + ) +- t = torch.tensor(item["data"], dtype=dtype).reshape( +- [int(d) for d in item["shape"]] +- ) +- _write(str(item["name"]), t) +- except Exception: +- for key in written: +- self._remove_quiet(key) +- raise +- +- return { +- "sample_id": sample_id, +- "store_id": store_id, +- "gen": gen, +- "aux_layer_ids": self.aux_layer_ids, +- "features": result_feats, +- } ++ t = torch.tensor(item["data"], dtype=dtype).reshape( ++ [int(d) for d in item["shape"]] ++ ) ++ _write(str(item["name"]), t) ++ except Exception: ++ for key in written: ++ self._remove_exact(key, force=True) ++ self._mark_lifecycle_cleaned(sample_id, gen) ++ raise ++ ++ result = { ++ "sample_id": sample_id, ++ "store_id": store_id, ++ "gen": gen, ++ "aux_layer_ids": self.aux_layer_ids, ++ "features": result_feats, ++ } ++ self._commit_write(sample_id, gen, result) ++ self._record_lifecycle( ++ sample_id, gen, artifact_names, total_bytes, "resident" ++ ) ++ return result + + + _SINK: Optional[SpecCaptureSink] = None +@@ -234,7 +620,31 @@ def maybe_init_sink(server_args) -> None: + """ + global _SINK + if getattr(server_args, "enable_spec_capture", False) and _SINK is None: ++ store_id = getattr(server_args, "spec_capture_store_id", None) ++ if not store_id: ++ raise ValueError( ++ "--enable-spec-capture requires --spec-capture-store-id" ++ ) ++ auth_token = os.environ.get("SGLANG_SPEC_CAPTURE_TOKEN") ++ if not auth_token: ++ raise ValueError( ++ "--enable-spec-capture requires SGLANG_SPEC_CAPTURE_TOKEN" ++ ) ++ lifecycle_db_path = getattr(server_args, "spec_capture_lifecycle_db", None) ++ if not lifecycle_db_path: ++ raise ValueError( ++ "--enable-spec-capture requires --spec-capture-lifecycle-db" ++ ) + _SINK = SpecCaptureSink( ++ store_id=store_id, ++ auth_token=auth_token, ++ max_sample_bytes=getattr( ++ server_args, "spec_capture_max_sample_bytes", 1 << 30 ++ ), ++ inventory_db_path=getattr( ++ server_args, "spec_capture_inventory_db", None ++ ), ++ lifecycle_db_path=lifecycle_db_path, + aux_layer_ids=getattr(server_args, "spec_capture_aux_layer_ids", None) + ) + diff --git a/patches/sglang/v0.5.14/spec-capture.patch b/patches/sglang/v0.5.14/spec-capture.patch index 15bcfa45d..176802cfe 100644 --- a/patches/sglang/v0.5.14/spec-capture.patch +++ b/patches/sglang/v0.5.14/spec-capture.patch @@ -1,5 +1,5 @@ diff --git a/python/sglang/srt/layers/logits_processor.py b/python/sglang/srt/layers/logits_processor.py -index a99d25267..f5d0b35e7 100644 +index a99d252..f5d0b35 100644 --- a/python/sglang/srt/layers/logits_processor.py +++ b/python/sglang/srt/layers/logits_processor.py @@ -94,6 +94,9 @@ class LogitsProcessorOutput: @@ -9,7 +9,7 @@ index a99d25267..f5d0b35e7 100644 + # Spec-training capture: under FULL+aux, `hidden_states` is the aux + # concatenation, so the post-norm last hidden is exposed here separately. + last_hidden_states: Optional[torch.Tensor] = None - + ## Part 2: This part will be assigned in python/sglang/srt/layers/sampler.py::Sampler # he log probs of output tokens, if SGLANG_RETURN_ORIGINAL_LOGPROB = True, will get the log probs before applying temperature. If False, will get the log probs before applying temperature. @@ -361,6 +364,16 @@ class LogitsProcessor(nn.Module): @@ -27,7 +27,7 @@ index a99d25267..f5d0b35e7 100644 + else None + ) del hidden_states - + if not logits_metadata.extend_return_logprob: @@ -374,6 +387,7 @@ class LogitsProcessor(nn.Module): return LogitsProcessorOutput( @@ -38,7 +38,7 @@ index a99d25267..f5d0b35e7 100644 # workaround since ForwardBatch is local to forward_batch_generation(). # They should be moved to GenerationBatchResult to keep this class clean. diff --git a/python/sglang/srt/managers/detokenizer_manager.py b/python/sglang/srt/managers/detokenizer_manager.py -index b05334dea..2853c8231 100644 +index b05334d..2853c82 100644 --- a/python/sglang/srt/managers/detokenizer_manager.py +++ b/python/sglang/srt/managers/detokenizer_manager.py @@ -441,6 +441,7 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin): @@ -50,13 +50,13 @@ index b05334dea..2853c8231 100644 indexer_topk=indexer_topk, customized_info=recv_obj.customized_info, diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py -index 951f35495..2359c31b7 100644 +index 951f354..2359c31 100644 --- a/python/sglang/srt/managers/io_struct.py +++ b/python/sglang/srt/managers/io_struct.py @@ -284,6 +284,10 @@ class GenerateReqInput(BaseReq): # Batch-level: List[List[int]] (one per request). After __getitem__: List[int]. multi_item_delimiter_indices: Optional[Union[List[List[int]], List[int]]] = None - + + # Spec-training capture sink instructions (see spec_capture_sink.py). + # Batch-level: List[Optional[dict]]; per-request after __getitem__. + spec_capture: Optional[Union[List[Optional[Dict]], Dict]] = None @@ -79,17 +79,17 @@ index 951f35495..2359c31b7 100644 @@ -842,6 +851,9 @@ class TokenizedGenerateReqInput(BaseReq): # Pre-computed delimiter indices for multi-item scoring multi_item_delimiter_indices: Optional[List[int]] = None - + + # Spec-training capture sink instructions (see GenerateReqInput.spec_capture) + spec_capture: Optional[Dict] = None + # For observability time_stats: Optional[Union[APIServerReqTimeStats, DPControllerReqTimeStats]] = None - + @@ -1179,6 +1191,9 @@ class BatchTokenIDOutput(BaseBatchReq, SpeculativeDecodingMetricsMixin): # The trainer step id. Used to know which step's weights are used for sampling. token_steps: List[List[int]] = None - + + # Spec-training capture: one result dict per request (see spec_capture_sink). + spec_capture: Optional[List[Any]] = None + @@ -99,7 +99,7 @@ index 951f35495..2359c31b7 100644 @@ -1247,6 +1262,9 @@ class BatchStrOutput(BaseBatchReq, SpeculativeDecodingMetricsMixin): # The trainer step id. Used to know which step's weights are used for sampling. token_steps: List[List[int]] = None - + + # Spec-training capture: one result dict per request (see spec_capture_sink). + spec_capture: Optional[List[Any]] = None + @@ -107,7 +107,7 @@ index 951f35495..2359c31b7 100644 customized_info: Optional[Dict[str, List[Any]]] = None # Detailed breakdown of cached tokens by source (device/host/storage) diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py -index f1dc81d17..a93a466b7 100755 +index f1dc81d..e4e84cd 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -708,6 +708,7 @@ class Req(ReqDllmMixin): @@ -122,25 +122,36 @@ index f1dc81d17..a93a466b7 100755 } self.sampling_params = sampling_params self.custom_logit_processor = custom_logit_processor -- self.return_hidden_states = return_hidden_states -+ # Spec-training capture: piggyback the return_hidden_states path (fires -+ # CaptureHiddenMode.FULL); the sink consumes the slices, not the response. ++ # Keep the client response flag separate from internal spec capture. ++ # ScheduleBatch enables CaptureHiddenMode.FULL for either use case. + self.spec_capture = spec_capture -+ self.return_hidden_states = return_hidden_states or spec_capture is not None + self.return_hidden_states = return_hidden_states + self.spec_capture_aux: List[torch.Tensor] = [] + self.spec_capture_last_hidden: List[torch.Tensor] = [] + self.spec_capture_result = None # per-request sink result -> output field - + # extra key for classifying the request (e.g. cache_salt) if lora_id is not None: +@@ -1865,7 +1872,10 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): + has_grammar=any(req.grammar for req in reqs), + device=req_to_token_pool.device, + spec_algorithm=spec_algorithm, +- return_hidden_states=any(req.return_hidden_states for req in reqs), ++ return_hidden_states=any( ++ req.return_hidden_states or req.spec_capture is not None ++ for req in reqs ++ ), + is_prefill_only=all(req.is_prefill_only for req in reqs), + chunked_req=chunked_req, + chunked_req_next_prompt_token=_compute_chunked_req_next_prompt_token( diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py -index abba37441..3018b8247 100644 +index abba374..701ee17 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -576,6 +576,25 @@ class Scheduler( - + self.init_batch_result_processor() - + + if server_args.enable_spec_capture: + # Capture needs single-pass prefill; chunking would drop all but the + # final chunk's hidden rows. @@ -161,23 +172,23 @@ index abba37441..3018b8247 100644 + spec_capture_sink.maybe_init_sink(server_args) + self.is_initializing = False - + def init_zbal_on_npu(self): -@@ -2054,6 +2067,7 @@ class Scheduler( +@@ -2054,6 +2073,7 @@ class Scheduler( dllm_config=self.dllm_config, time_stats=recv_req.time_stats, multi_item_delimiter_indices=recv_req.multi_item_delimiter_indices, + spec_capture=recv_req.spec_capture, ) req.tokenizer = self.tokenizer - + diff --git a/python/sglang/srt/managers/scheduler_components/batch_result_processor.py b/python/sglang/srt/managers/scheduler_components/batch_result_processor.py -index a9d5f0c28..001eff804 100644 +index a9d5f0c..001eff8 100644 --- a/python/sglang/srt/managers/scheduler_components/batch_result_processor.py +++ b/python/sglang/srt/managers/scheduler_components/batch_result_processor.py @@ -257,6 +257,18 @@ class SchedulerBatchResultProcessor: ) - + if ( + req.spec_capture is not None + and logits_output.hidden_states is not None @@ -194,10 +205,10 @@ index a9d5f0c28..001eff804 100644 req.return_hidden_states and logits_output.hidden_states is not None ): -@@ -462,6 +474,62 @@ class SchedulerBatchResultProcessor: +@@ -462,6 +474,66 @@ class SchedulerBatchResultProcessor: f"Placeholder zeros would be appended to output_ids." ) - + + def _append_spec_capture_states( + self, + *, @@ -208,15 +219,19 @@ index a9d5f0c28..001eff804 100644 + """Accumulate captured rows as CPU tensors for the Mooncake sink. + + Same offset arithmetic as ``_append_prefill_hidden_states`` but keeps -+ tensor slices (aux concat in ``hidden_states``, post-norm last in -+ ``last_hidden_states``) rather than the JSON-able response payload. ++ only the artifacts requested by this capture as CPU tensor slices. + """ + start = hidden_state_offset + end = start + len(req.origin_input_ids) -+ req.spec_capture_aux.append( -+ logits_output.hidden_states[start:end].cpu().clone() -+ ) -+ if logits_output.last_hidden_states is not None: ++ requested_artifacts = req.spec_capture.get("features") or {} ++ if "aux" in requested_artifacts: ++ req.spec_capture_aux.append( ++ logits_output.hidden_states[start:end].cpu().clone() ++ ) ++ if ( ++ "last_hidden" in requested_artifacts ++ and logits_output.last_hidden_states is not None ++ ): + req.spec_capture_last_hidden.append( + logits_output.last_hidden_states[start:end].cpu().clone() + ) @@ -258,7 +273,7 @@ index a9d5f0c28..001eff804 100644 self, *, diff --git a/python/sglang/srt/managers/scheduler_components/output_streamer.py b/python/sglang/srt/managers/scheduler_components/output_streamer.py -index f95c59f6f..9e64326a1 100644 +index f95c59f..9e64326 100644 --- a/python/sglang/srt/managers/scheduler_components/output_streamer.py +++ b/python/sglang/srt/managers/scheduler_components/output_streamer.py @@ -273,6 +273,7 @@ class _GenerationStreamAccumulator: @@ -287,7 +302,7 @@ index f95c59f6f..9e64326a1 100644 indexer_topk=self.indexer_topk, customized_info=self.customized_info, diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py -index bf932611a..f0e821ec4 100644 +index bf93261..f0e821e 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -1172,6 +1172,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): @@ -310,13 +325,13 @@ index bf932611a..f0e821ec4 100644 val = recv_obj.routed_experts[i] if val is not None: diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py -index 1cff5c983..a5935fa3c 100644 +index 1cff5c9..a5935fa 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -471,6 +471,19 @@ class ModelRunner(ModelRunnerKVCacheMixin): # if there is no aux layer, set to None self.eagle_aux_hidden_state_layer_ids = None - + + if server_args.enable_spec_capture and not self.is_draft_worker: + # Aux capture without a draft worker, routed to the strategy's own + # capture method (they wire different submodules — e.g. VL models @@ -332,12 +347,28 @@ index 1cff5c983..a5935fa3c 100644 + if self.spec_algorithm.is_dflash() and not self.is_draft_worker: from sglang.srt.speculative.dflash_utils import parse_dflash_draft_config - + +diff --git a/python/sglang/srt/models/qwen2.py b/python/sglang/srt/models/qwen2.py +--- a/python/sglang/srt/models/qwen2.py ++++ b/python/sglang/srt/models/qwen2.py +@@ -386,6 +386,12 @@ class Qwen2Model(nn.Module): + } + ) + else: ++ # Capture after the final transformer layer but before final norm. ++ # The ordinary pre-layer hook cannot observe end_layer itself. ++ if self.end_layer in self.layers_to_capture: ++ aux_hidden_states.append( ++ hidden_states + residual if residual is not None else hidden_states ++ ) + if hidden_states.shape[0] != 0: + if residual is None: + hidden_states = self.norm(hidden_states) diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py -index c7162c16d..d515b3c69 100644 +index c7162c1..d7c98f1 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py -@@ -2127,6 +2127,25 @@ class ServerArgs: +@@ -2127,6 +2127,44 @@ class ServerArgs: bool, "Enable returning hidden states with responses.", ] = False @@ -360,15 +391,34 @@ index c7162c16d..d515b3c69 100644 + "match the draft strategy being trained; they wire capture onto " + "different submodules (VL targets only populate the dflash path).", + ] = "eagle3" ++ spec_capture_store_id: A[ ++ Optional[str], ++ "Server-owned Mooncake namespace for spec capture. Required with " ++ "--enable-spec-capture; requests cannot override it.", ++ ] = None ++ spec_capture_max_sample_bytes: A[ ++ int, ++ "Maximum total bytes a single spec-capture request may hard-pin.", ++ ] = 1 << 30 ++ spec_capture_inventory_db: A[ ++ Optional[str], ++ "Server-side durable capture transaction inventory. Required with " ++ "--enable-spec-capture so response-loss retries can reclaim old keys.", ++ ] = None ++ spec_capture_lifecycle_db: A[ ++ Optional[str], ++ "Shared SpecForge owner lifecycle DB. Every planned Mooncake key is " ++ "recorded here before its first hard-pinned write.", ++ ] = None enable_return_routed_experts: A[ bool, "Enable returning routed experts of each layer with responses.", diff --git a/python/sglang/srt/spec_capture_sink.py b/python/sglang/srt/spec_capture_sink.py new file mode 100644 -index 000000000..d317b2801 +index 0000000..a633c56 --- /dev/null +++ b/python/sglang/srt/spec_capture_sink.py -@@ -0,0 +1,243 @@ +@@ -0,0 +1,653 @@ +# Copyright 2024 SGLang Team +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. @@ -377,9 +427,9 @@ index 000000000..d317b2801 +# http://www.apache.org/licenses/LICENSE-2.0 +"""Server-side spec-training capture sink (SpecForge DataFlow transport). + -+Under ``--enable-spec-capture``, a request's ``spec_capture`` dict tells this -+sink to write the prefill's captured tensors straight into a Mooncake store -+(one hard-pinned object per tensor at ``{store_id}/{sample_id}/g{gen}/{name}``, ++Under ``--enable-spec-capture``, an authenticated request's ``spec_capture`` ++dict tells this sink to write captured tensors into the server-owned Mooncake ++namespace (one object per tensor at ``{store_id}/{sample_id}/g{gen}/{name}``, +raw bytes — shape/dtype travel on the returned spec). Feature tensors never +touch the response path; ``meta_info["spec_capture"]`` returns only keys + +shapes/dtypes. Strategy naming is the client's (the ``features`` mapping); the @@ -388,7 +438,8 @@ index 000000000..d317b2801 + +Request schema:: + -+ {"store_id", "sample_id", "gen", "replace", # key namespace / retry policy ++ {"auth_token", "store_id", "sample_id", "gen", "replace", ++ # capability / namespace / retry + "features": {"aux": , "last_hidden": }, # artifact -> feature + "passthrough": [{"name", "data", "shape", "dtype"}]} # client tensors verbatim + @@ -401,9 +452,15 @@ index 000000000..d317b2801 + +from __future__ import annotations + ++import hmac ++import json +import logging ++import math +import os ++import re ++import sqlite3 +import threading ++import time +from typing import Any, Dict, List, Optional + +import torch @@ -428,18 +485,94 @@ index 000000000..d317b2801 + +_ARTIFACT_AUX = "aux" +_ARTIFACT_LAST_HIDDEN = "last_hidden" ++_ALLOWED_FEATURE_NAMES = { ++ "attention_mask", ++ "hidden_state", ++ "hidden_states", ++ "input_ids", ++ "loss_mask", ++ "target", ++} ++_SAFE_SAMPLE_ID = re.compile(r"^[A-Za-z0-9_.:-]{1,256}$") ++_MAX_PASSTHROUGH_FEATURES = 8 + + +class SpecCaptureSink: + """Writes captured per-request tensors into Mooncake in SpecForge layout.""" + -+ def __init__(self, aux_layer_ids: Optional[List[int]] = None) -> None: ++ def __init__( ++ self, ++ *, ++ store_id: str, ++ auth_token: str, ++ max_sample_bytes: int, ++ inventory_db_path: str, ++ lifecycle_db_path: str, ++ aux_layer_ids: Optional[List[int]] = None, ++ ) -> None: ++ if not store_id or "/" in store_id or store_id in {".", ".."}: ++ raise ValueError("spec-capture store id must be one safe path segment") ++ if not auth_token: ++ raise ValueError("SGLANG_SPEC_CAPTURE_TOKEN must be non-empty") ++ if max_sample_bytes <= 0: ++ raise ValueError("spec-capture max sample bytes must be positive") ++ if not inventory_db_path: ++ raise ValueError("spec-capture inventory DB path must be non-empty") ++ if not lifecycle_db_path: ++ raise ValueError("spec-capture lifecycle DB path must be non-empty") ++ self.store_id = store_id ++ self.auth_token = auth_token ++ self.max_sample_bytes = int(max_sample_bytes) + self.aux_layer_ids = list(aux_layer_ids) if aux_layer_ids else None ++ self.inventory_db_path = os.path.abspath(inventory_db_path) ++ os.makedirs(os.path.dirname(self.inventory_db_path), exist_ok=True) ++ self._inventory = sqlite3.connect( ++ self.inventory_db_path, check_same_thread=False, timeout=30.0 ++ ) ++ self._inventory.execute("PRAGMA journal_mode=WAL") ++ self._inventory.execute("PRAGMA synchronous=FULL") ++ self._inventory.execute("PRAGMA busy_timeout=30000") ++ self._inventory.execute( ++ "CREATE TABLE IF NOT EXISTS captures (" ++ "sample_id TEXT PRIMARY KEY, generation INTEGER NOT NULL, " ++ "keys_json TEXT NOT NULL, result_json TEXT, state TEXT NOT NULL, " ++ "prior_generation INTEGER, prior_keys_json TEXT)" ++ ) ++ capture_columns = { ++ row[1] for row in self._inventory.execute("PRAGMA table_info(captures)") ++ } ++ if "prior_generation" not in capture_columns: ++ self._inventory.execute( ++ "ALTER TABLE captures ADD COLUMN prior_generation INTEGER" ++ ) ++ if "prior_keys_json" not in capture_columns: ++ self._inventory.execute( ++ "ALTER TABLE captures ADD COLUMN prior_keys_json TEXT" ++ ) ++ self._inventory.commit() ++ self.lifecycle_db_path = os.path.abspath(lifecycle_db_path) ++ os.makedirs(os.path.dirname(self.lifecycle_db_path), exist_ok=True) ++ self._lifecycle = sqlite3.connect( ++ self.lifecycle_db_path, check_same_thread=False, timeout=30.0 ++ ) ++ self._lifecycle.execute("PRAGMA journal_mode=WAL") ++ self._lifecycle.execute("PRAGMA synchronous=FULL") ++ self._lifecycle.execute("PRAGMA busy_timeout=30000") ++ self._lifecycle.execute( ++ "CREATE TABLE IF NOT EXISTS mooncake_objects (" ++ "store_id TEXT NOT NULL, sample_id TEXT NOT NULL, " ++ "generation INTEGER NOT NULL, feature_names_json TEXT NOT NULL, " ++ "estimated_bytes INTEGER NOT NULL, state TEXT NOT NULL, " ++ "reason TEXT, updated_at REAL NOT NULL, " ++ "PRIMARY KEY (store_id, sample_id, generation))" ++ ) ++ self._lifecycle.commit() + self._store = None + self._put_config = None + self._lock = threading.Lock() + # Retried HTTP requests reuse deterministic keys. Striped locks keep -+ # replacement atomic per key without retaining one lock per sample. ++ # replacement atomic without retaining one lock per sample or key. ++ self._sample_locks = [threading.RLock() for _ in range(256)] + self._write_locks = [threading.Lock() for _ in range(256)] + + # -- connection --------------------------------------------------------- @@ -498,7 +631,7 @@ index 000000000..d317b2801 + if replace: + # Do not probe with is_exist(): Mooncake existence checks can + # acquire a read lease that prevents the following removal. -+ self._remove_quiet(key) ++ self._remove_exact(key, force=True) + try: + store.register_buffer(t.data_ptr(), nbytes) + except Exception: @@ -513,11 +646,281 @@ index 000000000..d317b2801 + if rc is not None and int(rc) < 0: + raise RuntimeError(f"spec-capture put_from failed (status {rc}) for {key}") + -+ def _remove_quiet(self, key: str) -> None: -+ try: -+ self._connect().remove(key) -+ except Exception: -+ pass ++ def _remove_exact(self, key: str, *, force: bool) -> None: ++ store = self._connect() ++ rc = store.remove(key, force) ++ if rc is not None and int(rc) not in (0, -704): ++ raise RuntimeError( ++ f"spec-capture remove failed (status {rc}) for {key}" ++ ) ++ ++ def _inventory_row(self, sample_id: str): ++ return self._inventory.execute( ++ "SELECT generation, keys_json, result_json, state, " ++ "prior_generation, prior_keys_json FROM captures " ++ "WHERE sample_id=?", ++ (sample_id,), ++ ).fetchone() ++ ++ def _record_lifecycle( ++ self, ++ sample_id: str, ++ gen: int, ++ feature_names: List[str], ++ estimated_bytes: int, ++ state: str, ++ ) -> None: ++ names_json = json.dumps(sorted(feature_names), separators=(",", ":")) ++ row = self._lifecycle.execute( ++ "SELECT feature_names_json, estimated_bytes, state FROM " ++ "mooncake_objects WHERE store_id=? AND sample_id=? AND generation=?", ++ (self.store_id, sample_id, gen), ++ ).fetchone() ++ if row is not None: ++ if row[0] != names_json or int(row[1]) != int(estimated_bytes): ++ raise RuntimeError( ++ f"spec-capture lifecycle identity changed for {sample_id} " ++ f"generation {gen}" ++ ) ++ if row[2] in {"tombstoned", "cleaned"}: ++ raise RuntimeError( ++ f"spec-capture lifecycle refuses {sample_id} generation {gen} " ++ f"from state {row[2]!r}" ++ ) ++ if state == "planned" or row[2] == "resident": ++ return ++ self._lifecycle.execute( ++ "UPDATE mooncake_objects SET state='resident', updated_at=? WHERE " ++ "store_id=? AND sample_id=? AND generation=? AND state='planned'", ++ (time.time(), self.store_id, sample_id, gen), ++ ) ++ else: ++ self._lifecycle.execute( ++ "INSERT INTO mooncake_objects " ++ "(store_id, sample_id, generation, feature_names_json, " ++ "estimated_bytes, state, reason, updated_at) " ++ "VALUES (?, ?, ?, ?, ?, ?, NULL, ?)", ++ ( ++ self.store_id, ++ sample_id, ++ gen, ++ names_json, ++ int(estimated_bytes), ++ state, ++ time.time(), ++ ), ++ ) ++ self._lifecycle.commit() ++ ++ def _mark_lifecycle_cleaned(self, sample_id: str, gen: int) -> None: ++ self._lifecycle.execute( ++ "UPDATE mooncake_objects SET state='cleaned', updated_at=? WHERE " ++ "store_id=? AND sample_id=? AND generation=?", ++ (time.time(), self.store_id, sample_id, gen), ++ ) ++ self._lifecycle.commit() ++ ++ def _reject_stale_generation(self, sample_id: str, gen: int) -> None: ++ row = self._inventory_row(sample_id) ++ if row is not None and int(row[0]) > gen: ++ raise ValueError( ++ f"stale spec-capture generation {gen}; latest is {row[0]}" ++ ) ++ ++ def _finish_replacement( ++ self, ++ sample_id: str, ++ gen: int, ++ prior_gen: int, ++ prior_keys: List[str], ++ ) -> None: ++ for key in prior_keys: ++ self._remove_exact(key, force=False) ++ self._mark_lifecycle_cleaned(sample_id, prior_gen) ++ cursor = self._inventory.execute( ++ "UPDATE captures SET state='writing', prior_generation=NULL, " ++ "prior_keys_json=NULL WHERE sample_id=? AND generation=? " ++ "AND state='replacing'", ++ (sample_id, gen), ++ ) ++ if cursor.rowcount != 1: ++ raise RuntimeError( ++ f"spec-capture replacement journal changed for {sample_id} " ++ f"generation {gen}" ++ ) ++ self._inventory.commit() ++ ++ def _prepare_write(self, sample_id: str, gen: int, keys: List[str]): ++ row = self._inventory_row(sample_id) ++ if row is None: ++ self._inventory.execute( ++ "INSERT INTO captures " ++ "(sample_id, generation, keys_json, result_json, state, " ++ "prior_generation, prior_keys_json) " ++ "VALUES (?, ?, ?, NULL, 'writing', NULL, NULL)", ++ (sample_id, gen, json.dumps(keys, separators=(",", ":"))), ++ ) ++ self._inventory.commit() ++ return None ++ ++ current_gen, keys_json, result_json, state, prior_gen, prior_keys_json = row ++ current_gen = int(current_gen) ++ if current_gen > gen: ++ raise ValueError( ++ f"stale spec-capture generation {gen}; latest is {current_gen}" ++ ) ++ if current_gen == gen and state == "committed": ++ return json.loads(result_json) ++ if state == "replacing": ++ if prior_gen is None or prior_keys_json is None: ++ raise RuntimeError( ++ f"incomplete replacement journal for {sample_id} generation " ++ f"{current_gen}" ++ ) ++ self._finish_replacement( ++ sample_id, ++ current_gen, ++ int(prior_gen), ++ list(json.loads(prior_keys_json)), ++ ) ++ row = self._inventory_row(sample_id) ++ current_gen, keys_json, result_json, state, _, _ = row ++ current_gen = int(current_gen) ++ if current_gen == gen: ++ if state != "writing": ++ raise RuntimeError( ++ f"invalid spec-capture inventory state {state!r} for " ++ f"{sample_id}" ++ ) ++ for key in json.loads(keys_json): ++ self._remove_exact(key, force=True) ++ return None ++ if state not in {"writing", "committed"}: ++ raise RuntimeError( ++ f"invalid spec-capture inventory state {state!r} for {sample_id}" ++ ) ++ ++ prior_keys = list(json.loads(keys_json)) ++ self._inventory.execute( ++ "INSERT OR REPLACE INTO captures " ++ "(sample_id, generation, keys_json, result_json, state, " ++ "prior_generation, prior_keys_json) " ++ "VALUES (?, ?, ?, NULL, 'replacing', ?, ?)", ++ ( ++ sample_id, ++ gen, ++ json.dumps(keys, separators=(",", ":")), ++ current_gen, ++ json.dumps(prior_keys, separators=(",", ":")), ++ ), ++ ) ++ self._inventory.commit() ++ self._finish_replacement(sample_id, gen, current_gen, prior_keys) ++ return None ++ ++ def _commit_write(self, sample_id: str, gen: int, result: Dict[str, Any]): ++ cursor = self._inventory.execute( ++ "UPDATE captures SET result_json=?, state='committed' " ++ "WHERE sample_id=? AND generation=? AND state='writing'", ++ (json.dumps(result, separators=(",", ":")), sample_id, gen), ++ ) ++ if cursor.rowcount != 1: ++ raise RuntimeError( ++ f"spec-capture write transaction changed for {sample_id} " ++ f"generation {gen}" ++ ) ++ self._inventory.commit() ++ ++ @staticmethod ++ def _validate_feature_name(name: Any) -> str: ++ name = str(name) ++ if name not in _ALLOWED_FEATURE_NAMES: ++ raise ValueError(f"spec-capture feature name {name!r} is not allowed") ++ return name ++ ++ def _validate_request( ++ self, ++ spec: Dict[str, Any], ++ aux: Optional[torch.Tensor], ++ last_hidden: Optional[torch.Tensor], ++ ) -> tuple[ ++ str, str, int, bool, Dict[str, str], List[Dict[str, Any]], int ++ ]: ++ supplied = str(spec.get("auth_token") or "") ++ if not hmac.compare_digest(supplied, self.auth_token): ++ raise PermissionError("invalid spec-capture capability") ++ store_id = str(spec.get("store_id") or "") ++ if store_id != self.store_id: ++ raise ValueError( ++ "spec-capture store_id must match the server-owned namespace" ++ ) ++ sample_id = str(spec.get("sample_id") or "") ++ if not _SAFE_SAMPLE_ID.fullmatch(sample_id) or not sample_id.startswith( ++ self.store_id + ":" ++ ): ++ raise ValueError( ++ "spec-capture sample_id must be safe and prefixed by the " ++ "server-owned store id" ++ ) ++ gen = int(spec.get("gen", 1)) ++ if not 1 <= gen <= 2**31 - 1: ++ raise ValueError("spec-capture generation is out of range") ++ replace = spec.get("replace", False) ++ if type(replace) is not bool: ++ raise ValueError("spec-capture replace must be a boolean") ++ ++ features = dict(spec.get("features") or {}) ++ unknown_artifacts = set(features) - {_ARTIFACT_AUX, _ARTIFACT_LAST_HIDDEN} ++ if unknown_artifacts: ++ raise ValueError( ++ f"unknown spec-capture artifacts: {sorted(unknown_artifacts)}" ++ ) ++ features = { ++ artifact: self._validate_feature_name(name) ++ for artifact, name in features.items() ++ } ++ passthrough = list(spec.get("passthrough") or []) ++ if len(passthrough) > _MAX_PASSTHROUGH_FEATURES: ++ raise ValueError("too many spec-capture passthrough features") ++ ++ names = list(features.values()) ++ total_bytes = 0 ++ selected_artifacts = ( ++ (_ARTIFACT_AUX, aux), ++ (_ARTIFACT_LAST_HIDDEN, last_hidden), ++ ) ++ for artifact, tensor in selected_artifacts: ++ if artifact in features and tensor is not None: ++ total_bytes += tensor.element_size() * tensor.numel() ++ for item in passthrough: ++ name = self._validate_feature_name(item.get("name")) ++ names.append(name) ++ dtype = _STR_DTYPE.get(str(item.get("dtype", "int64"))) ++ if dtype is None: ++ raise ValueError( ++ f"spec-capture passthrough {name!r} has unsupported dtype " ++ f"{item.get('dtype')!r}" ++ ) ++ shape = [int(d) for d in item.get("shape") or []] ++ if not shape or len(shape) > 4 or any(d <= 0 for d in shape): ++ raise ValueError( ++ f"spec-capture passthrough {name!r} has invalid shape {shape}" ++ ) ++ numel = math.prod(shape) ++ if numel != len(item.get("data") or []): ++ raise ValueError( ++ f"spec-capture passthrough {name!r} shape requires {numel} " ++ f"values, got {len(item.get('data') or [])}" ++ ) ++ total_bytes += numel * torch.empty((), dtype=dtype).element_size() ++ if len(names) != len(set(names)): ++ raise ValueError("spec-capture feature names must be unique") ++ if total_bytes > self.max_sample_bytes: ++ raise ValueError( ++ f"spec-capture sample is {total_bytes} bytes, above the " ++ f"{self.max_sample_bytes}-byte limit" ++ ) ++ return store_id, sample_id, gen, replace, features, passthrough, total_bytes + + # -- the one entry point -------------------------------------------------- + def put_sample( @@ -533,65 +936,98 @@ index 000000000..d317b2801 + stored with a leading batch dim of 1. On any failure the keys already + written are best-effort removed (no partial sample is consumable). + """ -+ store_id = str(spec["store_id"]) -+ sample_id = str(spec["sample_id"]) -+ gen = int(spec.get("gen", 1)) -+ replace = bool(spec.get("replace", False)) -+ features: Dict[str, str] = dict(spec.get("features") or {}) ++ ( ++ store_id, ++ sample_id, ++ gen, ++ replace, ++ features, ++ passthrough, ++ total_bytes, ++ ) = self._validate_request(spec, aux, last_hidden) + ++ artifact_names = list(features.values()) + [ ++ str(item["name"]) for item in passthrough ++ ] ++ planned_keys = [ ++ self._tkey(store_id, sample_id, gen, name) ++ for name in artifact_names ++ ] + written: List[str] = [] + result_feats: Dict[str, Dict[str, Any]] = {} + + def _write(name: str, t: torch.Tensor) -> None: + key = self._tkey(store_id, sample_id, gen, name) -+ self._put_tensor(key, t, replace=replace) ++ # Include the attempted key in rollback: a negative put status may ++ # still follow a partial remote write. + written.append(key) ++ self._put_tensor(key, t, replace=replace) + result_feats[name] = { + "shape": list(t.shape), + "dtype": _DTYPE_STR.get(t.dtype, str(t.dtype).replace("torch.", "")), + } + -+ try: -+ aux_name = features.get(_ARTIFACT_AUX) -+ if aux_name is not None: -+ if aux is None: -+ raise RuntimeError( ++ sample_lock = self._sample_locks[ ++ hash((store_id, sample_id)) % len(self._sample_locks) ++ ] ++ with sample_lock: ++ self._reject_stale_generation(sample_id, gen) ++ self._record_lifecycle( ++ sample_id, gen, artifact_names, total_bytes, "planned" ++ ) ++ cached = self._prepare_write(sample_id, gen, planned_keys) ++ if cached is not None: ++ self._record_lifecycle( ++ sample_id, gen, artifact_names, total_bytes, "resident" ++ ) ++ return cached ++ try: ++ aux_name = features.get(_ARTIFACT_AUX) ++ if aux_name is not None: ++ if aux is None: ++ raise RuntimeError( + "spec_capture requested 'aux' but no aux hidden states were " + "captured — launch the server with --enable-spec-capture " + "(and optionally --spec-capture-aux-layer-ids)" + ) -+ _write(aux_name, aux.unsqueeze(0)) -+ lh_name = features.get(_ARTIFACT_LAST_HIDDEN) -+ if lh_name is not None: -+ if last_hidden is None: -+ raise RuntimeError( ++ _write(aux_name, aux.unsqueeze(0)) ++ lh_name = features.get(_ARTIFACT_LAST_HIDDEN) ++ if lh_name is not None: ++ if last_hidden is None: ++ raise RuntimeError( + "spec_capture requested 'last_hidden' but the logits " + "processor did not return it (is aux capture enabled?)" + ) -+ _write(lh_name, last_hidden.unsqueeze(0)) -+ for item in spec.get("passthrough") or []: -+ dtype = _STR_DTYPE.get(str(item.get("dtype", "int64"))) -+ if dtype is None: -+ raise RuntimeError( ++ _write(lh_name, last_hidden.unsqueeze(0)) ++ for item in passthrough: ++ dtype = _STR_DTYPE.get(str(item.get("dtype", "int64"))) ++ if dtype is None: ++ raise RuntimeError( + f"spec_capture passthrough {item.get('name')!r}: " + f"unsupported dtype {item.get('dtype')!r}" + ) -+ t = torch.tensor(item["data"], dtype=dtype).reshape( -+ [int(d) for d in item["shape"]] -+ ) -+ _write(str(item["name"]), t) -+ except Exception: -+ for key in written: -+ self._remove_quiet(key) -+ raise -+ -+ return { -+ "sample_id": sample_id, -+ "store_id": store_id, -+ "gen": gen, -+ "aux_layer_ids": self.aux_layer_ids, -+ "features": result_feats, -+ } ++ t = torch.tensor(item["data"], dtype=dtype).reshape( ++ [int(d) for d in item["shape"]] ++ ) ++ _write(str(item["name"]), t) ++ except Exception: ++ for key in written: ++ self._remove_exact(key, force=True) ++ self._mark_lifecycle_cleaned(sample_id, gen) ++ raise ++ ++ result = { ++ "sample_id": sample_id, ++ "store_id": store_id, ++ "gen": gen, ++ "aux_layer_ids": self.aux_layer_ids, ++ "features": result_feats, ++ } ++ self._commit_write(sample_id, gen, result) ++ self._record_lifecycle( ++ sample_id, gen, artifact_names, total_bytes, "resident" ++ ) ++ return result + + +_SINK: Optional[SpecCaptureSink] = None @@ -605,7 +1041,31 @@ index 000000000..d317b2801 + """ + global _SINK + if getattr(server_args, "enable_spec_capture", False) and _SINK is None: ++ store_id = getattr(server_args, "spec_capture_store_id", None) ++ if not store_id: ++ raise ValueError( ++ "--enable-spec-capture requires --spec-capture-store-id" ++ ) ++ auth_token = os.environ.get("SGLANG_SPEC_CAPTURE_TOKEN") ++ if not auth_token: ++ raise ValueError( ++ "--enable-spec-capture requires SGLANG_SPEC_CAPTURE_TOKEN" ++ ) ++ lifecycle_db_path = getattr(server_args, "spec_capture_lifecycle_db", None) ++ if not lifecycle_db_path: ++ raise ValueError( ++ "--enable-spec-capture requires --spec-capture-lifecycle-db" ++ ) + _SINK = SpecCaptureSink( ++ store_id=store_id, ++ auth_token=auth_token, ++ max_sample_bytes=getattr( ++ server_args, "spec_capture_max_sample_bytes", 1 << 30 ++ ), ++ inventory_db_path=getattr( ++ server_args, "spec_capture_inventory_db", None ++ ), ++ lifecycle_db_path=lifecycle_db_path, + aux_layer_ids=getattr(server_args, "spec_capture_aux_layer_ids", None) + ) + diff --git a/scripts/apply_sglang_spec_capture_patch.sh b/scripts/apply_sglang_spec_capture_patch.sh index cb59afced..c0627ef45 100755 --- a/scripts/apply_sglang_spec_capture_patch.sh +++ b/scripts/apply_sglang_spec_capture_patch.sh @@ -1,65 +1,330 @@ #!/usr/bin/env bash -# Apply the spec-capture patch to the INSTALLED sglang (site-packages). -# -# The patch file is authored against the sglang source tree (git-style paths -# a/python/sglang/srt/...); an installed package drops the python/ prefix, so -# strip TWO components (a/ + python/) and apply from the site-packages parent. -# -# Idempotence is CONTENT-aware, not existence-based: the applied patch text is -# recorded next to the tree, so a revised patch reverses the recorded one and -# re-applies instead of silently keeping a stale version alive on cached -# venvs/runners. A tree patched before this record existed is adopted only -# when a reverse dry-run proves it matches the current patch byte-for-byte; -# anything else fails loudly rather than testing against unknown server code. -# -# Usage: scripts/apply_sglang_spec_capture_patch.sh [--reverse] +# Strictly apply, verify, or reverse the server-capture patch in one isolated +# sglang==0.5.14 installation. The Python executable is mandatory so an active +# shell environment can never redirect this operation to another site-packages. set -euo pipefail -HERE="$(cd "$(dirname "$0")/.." && pwd)" -PATCH="$HERE/patches/sglang/v0.5.14/spec-capture.patch" +readonly EXPECTED_PATCH_SHA256="a07c015a584bfff053d1358b1cc3c28ee79f1983ed8298953732c5ebe40647fb" +readonly EXPECTED_BASE_TO_CURRENT_SHA256="2c3a713514508b15f8a84a902c4af07e5a73894a6bb35a3b3f8f83c014150542" +readonly REPO_ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)" +readonly PATCH_PATH="$REPO_ROOT/patches/sglang/v0.5.14/spec-capture.patch" +readonly BASE_TO_CURRENT_PATH="$REPO_ROOT/patches/sglang/v0.5.14/spec-capture-base-to-current.patch" -SGL_PARENT="$(python -c 'import sglang, os; print(os.path.dirname(os.path.dirname(sglang.__file__)))')" -SGL_VERSION="$(python -c 'import sglang; print(sglang.__version__)')" -APPLIED_COPY="$SGL_PARENT/sglang/.spec_capture_patch.applied" -SINK="$SGL_PARENT/sglang/srt/spec_capture_sink.py" +usage() { + cat >&2 <<'EOF' +Usage: + scripts/apply_sglang_spec_capture_patch.sh --python /path/to/python [--apply] + scripts/apply_sglang_spec_capture_patch.sh --python /path/to/python --check + scripts/apply_sglang_spec_capture_patch.sh --python /path/to/python --reverse -if [[ "$SGL_VERSION" != 0.5.14* ]]; then - echo "WARNING: installed sglang is $SGL_VERSION; the patch targets v0.5.14" >&2 -fi +Modes are idempotent. --apply upgrades the published base capture patch when +needed. --check performs no writes and succeeds only when the +exact current v0.5.14 patch is applied and its imports and launch-server flags +work. +EOF +} + +die() { + echo "spec-capture patch error: $*" >&2 + exit 2 +} + +mode="apply" +python_executable="" +mode_seen=0 +while (($#)); do + case "$1" in + --python) + (($# >= 2)) || die "--python requires an executable" + python_executable="$2" + shift 2 + ;; + --apply|--check|--reverse) + ((mode_seen == 0)) || die "choose exactly one mode" + mode="${1#--}" + mode_seen=1 + shift + ;; + -h|--help) + usage + exit 0 + ;; + *) + usage + die "unknown argument $1" + ;; + esac +done -if [[ "${1:-}" == "--reverse" ]]; then - patch --reverse -p2 --batch -N -d "$SGL_PARENT" < "$PATCH" - rm -f "$APPLIED_COPY" - echo "spec-capture patch --reverse at $SGL_PARENT/sglang (sglang $SGL_VERSION)" - exit 0 +[[ -n "$python_executable" ]] || { + usage + die "--python is required" +} +if [[ "$python_executable" != /* ]]; then + python_executable="$(command -v -- "$python_executable" || true)" fi +[[ -n "$python_executable" && -x "$python_executable" ]] \ + || die "Python executable is unavailable" +[[ -f "$PATCH_PATH" && ! -L "$PATCH_PATH" ]] || die "patch file is unavailable" +[[ -f "$BASE_TO_CURRENT_PATH" && ! -L "$BASE_TO_CURRENT_PATH" ]] \ + || die "patch migration file is unavailable: $BASE_TO_CURRENT_PATH" -if [[ -f "$APPLIED_COPY" ]]; then - if cmp -s "$APPLIED_COPY" "$PATCH"; then - echo "spec-capture patch already applied at $SGL_PARENT/sglang" - exit 0 - fi - echo "spec-capture patch changed; reversing the recorded version first" - patch --reverse -p2 --batch -d "$SGL_PARENT" < "$APPLIED_COPY" -elif [[ -f "$SINK" ]]; then - # Patched before the applied-copy record existed. Adopt only a tree that - # provably matches the current patch; otherwise demand a clean reinstall. - # git apply verifies exact content on reverse; BSD patch dry-runs do not. - if command -v git > /dev/null; then - matches() { git -C "$SGL_PARENT" apply --reverse --check -p2 "$PATCH" 2> /dev/null; } +actual_patch_sha256="$(sha256sum "$PATCH_PATH" | awk '{print $1}')" +[[ "$actual_patch_sha256" == "$EXPECTED_PATCH_SHA256" ]] || die \ + "patch SHA256 mismatch: expected $EXPECTED_PATCH_SHA256, got $actual_patch_sha256" +actual_base_to_current_sha256="$(sha256sum "$BASE_TO_CURRENT_PATH" | awk '{print $1}')" +[[ "$actual_base_to_current_sha256" == "$EXPECTED_BASE_TO_CURRENT_SHA256" ]] || die \ + "base-to-current SHA256 mismatch: expected $EXPECTED_BASE_TO_CURRENT_SHA256, got $actual_base_to_current_sha256" + +mapfile -t package_info < <("$python_executable" - <<'PY' +import importlib.metadata +import importlib.util +import os +import pathlib +import sys +import sysconfig + +try: + version = importlib.metadata.version("sglang") +except importlib.metadata.PackageNotFoundError as exc: + raise SystemExit(f"sglang distribution is missing: {exc}") +spec = importlib.util.find_spec("sglang") +if spec is None or spec.origin is None: + raise SystemExit("sglang package is not importable") +root = pathlib.Path(spec.origin).resolve().parent +if root.name != "sglang" or not root.is_dir(): + raise SystemExit(f"unexpected sglang package root: {root}") +site_roots = { + pathlib.Path(value).resolve() + for key, value in sysconfig.get_paths().items() + if key in {"purelib", "platlib"} and value +} +if not any(root == site / "sglang" for site in site_roots): + raise SystemExit( + f"refusing to patch non-installed source tree {root}; " + f"expected one of {[str(site / 'sglang') for site in site_roots]}" + ) +print(version) +print(root) +print(root.parent) +print(pathlib.Path(sys.executable).resolve()) +PY +) || die "cannot inspect sglang through $python_executable" + +((${#package_info[@]} == 4)) || die "incomplete sglang package inspection" +sglang_version="${package_info[0]}" +sglang_root="${package_info[1]}" +sglang_parent="${package_info[2]}" +python_realpath="${package_info[3]}" +[[ "$sglang_version" == "0.5.14" ]] \ + || die "sglang must be exactly 0.5.14, got $sglang_version" + +mapfile -t patch_targets < <( + sed -n 's#^+++ b/python/##p' "$PATCH_PATH" | LC_ALL=C sort -u +) +((${#patch_targets[@]} == 12)) \ + || die "patch inventory must contain exactly 12 target files" +for target in "${patch_targets[@]}"; do + [[ "$target" == sglang/srt/* && "$target" != *".."* ]] \ + || die "unsafe patch target $target" +done +mapfile -t migration_targets < <( + sed -n 's#^+++ b/python/##p' "$BASE_TO_CURRENT_PATH" | LC_ALL=C sort -u +) +expected_migration_targets=( + "sglang/srt/managers/schedule_batch.py" + "sglang/srt/managers/scheduler_components/batch_result_processor.py" + "sglang/srt/models/qwen2.py" + "sglang/srt/server_args.py" + "sglang/srt/spec_capture_sink.py" +) +[[ "${migration_targets[*]}" == "${expected_migration_targets[*]}" ]] \ + || die "base-to-current migration has an unexpected target inventory" + +patch_probe() { + local direction="$1" + local output_file="$2" + local patch_path="${3:-$PATCH_PATH}" + local target_root="${4:-$sglang_parent}" + local -a args=(--dry-run --batch --silent -p2 -d "$target_root") + if [[ "$direction" == "reverse" ]]; then + args+=(--reverse) else - matches() { patch --reverse --dry-run -p2 --batch -d "$SGL_PARENT" < "$PATCH" > /dev/null; } + args+=(--forward) + fi + if ! patch "${args[@]}" < "$patch_path" >"$output_file" 2>&1; then + return 1 + fi + # GNU patch can return zero in --batch mode after ignoring every hunk in + # the wrong direction. Treat any skip/ignore diagnostic as a failed probe. + if grep -Eqi \ + '(^|[[:space:]])(ignoring|skipping)([[:space:]]|$)|FAILED|does not exist|malformed patch' \ + "$output_file"; then + return 1 fi - if matches; then - cp "$PATCH" "$APPLIED_COPY" - echo "spec-capture patch already applied at $SGL_PARENT/sglang (adopted)" - exit 0 +} + +tmpdir="$(mktemp -d "${TMPDIR:-/tmp}/spec-capture-patch.XXXXXX")" +trap 'rm -rf "$tmpdir"' EXIT + +apply_probe_migration() { + local patch_path="$1" + local probe_root="$2" + local output_file="$3" + if ! patch --batch --forward -p2 -d "$probe_root" \ + < "$patch_path" >"$output_file" 2>&1; then + return 1 fi - echo "ERROR: $SGL_PARENT/sglang carries an unknown spec-capture patch state" >&2 - echo "reinstall sglang (or clear the cached venv) and re-run this script" >&2 - exit 1 + ! grep -Eqi \ + '(^|[[:space:]])(ignoring|skipping)([[:space:]]|$)|FAILED|does not exist|malformed patch' \ + "$output_file" +} + +probe_base_upgrade() { + local probe_root="$tmpdir/published-base" + local target source destination + for target in "${patch_targets[@]}"; do + source="$sglang_parent/$target" + destination="$probe_root/$target" + [[ -f "$source" && ! -L "$source" ]] || return 1 + mkdir -p "$(dirname "$destination")" + cp -- "$source" "$destination" + done + apply_probe_migration \ + "$BASE_TO_CURRENT_PATH" "$probe_root" "$tmpdir/base-to-current.log" \ + || return 1 + patch_probe \ + reverse "$tmpdir/published-base-current.log" \ + "$PATCH_PATH" "$probe_root" +} + +forward_ok=0 +reverse_ok=0 +base_patch_applied=0 +patch_probe forward "$tmpdir/forward.log" && forward_ok=1 +patch_probe reverse "$tmpdir/reverse.log" && reverse_ok=1 +if ((forward_ok == 0 && reverse_ok == 0)); then + probe_base_upgrade && base_patch_applied=1 +fi + +if ((forward_ok == reverse_ok && base_patch_applied == 0)); then + echo "forward dry-run:" >&2 + sed -n '1,80p' "$tmpdir/forward.log" >&2 + echo "reverse dry-run:" >&2 + sed -n '1,80p' "$tmpdir/reverse.log" >&2 + die "target tree is neither a clean v0.5.14 tree nor the exact patched state" fi -patch -p2 --batch -N -d "$SGL_PARENT" < "$PATCH" -cp "$PATCH" "$APPLIED_COPY" -echo "spec-capture patch applied at $SGL_PARENT/sglang (sglang $SGL_VERSION)" +verify_patched_tree() { + SGLANG_CAPTURE_ROOT="$sglang_root" "$python_executable" - <<'PY' +import hashlib +import importlib +import os +from pathlib import Path + +root = Path(os.environ["SGLANG_CAPTURE_ROOT"]) +required = { + "srt/spec_capture_sink.py": ( + "class SpecCaptureSink", + "def get_sink", + "def maybe_init_sink", + ), + "srt/server_args.py": ( + "enable_spec_capture", + "spec_capture_aux_layer_ids", + "spec_capture_method", + ), + "srt/managers/scheduler.py": ("spec_capture_sink.maybe_init_sink",), + "srt/managers/schedule_batch.py": ( + "self.return_hidden_states = return_hidden_states", + "req.return_hidden_states or req.spec_capture is not None", + ), + "srt/model_executor/model_runner.py": ("spec_capture_method",), + "srt/models/qwen2.py": ("Capture after the final transformer layer",), + "srt/managers/scheduler_components/batch_result_processor.py": ( + "_sink_spec_capture", + "requested_artifacts = req.spec_capture.get", + ), +} +for relative, markers in required.items(): + path = root / relative + if not path.is_file() or path.is_symlink(): + raise SystemExit(f"missing or unsafe patched module: {path}") + text = path.read_text(encoding="utf-8") + missing = [marker for marker in markers if marker not in text] + if missing: + raise SystemExit(f"{path} lacks patch markers {missing}") + print(f"{relative}\t{hashlib.sha256(path.read_bytes()).hexdigest()}") +importlib.import_module("sglang.srt.spec_capture_sink") +importlib.import_module("sglang.srt.server_args") +PY + + CUDA_VISIBLE_DEVICES='' FLASHINFER_DISABLE_VERSION_CHECK=1 \ + "$python_executable" -m sglang.launch_server --help \ + >"$tmpdir/launch-help.txt" 2>"$tmpdir/launch-help.err" || { + sed -n '1,120p' "$tmpdir/launch-help.err" >&2 + die "patched sglang launch-server CLI is not importable" + } + local flag + for flag in \ + --enable-spec-capture \ + --spec-capture-method \ + --spec-capture-aux-layer-ids \ + --spec-capture-store-id \ + --spec-capture-max-sample-bytes \ + --spec-capture-inventory-db \ + --spec-capture-lifecycle-db \ + --disable-radix-cache \ + --chunked-prefill-size; do + grep -Fq -- "$flag" "$tmpdir/launch-help.txt" \ + || die "patched launch-server help lacks $flag" + done +} + +upgrade_base_tree() { + if ! patch --batch --forward -p2 -d "$sglang_parent" \ + < "$BASE_TO_CURRENT_PATH" >"$tmpdir/base-to-current-apply.log" 2>&1; then + sed -n '1,80p' "$tmpdir/base-to-current-apply.log" >&2 + die "base-to-current migration failed" + fi + patch_probe reverse "$tmpdir/post-migration.log" \ + || die "base migration did not produce the exact current patch state" + forward_ok=0 + reverse_ok=1 + base_patch_applied=0 +} + +case "$mode" in + check) + ((base_patch_applied == 0)) \ + || die "published base spec-capture patch is applied; run --apply to upgrade it" + ((reverse_ok == 1)) || die "spec-capture patch is not applied" + verify_patched_tree + echo "verified spec-capture patch $EXPECTED_PATCH_SHA256 in $sglang_root" + ;; + apply) + if ((base_patch_applied == 1)); then + upgrade_base_tree + elif ((forward_ok == 1)); then + patch --batch --forward -p2 -d "$sglang_parent" \ + < "$PATCH_PATH" >"$tmpdir/apply.log" + patch_probe reverse "$tmpdir/post-apply.log" \ + || die "post-apply reverse dry-run failed" + fi + verify_patched_tree + echo "applied and verified spec-capture patch $EXPECTED_PATCH_SHA256" + echo "python=$python_realpath sglang=$sglang_version root=$sglang_root" + ;; + reverse) + if ((base_patch_applied == 1)); then + upgrade_base_tree + fi + if ((reverse_ok == 1)); then + patch --batch --reverse -p2 -d "$sglang_parent" \ + < "$PATCH_PATH" >"$tmpdir/reverse-apply.log" + patch_probe forward "$tmpdir/post-reverse.log" \ + || die "post-reverse forward dry-run failed" + fi + echo "verified clean unpatched sglang $sglang_version at $sglang_root" + ;; +esac diff --git a/specforge/cli.py b/specforge/cli.py index 53b7d8a37..f95ec96fc 100644 --- a/specforge/cli.py +++ b/specforge/cli.py @@ -146,7 +146,9 @@ def _train(resolved) -> int: destroy_distributed() -def _config_for_role(cfg: Config, role: str) -> Config: +def _config_for_role( + cfg: Config, role: str, consumer_id: Optional[str] = None +) -> Config: """Resolve a launch role without changing the persisted run config. A shared disaggregated config may contain trainer-only state used by the @@ -156,11 +158,43 @@ def _config_for_role(cfg: Config, role: str) -> Config: raw = cfg.model_dump() raw["training"]["role"] = role disaggregated = raw["deployment"].get("disaggregated") + fanout = disaggregated.get("windowed_fanout") if disaggregated else None + managed_local = disaggregated.get("managed_local") if disaggregated else None + if fanout is not None and managed_local is not None: + devices = managed_local["trainer_cuda_visible_devices"] + for consumer, device in zip(fanout["consumers"], devices): + consumer["cuda_visible_device"] = device if disaggregated is not None and disaggregated.get("managed_local") is not None: # This field describes services owned by the parent supervisor. A role # child consumes the already-derived environment and must not attempt to # validate or own that stack again. disaggregated["managed_local"] = None + if role == "consumer" and fanout is not None: + consumer_id = consumer_id or os.environ.get("SPECFORGE_FANOUT_CONSUMER_ID") + matches = [ + consumer + for consumer in fanout["consumers"] + if consumer["consumer_id"] == consumer_id + ] + if len(matches) != 1: + raise ValueError( + f"unknown or missing windowed fanout consumer {consumer_id!r}" + ) + consumer = matches[0] + raw["training"].update( + { + "seed": consumer["seed"], + "loss_type": consumer["loss_type"], + "loss_decay_gamma": consumer["loss_decay_gamma"], + "dpace_alpha": consumer["dpace_alpha"], + "num_anchors": consumer["num_anchors"], + "learning_rate": consumer["learning_rate"], + "warmup_ratio": consumer["warmup_ratio"], + } + ) + if consumer["draft_block_size"] is not None: + raw["model"]["draft_block_size"] = consumer["draft_block_size"] + raw["output_dir"] = os.path.join(raw["output_dir"], consumer_id) if role == "producer": raw["profiling"]["enabled"] = False return Config.model_validate(raw) @@ -186,6 +220,11 @@ def main(argv: Optional[List[str]] = None) -> int: default=None, help="node-local rank for an explicit multi-node trainer launch", ) + train.add_argument( + "--consumer-id", + default=None, + help="select one independent consumer from windowed_fanout", + ) train.add_argument( "--plan", action="store_true", @@ -224,6 +263,7 @@ def main(argv: Optional[List[str]] = None) -> int: config_path=args.config, overrides=args.overrides, requested_role=args.role, + consumer_id=args.consumer_id, node_rank=args.node_rank, ) if args.plan: @@ -231,7 +271,7 @@ def main(argv: Optional[List[str]] = None) -> int: return 0 if plan.kind == "worker": os.environ.update(plan.worker_env) - role_config = _config_for_role(resolved.config, plan.role) + role_config = _config_for_role(resolved.config, plan.role, args.consumer_id) try: with _worker_signal_unwind(): _train(bind_run(role_config, resolved.algorithm)) diff --git a/specforge/config/__init__.py b/specforge/config/__init__.py index 70a885798..929170cfe 100644 --- a/specforge/config/__init__.py +++ b/specforge/config/__init__.py @@ -16,6 +16,8 @@ TrackingConfig, TrainerDeploymentConfig, TrainingConfig, + WindowedConsumerConfig, + WindowedFanoutConfig, apply_overrides, load_config, migrate_legacy_config, @@ -36,6 +38,8 @@ "RuntimeConfig", "SGLANG_CAPTURE_CONTEXT_HEADROOM", "TrainerDeploymentConfig", + "WindowedConsumerConfig", + "WindowedFanoutConfig", "load_config", "apply_overrides", "migrate_legacy_config", diff --git a/specforge/config/schema.py b/specforge/config/schema.py index 331db80ae..329fc14e8 100644 --- a/specforge/config/schema.py +++ b/specforge/config/schema.py @@ -19,14 +19,16 @@ import copy import json import os +import re from typing import List, Literal, Optional -from pydantic import BaseModel, ConfigDict, Field, model_validator +from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator # SGLang reserves one generated-token slot plus five internal slots, and its # request validator rejects ``input_len >= context_len - 6``. Accepting a # prompt whose length is exactly ``data.max_length`` therefore needs 7 slots. SGLANG_CAPTURE_CONTEXT_HEADROOM = 7 +_SAFE_FANOUT_ID = re.compile(r"^[A-Za-z0-9][A-Za-z0-9_.-]{0,127}$") class StrictConfigModel(BaseModel): @@ -355,6 +357,181 @@ def _validate_local_resources(self): return self +class WindowedConsumerConfig(StrictConfigModel): + """One independently scheduled single-GPU fanout consumer.""" + + consumer_id: str = Field(min_length=1, max_length=128) + seed: int = Field(ge=0) + loss_type: Literal[ + "dflash", + "dpace", + "dpace-cumulative-confidence-only", + "dpace-continuation-value-only", + ] + loss_decay_gamma: Optional[float] = Field(default=None, ge=0.0) + dpace_alpha: float = Field(ge=0.0, le=1.0) + draft_block_size: Optional[int] = Field(default=None, gt=0) + num_anchors: int = Field(gt=0) + learning_rate: float = Field(gt=0.0) + warmup_ratio: float = Field(ge=0.0, le=1.0) + cuda_visible_device: Optional[str] = None + resume_from: Optional[str] = None + window_lookbehind: Optional[int] = Field(default=None, ge=0) + window_lookahead: Optional[int] = Field(default=None, ge=0) + max_prefetch: Optional[int] = Field(default=None, ge=0) + + @field_validator("consumer_id") + @classmethod + def _safe_consumer_id(cls, value: str) -> str: + if _SAFE_FANOUT_ID.fullmatch(value) is None: + raise ValueError( + "windowed fanout consumer_id must contain only letters, digits, " + "dots, underscores, and hyphens" + ) + return value + + @field_validator("cuda_visible_device") + @classmethod + def _single_cuda_device(cls, value: Optional[str]) -> Optional[str]: + if value is None: + return value + _validate_cuda_devices( + [value], field_name="windowed_fanout.consumers[].cuda_visible_device" + ) + return value + + @field_validator("resume_from") + @classmethod + def _non_empty_resume(cls, value: Optional[str]) -> Optional[str]: + if value is not None and (not value or value.strip() != value): + raise ValueError( + "windowed fanout resume_from must be non-empty without " + "surrounding whitespace" + ) + return value + + +class WindowedFanoutConfig(StrictConfigModel): + """Bounded shared-capture policy for independent consumers.""" + + consumers: List[WindowedConsumerConfig] = Field(min_length=1) + window_lookbehind: int = Field(default=2, ge=0) + window_lookahead: int = Field(default=40, ge=0) + max_prefetch_per_consumer: int = Field(default=8, ge=0) + max_outstanding_per_consumer: int = Field(default=8, gt=0) + max_live_refs: int = Field(default=48, gt=0) + max_live_bytes: int = Field(gt=0) + capture_reservation_bytes: int = Field(default=128 << 20, gt=0) + capture_max_sample_bytes: int = Field(default=128 << 20, gt=0) + capture_batch_size: int = Field(default=8, gt=0) + capture_batch_wait_s: float = Field(default=0.002, ge=0.0) + registry_poll_s: float = Field(default=0.01, gt=0.0) + max_capture_retries: int = Field(default=2, ge=0) + capture_retry_backoff_s: float = Field(default=0.05, ge=0.0) + consumer_registration_timeout_s: float = Field(default=300.0, gt=0.0) + consumer_heartbeat_timeout_s: float = Field(default=30.0, gt=0.0) + consumer_heartbeat_interval_s: float = Field(default=5.0, gt=0.0) + consumer_idle_timeout_s: float = Field(default=1800.0, gt=0.0) + consumer_prefetch_batches: int = Field(default=1, ge=0) + + @model_validator(mode="after") + def _validate_window_contract(self): + ids = [consumer.consumer_id for consumer in self.consumers] + if len(ids) != len(set(ids)): + raise ValueError("windowed fanout consumer ids must be unique") + configured_devices = [ + consumer.cuda_visible_device + for consumer in self.consumers + if consumer.cuda_visible_device is not None + ] + if len(configured_devices) != len(set(configured_devices)): + raise ValueError("windowed fanout consumer CUDA devices must be unique") + if self.consumer_heartbeat_timeout_s <= self.consumer_heartbeat_interval_s: + raise ValueError( + "windowed fanout heartbeat timeout must exceed its interval" + ) + if self.max_prefetch_per_consumer > self.window_lookahead + 1: + raise ValueError( + "windowed fanout max_prefetch_per_consumer must not exceed " + "window_lookahead + 1" + ) + if self.max_live_refs < self.max_outstanding_per_consumer: + raise ValueError( + "windowed fanout max_live_refs must cover per-consumer " + "outstanding refs" + ) + if self.max_live_bytes < self.capture_reservation_bytes: + raise ValueError( + "windowed fanout max_live_bytes must cover one capture reservation" + ) + if self.capture_max_sample_bytes > self.capture_reservation_bytes: + raise ValueError( + "windowed fanout capture_max_sample_bytes must not exceed " + "capture_reservation_bytes" + ) + invalid_prefetch = { + consumer.consumer_id: ( + ( + self.max_prefetch_per_consumer + if consumer.max_prefetch is None + else consumer.max_prefetch + ), + ( + self.window_lookahead + if consumer.window_lookahead is None + else consumer.window_lookahead + ), + ) + for consumer in self.consumers + if ( + self.max_prefetch_per_consumer + if consumer.max_prefetch is None + else consumer.max_prefetch + ) + > ( + self.window_lookahead + if consumer.window_lookahead is None + else consumer.window_lookahead + ) + + 1 + } + if invalid_prefetch: + raise ValueError( + "windowed fanout consumer prefetch exceeds effective lookahead: " + f"{invalid_prefetch}" + ) + return self + + def consumer(self, consumer_id: str) -> WindowedConsumerConfig: + for consumer in self.consumers: + if consumer.consumer_id == consumer_id: + return consumer + raise ValueError( + f"unknown windowed fanout consumer {consumer_id!r}; expected one of " + f"{[consumer.consumer_id for consumer in self.consumers]}" + ) + + def window_for(self, consumer_id: str) -> tuple[int, int, int]: + consumer = self.consumer(consumer_id) + return ( + ( + self.window_lookbehind + if consumer.window_lookbehind is None + else consumer.window_lookbehind + ), + ( + self.window_lookahead + if consumer.window_lookahead is None + else consumer.window_lookahead + ), + ( + self.max_prefetch_per_consumer + if consumer.max_prefetch is None + else consumer.max_prefetch + ), + ) + + class DisaggregatedDeploymentConfig(StrictConfigModel): """Shared, non-secret topology for producer/consumer launch planning.""" @@ -388,6 +565,7 @@ class DisaggregatedDeploymentConfig(StrictConfigModel): #: managed_local supervisors use managed_local.shutdown_grace_s instead. shutdown_grace_s: float = Field(default=30.0, gt=0) managed_local: Optional[ManagedLocalStackConfig] = None + windowed_fanout: Optional[WindowedFanoutConfig] = None @model_validator(mode="after") def _validate_store(self): @@ -405,6 +583,8 @@ def _validate_store(self): raise ValueError( "deployment.disaggregated.store_root is required for shared_dir" ) + if self.windowed_fanout is not None and self.backend != "mooncake": + raise ValueError("windowed_fanout requires backend=mooncake") if self.managed_local is not None: if self.backend != "mooncake": raise ValueError("managed_local requires backend=mooncake") @@ -721,6 +901,11 @@ def _validate_run_structure(self): if self.deployment.disaggregated is not None else None ) + windowed_fanout = ( + self.deployment.disaggregated.windowed_fanout + if self.deployment.disaggregated is not None + else None + ) consumer_state_dir = ( self.deployment.disaggregated.consumer_state_dir if self.deployment.disaggregated is not None @@ -743,6 +928,95 @@ def _validate_run_structure(self): "multi-node online consumers require an explicit node-local " "deployment.disaggregated.consumer_state_dir for SQLite/WAL" ) + if windowed_fanout is not None: + if mode != "online" or deployment != "disaggregated": + raise ValueError( + "deployment.disaggregated.windowed_fanout requires online " + "disaggregated training" + ) + if self.training.strategy != "dflash": + raise ValueError("windowed_fanout currently supports DFlash only") + if self.training.num_epochs != 1: + raise ValueError("windowed_fanout supports exactly one prompt epoch") + if ( + self.deployment.trainer.nnodes != 1 + or self.deployment.trainer.nproc_per_node != 1 + ): + raise ValueError( + "windowed_fanout consumers are independent single-process " + "workers; keep deployment.trainer at 1x1" + ) + configured_capture_servers = ( + managed_local.capture_servers + if managed_local is not None + else self.deployment.disaggregated.server_urls + ) + if configured_capture_servers and len(configured_capture_servers) != 1: + raise ValueError( + "windowed_fanout currently requires exactly one capture server" + ) + if self.data.max_prompts is None or self.data.max_prompts < 1: + raise ValueError( + "windowed_fanout requires a positive data.max_prompts " + "for a fixed capture inventory" + ) + effective_batch = ( + self.training.batch_size * self.training.accumulation_steps + ) + if self.data.max_prompts % effective_batch: + raise ValueError( + "windowed_fanout data.max_prompts must be divisible by " + "training.batch_size * training.accumulation_steps" + ) + total_steps = self.data.max_prompts // effective_batch + if ( + self.training.total_steps is not None + and self.training.total_steps != total_steps + ): + raise ValueError( + "windowed_fanout training.total_steps must match the fixed " + f"prompt inventory ({total_steps})" + ) + if windowed_fanout.max_outstanding_per_consumer < effective_batch: + raise ValueError( + "windowed_fanout max_outstanding_per_consumer must cover " + "training.batch_size * training.accumulation_steps" + ) + configured_store_id = self.deployment.disaggregated.store_id + if configured_store_id not in (None, self.run_id): + raise ValueError( + "windowed_fanout requires deployment.disaggregated.store_id " + "to equal run_id" + ) + if self.training.resume_from is not None: + raise ValueError( + "windowed_fanout resume checkpoints belong to individual " + "consumer entries, not training.resume_from" + ) + if managed_local is None: + missing_devices = [ + consumer.consumer_id + for consumer in windowed_fanout.consumers + if consumer.cuda_visible_device is None + ] + if missing_devices: + raise ValueError( + "non-managed windowed_fanout consumers require " + "cuda_visible_device: " + f"{missing_devices}" + ) + else: + explicit_devices = [ + consumer.consumer_id + for consumer in windowed_fanout.consumers + if consumer.cuda_visible_device is not None + ] + if explicit_devices: + raise ValueError( + "managed_local windowed_fanout derives consumer devices " + "from trainer_cuda_visible_devices; remove per-consumer " + f"devices from {explicit_devices}" + ) if managed_local is not None: if mode != "online": raise ValueError("managed_local supports online capture only") @@ -754,6 +1028,17 @@ def _validate_run_structure(self): ) if self.training.resume_from is not None: raise ValueError("managed_local does not support resume") + if windowed_fanout is not None: + resumed_consumers = [ + consumer.consumer_id + for consumer in windowed_fanout.consumers + if consumer.resume_from is not None + ] + if resumed_consumers: + raise ValueError( + "managed_local requires a fresh fanout attempt; consumer " + f"resume is configured for {resumed_consumers}" + ) minimum_context_length = ( self.data.max_length + SGLANG_CAPTURE_CONTEXT_HEADROOM ) @@ -781,13 +1066,22 @@ def _validate_run_structure(self): "managed_local capture servers do not support SGLang DP " f"options: {unsupported_dp_options}" ) + expected_trainer_devices = ( + len(windowed_fanout.consumers) + if windowed_fanout is not None + else self.deployment.trainer.nproc_per_node + ) if ( len(managed_local.trainer_cuda_visible_devices) - != self.deployment.trainer.nproc_per_node + != expected_trainer_devices ): raise ValueError( "managed_local trainer_cuda_visible_devices count must equal " - "deployment.trainer.nproc_per_node" + + ( + "the windowed_fanout consumer count" + if windowed_fanout is not None + else "deployment.trainer.nproc_per_node" + ) ) ep_size = self.model.sglang_ep_size incompatible_tp_sizes = sorted( diff --git a/specforge/inference/adapters/__init__.py b/specforge/inference/adapters/__init__.py index f62a766eb..286d1074e 100644 --- a/specforge/inference/adapters/__init__.py +++ b/specforge/inference/adapters/__init__.py @@ -1,2 +1,2 @@ # coding=utf-8 -"""Feature-source adapters for external server transport.""" +"""Feature-source adapters for external transport and windowed materialization.""" diff --git a/specforge/inference/adapters/materializing.py b/specforge/inference/adapters/materializing.py new file mode 100644 index 000000000..d444ea470 --- /dev/null +++ b/specforge/inference/adapters/materializing.py @@ -0,0 +1,175 @@ +# coding=utf-8 +# Copyright 2024 The SpecForge team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Materialize a local ``FeatureSource`` into order-aligned ``SampleRef``s.""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any, Dict, List, Optional, Union + +from specforge.inference.capture import ( + CaptureConfig, + CaptureMismatchError, + verify_capture, +) +from specforge.runtime.contracts import PromptTask, SampleRef + + +@dataclass(frozen=True) +class MaterializationFailure: + """One task failed validation or persistence without failing its batch.""" + + task_id: str + reason: str + retryable: bool = True + + +class MaterializingRefSource: + """Adapt ``generate_features`` plus a store to the ``RefSource`` protocol. + + Windowed capture consumes order-aligned refs because it owns scheduling and + retry decisions. Local target engines instead expose ``generate_features``. + This adapter is the transport-neutral boundary between those APIs: it + verifies each capture before persistence and returns a typed per-task + failure when one record can be retried or rejected independently. + """ + + def __init__( + self, + feature_source: Any, + feature_store: Any, + *, + run_id: str, + strategy: str = "eagle3", + target_model_version: str = "unknown", + tokenizer_version: str = "unknown", + draft_weight_version: Optional[str] = None, + ) -> None: + if not callable(getattr(feature_source, "generate_features", None)): + raise TypeError( + "feature_source must expose generate_features(tasks, capture=...)" + ) + if not callable(getattr(feature_store, "put", None)): + raise TypeError("feature_store must expose put(tensors, ...)") + if not callable(getattr(feature_store, "abort", None)): + raise TypeError("feature_store must expose abort(sample_id, reason=...)") + self.feature_source = feature_source + self.feature_store = feature_store + self.run_id = run_id + self.strategy = strategy + self.target_model_version = target_model_version + self.tokenizer_version = tokenizer_version + self.draft_weight_version = draft_weight_version + + def _sample_id(self, task: PromptTask) -> str: + return f"{self.run_id}:{task.task_id}" + + def _put_metadata(self, task: PromptTask, capture: CaptureConfig) -> Dict[str, Any]: + return { + "run_id": self.run_id, + "source_task_id": task.task_id, + "strategy": self.strategy, + "target_repr": capture.target_repr, + "vocab_map_version": capture.vocab_map_version, + "ttt_length": capture.extra.get("ttt_length"), + "target_model_version": self.target_model_version, + "tokenizer_version": self.tokenizer_version, + "draft_weight_version": self.draft_weight_version, + "num_tokens": int(task.metadata.get("num_tokens", 0)), + } + + def produce_refs( + self, tasks: List[PromptTask], *, capture: CaptureConfig + ) -> List[Union[SampleRef, MaterializationFailure]]: + """Generate, validate, and persist exactly one aligned result per task.""" + features = self.feature_source.generate_features(tasks, capture=capture) + if len(features) != len(tasks): + reason = ( + f"generate_features returned {len(features)} feature records " + f"for {len(tasks)} tasks" + ) + return [ + MaterializationFailure( + task_id=task.task_id, + reason=reason, + retryable=False, + ) + for task in tasks + ] + + results: List[Union[SampleRef, MaterializationFailure]] = [] + for task, generated in zip(tasks, features): + sample_id = self._sample_id(task) + try: + tensors = dict(generated) + except (TypeError, ValueError) as exc: + results.append( + MaterializationFailure( + task_id=task.task_id, + reason=f"invalid feature record: {exc}", + retryable=False, + ) + ) + continue + recorded = tensors.pop("__aux_layer_ids__", None) + try: + if capture.aux_hidden_state_layer_ids and recorded is None: + raise CaptureMismatchError( + f"[{sample_id}] capture omitted aux-layer ids; cannot " + "verify requested layers " + f"{capture.aux_hidden_state_layer_ids}" + ) + verify_capture( + tensors, + capture, + sample_id=sample_id, + recorded_aux_layer_ids=recorded, + ) + except CaptureMismatchError as exc: + results.append( + MaterializationFailure( + task_id=task.task_id, + reason=str(exc), + retryable=False, + ) + ) + continue + + try: + ref = self.feature_store.put( + tensors, + sample_id=sample_id, + metadata=self._put_metadata(task, capture), + ) + except Exception as exc: + try: + self.feature_store.abort(sample_id, reason=f"put_failed:{exc}") + except Exception as cleanup_error: + exc.add_note( + f"failed to abort partial materialization: {cleanup_error}" + ) + results.append( + MaterializationFailure( + task_id=task.task_id, + reason=f"put_failed:{exc}", + retryable=True, + ) + ) + continue + results.append(ref) + return results + + +__all__ = ["MaterializationFailure", "MaterializingRefSource"] diff --git a/specforge/inference/adapters/server_capture.py b/specforge/inference/adapters/server_capture.py index 4cb3fc192..2dd95162f 100644 --- a/specforge/inference/adapters/server_capture.py +++ b/specforge/inference/adapters/server_capture.py @@ -8,25 +8,27 @@ # http://www.apache.org/licenses/LICENSE-2.0 """Server-side spec-capture rollout source (zero-copy Mooncake transport). -An external SGLang server patched with -``patches/sglang/v0.5.14/spec-capture.patch`` runs +The *server* transport (vs the in-process ``PolicyFeatureAdapter``): a live +SGLang server patched with ``patches/sglang/v0.5.14/spec-capture.patch`` runs the prefill and writes captured features straight into Mooncake in :class:`MooncakeFeatureStore`'s key layout. Tensors never pass through this process — the ``/generate`` response's ``meta_info["spec_capture"]`` carries only key/shape/dtype, from which :meth:`SGLangServerCaptureAdapter.produce_refs` builds committed-ready ``SampleRef``s. -The server knows only generic artifacts (``aux`` = capture layers -concatenated, ``last_hidden`` = post-norm final hidden) plus passthrough -tensors. The application composition root injects an algorithm-owned -:class:`ServerCaptureSchema`; this transport never resolves an algorithm name. +The server knows only generic artifacts (``aux`` = capture layers concatenated, +``last_hidden`` = post-norm final hidden) plus passthrough tensors. The +application composition root injects the algorithm-owned +:class:`ServerCaptureSchema`; this transport does not resolve algorithms. """ from __future__ import annotations import logging +import os +import time from dataclasses import dataclass -from typing import Any, Callable, Dict, List, Mapping, Optional, Tuple, Union +from typing import Any, Callable, Dict, Iterable, List, Mapping, Optional, Tuple, Union from specforge.inference.capture import ( CaptureConfig, @@ -120,7 +122,8 @@ class SGLangServerCaptureAdapter: ``store`` must be the run's :class:`MooncakeFeatureStore` (its ``store_id`` namespaces the keys and ``adopt()`` registers each ref so a later - ``abort()``/``gc()`` on the producer side can free server-written objects). + generation-aware ``reclaim()``/``gc()`` on the producer side can free + server-written objects). ``post_fn`` is injectable for tests. """ @@ -136,10 +139,12 @@ def __init__( timeout_s: float = 300.0, post_fn: Optional[Callable[..., Any]] = None, target_model_version: str = "unknown", + capture_token: Optional[str] = None, ) -> None: required_store_api = ( "adopt", "discard_external_attempts", + "reclaim", "store_id", "track_external_attempt", ) @@ -151,6 +156,11 @@ def __init__( "SGLangServerCaptureAdapter needs a MooncakeFeatureStore-like " f"store; missing {missing_store_api}" ) + if run_id != store.store_id: + raise ValueError( + "server capture requires run_id == store.store_id; got " + f"{run_id!r} != {store.store_id!r}" + ) self.base_url = base_url.rstrip("/") self.store = store self.run_id = run_id @@ -170,14 +180,26 @@ def __init__( self.timeout_s = timeout_s self.post_fn = post_fn or _default_post self.target_model_version = target_model_version + self.capture_token = capture_token or os.environ.get( + "SGLANG_SPEC_CAPTURE_TOKEN" + ) + if not self.capture_token: + raise ValueError( + "server capture requires capture_token or SGLANG_SPEC_CAPTURE_TOKEN" + ) self._healthy = True + self._rpc_calls = 0 + self._rpc_tasks = 0 + self._rpc_time_s = 0.0 + self._rpc_failures = 0 + self._result_failures = 0 # -- request construction ------------------------------------------------- def _sample_id(self, task: PromptTask) -> str: return f"{self.run_id}:{task.task_id}" def _request_inputs(self, tasks: List[PromptTask]) -> Dict[str, Any]: - """Build only model-input fields, keeping transport keys runtime-owned.""" + """Build model inputs while keeping transport fields runtime-owned.""" if self.request_input_adapter is None: return {"input_ids": [list(task.payload["input_ids"]) for task in tasks]} @@ -186,8 +208,7 @@ def _request_inputs(self, tasks: List[PromptTask]) -> Dict[str, Any]: raise TypeError( "ServerInputAdapter.build_request_inputs must return a mapping" ) - reserved = {"sampling_params", "spec_capture"} - conflicts = sorted(reserved & set(request_inputs)) + conflicts = sorted({"sampling_params", "spec_capture"} & set(request_inputs)) if conflicts: raise ValueError( "ServerInputAdapter cannot set runtime-owned request fields: " @@ -202,6 +223,16 @@ def _request_inputs(self, tasks: List[PromptTask]) -> Dict[str, Any]: def _spec_capture_payload(self, task: PromptTask) -> Dict[str, Any]: input_ids = list(task.payload["input_ids"]) length = len(input_ids) + generation = task.metadata.get("capture_generation", 1) + if ( + isinstance(generation, bool) + or not isinstance(generation, int) + or not 1 <= generation <= 2**31 - 1 + ): + raise ValueError( + f"task {task.task_id}: capture_generation must be an integer " + "in [1, 2**31 - 1]" + ) features: Dict[str, str] = {} if self.schema.aux_feature is not None: features["aux"] = self.schema.aux_feature @@ -241,15 +272,14 @@ def _spec_capture_payload(self, task: PromptTask) -> Dict[str, Any]: } ) return { + "auth_token": self.capture_token, "store_id": self.store.store_id, "sample_id": self._sample_id(task), - # A task id is unique within a run, and retries happen only before - # its ref is committed. Keep the generation stable so a response - # lost after the server write cannot strand gN when the retry writes - # gN+1. The server patch replaces these deterministic keys on a - # retry (``replace``), bounding the attempt to one namespace. - "gen": 1, - "replace": int(task.attempt) > 0, + # Ordinary response-loss retries keep g1. Windowed recapture supplies + # its durable registry generation so a cleaned sample gets a fresh + # physical namespace instead of attempting to resurrect g1. + "gen": generation, + "replace": int(task.attempt) > 0 or generation > 1, "features": features, "passthrough": passthrough, } @@ -266,8 +296,21 @@ def _ref_from_result( specs: Dict[str, FeatureSpec] = {} nbytes = 0 for name, meta in feats.items(): - shape = tuple(int(d) for d in meta["shape"]) - dtype = str(meta["dtype"]) + if not isinstance(meta, dict): + raise TypeError(f"feature {name!r} metadata must be an object") + raw_shape = meta["shape"] + if not isinstance(raw_shape, (list, tuple)): + raise TypeError(f"feature {name!r} shape must be a list or tuple") + if not raw_shape: + raise ValueError(f"feature {name!r} shape must not be empty") + if any(type(d) is not int or d <= 0 for d in raw_shape): + raise ValueError( + f"feature {name!r} shape must contain positive integers" + ) + shape = tuple(raw_shape) + dtype = meta["dtype"] + if not isinstance(dtype, str) or dtype not in _DTYPE_BYTES: + raise ValueError(f"feature {name!r} has unsupported dtype {dtype!r}") extra: Dict[str, Any] = {} if name == self.schema.last_hidden_feature: extra["target_repr"] = capture.target_repr @@ -305,6 +348,46 @@ def _ref_from_result( }, ) + def _cleanup_ref( + self, + task: PromptTask, + *, + sample_id: str, + store_id: str, + gen: int, + feature_names: Iterable[str], + ) -> "SampleRef": # noqa: F821 + """Build an exact-generation ref without trusting response tensor metadata.""" + from specforge.runtime.contracts import SampleRef + + names = tuple(sorted(feature_names)) + return SampleRef( + sample_id=sample_id, + run_id=self.run_id, + source_task_id=task.task_id, + feature_store_uri=f"mooncake://{store_id}/{sample_id}", + feature_keys={name: f"{sample_id}/{name}" for name in names}, + feature_specs={}, + strategy=self.strategy, + schema_version=SCHEMA_VERSION, + target_model_version=self.target_model_version, + tokenizer_version=str(task.metadata.get("tokenizer_version", "unknown")), + num_tokens=0, + estimated_bytes=0, + metadata={"generation": gen}, + ) + + @staticmethod + def _aux_layer_ids_from_result( + result: Dict[str, Any], + ) -> Optional[Tuple[int, ...]]: + raw = result.get("aux_layer_ids") + if raw is None: + return None + if not isinstance(raw, (list, tuple)) or any(type(v) is not int for v in raw): + raise ValueError("aux_layer_ids must be a list of integers") + return tuple(raw) + # -- the RefSource entry point ---------------------------------------------- def produce_refs( self, tasks: List[PromptTask], *, capture: CaptureConfig @@ -330,9 +413,18 @@ def produce_refs( generation=int(payload["gen"]), feature_names=feature_names, ) - rows = self.post_fn( - f"{self.base_url}/generate", json_body=body, timeout=self.timeout_s - ) + self._rpc_calls += 1 + self._rpc_tasks += len(tasks) + rpc_started = time.perf_counter() + try: + rows = self.post_fn( + f"{self.base_url}/generate", json_body=body, timeout=self.timeout_s + ) + except BaseException: + self._rpc_failures += 1 + raise + finally: + self._rpc_time_s += time.perf_counter() - rpc_started rows = _flatten_list_wrappers(rows) if len(rows) != len(tasks): raise RuntimeError( @@ -340,8 +432,8 @@ def produce_refs( f"{len(tasks)} tasks" ) out: List[Union[Any, ServerCaptureFailure]] = [] - successful_refs = [] - for task, row in zip(tasks, rows): + successful_refs: List["SampleRef"] = [] # noqa: F821 + for task, request_spec, row in zip(tasks, capture_payloads, rows): if not isinstance(row, dict): raise RuntimeError( "spec-capture server returned a non-object row for task " @@ -355,45 +447,81 @@ def produce_refs( task_id=task.task_id, expected_sample_id=self._sample_id(task), ) + expected_identity = { + "sample_id": self._sample_id(task), + "store_id": str(self.store.store_id), + "gen": int(request_spec["gen"]), + } + expected_features = set(request_spec["features"].values()) | { + item["name"] for item in request_spec["passthrough"] + } + cleanup_ref = self._cleanup_ref( + task, + sample_id=expected_identity["sample_id"], + store_id=expected_identity["store_id"], + gen=expected_identity["gen"], + feature_names=expected_features, + ) + + def reject(reason: str, *, retryable: bool) -> ServerCaptureFailure: + if not retryable: + self.store.adopt(cleanup_ref) + self.store.reclaim(cleanup_ref, reason=reason) + return ServerCaptureFailure( + task_id=task.task_id, + reason=f"server_capture:{reason}", + retryable=retryable, + ) + if not result: out.append( - ServerCaptureFailure( - task_id=task.task_id, - reason=( - "server_capture: response carries no spec_capture " - "result — is the server patched and launched with " - "--enable-spec-capture?" - ), + reject( + "response carries no spec_capture result; is the server " + "patched and launched with --enable-spec-capture?", retryable=False, ) ) continue if result.get("error"): - out.append( - ServerCaptureFailure( - task_id=task.task_id, - reason=f"server_capture:{result['error']}", - retryable=True, - ) - ) + out.append(reject(str(result["error"]), retryable=True)) continue - expected_identity = { - "sample_id": self._sample_id(task), - "store_id": str(self.store.store_id), - "gen": 1, - } - actual_identity = { - "sample_id": str(result.get("sample_id")), - "store_id": str(result.get("store_id")), - "gen": int(result.get("gen", -1)), - } + try: + actual_identity = { + "sample_id": str(result.get("sample_id")), + "store_id": str(result.get("store_id")), + "gen": int(result.get("gen", -1)), + } + except (TypeError, ValueError) as exc: + raise RuntimeError( + "spec-capture server returned malformed object identity for " + f"task {task.task_id}: {exc}" + ) from exc if actual_identity != expected_identity: raise RuntimeError( "spec-capture server returned the wrong object identity for " f"task {task.task_id}: {actual_identity} != " f"{expected_identity}" ) - ref = self._ref_from_result(task, result, capture) + features = result.get("features") + actual_features = set(features) if isinstance(features, dict) else set() + if actual_features != expected_features: + out.append( + reject( + "response feature set mismatch " + f"expected={sorted(expected_features)}, " + f"actual={sorted(actual_features)}", + retryable=False, + ) + ) + continue + try: + ref = self._ref_from_result(task, result, capture) + recorded_aux_layer_ids = self._aux_layer_ids_from_result(result) + except (KeyError, TypeError, ValueError) as exc: + out.append( + reject(f"malformed feature metadata: {exc}", retryable=False) + ) + continue # A capture shorter than the prompt is corrupt (classic cause: a # radix-cache prefix hit skips prefilling — and capturing — the # cached tokens; the patched scheduler refuses that config). @@ -404,22 +532,16 @@ def produce_refs( if len(spec.shape) >= 2 and spec.shape[1] != expected_len } if short: - self.store.adopt(ref) - self.store.abort(ref.sample_id, reason="seq-len-mismatch") out.append( - ServerCaptureFailure( - task_id=task.task_id, - reason=( - f"server_capture: captured seq len != prompt len " - f"{expected_len} for {short} — was the server " - f"started without --disable-radix-cache?" - ), + reject( + f"captured seq len != prompt len {expected_len} for " + f"{short}; was the server started without " + "--disable-radix-cache?", retryable=False, ) ) continue try: - recorded_aux_layer_ids = result.get("aux_layer_ids") if ( capture.aux_hidden_state_layer_ids and recorded_aux_layer_ids is None @@ -433,24 +555,14 @@ def produce_refs( ref.feature_specs, capture, sample_id=ref.sample_id, - recorded_aux_layer_ids=( - tuple(recorded_aux_layer_ids) - if recorded_aux_layer_ids is not None - else None - ), + recorded_aux_layer_ids=recorded_aux_layer_ids, aux_feature_name=self.schema.aux_feature or "hidden_state", target_feature_name=self.schema.last_hidden_feature or "target", ) except CaptureMismatchError as exc: # Loud boundary failure; free the server-written keys so a # mismatched sample is never consumable. - self.store.adopt(ref) - self.store.abort(ref.sample_id, reason=f"contract:{exc}") - out.append( - ServerCaptureFailure( - task_id=task.task_id, reason=str(exc), retryable=False - ) - ) + out.append(reject(f"capture contract mismatch: {exc}", retryable=False)) continue # Do not transfer a successful row out of provisional ownership # until every row in the response has passed structural and @@ -461,14 +573,27 @@ def produce_refs( out.append(ref) for ref in successful_refs: self.store.adopt(ref) + self._result_failures += sum( + isinstance(result, ServerCaptureFailure) for result in out + ) return out def health(self) -> Dict[str, Any]: + mean_batch_size = self._rpc_tasks / self._rpc_calls if self._rpc_calls else 0.0 return { "healthy": self._healthy, "backend": "sglang_server_capture", "base_url": self.base_url, "strategy": self.strategy, + "rpc_calls": self._rpc_calls, + "rpc_tasks": self._rpc_tasks, + "rpc_time_s": self._rpc_time_s, + "rpc_failures": self._rpc_failures, + "result_failures": self._result_failures, + "mean_batch_size": mean_batch_size, + "tasks_per_rpc_second": ( + self._rpc_tasks / self._rpc_time_s if self._rpc_time_s else 0.0 + ), } diff --git a/specforge/launch.py b/specforge/launch.py index 7ad9b7582..3a6238ce8 100644 --- a/specforge/launch.py +++ b/specforge/launch.py @@ -11,7 +11,9 @@ from __future__ import annotations import logging -from typing import Any, Callable, List, Mapping, Optional, Tuple +import os +from dataclasses import dataclass +from typing import Any, Callable, Dict, List, Mapping, Optional, Tuple from specforge.algorithms.registry import AlgorithmRegistration from specforge.runtime.contracts import SampleRef @@ -1651,9 +1653,548 @@ def stop_distributor_and_drain() -> None: raise +def _windowed_prompt_tasks(run_id: str, prompts) -> List[Any]: + """Normalize a restart-stable canonical prompt stream.""" + from specforge.runtime.contracts import PromptTask, assert_no_tensors + + tasks = [] + for prompt in prompts: + if isinstance(prompt, PromptTask): + task = prompt + if task.run_id != run_id: + raise ValueError( + f"prompt {task.task_id!r} run_id={task.run_id!r} does not " + f"match capture run {run_id!r}" + ) + else: + assert_no_tensors(prompt) + task_id = prompt.get("task_id") + if not isinstance(task_id, str) or not task_id: + raise ValueError( + "windowed capture requires an explicit stable task_id on " + "every prompt" + ) + task = PromptTask( + task_id=task_id, + run_id=run_id, + source_id=str(prompt.get("source_id", "prompt_source")), + payload=dict(prompt.get("payload", prompt)), + max_length=int(prompt.get("max_length", 2048)), + chat_template=prompt.get("chat_template"), + loss_mask_policy=dict(prompt.get("loss_mask_policy", {})), + target_model_version=str(prompt.get("target_model_version", "unknown")), + draft_weight_version=prompt.get("draft_weight_version"), + metadata=dict(prompt.get("metadata", {})), + ) + assert_no_tensors(task) + tasks.append(task) + if not tasks: + raise ValueError("windowed capture prompts must not be empty") + ids = [task.task_id for task in tasks] + if len(ids) != len(set(ids)): + raise ValueError("windowed capture task_ids must be unique") + return tasks + + +def _resolve_window_algorithm( + strategy: str | AlgorithmRegistration, +) -> AlgorithmRegistration: + if isinstance(strategy, AlgorithmRegistration): + return strategy + from specforge.algorithms.builtin import builtin_algorithm_registry + + try: + return builtin_algorithm_registry().resolve(strategy) + except KeyError as exc: + raise ValueError(str(exc)) from exc + + +def build_disagg_windowed_capture_contract( + *, + strategy: str | AlgorithmRegistration, + modality: str = "text", + target_hidden_size: int, + target_model_version: str, + tokenizer_version: str, + target_vocab_size: Optional[int] = None, + draft_vocab_size: Optional[int] = None, + target_repr: Optional[str] = None, + aux_hidden_state_layer_ids=None, + vocab_map_version: Optional[str] = None, +): + """Build the typed capture request and its cross-process identity digest.""" + from specforge.inference.capture import CaptureConfig + from specforge.runtime.data_plane.windowed_capture import capture_contract_digest + + algorithm = _resolve_window_algorithm(strategy) + feature_contract = algorithm.spec.feature_contract("streaming", modality) + capture = CaptureConfig.from_strategy( + required_features=feature_contract.required_tensors, + aux_hidden_state_layer_ids=tuple(aux_hidden_state_layer_ids or ()), + target_repr=target_repr, + target_hidden_size=target_hidden_size, + target_vocab_size=target_vocab_size, + draft_vocab_size=draft_vocab_size, + vocab_map_version=vocab_map_version, + ) + digest = capture_contract_digest( + { + "strategy": algorithm.name, + "capture": capture, + "target_model_version": target_model_version, + "tokenizer_version": tokenizer_version, + } + ) + return capture, digest + + +@dataclass +class DisaggWindowedProducerRuntime: + registry: Any + service: Any + contract_digest: str + + def drive(self, max_rounds: int = 10_000_000, *, should_stop=None) -> int: + return self.service.drive(should_stop=should_stop, max_rounds=max_rounds) + + def accounting_snapshot(self) -> Dict[str, Any]: + snapshot = self.service.snapshot() + snapshot["contract_digest"] = self.contract_digest + return snapshot + + def close(self) -> None: + self.registry.close() + + +def build_disagg_online_windowed_producer( + *, + prompts, + feature_store: FeatureStore, + feature_source: Any, + run_id: str, + consumer_ids, + registry_db_path: str, + max_live_refs: int, + target_hidden_size: int, + target_model_version: str, + tokenizer_version: str, + strategy: str = "eagle3", + modality: str = "text", + target_vocab_size: Optional[int] = None, + draft_vocab_size: Optional[int] = None, + target_repr: Optional[str] = None, + aux_hidden_state_layer_ids=None, + vocab_map_version: Optional[str] = None, + max_live_bytes: Optional[int] = None, + capture_reservation_bytes: Optional[int] = None, + capture_batch_size: int = 8, + capture_batch_wait_s: float = 0.002, + max_capture_retries: int = 2, + retry_backoff_s: float = 0.05, + consumer_registration_timeout_s: float = 600.0, + consumer_heartbeat_timeout_s: float = 120.0, + registry_poll_s: float = 0.01, + recover: bool = False, +) -> DisaggWindowedProducerRuntime: + """Build one demand-driven producer for fixed independent consumers.""" + from specforge.runtime.data_plane.windowed_capture import ( + SQLiteWindowedCaptureRegistry, + ) + from specforge.runtime.data_plane.windowed_capture_runtime import ( + WindowedCaptureService, + ) + + tasks = _windowed_prompt_tasks(run_id, prompts) + incompatible_prompts = [ + task.task_id + for task in tasks + if task.target_model_version not in ("unknown", target_model_version) + ] + if incompatible_prompts: + raise ValueError( + "windowed prompts target a different model version: " + f"{incompatible_prompts[:8]}" + ) + capture, digest = build_disagg_windowed_capture_contract( + strategy=strategy, + modality=modality, + target_hidden_size=target_hidden_size, + target_model_version=target_model_version, + tokenizer_version=tokenizer_version, + target_vocab_size=target_vocab_size, + draft_vocab_size=draft_vocab_size, + target_repr=target_repr, + aux_hidden_state_layer_ids=aux_hidden_state_layer_ids, + vocab_map_version=vocab_map_version, + ) + registry = SQLiteWindowedCaptureRegistry( + registry_db_path, + max_live_refs=max_live_refs, + max_live_bytes=max_live_bytes, + capture_reservation_bytes=capture_reservation_bytes, + poll_s=registry_poll_s, + ) + try: + registry.initialize_run( + run_id=run_id, + contract_digest=digest, + source_sample_ids=[task.task_id for task in tasks], + expected_consumers=tuple(consumer_ids), + recover_inflight=recover, + recovery_store=feature_store if recover else None, + ) + service = WindowedCaptureService( + registry, + prompts=tasks, + feature_source=feature_source, + capture=capture, + owner_store=feature_store, + capture_batch_size=capture_batch_size, + batch_wait_s=capture_batch_wait_s, + max_capture_retries=max_capture_retries, + retry_backoff_s=retry_backoff_s, + consumer_registration_timeout_s=consumer_registration_timeout_s, + consumer_heartbeat_timeout_s=consumer_heartbeat_timeout_s, + poll_s=registry_poll_s, + ) + except BaseException: + registry.close() + raise + return DisaggWindowedProducerRuntime(registry, service, digest) + + +@dataclass +class DisaggWindowedConsumerRuntime: + trainer: Any + loader: Any + queue: Any + control: Any + controller: DataFlowController + registry: Any + max_steps: Optional[int] + + def run(self) -> int: + """Train to EOF or to this consumer's explicit independent step cap.""" + try: + self.control.mark_ready() + step = self.trainer.fit() + self.control.ensure_healthy() + if self.queue.drained(): + self.control.complete() + elif self.max_steps is not None and step >= self.max_steps: + self.queue.close() + self.control.complete(allow_partial=True) + else: + raise RuntimeError( + "windowed consumer stopped before EOF without reaching its " + "configured max_steps" + ) + return step + except BaseException as exc: + try: + self.queue.close() + except BaseException as cleanup_error: + exc.add_note(f"failed to close windowed queue: {cleanup_error!r}") + try: + state = self.registry.snapshot()["consumers"][self.control.consumer_id][ + "state" + ] + if state != "completed": + self.control.fail(exc) + except BaseException as cleanup_error: + exc.add_note(f"failed to report consumer failure: {cleanup_error!r}") + raise + + def accounting_snapshot(self) -> Dict[str, Any]: + marker = self.controller.store.durable_marker() + return { + "consumer_id": self.control.consumer_id, + "window": self.registry.snapshot()["consumers"][self.control.consumer_id], + "committed": self.controller.store.committed_count(), + "acked": len(marker["acked"]), + "global_step": marker["global_step"], + "queue": self.queue.metrics(), + } + + def close(self) -> None: + self.control.close() + close_store = getattr(self.controller.store, "close", None) + if callable(close_store): + close_store() + self.registry.close() + + +def _durable_window_cursor(metadata_store: MetadataStore) -> int: + """Return the contiguous optimizer-durable prefix in canonical fetch order.""" + marker = metadata_store.durable_marker() + acked = marker["acked"] + cursor = 0 + for sample_id in metadata_store.all_committed_ids(): + if sample_id not in acked: + break + cursor += 1 + return cursor + + +def build_disagg_online_windowed_consumer( + *, + consumer_id: str, + registry_db_path: str, + max_live_refs: int, + contract_digest: str, + total_samples: int, + feature_store: FeatureStore, + draft_model, + optimizer_factory, + run_id: str, + output_dir: str, + metadata_db_path: str, + lookbehind: int = 0, + lookahead: int = 0, + prefetch_depth: int = 0, + max_outstanding: int = 1, + strategy: str | AlgorithmRegistration = "eagle3", + modality: str = "text", + batch_size: int = 1, + accumulation_steps: int = 1, + num_epochs: int = 1, + max_steps: Optional[int] = None, + total_steps: Optional[int] = None, + save_interval: int = 0, + eval_interval: int = 0, + collate_fn=None, + idle_timeout_s: Optional[float] = 1800.0, + logger=None, + log_interval: int = 50, + strategy_kwargs: Optional[dict] = None, + resume: bool = False, + resume_from: Optional[str] = None, + max_checkpoints: int = 0, + max_live_bytes: Optional[int] = None, + capture_reservation_bytes: Optional[int] = None, + heartbeat_interval_s: float = 5.0, + initialization_timeout_s: float = 600.0, + registry_poll_s: float = 0.01, + loader_prefetch_batches: int = 0, + consumer_control=None, +) -> DisaggWindowedConsumerRuntime: + """Build one single-GPU trainer over an independent windowed cursor.""" + from specforge.runtime.control_plane.metadata_store import SQLiteMetadataStore + from specforge.runtime.data_plane.windowed_capture import ( + SQLiteWindowedCaptureRegistry, + WindowedCaptureQueue, + ) + from specforge.runtime.data_plane.windowed_capture_runtime import ( + start_windowed_consumer_control, + ) + + if num_epochs != 1: + raise ValueError("windowed canonical streams support exactly one epoch") + if max_steps is not None and max_steps < 1: + raise ValueError("max_steps must be >= 1 or None") + if resume_from is not None and not resume: + raise ValueError("resume_from requires resume=True") + if max_outstanding < batch_size * accumulation_steps: + raise ValueError( + "max_outstanding must cover batch_size * accumulation_steps so " + "optimizer-boundary ACKs cannot deadlock the queue" + ) + if hasattr(feature_store, "lifetime_owner") and feature_store.lifetime_owner: + raise ValueError("windowed consumers must not own shared payload lifetime") + + owns_registry = consumer_control is None + registry = ( + consumer_control.registry + if not owns_registry + else SQLiteWindowedCaptureRegistry( + registry_db_path, + max_live_refs=max_live_refs, + max_live_bytes=max_live_bytes, + capture_reservation_bytes=capture_reservation_bytes, + poll_s=registry_poll_s, + ) + ) + metadata_store = None + try: + initialized = registry.wait_initialized(initialization_timeout_s) + observed = ( + initialized["run_id"], + initialized["contract_digest"], + initialized["total_samples"], + initialized["max_live_refs"], + initialized["max_live_bytes"], + initialized["capture_reservation_bytes"], + ) + expected = ( + run_id, + contract_digest, + total_samples, + max_live_refs, + max_live_bytes, + capture_reservation_bytes or 0, + ) + if observed != expected: + raise RuntimeError( + f"windowed registry identity mismatch: expected={expected!r}, " + f"observed={observed!r}" + ) + + os.makedirs(os.path.dirname(os.path.abspath(metadata_db_path)), exist_ok=True) + metadata_store = SQLiteMetadataStore(metadata_db_path) + if not resume and metadata_store.committed_count(): + raise ValueError( + "windowed consumer metadata store is not fresh; pass resume=True " + "with a matching checkpoint or use a new run directory" + ) + durable_cursor = _durable_window_cursor(metadata_store) if resume else 0 + if durable_cursor and resume_from is None: + raise ValueError( + "windowed consumer resume found an acknowledged prefix but no " + "resume_from checkpoint; skipping trained samples without restoring " + "their weight updates would lose data" + ) + if resume_from is not None: + marker_step = metadata_store.durable_marker()["global_step"] + checkpoint_step = _checkpoint_global_step(resume_from) + if marker_step is not None and marker_step > checkpoint_step: + raise RuntimeError( + f"durable marker global_step={marker_step} is ahead of " + f"checkpoint global_step={checkpoint_step}" + ) + if consumer_control is None: + existing = registry.snapshot()["consumers"].get(consumer_id) + if existing is not None and not resume: + raise RuntimeError( + f"consumer {consumer_id!r} already exists; pass resume=True" + ) + consumer_control = start_windowed_consumer_control( + registry, + consumer_id, + lookbehind=lookbehind, + lookahead=lookahead, + prefetch_depth=prefetch_depth, + max_outstanding=max_outstanding, + heartbeat_interval_s=heartbeat_interval_s, + durable_cursor=durable_cursor, + ) + else: + if consumer_control.consumer_id != consumer_id: + raise ValueError("consumer_control identity mismatch") + existing = registry.snapshot()["consumers"][consumer_id] + expected_window = ( + lookbehind, + lookahead, + prefetch_depth, + max_outstanding, + ) + observed_window = tuple( + int(existing[name]) + for name in ( + "lookbehind", + "lookahead", + "prefetch_depth", + "max_outstanding", + ) + ) + if observed_window != expected_window: + raise ValueError( + f"consumer_control window mismatch: expected={expected_window}, " + f"observed={observed_window}" + ) + if resume: + registry.resume_consumer(consumer_id, durable_cursor=durable_cursor) + except BaseException as exc: + if owns_registry and consumer_control is not None: + try: + consumer_control.fail(exc) + except BaseException as cleanup_error: + exc.add_note( + f"failed to report windowed consumer setup failure: " + f"{cleanup_error!r}" + ) + if metadata_store is not None: + metadata_store.close() + if owns_registry: + registry.close() + raise + + queue = None + try: + controller = DataFlowController( + run_id, + metadata_store=metadata_store, + enable_sample_queue=False, + ) + queue = WindowedCaptureQueue( + registry, + consumer_id, + idle_timeout_s=idle_timeout_s, + record_refs=lambda refs: controller.record_external_refs(list(refs)), + ) + algorithm = _resolve_window_algorithm(strategy) + trainer = _assemble_trainer( + algorithm=algorithm, + controller=controller, + store=feature_store, + ref_source={"queue": queue, "defer_ack_until_durable": True}, + model=draft_model, + target_head=None, + optimizer_factory=optimizer_factory, + run_id=run_id, + output_dir=output_dir, + batch_size=batch_size, + accumulation_steps=accumulation_steps, + num_epochs=1, + max_steps=max_steps, + total_steps=total_steps, + save_interval=save_interval, + eval_interval=eval_interval, + tp_size=1, + sp_ulysses_size=1, + sp_ring_size=1, + logger=logger, + log_interval=log_interval, + collate_fn=_streaming_collate(algorithm, modality, collate_fn), + strategy_kwargs=strategy_kwargs, + per_sample_transform=None, + durable_ack=True, + resume_from=resume_from, + max_checkpoints=max_checkpoints, + dataloader_num_workers=loader_prefetch_batches, + ) + except BaseException as exc: + if queue is not None: + try: + queue.close() + except BaseException as cleanup_error: + exc.add_note(f"failed to close windowed queue: {cleanup_error!r}") + try: + consumer_control.fail(exc) + except BaseException as cleanup_error: + exc.add_note( + f"failed to report windowed consumer setup failure: " + f"{cleanup_error!r}" + ) + metadata_store.close() + if owns_registry: + registry.close() + raise + return DisaggWindowedConsumerRuntime( + trainer, + trainer.loader, + queue, + consumer_control, + controller, + registry, + max_steps, + ) + + __all__ = [ "build_offline_runtime", "build_disagg_offline_runtime", "build_disagg_online_producer", "build_disagg_online_consumer", + "build_disagg_online_windowed_producer", + "build_disagg_online_windowed_consumer", ] diff --git a/specforge/launch_plan.py b/specforge/launch_plan.py index 98f852f49..d6d07e29e 100644 --- a/specforge/launch_plan.py +++ b/specforge/launch_plan.py @@ -6,6 +6,7 @@ import importlib.util import json import os +import secrets import shutil import signal import socket @@ -27,11 +28,14 @@ LaunchRole = Literal["auto", "all", "producer", "consumer", "both"] PlanKind = Literal["worker", "command", "supervisor", "managed_supervisor"] ReadinessKind = Literal["http", "mooncake"] +_HTTP_READINESS_PROBE_TIMEOUT_S = 5.0 _DIST_ENV = ("RANK", "WORLD_SIZE", "LOCAL_RANK", "MASTER_ADDR", "MASTER_PORT") _MANAGED_CHILD_ENV = "SPECFORGE_MANAGED_LOCAL_CHILD" +_FANOUT_CONSUMER_ENV = "SPECFORGE_FANOUT_CONSUMER_ID" _SECRET_NAMES = ( "auth_token", + "capture_token", "password", "secret", "credential", @@ -271,23 +275,34 @@ def _disaggregated_env( # Online feature objects are allocated by the external capture server. # SpecForge roles only read or publish references to those objects. values["DISAGG_CLIENT_SEGMENT_SIZE"] = "0" - values.update( - { - "DISAGG_REF_CHANNEL": str(control_dir / "refs.jsonl"), - "DISAGG_DB": str(consumer_state_dir / "consumer.sqlite"), - # SQLite/WAL stays on rank 0's local filesystem. Inboxes are - # ordinary append-only channels and must remain visible to - # ranks on every trainer node. - "DISAGG_INBOX_DIR": str( - ( - control_dir - if cfg.deployment.trainer.nnodes > 1 - else consumer_state_dir - ) - / "inboxes" - ), - } - ) + if deployment.windowed_fanout is not None: + values.update( + { + "DISAGG_WINDOW_REGISTRY": str( + control_dir / "windowed-capture.sqlite" + ), + "DISAGG_CAPTURE_LIFECYCLE_DB": str( + control_dir / "capture-lifecycle.sqlite" + ), + } + ) + else: + values.update( + { + "DISAGG_REF_CHANNEL": str(control_dir / "refs.jsonl"), + "DISAGG_DB": str(consumer_state_dir / "consumer.sqlite"), + # SQLite/WAL stays on rank 0's local filesystem. Inboxes + # must remain visible to ranks on every trainer node. + "DISAGG_INBOX_DIR": str( + ( + control_dir + if cfg.deployment.trainer.nnodes > 1 + else consumer_state_dir + ) + / "inboxes" + ), + } + ) else: values["DISAGG_MANIFEST"] = str(control_dir / "manifest.json") if deployment.backend == "mooncake": @@ -318,6 +333,7 @@ def _disaggregated_env( "DISAGG_IDLE_TIMEOUT": deployment.idle_timeout_s, "DISAGG_PEER_WAIT_TIMEOUT": deployment.peer_wait_timeout_s, "DISAGG_PRODUCER_HOLD_S": deployment.producer_hold_s, + "SGLANG_SPEC_CAPTURE_TOKEN": base_env.get("SGLANG_SPEC_CAPTURE_TOKEN"), } for name, configured in optional_values.items(): value = base_env.get(name, configured) @@ -337,7 +353,9 @@ def _disaggregated_env( return values -def _managed_local_environment(cfg: Config) -> dict[str, str]: +def _managed_local_environment( + cfg: Config, *, capture_token: Optional[str] = None +) -> dict[str, str]: deployment = cfg.deployment.disaggregated assert deployment is not None and deployment.managed_local is not None managed = deployment.managed_local @@ -356,6 +374,8 @@ def _managed_local_environment(cfg: Config) -> dict[str, str]: } if mooncake.rdma_devices: values["MOONCAKE_RDMA_DEVICES"] = mooncake.rdma_devices + if capture_token is not None: + values["SGLANG_SPEC_CAPTURE_TOKEN"] = capture_token return values @@ -363,6 +383,7 @@ def _managed_local_services( cfg: Config, *, algorithm: "AlgorithmRegistration", + shared_env: Optional[Mapping[str, str]] = None, ) -> tuple[ServiceSpec, ...]: from specforge.training.capture_contract import resolve_server_capture_contract @@ -372,7 +393,12 @@ def _managed_local_services( mooncake = managed.mooncake control_dir = Path(deployment.control_dir) log_dir = control_dir / "logs" - shared_env = _managed_local_environment(cfg) + shared_env = dict(shared_env or _managed_local_environment(cfg)) + fanout = deployment.windowed_fanout + capture_max_sample_bytes = ( + fanout.capture_max_sample_bytes if fanout is not None else 1 << 30 + ) + capture_lifecycle_db = control_dir / "capture-lifecycle.sqlite" capture_context_length = cfg.model.sglang_context_length or ( cfg.data.max_length + SGLANG_CAPTURE_CONTEXT_HEADROOM ) @@ -436,6 +462,14 @@ def _managed_local_services( "-1", "--disable-radix-cache", "--enable-spec-capture", + "--spec-capture-store-id", + deployment.store_id or cfg.run_id, + "--spec-capture-max-sample-bytes", + str(capture_max_sample_bytes), + "--spec-capture-inventory-db", + str(control_dir / f"capture-server-{index}-inventory.sqlite"), + "--spec-capture-lifecycle-db", + str(capture_lifecycle_db), "--spec-capture-method", contract.method, "--spec-capture-aux-layer-ids", @@ -490,16 +524,21 @@ def _worker_argv( config_path: str, role: Literal["all", "producer", "consumer"], overrides: Sequence[str], + *, + consumer_id: Optional[str] = None, ) -> list[str]: - return [ + argv = [ *command_prefix, "train", "--config", config_path, "--role", role, - *overrides, ] + if consumer_id is not None: + argv.extend(("--consumer-id", consumer_id)) + argv.extend(overrides) + return argv def _trainer_command( @@ -513,13 +552,20 @@ def _trainer_command( distributed_entry: Sequence[str], node_rank: Optional[int], env: Mapping[str, str], + consumer_id: Optional[str] = None, ) -> CommandSpec: topology = cfg.deployment.trainer - worker = _worker_argv(worker_prefix, config_path, role, overrides) + worker = _worker_argv( + worker_prefix, + config_path, + role, + overrides, + consumer_id=consumer_id, + ) world_size = topology.nnodes * topology.nproc_per_node cfg.validate_world_size(world_size) if world_size == 1: - return CommandSpec(role, tuple(worker), env) + return CommandSpec(consumer_id or role, tuple(worker), env) if topology.nnodes == 1: launch = [ *torchrun_prefix, @@ -546,7 +592,67 @@ def _trainer_command( str(topology.nproc_per_node), ] worker_args = worker[len(worker_prefix) :] - return CommandSpec(role, tuple([*launch, *distributed_entry, *worker_args]), env) + return CommandSpec( + consumer_id or role, + tuple([*launch, *distributed_entry, *worker_args]), + env, + ) + + +def _fanout_consumer_env( + cfg: Config, + consumer_id: str, + base_env: Mapping[str, str], +) -> dict[str, str]: + deployment = cfg.deployment.disaggregated + assert deployment is not None and deployment.windowed_fanout is not None + fanout = deployment.windowed_fanout + consumer = fanout.consumer(consumer_id) + env = dict(base_env) + env[_FANOUT_CONSUMER_ENV] = consumer_id + state_root = Path(deployment.consumer_state_dir or deployment.control_dir) + env["DISAGG_DB"] = str(state_root / "consumers" / consumer_id / "consumer.sqlite") + if deployment.managed_local is not None: + index = [item.consumer_id for item in fanout.consumers].index(consumer_id) + device = deployment.managed_local.trainer_cuda_visible_devices[index] + else: + assert consumer.cuda_visible_device is not None + device = consumer.cuda_visible_device + env["CUDA_VISIBLE_DEVICES"] = device + return env + + +def _validate_fanout_consumer_database( + cfg: Config, + consumer_id: str, + launch_env: Mapping[str, str], + *, + base_env: Mapping[str, str], + distributed: bool, + node_rank: Optional[int], +) -> None: + deployment = cfg.deployment.disaggregated + assert deployment is not None and deployment.windowed_fanout is not None + consumer = deployment.windowed_fanout.consumer(consumer_id) + database = launch_env["DISAGG_DB"] + state_owner = int(base_env["RANK"]) == 0 if distributed else node_rank in (None, 0) + if consumer.resume_from is not None: + if state_owner and not os.path.exists(database): + raise ValueError( + f"fanout consumer {consumer_id!r} resume requires retained " + f"metadata database: {database}" + ) + return + stale = [ + path + for path in (database, f"{database}-wal", f"{database}-shm") + if os.path.exists(path) + ] + if state_owner and stale: + raise ValueError( + f"fanout consumer {consumer_id!r} requires a fresh metadata path; " + f"found {stale}" + ) def _validate_consumer_database( @@ -564,8 +670,10 @@ def _validate_consumer_database( or role not in ("consumer", "both") ): return - database = launch_env.get("DISAGG_DB") or base_env.get("DISAGG_DB") deployment = cfg.deployment.disaggregated + if deployment is not None and deployment.windowed_fanout is not None: + return + database = launch_env.get("DISAGG_DB") or base_env.get("DISAGG_DB") if not database and deployment is not None: state_dir = deployment.consumer_state_dir or deployment.control_dir database = str(Path(state_dir) / "consumer.sqlite") @@ -605,15 +713,37 @@ def _validate_capture_urls( return deployment = cfg.deployment.disaggregated if deployment is not None and deployment.managed_local is not None: + if ( + deployment.windowed_fanout is not None + and len(deployment.managed_local.capture_servers) != 1 + ): + raise ValueError( + "windowed_fanout currently requires exactly one capture server" + ) return - if deployment is not None and deployment.server_urls: - return - if base_env.get("DISAGG_SERVER_URLS") or base_env.get("DISAGG_SERVER_URL"): - return - raise ValueError( - "online disaggregated producer requires server URLs in " - "deployment.disaggregated.server_urls or DISAGG_SERVER_URL(S)" - ) + urls = list(deployment.server_urls) if deployment is not None else [] + if not urls: + raw_urls = base_env.get("DISAGG_SERVER_URLS") or base_env.get( + "DISAGG_SERVER_URL" + ) + urls = ( + [item.strip() for item in raw_urls.split(",") if item.strip()] + if raw_urls + else [] + ) + if not urls: + raise ValueError( + "online disaggregated producer requires server URLs in " + "deployment.disaggregated.server_urls or DISAGG_SERVER_URL(S)" + ) + if ( + deployment is not None + and deployment.windowed_fanout is not None + and len(urls) != 1 + ): + raise ValueError( + "windowed_fanout currently requires exactly one capture server" + ) def build_launch_plan( @@ -623,6 +753,7 @@ def build_launch_plan( config_path: str, overrides: Sequence[str] = (), requested_role: LaunchRole = "auto", + consumer_id: Optional[str] = None, node_rank: Optional[int] = None, env: Optional[Mapping[str, str]] = None, worker_prefix: Optional[Sequence[str]] = None, @@ -643,6 +774,7 @@ def build_launch_plan( distributed = _distributed_state(base_env) deployment = cfg.deployment.disaggregated managed_local = deployment.managed_local if deployment is not None else None + fanout = deployment.windowed_fanout if deployment is not None else None if managed_local is not None and algorithm is None: raise ValueError( "managed_local launch planning requires a resolved algorithm registration" @@ -675,6 +807,22 @@ def build_launch_plan( f"exists: {deployment.control_dir}" ) role = _resolve_role(cfg, requested_role, distributed=distributed) + selected_consumer_id: Optional[str] = None + if fanout is None: + if consumer_id is not None: + raise ValueError("--consumer-id requires windowed_fanout") + elif role == "consumer": + if consumer_id is None: + if len(fanout.consumers) != 1: + raise ValueError( + "--role consumer requires --consumer-id when windowed_fanout " + "defines multiple consumers" + ) + consumer_id = fanout.consumers[0].consumer_id + fanout.consumer(consumer_id) + selected_consumer_id = consumer_id + elif consumer_id is not None: + raise ValueError("--consumer-id is valid only with --role consumer") topology = cfg.deployment.trainer if role == "both" and topology.nnodes > 1: raise ValueError( @@ -688,25 +836,52 @@ def build_launch_plan( ) producer_env: dict[str, str] = {} consumer_env: dict[str, str] = {} + fanout_consumer_envs: dict[str, dict[str, str]] = {} + managed_environment: dict[str, str] = {} if cfg.deployment.mode == "disaggregated": role_base_env = base_env - managed_environment: dict[str, str] = {} if managed_local is not None: - managed_environment = _managed_local_environment(cfg) + capture_token = base_env.get("SGLANG_SPEC_CAPTURE_TOKEN") + if capture_token is None: + if managed_child: + raise ValueError( + "managed_local child is missing SGLANG_SPEC_CAPTURE_TOKEN" + ) + capture_token = secrets.token_urlsafe(32) + managed_environment = _managed_local_environment( + cfg, capture_token=capture_token + ) role_base_env = {**base_env, **managed_environment} if role in ("producer", "both"): producer_env = _disaggregated_env(cfg, role_base_env, role="producer") if role in ("consumer", "both"): - consumer_env = _disaggregated_env(cfg, role_base_env, role="consumer") + base_consumer_env = _disaggregated_env(cfg, role_base_env, role="consumer") if managed_local is not None: producer_env.update(managed_environment) producer_env["CUDA_VISIBLE_DEVICES"] = "" producer_env[_MANAGED_CHILD_ENV] = "1" - consumer_env.update(managed_environment) - consumer_env["CUDA_VISIBLE_DEVICES"] = ",".join( - managed_local.trainer_cuda_visible_devices - ) - consumer_env[_MANAGED_CHILD_ENV] = "1" + if role in ("consumer", "both"): + base_consumer_env.update(managed_environment) + base_consumer_env[_MANAGED_CHILD_ENV] = "1" + if role in ("consumer", "both"): + if fanout is None: + consumer_env = base_consumer_env + if managed_local is not None: + consumer_env["CUDA_VISIBLE_DEVICES"] = ",".join( + managed_local.trainer_cuda_visible_devices + ) + else: + ids = ( + [selected_consumer_id] + if selected_consumer_id is not None + else [consumer.consumer_id for consumer in fanout.consumers] + ) + fanout_consumer_envs = { + item: _fanout_consumer_env(cfg, item, base_consumer_env) + for item in ids + } + if selected_consumer_id is not None: + consumer_env = fanout_consumer_envs[selected_consumer_id] _validate_capture_urls(cfg, role=role, base_env=base_env) _validate_consumer_database( @@ -717,6 +892,15 @@ def build_launch_plan( distributed=distributed, node_rank=resolved_rank, ) + for item, item_env in fanout_consumer_envs.items(): + _validate_fanout_consumer_database( + cfg, + item, + item_env, + base_env=base_env, + distributed=distributed, + node_rank=resolved_rank, + ) if distributed: assert role != "both" if role == "producer" and int(base_env["WORLD_SIZE"]) > 1: @@ -759,6 +943,7 @@ def build_launch_plan( distributed_entry=distributed_entry, node_rank=resolved_rank, env=consumer_env, + consumer_id=selected_consumer_id, ) if command.argv[: len(worker_prefix)] == worker_prefix: return LaunchPlan("worker", role, worker_env=consumer_env) @@ -776,6 +961,49 @@ def build_launch_plan( tuple(_worker_argv(worker_prefix, config_path, "producer", overrides)), producer_env, ) + if fanout is not None: + consumers = tuple( + _trainer_command( + cfg, + config_path=config_path, + role="consumer", + overrides=overrides, + worker_prefix=worker_prefix, + torchrun_prefix=torchrun_prefix, + distributed_entry=distributed_entry, + node_rank=resolved_rank, + env=fanout_consumer_envs[item.consumer_id], + consumer_id=item.consumer_id, + ) + for item in fanout.consumers + ) + commands = (*consumers, producer) + if managed_local is not None: + return LaunchPlan( + "managed_supervisor", + "both", + commands=commands, + services=_managed_local_services( + cfg, + algorithm=algorithm, + shared_env=managed_environment, + ), + managed_root=deployment.control_dir, + managed_ports=( + managed_local.mooncake.rpc_port, + managed_local.mooncake.metadata_port, + managed_local.mooncake.metrics_port, + *[server.port for server in managed_local.capture_servers], + ), + shutdown_grace_s=managed_local.shutdown_grace_s, + ) + return LaunchPlan( + "supervisor", + "both", + commands=commands, + shutdown_grace_s=deployment.shutdown_grace_s, + ) + consumer = _trainer_command( cfg, config_path=config_path, @@ -792,7 +1020,9 @@ def build_launch_plan( "managed_supervisor", "both", commands=(producer, consumer), - services=_managed_local_services(cfg, algorithm=algorithm), + services=_managed_local_services( + cfg, algorithm=algorithm, shared_env=managed_environment + ), managed_root=deployment.control_dir, managed_ports=( managed_local.mooncake.rpc_port, @@ -912,7 +1142,10 @@ def _managed_preflight(plan: LaunchPlan) -> None: def _http_ready(readiness: ReadinessSpec) -> bool: try: - with urllib_request.urlopen(readiness.url, timeout=1.0) as response: + with urllib_request.urlopen( + readiness.url, + timeout=min(_HTTP_READINESS_PROBE_TIMEOUT_S, readiness.timeout_s), + ) as response: status = getattr(response, "status", 200) return readiness.kind == "mooncake" or 200 <= status < 300 except urllib_error.HTTPError as exc: diff --git a/specforge/runtime/ARCHITECTURE.md b/specforge/runtime/ARCHITECTURE.md index 25a3d238a..655802b97 100644 --- a/specforge/runtime/ARCHITECTURE.md +++ b/specforge/runtime/ARCHITECTURE.md @@ -2,12 +2,15 @@ SpecForge has one public training entry point, `specforge train`. A typed run configuration selects an algorithm and a topology; it does not select a second -trainer. The launch layer exposes exactly four topology builders: +trainer. The launch layer exposes canonical builders for fixed-stream and +windowed-stream roles: - `build_offline_runtime` - `build_disagg_offline_runtime` - `build_disagg_online_producer` - `build_disagg_online_consumer` +- `build_disagg_online_windowed_producer` +- `build_disagg_online_windowed_consumer` All trainer-bearing builders converge on the same `Trainer -> FeatureDataLoader -> TrainerController -> TrainerCore` path. Only @@ -20,12 +23,17 @@ the reference source and feature-store backend change. | Colocated offline | Precomputed feature files | Fixed `SampleRef` list | `LocalFeatureStore` reads `file://` refs | Re-iterable; epochs and checkpoint resume are supported | | Disaggregated offline | `CONFIG=path/to/offline-disagg.yaml run_offline.sh --role producer` ingests existing files and writes a static manifest | Fixed manifest refs | Shared directory or Mooncake | Re-iterable; DP/multi-node epochs and checkpoint resume are supported | | Online | Patched SGLang server writes tensors; producer publishes refs | Per-rank `StreamingRefQueue` inbox | Mooncake | Consume once; consumer-only recovery reconciles retained state; no producer resume or second pass | +| Online independent fanout | One producer serves demand from independently scheduled consumers | Per-consumer `WindowedCaptureQueue` | Mooncake with a bounded SQLite capture registry | One fixed prompt pass; consumers may differ and advance independently; compatible live captures are shared | `training.num_epochs` on an online run controls how many prompt passes the producer creates. Each pass receives new task and sample ids. The consumer still iterates one consume-once stream exactly once; it never replays a prior stream as a second trainer epoch. +Windowed fanout is stricter: it requires exactly one prompt pass and a fixed +`data.max_prompts` inventory so every independent consumer has the same stable +source index space. + ## Cross-plane contracts - The control plane carries `PromptTask` and `SampleRef` metadata only. @@ -114,6 +122,38 @@ counter is settled. Every inbox then closes normally after the aligned prefix; an actual cleanup failure still poisons the inboxes. A partial global optimizer step is never dispatched. +## Independent windowed fanout + +Windowed fanout is not a data-parallel consumer. `launch_plan` starts one +producer command and one single-process `specforge train` command per consumer. +Every consumer uses the normal model, optimizer, loader, controller, and trainer +assembly, but has its own cursor, checkpoint, output directory, and projected +training parameters. + +```mermaid +flowchart LR + P[producer] --> R[(SQLite window registry)] + R --> S[patched SGLang capture] + S --> M[Mooncake payloads] + R --> Q1[consumer A queue] + R --> Q2[consumer B queue] + R --> Q3[consumer C queue] + M -.-> Q1 + M -.-> Q2 + M -.-> Q3 + Q1 --> T1[canonical trainer A] + Q2 --> T2[canonical trainer B] + Q3 --> T3[canonical trainer C] +``` + +Consumers register and begin heartbeats before expensive model construction. +Moving a cursor updates its legal lookbehind/lookahead interests and schedules +bounded prefetch. Demand, active capture, and read leases are hard interests; +window-only entries may be reclaimed under slot or byte pressure and recaptured +later with a new generation. The producer owns Mooncake lifetime cleanup, while +consumer stores are non-owners. One consumer failure releases only its interests; +the remaining consumers continue independently. + The terminal tail remains committed-but-unacknowledged in the attempt's metadata ledger even though its feature objects are removed. A successfully completed attempt with such a tail must therefore start any later run with a diff --git a/specforge/runtime/control_plane/controller.py b/specforge/runtime/control_plane/controller.py index 53a577f50..804b3a4a1 100644 --- a/specforge/runtime/control_plane/controller.py +++ b/specforge/runtime/control_plane/controller.py @@ -200,7 +200,14 @@ def commit_samples(self, worker_id: str, refs: List[SampleRef]) -> List[SampleRe self.sample_queue.put(fresh) return fresh - # -- offline ingest ---------------------------------------------------- + def record_external_refs(self, refs: List[SampleRef]) -> int: + """Ledger refs supplied by an external queue without enqueueing them.""" + committed = 0 + for ref in refs: + assert_no_tensors(ref) + committed += int(self.store.commit_sample(ref)) + return committed + # -- train-side durable ack ------------------------------------------- def ack_train_refs( self, diff --git a/specforge/runtime/data_plane/DESIGN.md b/specforge/runtime/data_plane/DESIGN.md index 1caf60275..2d59d880c 100644 --- a/specforge/runtime/data_plane/DESIGN.md +++ b/specforge/runtime/data_plane/DESIGN.md @@ -14,6 +14,8 @@ See [`../ARCHITECTURE.md`](../ARCHITECTURE.md) for the complete topology. - Offline refs are fixed and re-iterable. Online refs are consume-once and are never replayed as a second consumer epoch; resume rebuilds or reconciles only the untrained suffix. +- Independent fanout keeps one fixed source-index inventory, bounds both live + refs and live bytes, and never requires consumers to advance in lockstep. ## Topology map @@ -151,6 +153,35 @@ private inbox advances its exact acknowledged prefix. This ordering keeps producer backpressure and restart accounting aligned with durable optimizer progress. +### Windowed independent consumers + +The original `windowed_capture` import is a compatibility facade over three +ownership-focused modules: + +- `windowed_capture_contracts.py` defines capture keys, requests, leases, + failures, and the canonical contract digest; +- `windowed_capture_registry.py` owns all SQLite transactions, generations, + consumer cursors, interest accounting, reservations, and reclamation; +- `windowed_capture_queue.py` adapts one consumer cursor to the loader queue + protocol. + +`windowed_capture_runtime.py` adds producer batching/retry and consumer +heartbeat control. A consumer registers before model loading so producer +startup does not mistake long CUDA/model initialization for a missing worker. +Its queue acquires refs in source order and advances the registry cursor only +after the optimizer-durable ACK reaches the queue. + +The registry distinguishes soft window interest from hard demand/read interest. +It can evict a ready entry used only for lookahead, but cannot evict an acquired +lease. Capture slots and bytes are reserved in the same transaction that claims +work. A failed or reclaimed attempt fences its generation, so a delayed capture +completion cannot publish stale payload metadata. + +The producer is the Mooncake lifetime owner and drains every retained object on +success, failure, or signal unwind. Consumers create non-owner store clients and +may finish or fail independently; the registry removes only that consumer's +waiters, leases, and window interests. + ## Attempt lifecycle An online producer always claims a fresh attempt and cannot resume. A fresh diff --git a/specforge/runtime/data_plane/__init__.py b/specforge/runtime/data_plane/__init__.py index b0592f482..e780ef5f1 100644 --- a/specforge/runtime/data_plane/__init__.py +++ b/specforge/runtime/data_plane/__init__.py @@ -21,6 +21,16 @@ "SharedDirFeatureStore", "MooncakeFeatureStore", "AuthPolicy", + "CaptureFailedError", + "CapturePriority", + "CaptureReadLease", + "CaptureRequest", + "SQLiteWindowedCaptureRegistry", + "WindowedCaptureQueue", + "capture_contract_digest", + "WindowedCaptureService", + "WindowedConsumerControl", + "start_windowed_consumer_control", ] _EXPORT_MODULE = { @@ -37,6 +47,16 @@ "SharedDirFeatureStore": "disaggregated", "MooncakeFeatureStore": "mooncake_store", "AuthPolicy": "disaggregated", + "CaptureFailedError": "windowed_capture", + "CapturePriority": "windowed_capture", + "CaptureReadLease": "windowed_capture", + "CaptureRequest": "windowed_capture", + "SQLiteWindowedCaptureRegistry": "windowed_capture", + "WindowedCaptureQueue": "windowed_capture", + "capture_contract_digest": "windowed_capture", + "WindowedCaptureService": "windowed_capture_runtime", + "WindowedConsumerControl": "windowed_capture_runtime", + "start_windowed_consumer_control": "windowed_capture_runtime", } diff --git a/specforge/runtime/data_plane/disaggregated.py b/specforge/runtime/data_plane/disaggregated.py index 39e642d11..0bcfe1cf1 100644 --- a/specforge/runtime/data_plane/disaggregated.py +++ b/specforge/runtime/data_plane/disaggregated.py @@ -78,6 +78,7 @@ def __init__( credential: Optional[str] = None, max_hold_age_s: Optional[float] = None, retain_on_release: bool = False, + lifetime_owner: bool = True, clock: Callable[[], float] = time.monotonic, ) -> None: self.auth = auth or AuthPolicy() @@ -92,6 +93,9 @@ def __init__( # LocalFeatureStore's file:// no-op release). Cleanup is whole-store at run # end; consume-once free (retain_on_release=False) is for online rollout. self.retain_on_release = retain_on_release + # Independent fan-out readers release only their local leases. The + # producer is the sole owner allowed to reclaim shared payload files. + self.lifetime_owner = lifetime_owner self._clock = clock # in-process liveness index (generation / put-time / active leases) self._generation: Dict[str, int] = {} @@ -239,6 +243,8 @@ def release(self, handle: FeatureHandle, *, reason: str = "consumed") -> None: # never delete the freshly re-put current generation (different filename). with self._lock: self._active_leases.pop(handle.lease_token, None) + if not self.lifetime_owner: + return if self.retain_on_release: return # offline re-iterable set: keep the file for the next epoch sid, gen = handle.sample_id, handle.generation @@ -251,17 +257,53 @@ def release(self, handle: FeatureHandle, *, reason: str = "consumed") -> None: def abort(self, sample_id: str, *, reason: str = "aborted") -> None: with self._lock: + if not self.lifetime_owner: + self._active_leases = { + token: handle + for token, handle in self._active_leases.items() + if handle.sample_id != sample_id + } + return for gen in self._disk_gens(sample_id): self._free_gen_locked(sample_id, gen) self._generation.pop(sample_id, None) self._put_time.pop(sample_id, None) + def reclaim(self, sample_ref: SampleRef, *, reason: str = "consumed") -> None: + """Free only the generation named by a globally consumed fan-out ref.""" + gen = sample_ref.metadata.get("generation") + if gen is None: + raise ValueError( + f"cannot reclaim {sample_ref.sample_id}: ref carries no generation" + ) + gen = int(gen) + with self._lock: + if not self.lifetime_owner: + raise RuntimeError( + "only the shared-directory lifetime owner may reclaim" + ) + if any( + handle.sample_id == sample_ref.sample_id and handle.generation == gen + for handle in self._active_leases.values() + ): + raise RuntimeError( + f"cannot reclaim leased sample {sample_ref.sample_id} " + f"generation {gen}" + ) + self._free_gen_locked(sample_ref.sample_id, gen) + def gc(self, *, now: Optional[float] = None) -> Dict[str, int]: # Max-hold force-free tracks puts made by this process. Other processes # observe disk residency through health(), but do not share lease ages. now = self._clock() if now is None else now freed = freed_bytes = 0 with self._lock: + if not self.lifetime_owner: + return { + "force_freed": 0, + "force_freed_bytes": 0, + "release_pending": 0, + } if self.max_hold_age_s is not None: stale = [ sid @@ -313,6 +355,7 @@ def health(self) -> Dict[str, Any]: "active_leases": active_leases, "resident_bytes": resident_bytes, "auth_required": self.auth.required, + "lifetime_owner": self.lifetime_owner, "oldest_age_s": max(ages) if ages else 0.0, "avg_age_s": (sum(ages) / len(ages)) if ages else 0.0, "force_freed_total": force_freed, diff --git a/specforge/runtime/data_plane/mooncake_lifecycle.py b/specforge/runtime/data_plane/mooncake_lifecycle.py new file mode 100644 index 000000000..95044cbad --- /dev/null +++ b/specforge/runtime/data_plane/mooncake_lifecycle.py @@ -0,0 +1,239 @@ +# coding=utf-8 +"""Durable single-host lifecycle index for Mooncake-owned feature objects.""" + +from __future__ import annotations + +import json +import os +import sqlite3 +import threading +import time +from dataclasses import dataclass +from typing import Iterable, Optional + + +@dataclass(frozen=True) +class LifecycleRecord: + sample_id: str + generation: int + feature_names: tuple[str, ...] + estimated_bytes: int + state: str + + +class SQLiteMooncakeLifecycleIndex: + """Persistent owner inventory and cross-process tombstone authority. + + The database contains metadata only. Tensor payloads remain exclusively in + Mooncake. SQLite is the single-host production tier used by the manifest + supervisor; a cross-node deployment must provide an equivalent shared + metadata service before sharing this contract across hosts. + """ + + _PLANNED = "planned" + _LIVE = "resident" + _TOMBSTONED = "tombstoned" + _CLEANED = "cleaned" + + def __init__(self, path: str, *, store_id: str) -> None: + self.path = os.path.abspath(path) + self.store_id = store_id + os.makedirs(os.path.dirname(self.path), exist_ok=True) + self._conn = sqlite3.connect(self.path, check_same_thread=False, timeout=30.0) + self._conn.execute("PRAGMA busy_timeout=30000") + self._conn.execute("PRAGMA journal_mode=WAL") + self._conn.execute("PRAGMA synchronous=FULL") + self._lock = threading.RLock() + with self._lock: + self._conn.execute( + "CREATE TABLE IF NOT EXISTS mooncake_objects (" + "store_id TEXT NOT NULL, sample_id TEXT NOT NULL, " + "generation INTEGER NOT NULL, feature_names_json TEXT NOT NULL, " + "estimated_bytes INTEGER NOT NULL, state TEXT NOT NULL, " + "reason TEXT, updated_at REAL NOT NULL, " + "PRIMARY KEY (store_id, sample_id, generation))" + ) + self._conn.commit() + + def record_planned( + self, + sample_id: str, + generation: int, + feature_names: Iterable[str], + estimated_bytes: int, + ) -> None: + """Durably declare every key before the first hard-pinned write.""" + names_json = json.dumps(sorted(feature_names), separators=(",", ":")) + with self._lock: + row = self._conn.execute( + "SELECT feature_names_json, estimated_bytes, state FROM " + "mooncake_objects WHERE store_id=? AND sample_id=? AND generation=?", + (self.store_id, sample_id, generation), + ).fetchone() + if row is not None: + prior_names, prior_bytes, state = row + if state in {self._TOMBSTONED, self._CLEANED}: + raise RuntimeError( + f"refusing to plan {sample_id} generation {generation} " + f"from lifecycle state {state!r}" + ) + if prior_names != names_json or int(prior_bytes) != int( + estimated_bytes + ): + raise RuntimeError( + f"lifecycle identity changed for {sample_id} generation " + f"{generation}" + ) + return + self._conn.execute( + "INSERT INTO mooncake_objects " + "(store_id, sample_id, generation, feature_names_json, " + "estimated_bytes, state, reason, updated_at) " + "VALUES (?, ?, ?, ?, ?, ?, NULL, ?)", + ( + self.store_id, + sample_id, + generation, + names_json, + int(estimated_bytes), + self._PLANNED, + time.time(), + ), + ) + self._conn.commit() + + def record_resident( + self, + sample_id: str, + generation: int, + feature_names: Iterable[str], + estimated_bytes: int, + ) -> None: + names_json = json.dumps(sorted(feature_names), separators=(",", ":")) + with self._lock: + row = self._conn.execute( + "SELECT feature_names_json, estimated_bytes, state FROM " + "mooncake_objects WHERE store_id=? AND " + "sample_id=? AND generation=?", + (self.store_id, sample_id, generation), + ).fetchone() + if row is not None and row[2] not in {self._PLANNED, self._LIVE}: + raise RuntimeError( + f"refusing to resurrect {sample_id} generation {generation} " + f"from lifecycle state {row[2]!r}" + ) + if row is not None and ( + row[0] != names_json or int(row[1]) != int(estimated_bytes) + ): + raise RuntimeError( + f"lifecycle identity changed for {sample_id} generation " + f"{generation}" + ) + self._conn.execute( + "INSERT OR REPLACE INTO mooncake_objects " + "(store_id, sample_id, generation, feature_names_json, " + "estimated_bytes, state, reason, updated_at) " + "VALUES (?, ?, ?, ?, ?, ?, NULL, ?)", + ( + self.store_id, + sample_id, + generation, + names_json, + int(estimated_bytes), + self._LIVE, + time.time(), + ), + ) + self._conn.commit() + + def tombstone(self, sample_id: str, generation: int, reason: str) -> None: + with self._lock: + cur = self._conn.execute( + "UPDATE mooncake_objects SET state=?, reason=?, updated_at=? " + "WHERE store_id=? AND sample_id=? AND generation=?", + ( + self._TOMBSTONED, + reason, + time.time(), + self.store_id, + sample_id, + generation, + ), + ) + if cur.rowcount != 1: + raise KeyError( + f"no lifecycle inventory for {sample_id} generation {generation}" + ) + self._conn.commit() + + def mark_cleaned(self, sample_id: str, generation: int) -> None: + with self._lock: + cur = self._conn.execute( + "UPDATE mooncake_objects SET state=?, updated_at=? WHERE " + "store_id=? AND sample_id=? AND generation=?", + ( + self._CLEANED, + time.time(), + self.store_id, + sample_id, + generation, + ), + ) + if cur.rowcount != 1: + raise KeyError( + f"no lifecycle inventory for {sample_id} generation {generation}" + ) + self._conn.commit() + + def state(self, sample_id: str, generation: int) -> Optional[str]: + with self._lock: + row = self._conn.execute( + "SELECT state FROM mooncake_objects WHERE store_id=? AND " + "sample_id=? AND generation=?", + (self.store_id, sample_id, generation), + ).fetchone() + return row[0] if row is not None else None + + def record(self, sample_id: str, generation: int) -> Optional[LifecycleRecord]: + """Return one exact-generation inventory row, including cleaned state.""" + with self._lock: + row = self._conn.execute( + "SELECT feature_names_json, estimated_bytes, state FROM " + "mooncake_objects WHERE store_id=? AND sample_id=? AND generation=?", + (self.store_id, sample_id, generation), + ).fetchone() + if row is None: + return None + return LifecycleRecord( + sample_id=sample_id, + generation=generation, + feature_names=tuple(json.loads(row[0])), + estimated_bytes=int(row[1]), + state=row[2], + ) + + def pending(self) -> tuple[LifecycleRecord, ...]: + with self._lock: + rows = self._conn.execute( + "SELECT sample_id, generation, feature_names_json, " + "estimated_bytes, state FROM mooncake_objects WHERE store_id=? " + "AND state != ? ORDER BY rowid", + (self.store_id, self._CLEANED), + ).fetchall() + return tuple( + LifecycleRecord( + sample_id=row[0], + generation=int(row[1]), + feature_names=tuple(json.loads(row[2])), + estimated_bytes=int(row[3]), + state=row[4], + ) + for row in rows + ) + + def close(self) -> None: + with self._lock: + self._conn.close() + + +__all__ = ["LifecycleRecord", "SQLiteMooncakeLifecycleIndex"] diff --git a/specforge/runtime/data_plane/mooncake_store.py b/specforge/runtime/data_plane/mooncake_store.py index 5ef4e75d3..f7276cd14 100644 --- a/specforge/runtime/data_plane/mooncake_store.py +++ b/specforge/runtime/data_plane/mooncake_store.py @@ -6,7 +6,7 @@ # You may obtain a copy of the License at # # http://www.apache.org/licenses/LICENSE-2.0 -"""Mooncake-backed tensor store for the disaggregated runtime. +"""Mooncake-backed FeatureStore: the M6 *fast path* for disaggregated EAGLE3. ``SharedDirFeatureStore`` (``disaggregated.py``) locked down the disaggregation *contract* over a shared POSIX directory. ``MooncakeFeatureStore`` swaps that @@ -17,11 +17,13 @@ another, peer-to-peer, with no shared filesystem. That is what makes this a genuine network object store rather than a shared mount. -The wire contract is intentionally singular: every tensor is transferred as a -raw buffer with ``put_from``/``get_into``. Shape and dtype travel in the -metadata-only :class:`SampleRef`; no serialized tensor blob is accepted or -produced. Construction fails immediately when the installed Mooncake client -does not expose that API. +Lifetime roles are explicit. The producer-side owner keeps the generation index +and calls ``reclaim(ref)`` after the fan-out controller advances its global ACK +prefix. Independent readers use ``lifetime_owner=False``: their releases drop +only local lease bookkeeping and never remove the shared object. Offline +re-iterable owners continue to use ``retain_on_release``. A durable lifecycle +index records planned, resident, tombstoned, and cleaned generations so owner +restart can finish cleanup without trusting process-local bookkeeping. Contract carried from the reference backend: @@ -41,14 +43,16 @@ calls :meth:`drain_pending_removals`, a separate bounded retry that raises if physical removal never succeeds; failed hard-pinned objects are never silently dropped from bookkeeping. +In fan-out mode, trainer ``release(handle)`` is local-only and the owner calls +``reclaim(ref)`` after every subscriber acknowledges it. Required fan-out +reclaims remain tracked and raise when retries are exhausted. Concurrency: ``release``/``abort``/``gc`` hold ``self._lock`` across the ``remove()`` RPC. The lock is what makes consume-once free race-free against a concurrent ``get()`` (it prevents a re-lease between "decide to free" and the remote delete), exactly as ``SharedDirFeatureStore`` holds its lock across -``os.remove``. For the offline single-consumer path this is fine; a high-fanout -online deployment that wants the ``remove()`` RPC off the critical section needs -a tombstone-then-free protocol — a follow-up tied to the shared metadata index. +``os.remove``. Fan-out performs this RPC only once on the owner, after all readers +have released and acknowledged the ordered prefix. """ from __future__ import annotations @@ -58,12 +62,17 @@ import time import uuid from typing import Any, Callable, Dict, List, Optional, Tuple +from urllib.parse import urlparse import torch from specforge.runtime.contracts import SCHEMA_VERSION, FeatureHandle, SampleRef from specforge.runtime.data_plane.disaggregated import AuthPolicy from specforge.runtime.data_plane.feature_store import FeatureStore, spec_from_tensor +from specforge.runtime.data_plane.mooncake_lifecycle import ( + LifecycleRecord, + SQLiteMooncakeLifecycleIndex, +) logger = logging.getLogger(__name__) @@ -75,6 +84,8 @@ "rdma_devices": "", } +_MOONCAKE_MISSING_OBJECT = -704 + class _InjectedReplicateConfig: """Minimal config object for an explicitly injected test backend.""" @@ -129,7 +140,7 @@ def _require_store_api(store: Any) -> None: ) -# FeatureSpec.dtype (a string) -> torch dtype, for allocating receive +# FeatureSpec.dtype (a string) -> torch dtype, for allocating zero-copy receive # tensors from the ref alone (the ref carries shape+dtype, so get() needs no # serialized header). _TORCH_DTYPES = { @@ -146,12 +157,14 @@ def _require_store_api(store: Any) -> None: } -def _alloc_from_spec(spec) -> torch.Tensor: - """Allocate a fresh contiguous receive tensor matching a FeatureSpec.""" +def _alloc_from_spec(spec, *, pin_memory: bool = False) -> torch.Tensor: + """A fresh contiguous tensor matching a FeatureSpec (the zero-copy dst).""" dtype = _TORCH_DTYPES.get(spec.dtype) if dtype is None: - raise KeyError(f"unsupported feature dtype {spec.dtype!r} for Mooncake get") - return torch.empty(tuple(int(d) for d in spec.shape), dtype=dtype) + raise KeyError(f"unsupported feature dtype {spec.dtype!r} for zero-copy get") + return torch.empty( + tuple(int(d) for d in spec.shape), dtype=dtype, pin_memory=pin_memory + ) def _nbytes(t: torch.Tensor) -> int: @@ -177,6 +190,11 @@ class MooncakeFeatureStore(FeatureStore): during construction rather than selected as a different transport. """ + @property + def get_returns_fresh_tensors(self) -> bool: + # Raw reads allocate fresh storage from FeatureSpec. + return True + def __init__( self, *, @@ -185,9 +203,11 @@ def __init__( setup_kwargs: Optional[Dict[str, Any]] = None, auth: Optional[AuthPolicy] = None, credential: Optional[str] = None, + lifecycle_db_path: Optional[str] = None, max_resident_bytes: Optional[int] = None, max_hold_age_s: Optional[float] = None, retain_on_release: bool = False, + lifetime_owner: bool = True, max_release_attempts: int = 3, replica_num: int = 1, hard_pin: bool = True, @@ -216,6 +236,9 @@ def __init__( # Offline re-iterable mode: release() must NOT free (multi-epoch); mirrors # SharedDirFeatureStore / LocalFeatureStore file:// no-op release. self.retain_on_release = retain_on_release + # Exactly one store instance owns remote object lifetime. The default is + # the legacy consume-once behavior; fan-out trainers opt out explicitly. + self.lifetime_owner = lifetime_owner self.max_release_attempts = max_release_attempts self._clock = clock # in-process liveness index (single-host; see module docstring) @@ -223,9 +246,8 @@ def __init__( self._put_time: Dict[str, float] = {} self._sample_bytes: Dict[str, int] = {} # feature names per resident sample -> the per-tensor keys to remove on - # free. Cached on both put() (producer) and get() - # (consumer) so each side can free the sample it owns/consumed without the - # ref in hand at release() time. + # free (zero-copy mode). Owner instances cache these on put/adopt/get; + # non-owner readers intentionally do not retain a per-sample index. self._sample_names: Dict[str, List[str]] = {} # Server capture registers deterministic keys before issuing HTTP. If # the response is lost, no SampleRef exists to adopt/abort them. Keep a @@ -236,18 +258,111 @@ def __init__( # Samples whose remote remove() failed. gc() performs bounded # steady-state retries; lifecycle drain either removes them or raises. self._release_pending: Dict[str, int] = {} - # (sample_id, generation) logically freed in THIS process. Mooncake's - # remove() is lease-deferred (an object keeps a short read-lease), so the - # bytes can linger after release/abort; this makes the B5 "no - # use-after-free" guarantee immediate — get() of a freed ref raises even - # while physical reclamation is still pending. Grows with consume-once - # frees within a run (empty in retain_on_release/offline mode); a durable - # shared index would own this in the online multi-node follow-up. - self._freed: set = set() + self._external_release_pending: Dict[Tuple[str, int], int] = {} + self._external_release_records: Dict[Tuple[str, int], LifecycleRecord] = {} + self._force_release_pending: set[str] = set() + self._external_force_release_pending: set[Tuple[str, int]] = set() + # Fan-out reclaim is correctness-critical: unlike legacy best-effort + # cleanup, these entries must remain visible if retries are exhausted. + self._required_reclaims: set = set() + # (sample_id, generation) logically freed in THIS process. Stores without + # a lifecycle index need this guard while Mooncake deletion is lease- + # deferred. Production owners use the durable lifecycle state instead, so + # this set remains empty rather than growing once per consumed sample. + self._freed: set[Tuple[str, int]] = set() self._lock = threading.RLock() self._counter = 0 self._gen_counter = 0 self._stats = {"force_freed": 0, "force_freed_bytes": 0} + self._lifecycle = ( + SQLiteMooncakeLifecycleIndex(lifecycle_db_path, store_id=self.store_id) + if lifecycle_db_path is not None + else None + ) + if self._lifecycle is not None and self.lifetime_owner: + with self._lock: + self._sync_lifecycle_pending_locked() + + def _sync_lifecycle_pending_locked(self) -> None: + """Import server/local planned writes that appeared after owner startup.""" + if self._lifecycle is None or not self.lifetime_owner: + return + records = self._lifecycle.pending() + pending_identities = { + (record.sample_id, record.generation) for record in records + } + for sample_id, generation in list(self._generation.items()): + if (sample_id, generation) in pending_identities: + continue + self._generation.pop(sample_id, None) + self._sample_names.pop(sample_id, None) + self._sample_bytes.pop(sample_id, None) + self._put_time.pop(sample_id, None) + self._release_pending.pop(sample_id, None) + self._force_release_pending.discard(sample_id) + self._required_reclaims.discard(sample_id) + latest: Dict[str, Any] = {} + for record in records: + prior = latest.get(record.sample_id) + if prior is None or record.generation > prior.generation: + latest[record.sample_id] = record + self._gen_counter = max(self._gen_counter, record.generation) + for sample_id, record in latest.items(): + current = self._generation.get(sample_id) + if current is not None and current > record.generation: + continue + self._generation[sample_id] = record.generation + self._sample_names[sample_id] = list(record.feature_names) + self._sample_bytes[sample_id] = record.estimated_bytes + self._put_time.setdefault(sample_id, self._clock()) + if record.state == "tombstoned": + self._release_pending.setdefault(sample_id, 0) + self._required_reclaims.add(sample_id) + for record in records: + if latest[record.sample_id].generation == record.generation: + continue + identity = (record.sample_id, record.generation) + if record.state != "tombstoned": + self._lifecycle.tombstone( + record.sample_id, + record.generation, + "superseded-generation", + ) + self._external_release_pending.setdefault(identity, 0) + + def _external_records_locked(self) -> Dict[Tuple[str, int], LifecycleRecord]: + records = dict(self._external_release_records) + if self._lifecycle is not None: + records.update( + { + (record.sample_id, record.generation): record + for record in self._lifecycle.pending() + } + ) + return records + + def _require_lifetime_owner(self, operation: str) -> None: + if not self.lifetime_owner: + raise PermissionError( + f"MooncakeFeatureStore {self.store_id} is a non-owner reader; " + f"{operation} is reserved for the lifetime owner" + ) + + def _validate_ref_namespace(self, sample_ref: SampleRef) -> None: + parsed = urlparse(sample_ref.feature_store_uri) + if ( + parsed.scheme != "mooncake" + or parsed.netloc != self.store_id + or parsed.path != f"/{sample_ref.sample_id}" + or parsed.params + or parsed.query + or parsed.fragment + ): + raise ValueError( + f"ref {sample_ref.sample_id} points to " + f"{sample_ref.feature_store_uri!r}, not Mooncake store " + f"{self.store_id!r}" + ) # -- keys -------------------------------------------------------------- def _tkey(self, sample_id: str, gen: int, name: str) -> str: @@ -263,12 +378,12 @@ def _store_exists(self, key: str) -> bool: def _store_put_tensor(self, key: str, t: torch.Tensor) -> None: """Zero-copy publish: DMA straight from the tensor's storage, hard-pinned. - ``t`` must be contiguous + CPU (caller stages it). The bytes are the raw - tensor buffer; shape/dtype travel on the ref's FeatureSpec, so get() - needs no header. The source is registered with the transfer engine for - the duration of the put -- RDMA transfers it by DMA and rejects an - unregistered address (AddressNotRegistered); TCP ignores the - registration. + ``t`` must be contiguous + CPU (caller stages it). No torch.save: the + bytes are the raw tensor buffer; shape/dtype travel on the ref's + FeatureSpec, so get() needs no header. The source is registered with the + transfer engine for the duration of the put -- RDMA transfers it by DMA + and rejects an unregistered address (AddressNotRegistered); TCP ignores + the registration. """ nb = _nbytes(t) try: @@ -307,20 +422,26 @@ def _store_get_tensor(self, key: str, out: torch.Tensor) -> None: raise KeyError(f"mooncake get_into failed (status {rc}) for {key}") # get_into returns the number of bytes read; a full read returns exactly # nb. A short read (0 <= rc < nb) would leave the tail of this freshly - # allocated buffer as uninitialized garbage. Reject it rather than hand - # the trainer silently-corrupt data (B5: never serve wrong bytes). + # torch.empty'd buffer as uninitialized garbage. Reject it rather than + # hand the trainer silently-corrupt data (B5: never serve wrong bytes). if int(rc) != nb: raise KeyError( f"mooncake get_into short read for {key}: got {rc} of {nb} bytes" ) - def _store_remove(self, key: str) -> bool: + def _store_remove(self, key: str, *, force: bool = False) -> bool: """Best-effort physical free. Returns True on confirmed removal.""" try: - rc = self._store.remove(key) + try: + rc = self._store.remove(key, force) + except TypeError: + # Older Mooncake bindings and injected test backends expose only + # remove(key). They do not support force removal, but ordinary + # cleanup must remain compatible with that API. + rc = self._store.remove(key) except Exception: # pragma: no cover - transient RPC failure return False - return rc is None or int(rc) == 0 + return rc is None or int(rc) in (0, _MOONCAKE_MISSING_OBJECT) # -- write ------------------------------------------------------------- def put( @@ -330,6 +451,7 @@ def put( sample_id: str, metadata: Dict[str, Any], ) -> SampleRef: + self._require_lifetime_owner("put") self.auth.check(self._credential) if not tensors: raise ValueError("put requires at least one tensor") @@ -337,6 +459,11 @@ def put( specs = {k: spec_from_tensor(k, v) for k, v in staged.items()} nbytes = sum(_nbytes(t) for t in staged.values()) with self._lock: + if sample_id in self._release_pending: + raise RuntimeError( + f"cannot put {sample_id}: its prior generation is still " + "pending remote removal; run gc() first" + ) if ( self.max_resident_bytes is not None and sum(self._sample_bytes.values()) + nbytes > self.max_resident_bytes @@ -349,17 +476,51 @@ def put( gen = self._gen_counter prior_gen = self._generation.get(sample_id) prior_names = self._sample_names.get(sample_id, []) + if self._lifecycle is not None: + self._lifecycle.record_planned(sample_id, gen, staged, nbytes) # One hard-pinned object per tensor, DMA'd straight from its storage. # staged keeps the source tensors alive across the synchronous puts. - for name, t in staged.items(): - self._store_put_tensor(self._tkey(sample_id, gen, name), t) + attempted: List[str] = [] + try: + for name, t in staged.items(): + attempted.append(name) + self._store_put_tensor(self._tkey(sample_id, gen, name), t) + except BaseException as error: + leaked = [ + name + for name in attempted + if not self._store_remove(self._tkey(sample_id, gen, name), force=False) + ] + identity = (sample_id, gen) + if self._lifecycle is not None: + self._lifecycle.tombstone(sample_id, gen, "partial-put-failure") + if not leaked: + self._lifecycle.mark_cleaned(sample_id, gen) + if leaked: + record = LifecycleRecord( + sample_id=sample_id, + generation=gen, + feature_names=tuple(leaked), + estimated_bytes=sum(_nbytes(staged[name]) for name in leaked), + state="tombstoned", + ) + with self._lock: + self._external_release_records[identity] = record + self._external_release_pending.setdefault(identity, 0) + error.add_note( + "partial Mooncake put cleanup is pending for " + f"{sample_id!r} generation {gen}: {leaked}" + ) + raise # Overwrite-safe: drop the prior generation's tensor keys so a stale # ref's keys are gone (its get() then raises -> no use-after-free). if prior_gen is not None and prior_gen != gen: leaked = [ name for name in prior_names - if not self._store_remove(self._tkey(sample_id, prior_gen, name)) + if not self._store_remove( + self._tkey(sample_id, prior_gen, name), force=False + ) ] if leaked: logger.warning( @@ -371,12 +532,9 @@ def put( prior_gen, leaked, ) - with self._lock: - self._generation[sample_id] = gen - self._put_time[sample_id] = self._clock() - self._sample_bytes[sample_id] = nbytes - self._sample_names[sample_id] = list(staged) - return SampleRef( + elif self._lifecycle is not None: + self._lifecycle.mark_cleaned(sample_id, prior_gen) + ref = SampleRef( sample_id=sample_id, run_id=str(metadata.get("run_id", "unknown")), source_task_id=metadata.get("source_task_id"), @@ -395,6 +553,14 @@ def put( "generation": gen, # travels with the ref for the staleness guard }, ) + if self._lifecycle is not None: + self._lifecycle.record_resident(sample_id, gen, staged, nbytes) + with self._lock: + self._generation[sample_id] = gen + self._put_time[sample_id] = self._clock() + self._sample_bytes[sample_id] = nbytes + self._sample_names[sample_id] = list(staged) + return ref def adopt(self, sample_ref: SampleRef) -> None: """Register an externally-produced sample for lifecycle management. @@ -405,22 +571,78 @@ def adopt(self, sample_ref: SampleRef) -> None: ref's generation / feature names / size so ``release``/``abort``/``gc`` can free the server-written objects exactly like locally-put ones. """ + self._require_lifetime_owner("adopt") + self._validate_ref_namespace(sample_ref) gen = sample_ref.metadata.get("generation") if gen is None: raise ValueError( f"cannot adopt {sample_ref.sample_id}: ref carries no generation" ) - with self._lock: - gen = int(gen) - self._generation[sample_ref.sample_id] = gen - self._sample_names[sample_ref.sample_id] = list( - sample_ref.feature_keys.keys() + gen = int(gen) + sample_bytes = int(sample_ref.estimated_bytes or 0) + feature_names = list(sample_ref.feature_keys.keys()) + if self._lifecycle is not None and sample_bytes == 0: + inventory = self._lifecycle.record(sample_ref.sample_id, gen) + if inventory is not None: + sample_bytes = inventory.estimated_bytes + if sample_bytes == 0 and feature_names: + sizes: List[int] = [] + for name in feature_names: + try: + size = int( + self._store.get_size( + self._tkey(sample_ref.sample_id, gen, name) + ) + ) + except (AttributeError, TypeError, ValueError): + sizes = [] + break + if size < 0: + sizes = [] + break + sizes.append(size) + if sizes: + sample_bytes = sum(sizes) + if sample_bytes == 0 and self.max_resident_bytes is not None: + raise ValueError( + f"cannot adopt {sample_ref.sample_id}: payload size is unknown " + "while max_resident_bytes is enforced" ) - self._sample_bytes[sample_ref.sample_id] = int( - sample_ref.estimated_bytes or 0 + with self._lock: + if sample_ref.sample_id in self._release_pending: + raise RuntimeError( + f"cannot adopt {sample_ref.sample_id}: its prior generation " + "is still pending remote removal; run gc() first" + ) + if self._lifecycle is not None: + self._lifecycle.record_resident( + sample_ref.sample_id, gen, feature_names, sample_bytes ) + with self._lock: + previous_bytes = self._sample_bytes.get(sample_ref.sample_id, 0) + projected = sum(self._sample_bytes.values()) - previous_bytes + sample_bytes + self._generation[sample_ref.sample_id] = gen + self._sample_names[sample_ref.sample_id] = feature_names + self._sample_bytes[sample_ref.sample_id] = sample_bytes self._put_time[sample_ref.sample_id] = self._clock() self._external_provisional.pop((sample_ref.sample_id, gen), None) + over_budget = ( + self.max_resident_bytes is not None + and projected > self.max_resident_bytes + ) + if over_budget: + self._abort_locked( + sample_ref.sample_id, + reason="adopt-over-budget", + required_reclaim=True, + force=True, + ) + if over_budget: + raise MemoryError( + f"MooncakeFeatureStore {self.store_id} adopt would exceed " + f"{self.max_resident_bytes} bytes; server-written sample was " + "scheduled for required cleanup" + ) def track_external_attempt( self, @@ -430,6 +652,7 @@ def track_external_attempt( feature_names: List[str], ) -> None: """Track server-owned keys before an HTTP response makes a ref adoptable.""" + self._require_lifetime_owner("track external capture") names = list(dict.fromkeys(str(name) for name in feature_names)) if not names: raise ValueError("external capture attempt must name at least one feature") @@ -447,6 +670,7 @@ def discard_external_attempts( shutdown path. Physical-remove failures remain visible to :meth:`drain_pending_removals`. """ + self._require_lifetime_owner("discard external captures") with self._lock: attempts = list(self._external_provisional.items()) @@ -502,10 +726,21 @@ def get( *, device: "torch.device | str" = "cpu", names: Optional[List[str]] = None, + pin_memory: bool = False, ) -> Tuple[Dict[str, torch.Tensor], FeatureHandle]: self.auth.check(self._credential) + self._validate_ref_namespace(sample_ref) sid = sample_ref.sample_id ref_gen = sample_ref.metadata.get("generation") + if self._lifecycle is not None: + if ref_gen is None: + raise KeyError(f"sample {sid} ref carries no lifecycle generation") + state = self._lifecycle.state(sid, int(ref_gen)) + if state != "resident": + raise KeyError( + f"sample {sid} generation {ref_gen} lifecycle state is " + f"{state!r}; refusing read" + ) with self._lock: if ref_gen is not None and (sid, int(ref_gen)) in self._freed: # logically freed here; the remote bytes may still linger under @@ -515,17 +750,24 @@ def get( f"refusing use-after-free" ) wanted = names or list(sample_ref.feature_keys.keys()) - out, gen = self._get_tensors(sample_ref, wanted) + out, gen = self._get_tensors(sample_ref, wanted, pin_memory=pin_memory) if str(device) != "cpu": out = {k: v.to(device) for k, v in out.items()} + if self._lifecycle is not None: + state = self._lifecycle.state(sid, int(ref_gen)) + if state != "resident": + raise KeyError( + f"sample {sid} generation {ref_gen} was tombstoned while " + "materializing; refusing read" + ) with self._lock: self._counter += 1 - # Consumer-side cache: a process that only get()s a sample (never - # put() it) still needs gen + feature names so its release()/abort() - # can free the per-tensor keys. setdefault keeps the producer's own - # entries authoritative when producer and consumer are one instance. - self._generation.setdefault(sid, gen) - self._sample_names.setdefault(sid, list(sample_ref.feature_keys.keys())) + if self.lifetime_owner: + # A legacy consume-once instance that only get()s still needs + # generation + names so release() can remove the remote keys. + # Non-owner readers intentionally retain no per-sample index. + self._generation.setdefault(sid, gen) + self._sample_names.setdefault(sid, list(sample_ref.feature_keys.keys())) handle = FeatureHandle( sample_id=sid, generation=gen, @@ -535,7 +777,7 @@ def get( return out, handle def _get_tensors( - self, ref: SampleRef, wanted: List[str] + self, ref: SampleRef, wanted: List[str], *, pin_memory: bool = False ) -> Tuple[Dict[str, torch.Tensor], int]: """Read each feature straight into a spec-allocated tensor.""" sid = ref.sample_id @@ -555,7 +797,8 @@ def _get_tensors( f"sample {sid} gen {gen} feature {n!r} not available " f"(freed, stale, or never written)" ) - out[n] = _alloc_from_spec(spec) # fresh -> clone-on-fetch for free (B5) + # Fresh storage makes the loader's clone-on-fetch redundant (B5). + out[n] = _alloc_from_spec(spec, pin_memory=pin_memory) self._store_get_tensor(key, out[n]) return out, gen @@ -564,42 +807,81 @@ def _try_physical_free( self, sample_id: str, *, + force: bool = False, confirm_absent_on_failure: bool = True, ) -> bool: - """Remove all tensor objects. False on a retryable RPC failure. - - Order matters against Mooncake's lease semantics: an is_exist probe - GRANTS a read lease, and a remove during any live lease fails (-706). - So each key is removed FIRST. The optional exist probe runs only after a - failed remove, purely to classify "already gone" as freed. Retry loops - disable that probe because probing a still-live key would renew its - lease and make every following remove fail again. + """Remove the remote object(s). False on a (retryable) RPC failure. + + One object backs each tensor, so remove every per-tensor key of the + sample's current generation. + + ``force`` is reserved for callers with an external proof that no reader + may still use the generation, such as the fan-out global durable ACK. + Mooncake reports an already-missing key as -704, so no ``is_exist`` + probe is needed; that probe would grant a fresh read lease. + A failed remove may be classified as already absent, but retry loops + disable that probe because ``is_exist`` grants a fresh read lease. """ + self._require_lifetime_owner("remote removal") gen = self._generation.get(sample_id) if gen is None: return True # nothing tracked to remove (already freed) ok = True for name in self._sample_names.get(sample_id, []): key = self._tkey(sample_id, gen, name) - if self._store_remove(key): + if self._store_remove(key, force=force): continue if confirm_absent_on_failure and not self._store_exists(key): - continue # already gone (freed remotely) counts as freed + continue + ok = False + return ok + + def _try_physical_free_record(self, record, *, force: bool = False) -> bool: + """Remove one exact lifecycle generation not represented by current maps.""" + ok = True + for name in record.feature_names: + key = self._tkey(record.sample_id, record.generation, name) + if self._store_remove(key, force=force): + continue ok = False return ok + def _sample_exists(self, sample_id: str) -> bool: + """True if any object backing the sample's current generation is present.""" + gen = self._generation.get(sample_id) + if gen is None: + return False + return any( + self._store_exists(self._tkey(sample_id, gen, n)) + for n in self._sample_names.get(sample_id, []) + ) + def _free_bookkeeping_locked(self, sample_id: str) -> int: """Drop in-process tracking for a sample. Returns bytes accounted freed.""" generation = self._generation.get(sample_id) + if self._lifecycle is not None and generation is not None: + self._lifecycle.mark_cleaned(sample_id, generation) nbytes = self._sample_bytes.pop(sample_id, 0) self._generation.pop(sample_id, None) self._put_time.pop(sample_id, None) self._sample_names.pop(sample_id, None) self._release_pending.pop(sample_id, None) + self._force_release_pending.discard(sample_id) + self._required_reclaims.discard(sample_id) if generation is not None: self._external_provisional.pop((sample_id, generation), None) return nbytes + def _tombstone_locked(self, sample_id: str, reason: str) -> Optional[int]: + generation = self._generation.get(sample_id) + if generation is None: + return None + if self._lifecycle is not None: + self._lifecycle.tombstone(sample_id, generation, reason) + else: + self._freed.add((sample_id, generation)) + return generation + def _still_leased_locked(self, sample_id: str, generation: Optional[int]) -> bool: # generation-aware: a stale older-generation lease does not pin the # current generation (matches LocalFeatureStore's invariant). @@ -609,32 +891,139 @@ def _still_leased_locked(self, sample_id: str, generation: Optional[int]) -> boo ) def release(self, handle: FeatureHandle, *, reason: str = "consumed") -> None: + """End a materialization lease. + + On a non-owner reader this is strictly local: it drops the active handle + and never removes shared Mooncake data. The fan-out coordinator must use + :meth:`reclaim` after the global minimum acknowledgement advances. + """ with self._lock: self._active_leases.pop(handle.lease_token, None) - if self.retain_on_release: - return # offline re-iterable set: keep for the next epoch sid = handle.sample_id cur = self._generation.get(sid) + if not self.lifetime_owner: + if not self._still_leased_locked(sid, cur): + self._free_bookkeeping_locked(sid) + return + if self.retain_on_release: + return # offline re-iterable set: keep for the next epoch if cur is not None and handle.generation != cur: return # stale lease -> no-op if self._still_leased_locked(sid, cur): return - self._freed.add((sid, handle.generation)) # immediate logical free - if self._try_physical_free(sid): + self._tombstone_locked(sid, reason) + if self._try_physical_free(sid, force=False): self._free_bookkeeping_locked(sid) else: # remote free deferred (lease) / failed -> gc() retries self._release_pending.setdefault(sid, 0) - def abort(self, sample_id: str, *, reason: str = "aborted") -> None: + def _abort_locked( + self, + sample_id: str, + *, + reason: str, + required_reclaim: bool = False, + force: bool = False, + ) -> None: + self._tombstone_locked(sample_id, reason) + if self._try_physical_free(sample_id, force=force): + self._free_bookkeeping_locked(sample_id) + else: + self._release_pending.setdefault(sample_id, 0) + if force: + self._force_release_pending.add(sample_id) + if required_reclaim: + self._required_reclaims.add(sample_id) + + def reclaim( + self, sample_ref: SampleRef, *, reason: str = "globally-consumed" + ) -> None: + """Owner-only deletion of one exact, globally-consumed reference. + + ``release(handle)`` means one reader finished materializing tensors; + ``reclaim(ref)`` means *all* fan-out subscribers acknowledged the ref and + its remote objects may be deleted. Matching the ref generation before + removal prevents a delayed acknowledgement from deleting a newer sample + with the same ID. A lease/busy removal failure enters ``release_pending`` + and is retried by :meth:`gc`. + """ + self._require_lifetime_owner("reclaim") + self._validate_ref_namespace(sample_ref) + ref_gen = sample_ref.metadata.get("generation") + if ref_gen is None: + raise ValueError( + f"cannot reclaim {sample_ref.sample_id}: ref carries no generation" + ) + ref_gen = int(ref_gen) with self._lock: - gen = self._generation.get(sample_id) - if gen is not None: - self._freed.add((sample_id, gen)) # immediate logical free - if self._try_physical_free(sample_id): - self._free_bookkeeping_locked(sample_id) - else: - self._release_pending.setdefault(sample_id, 0) + current_gen = self._generation.get(sample_ref.sample_id) + if current_gen is None: + raise KeyError( + f"cannot reclaim untracked sample {sample_ref.sample_id}; " + "the lifetime owner must put() or adopt() the ref first" + ) + if current_gen != ref_gen: + raise KeyError( + f"refusing to reclaim stale sample {sample_ref.sample_id} " + f"generation {ref_gen}; current generation is {current_gen}" + ) + self._abort_locked( + sample_ref.sample_id, + reason=reason, + required_reclaim=True, + force=True, + ) + + def abort( + self, sample_id: str, *, reason: str = "aborted", force: bool = False + ) -> None: + """Owner-only terminal cleanup by sample ID. + + Error paths that already own/adopt a sample use this legacy API. Fan-out + acknowledgement cleanup must use :meth:`reclaim` so generation matching + is explicit. + """ + self._require_lifetime_owner("abort") + with self._lock: + self._abort_locked(sample_id, reason=reason, force=force) + + def abort_all(self, *, reason: str = "owner-failure", force: bool = False) -> int: + """Persist tombstones and attempt cleanup for every tracked sample.""" + self._require_lifetime_owner("abort_all") + with self._lock: + self._sync_lifecycle_pending_locked() + records = self._lifecycle.pending() if self._lifecycle is not None else () + handled = { + (sample_id, generation) + for sample_id, generation in self._generation.items() + } + sample_ids = list(self._generation) + for sample_id in sample_ids: + self._abort_locked( + sample_id, + reason=reason, + required_reclaim=True, + force=force, + ) + for record in records: + identity = (record.sample_id, record.generation) + if identity in handled: + continue + if record.state != "tombstoned": + self._lifecycle.tombstone( + record.sample_id, record.generation, reason + ) + if self._try_physical_free_record(record, force=force): + self._lifecycle.mark_cleaned(record.sample_id, record.generation) + self._external_release_pending.pop(identity, None) + self._external_force_release_pending.discard(identity) + self._external_release_records.pop(identity, None) + else: + self._external_release_pending.setdefault(identity, 0) + if force: + self._external_force_release_pending.add(identity) + return len(records) if self._lifecycle is not None else len(sample_ids) def drain_pending_removals( self, @@ -662,8 +1051,10 @@ def drain_pending_removals( for attempt in range(max_attempts): attempts_run = attempt + 1 with self._lock: + self._sync_lifecycle_pending_locked() pending = list(self._release_pending) - if not pending: + external_pending = list(self._external_release_pending) + if not pending and not external_pending: return { "removed": removed, "removed_bytes": removed_bytes, @@ -696,8 +1087,45 @@ def drain_pending_removals( self.max_release_attempts, self._release_pending.get(sample_id, 0) + 1, ) - remaining = list(self._release_pending) - if not remaining: + records = self._external_records_locked() + for identity in external_pending: + record = records.get(identity) + if record is None: + self._external_release_pending.pop(identity, None) + self._external_force_release_pending.discard(identity) + self._external_release_records.pop(identity, None) + continue + label = f"{record.sample_id}:g{record.generation}" + try: + physically_removed = self._try_physical_free_record( + record, + force=identity in self._external_force_release_pending, + ) + except Exception as exc: + last_errors[label] = f"{type(exc).__name__}: {exc}" + physically_removed = False + if physically_removed: + if self._lifecycle is not None: + lifecycle_record = self._lifecycle.record(*identity) + if lifecycle_record is not None: + self._lifecycle.mark_cleaned(*identity) + self._external_release_pending.pop(identity, None) + self._external_force_release_pending.discard(identity) + self._external_release_records.pop(identity, None) + removed += 1 + removed_bytes += record.estimated_bytes + self._stats["force_freed"] += 1 + self._stats["force_freed_bytes"] += record.estimated_bytes + last_errors.pop(label, None) + else: + self._external_release_pending[identity] = min( + self.max_release_attempts, + self._external_release_pending.get(identity, 0) + 1, + ) + remaining_count = len(self._release_pending) + len( + self._external_release_pending + ) + if not remaining_count: return { "removed": removed, "removed_bytes": removed_bytes, @@ -709,6 +1137,10 @@ def drain_pending_removals( with self._lock: remaining = list(self._release_pending) + remaining.extend( + f"{sample_id}:g{generation}" + for sample_id, generation in self._external_release_pending + ) preview = remaining[:16] detail = f"; last errors={last_errors}" if last_errors else "" raise RuntimeError( @@ -718,9 +1150,16 @@ def drain_pending_removals( ) def gc(self, *, now: Optional[float] = None) -> Dict[str, int]: + if not self.lifetime_owner: + return { + "force_freed": 0, + "force_freed_bytes": 0, + "release_pending": 0, + } now = self._clock() if now is None else now freed = freed_bytes = 0 with self._lock: + self._sync_lifecycle_pending_locked() # max-hold sweep: force-free abandoned samples (spare still-leased) if self.max_hold_age_s is not None: stale = [ @@ -730,13 +1169,16 @@ def gc(self, *, now: Optional[float] = None) -> Dict[str, int]: and not self._still_leased_locked(sid, self._generation.get(sid)) ] for sid in stale: - if self._try_physical_free(sid, confirm_absent_on_failure=False): + self._tombstone_locked(sid, "max-hold-age") + if self._try_physical_free( + sid, force=False, confirm_absent_on_failure=False + ): freed_bytes += self._free_bookkeeping_locked(sid) freed += 1 else: self._release_pending.setdefault(sid, 0) - # Reconcile release-pending without an exists probe: is_exist grants - # a read lease that would make the next remove fail (-706). + # Reconcile release-pending without an existence probe: is_exist + # grants a read lease, while remove(-704) already reports absence. for sid in list(self._release_pending): if self._release_pending[sid] >= self.max_release_attempts: # Keep the physical key metadata and surface the pending @@ -745,36 +1187,91 @@ def gc(self, *, now: Optional[float] = None) -> Dict[str, int]: # make a hard-pinned remote leak invisible. continue attempts = self._release_pending[sid] + 1 - if self._try_physical_free(sid, confirm_absent_on_failure=False): + if self._try_physical_free( + sid, + force=sid in self._force_release_pending, + confirm_absent_on_failure=False, + ): freed_bytes += self._free_bookkeeping_locked(sid) freed += 1 else: self._release_pending[sid] = attempts + records = self._external_records_locked() + for identity in list(self._external_release_pending): + record = records.get(identity) + if record is None: + self._external_release_pending.pop(identity, None) + self._external_force_release_pending.discard(identity) + self._external_release_records.pop(identity, None) + continue + if ( + self._external_release_pending[identity] + >= self.max_release_attempts + ): + continue + attempts = self._external_release_pending[identity] + 1 + if self._try_physical_free_record( + record, + force=identity in self._external_force_release_pending, + ): + if self._lifecycle is not None: + self._lifecycle.mark_cleaned(*identity) + self._external_release_pending.pop(identity, None) + self._external_force_release_pending.discard(identity) + self._external_release_records.pop(identity, None) + freed += 1 + freed_bytes += record.estimated_bytes + else: + self._external_release_pending[identity] = min( + self.max_release_attempts, attempts + ) self._stats["force_freed"] += freed self._stats["force_freed_bytes"] += freed_bytes - return { - "force_freed": freed, - "force_freed_bytes": freed_bytes, - "release_pending": len(self._release_pending), - } + report = { + "force_freed": freed, + "force_freed_bytes": freed_bytes, + "release_pending": len(self._release_pending) + + len(self._external_release_pending), + } + return report def health(self) -> Dict[str, Any]: with self._lock: + self._sync_lifecycle_pending_locked() now = self._clock() ages = [now - t for t in self._put_time.values()] + lifecycle_pending = ( + self._lifecycle.pending() + if self._lifecycle is not None and self.lifetime_owner + else () + ) # NOTE: resident_bytes is an in-process accounting sum, not a live # Mooncake pool-usage query (the Python API exposes only per-key # get_size). A cross-node pool-usage signal is a follow-up. return { "store_id": self.store_id, "backend": "mooncake", - "resident_samples": len(self._generation), + "lifetime_owner": self.lifetime_owner, + "resident_samples": ( + len(lifecycle_pending) + if self._lifecycle is not None and self.lifetime_owner + else len(self._generation) + ), "provisional_external": len(self._external_provisional), "active_leases": len(self._active_leases), - "resident_bytes": sum(self._sample_bytes.values()), + "resident_bytes": ( + sum(record.estimated_bytes for record in lifecycle_pending) + if self._lifecycle is not None and self.lifetime_owner + else sum(self._sample_bytes.values()) + ), "max_resident_bytes": self.max_resident_bytes, + "retain_on_release": self.retain_on_release, "auth_required": self.auth.required, - "release_pending": len(self._release_pending), + "release_pending": len(self._release_pending) + + len(self._external_release_pending), + "required_reclaims_pending": len(self._required_reclaims) + + len(self._external_release_pending), + "local_tombstones": len(self._freed), "oldest_age_s": max(ages) if ages else 0.0, "avg_age_s": (sum(ages) / len(ages)) if ages else 0.0, "force_freed_total": self._stats["force_freed"], diff --git a/specforge/runtime/data_plane/windowed_capture.py b/specforge/runtime/data_plane/windowed_capture.py new file mode 100644 index 000000000..33aa2c9de --- /dev/null +++ b/specforge/runtime/data_plane/windowed_capture.py @@ -0,0 +1,46 @@ +# coding=utf-8 +# Copyright 2024 The SpecForge team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +"""Public facade for transactional, consumer-driven capture windows. + +Payloads remain in ``FeatureStore``. The implementation is split by ownership: +serializable contracts, SQLite state transitions, and the consumer queue facade. +Imports from this original module remain stable for callers. +""" + +from specforge.runtime.data_plane.windowed_capture_contracts import ( + AcquireTicket, + CaptureFailedError, + CaptureKey, + CapturePriority, + CaptureReadLease, + CaptureRequest, + CaptureState, + ConsumerFailedError, + EvictionCandidate, + capture_contract_digest, +) +from specforge.runtime.data_plane.windowed_capture_queue import WindowedCaptureQueue +from specforge.runtime.data_plane.windowed_capture_registry import ( + SQLiteWindowedCaptureRegistry, +) + +__all__ = [ + "AcquireTicket", + "CaptureFailedError", + "CaptureKey", + "CapturePriority", + "CaptureReadLease", + "CaptureRequest", + "CaptureState", + "ConsumerFailedError", + "EvictionCandidate", + "SQLiteWindowedCaptureRegistry", + "WindowedCaptureQueue", + "capture_contract_digest", +] diff --git a/specforge/runtime/data_plane/windowed_capture_contracts.py b/specforge/runtime/data_plane/windowed_capture_contracts.py new file mode 100644 index 000000000..a6635fe40 --- /dev/null +++ b/specforge/runtime/data_plane/windowed_capture_contracts.py @@ -0,0 +1,141 @@ +# coding=utf-8 +# Copyright 2024 The SpecForge team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +"""Serializable identities and state shared by windowed capture components.""" + +from __future__ import annotations + +import dataclasses +import hashlib +import json +from dataclasses import dataclass +from enum import Enum, IntEnum +from typing import Any, Mapping + +from specforge.runtime.contracts import SampleRef + + +class CaptureState(str, Enum): + ABSENT = "absent" + QUEUED = "queued" + CAPTURING = "capturing" + COMMITTING = "committing" + READY = "ready" + EVICTING = "evicting" + FAILED = "failed" + + +class CapturePriority(IntEnum): + DEMAND = 0 + PREFETCH = 1 + + +@dataclass(frozen=True, order=True) +class CaptureKey: + source_sample_id: str + contract_digest: str + + def __post_init__(self) -> None: + if not isinstance(self.source_sample_id, str) or not self.source_sample_id: + raise ValueError("source_sample_id must be a non-empty string") + digest = self.contract_digest + if ( + not isinstance(digest, str) + or len(digest) != 64 + or any(char not in "0123456789abcdef" for char in digest) + ): + raise ValueError("contract_digest must be a lowercase SHA-256 digest") + + +@dataclass(frozen=True) +class AcquireTicket: + token: str + consumer_id: str + key: CaptureKey + source_index: int + ready_at_request: bool + + +@dataclass(frozen=True) +class CaptureReadLease: + token: str + consumer_id: str + key: CaptureKey + source_index: int + generation: int + ref: SampleRef + ready_at_request: bool + wait_s: float + + +@dataclass(frozen=True) +class CaptureRequest: + key: CaptureKey + source_index: int + generation: int + priority: CapturePriority + demand_consumers: tuple[str, ...] + reserved_bytes: int + + +@dataclass(frozen=True) +class EvictionCandidate: + key: CaptureKey + source_index: int + generation: int + ref: SampleRef + + +class CaptureFailedError(RuntimeError): + """A requested capture reached a terminal failure.""" + + +class ConsumerFailedError(RuntimeError): + """A consumer was failed or expired while waiting for a capture.""" + + +def _jsonable_contract(value: Any) -> Any: + if dataclasses.is_dataclass(value): + value = dataclasses.asdict(value) + if isinstance(value, Mapping): + return { + str(key): _jsonable_contract(item) + for key, item in sorted(value.items(), key=lambda pair: str(pair[0])) + } + if isinstance(value, (set, frozenset)): + return sorted(_jsonable_contract(item) for item in value) + if isinstance(value, (tuple, list)): + return [_jsonable_contract(item) for item in value] + if value is None or isinstance(value, (str, int, float, bool)): + return value + raise TypeError(f"capture contract contains unsupported {type(value).__name__}") + + +def capture_contract_digest(contract: Any) -> str: + """Return a canonical digest separating incompatible capture payloads.""" + encoded = json.dumps( + _jsonable_contract(contract), + sort_keys=True, + separators=(",", ":"), + ensure_ascii=True, + ).encode("ascii") + return hashlib.sha256(encoded).hexdigest() + + +__all__ = [ + "AcquireTicket", + "CaptureFailedError", + "CaptureKey", + "CapturePriority", + "CaptureReadLease", + "CaptureRequest", + "CaptureState", + "ConsumerFailedError", + "EvictionCandidate", + "capture_contract_digest", +] diff --git a/specforge/runtime/data_plane/windowed_capture_queue.py b/specforge/runtime/data_plane/windowed_capture_queue.py new file mode 100644 index 000000000..a38acaaeb --- /dev/null +++ b/specforge/runtime/data_plane/windowed_capture_queue.py @@ -0,0 +1,235 @@ +# coding=utf-8 +# Copyright 2024 The SpecForge team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +"""Ordered queue facade over one consumer's capture-window cursor.""" + +from __future__ import annotations + +import threading +from typing import Callable, Optional, Sequence + +from specforge.runtime.contracts import SampleRef +from specforge.runtime.data_plane.windowed_capture_contracts import CaptureReadLease +from specforge.runtime.data_plane.windowed_capture_registry import ( + SQLiteWindowedCaptureRegistry, +) + + +class WindowedCaptureQueue: + """Ordered ``SampleRefQueue`` facade for one logical consumer.""" + + def __init__( + self, + registry: SQLiteWindowedCaptureRegistry, + consumer_id: str, + *, + idle_timeout_s: Optional[float] = 1800.0, + record_refs: Optional[Callable[[Sequence[SampleRef]], None]] = None, + ) -> None: + if idle_timeout_s is not None and idle_timeout_s <= 0: + raise ValueError("idle_timeout_s must be > 0 or None") + snapshot = registry.snapshot() + if consumer_id not in snapshot["consumers"]: + raise KeyError(f"unknown consumer {consumer_id!r}") + self.registry = registry + self.consumer_id = consumer_id + self.total_samples = int(snapshot["total_samples"]) + self.idle_timeout_s = idle_timeout_s + if record_refs is not None and not callable(record_refs): + raise TypeError("record_refs must be callable or None") + self._record_refs = record_refs + self._next_fetch = int(snapshot["consumers"][consumer_id]["cursor"]) + self._leases: dict[str, CaptureReadLease] = {} + self._closed = False + self._get_lock = threading.Lock() + self._state_lock = threading.RLock() + self._metrics: dict[str, float | int] = { + "refs": 0, + "ready_at_request_refs": 0, + "demand_wait_s": 0.0, + "max_demand_wait_s": 0.0, + } + + def get(self, n: int, timeout_s: float = 0.0) -> list[SampleRef]: + del timeout_s # Registry waits use the explicit idle timeout. + if isinstance(n, bool) or not isinstance(n, int) or n < 1: + raise ValueError("n must be a positive integer") + with self._get_lock: + while True: + with self._state_lock: + if self._closed: + return [] + start = self._next_fetch + if start >= self.total_samples: + if not self._leases: + self.registry.complete_consumer(self.consumer_id) + self._closed = True + return [] + stop = min(self.total_samples, start + n) + tickets = self.registry.request_many( + self.consumer_id, range(start, stop) + ) + acquired: list[CaptureReadLease] = [] + try: + for ticket in tickets: + acquired.append( + self.registry.wait_ready( + ticket, timeout_s=self.idle_timeout_s + ) + ) + except BaseException: + acquired_indices = {lease.source_index for lease in acquired} + for ticket in tickets: + if ticket.source_index not in acquired_indices: + try: + self.registry.cancel_acquire(ticket) + except RuntimeError: + pass + self.registry.abandon_leases(self.consumer_id, acquired) + raise + refs = [lease.ref for lease in acquired] + if len({ref.sample_id for ref in refs}) != len(refs): + self.registry.abandon_leases(self.consumer_id, acquired) + raise RuntimeError( + "windowed capture batch contains duplicate sample IDs" + ) + if self._record_refs is not None: + try: + self._record_refs(refs) + except BaseException: + self.registry.abandon_leases(self.consumer_id, acquired) + raise + with self._state_lock: + if self._closed: + self.registry.abandon_leases(self.consumer_id, acquired) + return [] + if self._next_fetch < start: + self.registry.abandon_leases(self.consumer_id, acquired) + continue + self._leases.update( + (lease.ref.sample_id, lease) for lease in acquired + ) + self._next_fetch = max(self._next_fetch, stop) + for lease in acquired: + self._metrics["refs"] += 1 + self._metrics["ready_at_request_refs"] += int( + lease.ready_at_request + ) + self._metrics["demand_wait_s"] += lease.wait_s + self._metrics["max_demand_wait_s"] = max( + self._metrics["max_demand_wait_s"], lease.wait_s + ) + return refs + + def _resolve_leases(self, refs: Sequence[SampleRef]) -> list[CaptureReadLease]: + missing = [ref.sample_id for ref in refs if ref.sample_id not in self._leases] + if missing: + raise RuntimeError( + f"windowed queue references samples not leased: {missing}" + ) + return [self._leases[ref.sample_id] for ref in refs] + + def ack(self, refs: list[SampleRef]) -> None: + self.ack_ids([ref.sample_id for ref in refs]) + + def ack_ids(self, sample_ids: list[str]) -> None: + """Acknowledge the exact leased prefix after its durable train ACK.""" + if not sample_ids: + return + with self._state_lock: + expected = list(self._leases)[: len(sample_ids)] + if expected != list(sample_ids): + raise RuntimeError( + "windowed queue acknowledgement is not the leased prefix: " + f"expected={expected}, got={list(sample_ids)}" + ) + leases = [self._leases[sample_id] for sample_id in sample_ids] + self.registry.release_and_advance(self.consumer_id, leases) + for sample_id in sample_ids: + self._leases.pop(sample_id) + + def fail(self, refs: list[SampleRef], reason: str, retryable: bool) -> None: + del reason + with self._state_lock: + leases = self._resolve_leases(refs) + self.registry.abandon_leases(self.consumer_id, leases) + for ref in refs: + self._leases.pop(ref.sample_id) + if retryable and leases: + self._next_fetch = min( + self._next_fetch, + min(lease.source_index for lease in leases), + ) + + def depth(self) -> int: + with self._state_lock: + return max(0, self.total_samples - self._next_fetch) + + def in_flight(self) -> int: + with self._state_lock: + return len(self._leases) + + def metrics(self) -> dict[str, float | int]: + with self._state_lock: + refs = int(self._metrics["refs"]) + return { + **self._metrics, + "ready_at_request_ratio": ( + float(self._metrics["ready_at_request_refs"]) / refs + if refs + else 0.0 + ), + "mean_demand_wait_s": ( + float(self._metrics["demand_wait_s"]) / refs if refs else 0.0 + ), + "next_fetch": self._next_fetch, + "in_flight": len(self._leases), + } + + def drained(self) -> bool: + with self._state_lock: + return self._next_fetch == self.total_samples and not self._leases + + def finalize(self) -> None: + self.complete() + + def complete(self, *, allow_partial: bool = False) -> None: + with self._state_lock: + if not allow_partial and ( + self._next_fetch != self.total_samples or self._leases + ): + raise RuntimeError( + f"consumer {self.consumer_id!r} queue did not drain: " + f"next={self._next_fetch}/{self.total_samples}, " + f"leases={len(self._leases)}" + ) + if allow_partial and self._leases: + self.registry.abandon_leases( + self.consumer_id, list(self._leases.values()) + ) + self._leases.clear() + self.registry.complete_consumer( + self.consumer_id, allow_partial=allow_partial + ) + self._closed = True + + def close(self, error: Optional[BaseException | str] = None) -> None: + with self._state_lock: + if self._closed: + return + if self._leases: + self.registry.abandon_leases( + self.consumer_id, list(self._leases.values()) + ) + self._leases.clear() + if error is not None: + self.registry.fail_consumer(self.consumer_id, error) + self._closed = True + + +__all__ = ["WindowedCaptureQueue"] diff --git a/specforge/runtime/data_plane/windowed_capture_registry.py b/specforge/runtime/data_plane/windowed_capture_registry.py new file mode 100644 index 000000000..0e110306d --- /dev/null +++ b/specforge/runtime/data_plane/windowed_capture_registry.py @@ -0,0 +1,1865 @@ +# coding=utf-8 +# Copyright 2024 The SpecForge team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +"""SQLite state transitions for bounded, consumer-driven capture windows.""" + +from __future__ import annotations + +import hashlib +import json +import os +import re +import sqlite3 +import threading +import time +import uuid +from contextlib import contextmanager +from typing import TYPE_CHECKING, Any, Callable, Iterable, Iterator, Optional, Sequence + +from specforge.runtime.contracts import SampleRef, assert_no_tensors +from specforge.runtime.data_plane.ref_serialization import ref_from_dict, ref_to_dict +from specforge.runtime.data_plane.windowed_capture_contracts import ( + AcquireTicket, + CaptureFailedError, + CaptureKey, + CapturePriority, + CaptureReadLease, + CaptureRequest, + CaptureState, + ConsumerFailedError, + EvictionCandidate, +) + +if TYPE_CHECKING: + from specforge.runtime.data_plane.feature_store import FeatureStore + + +_CONSUMER_ID = re.compile(r"[A-Za-z0-9][A-Za-z0-9._-]{0,127}\Z") + + +class SQLiteWindowedCaptureRegistry: + """SQLite authority for reusable captures and independent cursors.""" + + _ACTIVE_CONSUMER_STATES = ("initializing", "ready", "eof") + _LIVE_CAPTURE_STATES = ( + CaptureState.CAPTURING.value, + CaptureState.COMMITTING.value, + CaptureState.READY.value, + CaptureState.EVICTING.value, + ) + + def __init__( + self, + path: str, + *, + max_live_refs: int, + max_live_bytes: Optional[int] = None, + capture_reservation_bytes: Optional[int] = None, + clock: Callable[[], float] = time.time, + poll_s: float = 0.01, + ) -> None: + if isinstance(max_live_refs, bool) or not isinstance(max_live_refs, int): + raise TypeError("max_live_refs must be an integer") + if max_live_refs < 1: + raise ValueError("max_live_refs must be >= 1") + if (max_live_bytes is None) != (capture_reservation_bytes is None): + raise ValueError( + "max_live_bytes and capture_reservation_bytes must be set together" + ) + if max_live_bytes is not None: + for name, value in ( + ("max_live_bytes", max_live_bytes), + ("capture_reservation_bytes", capture_reservation_bytes), + ): + if isinstance(value, bool) or not isinstance(value, int) or value < 1: + raise ValueError(f"{name} must be a positive integer") + if capture_reservation_bytes > max_live_bytes: + raise ValueError("capture_reservation_bytes exceeds max_live_bytes") + if poll_s <= 0: + raise ValueError("poll_s must be > 0") + + self.path = os.path.abspath(path) + parent = os.path.dirname(self.path) + if parent: + os.makedirs(parent, exist_ok=True) + self.max_live_refs = max_live_refs + self.max_live_bytes = max_live_bytes + self.capture_reservation_bytes = capture_reservation_bytes or 0 + self._clock = clock + self.poll_s = poll_s + self._lock = threading.RLock() + self._closed = False + self._conn = sqlite3.connect( + self.path, check_same_thread=False, timeout=30.0, isolation_level=None + ) + self._conn.row_factory = sqlite3.Row + self._conn.execute("PRAGMA busy_timeout=30000") + self._conn.execute("PRAGMA journal_mode=WAL") + self._conn.execute("PRAGMA synchronous=NORMAL") + self._conn.execute("PRAGMA foreign_keys=ON") + self._create_schema() + + def _create_schema(self) -> None: + schema = """ + CREATE TABLE IF NOT EXISTS run_config ( + singleton INTEGER PRIMARY KEY CHECK(singleton = 1), + run_id TEXT NOT NULL, + contract_digest TEXT NOT NULL, + source_digest TEXT NOT NULL, + total_samples INTEGER NOT NULL, + expected_consumers_json TEXT NOT NULL, + max_live_refs INTEGER NOT NULL, + max_live_bytes INTEGER, + capture_reservation_bytes INTEGER NOT NULL, + status TEXT NOT NULL, + created_at REAL NOT NULL + ); + CREATE TABLE IF NOT EXISTS sources ( + source_index INTEGER PRIMARY KEY, + source_sample_id TEXT NOT NULL UNIQUE + ); + CREATE TABLE IF NOT EXISTS consumers ( + consumer_id TEXT PRIMARY KEY, + cursor INTEGER NOT NULL, + total_samples INTEGER NOT NULL, + lookbehind INTEGER NOT NULL, + lookahead INTEGER NOT NULL, + prefetch_depth INTEGER NOT NULL, + max_outstanding INTEGER NOT NULL, + state TEXT NOT NULL, + heartbeat_at REAL NOT NULL, + failure TEXT, + updated_at REAL NOT NULL + ); + CREATE TABLE IF NOT EXISTS entries ( + source_index INTEGER PRIMARY KEY, + source_sample_id TEXT NOT NULL UNIQUE, + generation INTEGER NOT NULL DEFAULT 0, + materializations INTEGER NOT NULL DEFAULT 0, + state TEXT NOT NULL, + priority INTEGER, + queued_at REAL, + captured_at REAL, + reserved_bytes INTEGER NOT NULL DEFAULT 0, + estimated_bytes INTEGER NOT NULL DEFAULT 0, + ref_json TEXT, + retry_count INTEGER NOT NULL DEFAULT 0, + last_error TEXT, + prefetch_suppressed INTEGER NOT NULL DEFAULT 0, + FOREIGN KEY(source_index) REFERENCES sources(source_index) + ); + CREATE INDEX IF NOT EXISTS entries_schedule + ON entries(state, priority, queued_at, source_index); + CREATE TABLE IF NOT EXISTS interests ( + consumer_id TEXT NOT NULL, + source_index INTEGER NOT NULL, + kind TEXT NOT NULL CHECK(kind IN ('window', 'demand')), + created_at REAL NOT NULL, + PRIMARY KEY(consumer_id, source_index, kind), + FOREIGN KEY(consumer_id) REFERENCES consumers(consumer_id) + ON DELETE CASCADE, + FOREIGN KEY(source_index) REFERENCES sources(source_index) + ); + CREATE TABLE IF NOT EXISTS waiters ( + token TEXT PRIMARY KEY, + consumer_id TEXT NOT NULL, + source_index INTEGER NOT NULL, + created_at REAL NOT NULL, + ready_at_request INTEGER NOT NULL, + UNIQUE(consumer_id, source_index), + FOREIGN KEY(consumer_id) REFERENCES consumers(consumer_id) + ON DELETE CASCADE + ); + CREATE TABLE IF NOT EXISTS read_leases ( + token TEXT PRIMARY KEY, + consumer_id TEXT NOT NULL, + source_index INTEGER NOT NULL, + generation INTEGER NOT NULL, + created_at REAL NOT NULL, + UNIQUE(consumer_id, source_index), + FOREIGN KEY(consumer_id) REFERENCES consumers(consumer_id) + ON DELETE CASCADE + ); + CREATE TABLE IF NOT EXISTS completed_samples ( + consumer_id TEXT NOT NULL, + source_index INTEGER NOT NULL, + PRIMARY KEY(consumer_id, source_index), + FOREIGN KEY(consumer_id) REFERENCES consumers(consumer_id) + ON DELETE CASCADE + ); + CREATE TABLE IF NOT EXISTS registry_meta ( + key TEXT PRIMARY KEY, + value TEXT NOT NULL + ); + """ + with self._lock: + self._conn.executescript(schema) + + @contextmanager + def _transaction(self) -> Iterator[sqlite3.Connection]: + with self._lock: + self._conn.execute("BEGIN IMMEDIATE") + try: + yield self._conn + except BaseException: + self._conn.rollback() + raise + else: + self._conn.commit() + + @staticmethod + def _validate_consumer_id(consumer_id: str) -> None: + if ( + not isinstance(consumer_id, str) + or _CONSUMER_ID.fullmatch(consumer_id) is None + ): + raise ValueError("consumer_id must match [A-Za-z0-9][A-Za-z0-9._-]{0,127}") + + @staticmethod + def _validate_non_negative(name: str, value: int) -> None: + if isinstance(value, bool) or not isinstance(value, int) or value < 0: + raise ValueError(f"{name} must be a non-negative integer") + + def initialize_run( + self, + *, + run_id: str, + contract_digest: str, + source_sample_ids: Sequence[str], + expected_consumers: Sequence[str], + recover_inflight: bool = False, + recovery_store: Optional["FeatureStore"] = None, + ) -> None: + """Create or validate a run, optionally recovering interrupted work. + + ``recover_inflight`` is an owner-only operation. Captures that have no + published payload are fenced by advancing their generation. A + ``recovery_store`` is required for COMMITTING or EVICTING entries so the + old payload is generation-safely reclaimed before work is requeued. + """ + CaptureKey("validation", contract_digest) + if not isinstance(run_id, str) or not run_id: + raise ValueError("run_id must be a non-empty string") + sources = tuple(source_sample_ids) + if ( + not sources + or any(not isinstance(item, str) or not item for item in sources) + or len(sources) != len(set(sources)) + ): + raise ValueError("source_sample_ids must be non-empty, unique strings") + consumers = tuple(expected_consumers) + if not consumers or len(consumers) != len(set(consumers)): + raise ValueError("expected_consumers must be non-empty and unique") + for consumer_id in consumers: + self._validate_consumer_id(consumer_id) + source_digest = hashlib.sha256( + json.dumps(sources, separators=(",", ":")).encode("utf-8") + ).hexdigest() + expected_json = json.dumps(sorted(consumers), separators=(",", ":")) + identity = ( + run_id, + contract_digest, + source_digest, + len(sources), + expected_json, + self.max_live_refs, + self.max_live_bytes, + self.capture_reservation_bytes, + ) + recovery_candidates: tuple[EvictionCandidate, ...] = () + with self._transaction() as conn: + existing = conn.execute( + "SELECT * FROM run_config WHERE singleton=1" + ).fetchone() + if existing is not None: + observed = ( + existing["run_id"], + existing["contract_digest"], + existing["source_digest"], + int(existing["total_samples"]), + existing["expected_consumers_json"], + int(existing["max_live_refs"]), + existing["max_live_bytes"], + int(existing["capture_reservation_bytes"]), + ) + if observed != identity: + raise RuntimeError( + "windowed capture registry identity mismatch: " + f"expected={identity!r}, observed={observed!r}" + ) + interrupted = int( + conn.execute( + "SELECT COUNT(*) FROM entries WHERE state IN (?,?,?)", + ( + CaptureState.CAPTURING.value, + CaptureState.COMMITTING.value, + CaptureState.EVICTING.value, + ), + ).fetchone()[0] + ) + physical = int( + conn.execute( + "SELECT COUNT(*) FROM entries WHERE state IN (?,?)", + ( + CaptureState.COMMITTING.value, + CaptureState.EVICTING.value, + ), + ).fetchone()[0] + ) + if interrupted and not recover_inflight: + raise RuntimeError( + "registry has interrupted captures; owner must explicitly " + "set recover_inflight=True" + ) + if physical and recovery_store is None: + raise RuntimeError( + "recovery_store is required to resolve interrupted payloads" + ) + if interrupted: + recovery_candidates = self._recover_inflight_locked(conn) + else: + conn.execute( + "INSERT INTO run_config VALUES(1,?,?,?,?,?,?,?,?,?,?)", + (*identity, "active", self._clock()), + ) + conn.executemany( + "INSERT INTO sources(source_index,source_sample_id) VALUES(?,?)", + enumerate(sources), + ) + conn.executemany( + "INSERT INTO registry_meta(key,value) VALUES(?,?)", + ( + ("scheduler_cursor", ""), + ("capture_count", "0"), + ("recapture_count", "0"), + ("peak_live_refs", "0"), + ("peak_live_bytes", "0"), + ), + ) + + if recovery_candidates: + assert recovery_store is not None + completed: list[EvictionCandidate] = [] + try: + for candidate in recovery_candidates: + try: + recovery_store.reclaim( + candidate.ref, reason="interrupted-capture-recovery" + ) + except KeyError: + # Recovery is idempotent across a crash after physical reclaim. + pass + completed.append(candidate) + except BaseException: + if completed: + self.finish_evictions(completed) + raise + self.finish_evictions(completed) + + def _recover_inflight_locked( + self, conn: sqlite3.Connection + ) -> tuple[EvictionCandidate, ...]: + rows = conn.execute( + "SELECT * FROM entries WHERE state IN (?,?,?) ORDER BY source_index", + ( + CaptureState.CAPTURING.value, + CaptureState.COMMITTING.value, + CaptureState.EVICTING.value, + ), + ).fetchall() + candidates: list[EvictionCandidate] = [] + for row in rows: + state = CaptureState(row["state"]) + if state in (CaptureState.COMMITTING, CaptureState.EVICTING): + if not row["ref_json"]: + raise RuntimeError( + f"interrupted {state.value} capture has no payload metadata" + ) + if state == CaptureState.COMMITTING: + conn.execute( + "UPDATE entries SET state=? WHERE source_index=?", + (CaptureState.EVICTING.value, row["source_index"]), + ) + candidates.append( + EvictionCandidate( + key=self._key(conn, int(row["source_index"])), + source_index=int(row["source_index"]), + generation=int(row["generation"]), + ref=ref_from_dict(json.loads(row["ref_json"])), + ) + ) + continue + + demand = conn.execute( + "SELECT 1 FROM interests WHERE source_index=? AND kind='demand' " + "LIMIT 1", + (row["source_index"],), + ).fetchone() + interested = conn.execute( + "SELECT 1 FROM interests WHERE source_index=? LIMIT 1", + (row["source_index"],), + ).fetchone() + if interested is None: + conn.execute( + "UPDATE entries SET state=?,priority=NULL,queued_at=NULL," + "reserved_bytes=0,last_error=? WHERE source_index=?", + ( + CaptureState.ABSENT.value, + "producer interrupted", + row["source_index"], + ), + ) + continue + priority = CapturePriority.DEMAND if demand else CapturePriority.PREFETCH + conn.execute( + "UPDATE entries SET state=?,generation=generation+1,priority=?," + "queued_at=?,captured_at=NULL,reserved_bytes=0,retry_count=" + "retry_count+1,last_error=? WHERE source_index=?", + ( + CaptureState.QUEUED.value, + int(priority), + self._clock(), + "producer interrupted", + row["source_index"], + ), + ) + return tuple(candidates) + + def wait_initialized(self, timeout_s: float) -> dict[str, Any]: + if timeout_s <= 0: + raise ValueError("timeout_s must be > 0") + deadline = time.monotonic() + timeout_s + while True: + with self._lock: + row = self._conn.execute( + "SELECT * FROM run_config WHERE singleton=1" + ).fetchone() + if row is not None: + return { + "run_id": row["run_id"], + "contract_digest": row["contract_digest"], + "total_samples": int(row["total_samples"]), + "expected_consumers": tuple( + json.loads(row["expected_consumers_json"]) + ), + "status": row["status"], + "max_live_refs": int(row["max_live_refs"]), + "max_live_bytes": row["max_live_bytes"], + "capture_reservation_bytes": int(row["capture_reservation_bytes"]), + } + if time.monotonic() >= deadline: + raise TimeoutError( + f"registry {self.path!r} was not initialized within " + f"{timeout_s:.1f}s" + ) + time.sleep(min(self.poll_s, max(0.0, deadline - time.monotonic()))) + + def _run(self, conn: sqlite3.Connection) -> sqlite3.Row: + row = conn.execute("SELECT * FROM run_config WHERE singleton=1").fetchone() + if row is None: + raise RuntimeError("windowed capture run is not initialized") + return row + + def _source(self, conn: sqlite3.Connection, source_index: int) -> sqlite3.Row: + if isinstance(source_index, bool) or not isinstance(source_index, int): + raise TypeError("source_index must be an integer") + run = self._run(conn) + if source_index < 0 or source_index >= int(run["total_samples"]): + raise IndexError( + f"source_index {source_index} outside [0, {run['total_samples']})" + ) + row = conn.execute( + "SELECT * FROM sources WHERE source_index=?", (source_index,) + ).fetchone() + if row is None: + raise RuntimeError(f"missing source catalog entry {source_index}") + return row + + def _ensure_entry(self, conn: sqlite3.Connection, source_index: int) -> sqlite3.Row: + source = self._source(conn, source_index) + conn.execute( + "INSERT OR IGNORE INTO entries(source_index,source_sample_id,state) " + "VALUES(?,?,?)", + (source_index, source["source_sample_id"], CaptureState.ABSENT.value), + ) + return conn.execute( + "SELECT * FROM entries WHERE source_index=?", (source_index,) + ).fetchone() + + def _key(self, conn: sqlite3.Connection, source_index: int) -> CaptureKey: + source = self._source(conn, source_index) + run = self._run(conn) + return CaptureKey(source["source_sample_id"], run["contract_digest"]) + + def _queue_entry( + self, + conn: sqlite3.Connection, + row: sqlite3.Row, + priority: CapturePriority, + ) -> sqlite3.Row: + state = CaptureState(row["state"]) + if state == CaptureState.QUEUED: + if int(row["priority"]) > int(priority): + conn.execute( + "UPDATE entries SET priority=?,prefetch_suppressed=0 " + "WHERE source_index=?", + (int(priority), row["source_index"]), + ) + return conn.execute( + "SELECT * FROM entries WHERE source_index=?", (row["source_index"],) + ).fetchone() + if state not in (CaptureState.ABSENT, CaptureState.FAILED): + return row + if state == CaptureState.FAILED: + waiters = conn.execute( + "SELECT COUNT(*) FROM waiters WHERE source_index=?", + (row["source_index"],), + ).fetchone()[0] + if waiters: + return row + if priority == CapturePriority.PREFETCH and row["prefetch_suppressed"]: + return row + conn.execute( + "UPDATE entries SET state=?,generation=generation+1,priority=?," + "queued_at=?,captured_at=NULL,reserved_bytes=0,last_error=NULL," + "retry_count=0,prefetch_suppressed=0 WHERE source_index=?", + ( + CaptureState.QUEUED.value, + int(priority), + self._clock(), + row["source_index"], + ), + ) + return conn.execute( + "SELECT * FROM entries WHERE source_index=?", (row["source_index"],) + ).fetchone() + + def _window_indices(self, consumer: sqlite3.Row) -> range: + cursor = int(consumer["cursor"]) + low = max(0, cursor - int(consumer["lookbehind"])) + high = min( + int(consumer["total_samples"]), cursor + int(consumer["lookahead"]) + 1 + ) + return range(low, high) + + def _refresh_window_locked( + self, conn: sqlite3.Connection, consumer_id: str + ) -> None: + consumer = conn.execute( + "SELECT * FROM consumers WHERE consumer_id=?", (consumer_id,) + ).fetchone() + if consumer is None: + raise KeyError(f"unknown consumer {consumer_id!r}") + desired = set(self._window_indices(consumer)) + existing = { + int(row[0]) + for row in conn.execute( + "SELECT source_index FROM interests WHERE consumer_id=? " + "AND kind='window'", + (consumer_id,), + ).fetchall() + } + now = self._clock() + for source_index in sorted(desired - existing): + self._ensure_entry(conn, source_index) + conn.execute( + "INSERT INTO interests VALUES(?,?,'window',?)", + (consumer_id, source_index, now), + ) + conn.execute( + "UPDATE entries SET prefetch_suppressed=0 WHERE source_index=?", + (source_index,), + ) + for source_index in sorted(existing - desired): + conn.execute( + "DELETE FROM interests WHERE consumer_id=? AND source_index=? " + "AND kind='window'", + (consumer_id, source_index), + ) + self._prune_orphans_locked(conn) + self._top_up_prefetch_locked(conn, consumer_id) + + def _top_up_prefetch_locked( + self, conn: sqlite3.Connection, consumer_id: str + ) -> None: + consumer = conn.execute( + "SELECT * FROM consumers WHERE consumer_id=?", (consumer_id,) + ).fetchone() + if consumer is None or consumer["state"] not in self._ACTIVE_CONSUMER_STATES: + return + depth = int(consumer["prefetch_depth"]) + if depth == 0: + return + # READY rows are cached results, not occupied producer slots. Keep a + # rolling number of capture operations active until the legal window is full. + outstanding = conn.execute( + "SELECT COUNT(DISTINCT e.source_index) FROM interests i " + "JOIN entries e ON e.source_index=i.source_index " + "WHERE i.consumer_id=? AND i.kind='window' AND e.priority=? " + "AND e.state IN (?,?,?)", + ( + consumer_id, + int(CapturePriority.PREFETCH), + CaptureState.QUEUED.value, + CaptureState.CAPTURING.value, + CaptureState.COMMITTING.value, + ), + ).fetchone()[0] + remaining = max(0, depth - int(outstanding)) + if remaining == 0: + return + cursor = int(consumer["cursor"]) + stop = min( + int(consumer["total_samples"]), + cursor + int(consumer["lookahead"]) + 1, + ) + candidates = conn.execute( + "SELECT DISTINCT e.* FROM interests i JOIN entries e ON " + "e.source_index=i.source_index WHERE i.consumer_id=? AND " + "i.kind='window' AND e.source_index>=? AND e.source_index None: + outstanding = conn.execute( + "SELECT (SELECT COUNT(*) FROM waiters WHERE consumer_id=? AND " + "source_index=?) + (SELECT COUNT(*) FROM read_leases WHERE " + "consumer_id=? AND source_index=?)", + (consumer_id, source_index, consumer_id, source_index), + ).fetchone()[0] + if not outstanding: + conn.execute( + "DELETE FROM interests WHERE consumer_id=? AND source_index=? " + "AND kind='demand'", + (consumer_id, source_index), + ) + + def _prune_orphans_locked(self, conn: sqlite3.Connection) -> int: + rows = conn.execute( + "SELECT source_index FROM entries e WHERE e.state IN (?,?) AND " + "NOT EXISTS(SELECT 1 FROM interests i WHERE i.source_index=" + "e.source_index) AND NOT EXISTS(SELECT 1 FROM waiters w WHERE " + "w.source_index=e.source_index)", + (CaptureState.QUEUED.value, CaptureState.FAILED.value), + ).fetchall() + for row in rows: + conn.execute( + "UPDATE entries SET state=?,priority=NULL,queued_at=NULL," + "reserved_bytes=0,last_error=NULL WHERE source_index=?", + (CaptureState.ABSENT.value, row["source_index"]), + ) + return len(rows) + + def register_consumer( + self, + consumer_id: str, + *, + lookbehind: int = 0, + lookahead: int = 0, + prefetch_depth: int = 0, + max_outstanding: int = 1, + cursor: int = 0, + ) -> None: + """Register one stable logical consumer and its bounded cache window.""" + self._validate_consumer_id(consumer_id) + for name, value in ( + ("lookbehind", lookbehind), + ("lookahead", lookahead), + ("prefetch_depth", prefetch_depth), + ("max_outstanding", max_outstanding), + ("cursor", cursor), + ): + self._validate_non_negative(name, value) + if max_outstanding < 1: + raise ValueError("max_outstanding must be >= 1") + if prefetch_depth > lookahead + 1: + raise ValueError("prefetch_depth cannot exceed lookahead + 1") + + with self._transaction() as conn: + run = self._run(conn) + expected = set(json.loads(run["expected_consumers_json"])) + if consumer_id not in expected: + raise ValueError( + f"consumer {consumer_id!r} is not expected; expected={sorted(expected)}" + ) + total = int(run["total_samples"]) + if cursor > total: + raise ValueError("consumer cursor exceeds total samples") + identity = ( + cursor, + total, + lookbehind, + lookahead, + prefetch_depth, + max_outstanding, + ) + existing = conn.execute( + "SELECT * FROM consumers WHERE consumer_id=?", (consumer_id,) + ).fetchone() + if existing is not None: + observed = tuple( + int(existing[name]) + for name in ( + "cursor", + "total_samples", + "lookbehind", + "lookahead", + "prefetch_depth", + "max_outstanding", + ) + ) + if observed != identity or existing["state"] in ("completed", "failed"): + raise RuntimeError( + f"consumer {consumer_id!r} registration mismatch: " + f"expected={identity}, observed={observed}, " + f"state={existing['state']!r}" + ) + return + now = self._clock() + state = "eof" if cursor == total else "initializing" + conn.execute( + "INSERT INTO consumers VALUES(?,?,?,?,?,?,?,?,?,?,?)", + ( + consumer_id, + cursor, + total, + lookbehind, + lookahead, + prefetch_depth, + max_outstanding, + state, + now, + None, + now, + ), + ) + self._refresh_window_locked(conn, consumer_id) + + def resume_consumer(self, consumer_id: str, *, durable_cursor: int) -> None: + """Reconcile transient state to an optimizer-durable prefix cursor. + + The cursor may move in either direction. A rewind deliberately replays + captures acknowledged only by the crashed process; a fast-forward is + accepted only because the caller explicitly supplies durable progress. + """ + self._validate_non_negative("durable_cursor", durable_cursor) + with self._transaction() as conn: + consumer = conn.execute( + "SELECT * FROM consumers WHERE consumer_id=?", (consumer_id,) + ).fetchone() + if consumer is None: + raise KeyError(f"unknown consumer {consumer_id!r}") + if consumer["state"] == "failed": + raise RuntimeError(f"cannot resume failed consumer {consumer_id!r}") + total = int(consumer["total_samples"]) + if durable_cursor > total: + raise ValueError("durable_cursor exceeds total samples") + self._release_consumer_transients_locked(conn, consumer_id) + conn.execute( + "DELETE FROM completed_samples WHERE consumer_id=?", (consumer_id,) + ) + now = self._clock() + state = "eof" if durable_cursor == total else "initializing" + conn.execute( + "UPDATE consumers SET cursor=?,state=?,failure=NULL,heartbeat_at=?," + "updated_at=? WHERE consumer_id=?", + (durable_cursor, state, now, now, consumer_id), + ) + self._refresh_window_locked(conn, consumer_id) + + def heartbeat(self, consumer_id: str, *, ready: bool = False) -> None: + with self._transaction() as conn: + consumer = conn.execute( + "SELECT * FROM consumers WHERE consumer_id=?", (consumer_id,) + ).fetchone() + if consumer is None: + raise KeyError(f"unknown consumer {consumer_id!r}") + if consumer["state"] not in self._ACTIVE_CONSUMER_STATES: + raise RuntimeError(f"consumer {consumer_id!r} is {consumer['state']!r}") + state = consumer["state"] + if ready and state == "initializing": + state = "ready" + now = self._clock() + conn.execute( + "UPDATE consumers SET state=?,heartbeat_at=?,updated_at=? " + "WHERE consumer_id=?", + (state, now, now, consumer_id), + ) + + def consumer_cursor(self, consumer_id: str) -> int: + with self._lock: + row = self._conn.execute( + "SELECT cursor FROM consumers WHERE consumer_id=?", (consumer_id,) + ).fetchone() + if row is None: + raise KeyError(f"unknown consumer {consumer_id!r}") + return int(row["cursor"]) + + def request_acquire(self, consumer_id: str, source_index: int) -> AcquireTicket: + token = uuid.uuid4().hex + with self._transaction() as conn: + consumer = conn.execute( + "SELECT * FROM consumers WHERE consumer_id=?", (consumer_id,) + ).fetchone() + if consumer is None: + raise KeyError(f"unknown consumer {consumer_id!r}") + if consumer["state"] not in self._ACTIVE_CONSUMER_STATES: + raise RuntimeError(f"consumer {consumer_id!r} is {consumer['state']!r}") + if source_index not in self._window_indices(consumer): + window = self._window_indices(consumer) + raise ValueError( + f"source_index {source_index} outside consumer {consumer_id!r} " + f"window [{window.start}, {window.stop})" + ) + duplicate = conn.execute( + "SELECT (SELECT COUNT(*) FROM waiters WHERE consumer_id=? AND " + "source_index=?) + (SELECT COUNT(*) FROM read_leases WHERE " + "consumer_id=? AND source_index=?)", + (consumer_id, source_index, consumer_id, source_index), + ).fetchone()[0] + if duplicate: + raise RuntimeError( + f"consumer {consumer_id!r} already has source_index " + f"{source_index} outstanding" + ) + outstanding = conn.execute( + "SELECT (SELECT COUNT(*) FROM waiters WHERE consumer_id=?) + " + "(SELECT COUNT(*) FROM read_leases WHERE consumer_id=?)", + (consumer_id, consumer_id), + ).fetchone()[0] + if int(outstanding) >= int(consumer["max_outstanding"]): + raise RuntimeError( + f"consumer {consumer_id!r} reached max_outstanding=" + f"{consumer['max_outstanding']}" + ) + row = self._ensure_entry(conn, source_index) + ready = CaptureState(row["state"]) == CaptureState.READY + now = self._clock() + conn.execute( + "INSERT OR IGNORE INTO interests VALUES(?,?,'demand',?)", + (consumer_id, source_index, now), + ) + self._queue_entry(conn, row, CapturePriority.DEMAND) + conn.execute( + "INSERT INTO waiters VALUES(?,?,?,?,?)", + (token, consumer_id, source_index, now, int(ready)), + ) + key = self._key(conn, source_index) + return AcquireTicket(token, consumer_id, key, source_index, ready) + + def request_many( + self, consumer_id: str, source_indices: Iterable[int] + ) -> tuple[AcquireTicket, ...]: + tickets: list[AcquireTicket] = [] + try: + for source_index in source_indices: + tickets.append(self.request_acquire(consumer_id, source_index)) + except BaseException: + for ticket in tickets: + self.cancel_acquire(ticket) + raise + return tuple(tickets) + + def wait_ready( + self, ticket: AcquireTicket, *, timeout_s: Optional[float] + ) -> CaptureReadLease: + if timeout_s is not None and timeout_s <= 0: + raise ValueError("timeout_s must be > 0 or None") + started = time.monotonic() + deadline = None if timeout_s is None else started + timeout_s + while True: + with self._lock: + waiter = self._conn.execute( + "SELECT 1 FROM waiters WHERE token=?", (ticket.token,) + ).fetchone() + entry = self._conn.execute( + "SELECT state FROM entries WHERE source_index=?", + (ticket.source_index,), + ).fetchone() + consumer = self._conn.execute( + "SELECT state,failure FROM consumers WHERE consumer_id=?", + (ticket.consumer_id,), + ).fetchone() + if waiter is None: + if consumer is not None and consumer["state"] == "failed": + raise ConsumerFailedError( + f"consumer {ticket.consumer_id!r} failed: " + f"{consumer['failure'] or 'unknown failure'}" + ) + raise RuntimeError(f"acquire ticket {ticket.token} is no longer active") + if entry is None: + raise RuntimeError(f"capture entry disappeared for {ticket.key}") + if CaptureState(entry["state"]) in ( + CaptureState.READY, + CaptureState.FAILED, + ): + lease = self._finish_wait(ticket, started) + if lease is not None: + return lease + if deadline is not None and time.monotonic() >= deadline: + self.cancel_acquire(ticket) + raise TimeoutError( + f"capture {ticket.key.source_sample_id} did not become ready " + f"within {timeout_s:.1f}s" + ) + sleep_s = self.poll_s + if deadline is not None: + sleep_s = min(sleep_s, max(0.0, deadline - time.monotonic())) + time.sleep(sleep_s) + + def _finish_wait( + self, ticket: AcquireTicket, started: float + ) -> Optional[CaptureReadLease]: + failure: Optional[str] = None + lease: Optional[CaptureReadLease] = None + with self._transaction() as conn: + waiter = conn.execute( + "SELECT 1 FROM waiters WHERE token=?", (ticket.token,) + ).fetchone() + if waiter is None: + raise RuntimeError(f"acquire ticket {ticket.token} is no longer active") + row = conn.execute( + "SELECT * FROM entries WHERE source_index=?", (ticket.source_index,) + ).fetchone() + state = CaptureState(row["state"]) + if state not in (CaptureState.READY, CaptureState.FAILED): + return None + if state == CaptureState.FAILED: + conn.execute("DELETE FROM waiters WHERE token=?", (ticket.token,)) + self._drop_demand_if_idle_locked( + conn, ticket.consumer_id, ticket.source_index + ) + failure = ( + f"capture {ticket.key.source_sample_id} generation " + f"{row['generation']} failed: {row['last_error'] or 'unknown error'}" + ) + else: + if not row["ref_json"]: + raise RuntimeError("READY capture has no SampleRef metadata") + lease_token = uuid.uuid4().hex + conn.execute( + "INSERT INTO read_leases VALUES(?,?,?,?,?)", + ( + lease_token, + ticket.consumer_id, + ticket.source_index, + int(row["generation"]), + self._clock(), + ), + ) + conn.execute("DELETE FROM waiters WHERE token=?", (ticket.token,)) + lease = CaptureReadLease( + token=lease_token, + consumer_id=ticket.consumer_id, + key=ticket.key, + source_index=ticket.source_index, + generation=int(row["generation"]), + ref=ref_from_dict(json.loads(row["ref_json"])), + ready_at_request=ticket.ready_at_request, + wait_s=time.monotonic() - started, + ) + if failure is not None: + raise CaptureFailedError(failure) + return lease + + def cancel_acquire(self, ticket: AcquireTicket) -> None: + with self._transaction() as conn: + conn.execute("DELETE FROM waiters WHERE token=?", (ticket.token,)) + self._drop_demand_if_idle_locked( + conn, ticket.consumer_id, ticket.source_index + ) + self._prune_orphans_locked(conn) + + def _live_usage_locked(self, conn: sqlite3.Connection) -> tuple[int, int]: + placeholders = ",".join("?" for _ in self._LIVE_CAPTURE_STATES) + row = conn.execute( + f"SELECT COUNT(*) AS refs,COALESCE(SUM(CASE WHEN state IN (?,?) " + f"THEN estimated_bytes ELSE reserved_bytes END),0) AS bytes " + f"FROM entries WHERE state IN ({placeholders})", + ( + CaptureState.READY.value, + CaptureState.EVICTING.value, + *self._LIVE_CAPTURE_STATES, + ), + ).fetchone() + return int(row["refs"]), int(row["bytes"]) + + def _claim_capacity_locked(self, conn: sqlite3.Connection) -> int: + live_refs, live_bytes = self._live_usage_locked(conn) + capacity = max(0, self.max_live_refs - live_refs) + if self.max_live_bytes is not None: + byte_capacity = max(0, self.max_live_bytes - live_bytes) + capacity = min(capacity, byte_capacity // self.capture_reservation_bytes) + return capacity + + @staticmethod + def _meta_int(conn: sqlite3.Connection, key: str) -> int: + row = conn.execute( + "SELECT value FROM registry_meta WHERE key=?", (key,) + ).fetchone() + if row is None: + raise RuntimeError(f"missing registry metadata {key!r}") + return int(row["value"]) + + @staticmethod + def _set_meta_int(conn: sqlite3.Connection, key: str, value: int) -> None: + conn.execute("UPDATE registry_meta SET value=? WHERE key=?", (str(value), key)) + + def _record_peaks_locked(self, conn: sqlite3.Connection) -> None: + live_refs, live_bytes = self._live_usage_locked(conn) + if live_refs > self._meta_int(conn, "peak_live_refs"): + self._set_meta_int(conn, "peak_live_refs", live_refs) + if live_bytes > self._meta_int(conn, "peak_live_bytes"): + self._set_meta_int(conn, "peak_live_bytes", live_bytes) + + def claim_batch(self, max_requests: int) -> tuple[CaptureRequest, ...]: + """Claim demand-first work while reserving authoritative capacity.""" + if isinstance(max_requests, bool) or not isinstance(max_requests, int): + raise TypeError("max_requests must be an integer") + if max_requests < 1: + raise ValueError("max_requests must be >= 1") + with self._transaction() as conn: + limit = min(max_requests, self._claim_capacity_locked(conn)) + if limit == 0: + return () + rows = conn.execute( + "SELECT * FROM entries WHERE state=? ORDER BY priority,queued_at," + "source_index", + (CaptureState.QUEUED.value,), + ).fetchall() + if not rows: + return () + demand_rows = [ + row + for row in rows + if int(row["priority"]) == int(CapturePriority.DEMAND) + ] + prefetch_rows = [ + row + for row in rows + if int(row["priority"]) == int(CapturePriority.PREFETCH) + ] + owners: dict[int, tuple[str, ...]] = {} + by_consumer: dict[str, list[sqlite3.Row]] = {} + for row in demand_rows: + source_index = int(row["source_index"]) + demanders = tuple( + owner["consumer_id"] + for owner in conn.execute( + "SELECT consumer_id FROM interests WHERE source_index=? " + "AND kind='demand' ORDER BY consumer_id", + (source_index,), + ).fetchall() + ) + owners[source_index] = demanders + for consumer_id in demanders: + by_consumer.setdefault(consumer_id, []).append(row) + + cursor_row = conn.execute( + "SELECT value FROM registry_meta WHERE key='scheduler_cursor'" + ).fetchone() + last_consumer = cursor_row["value"] if cursor_row else "" + consumer_order = sorted(by_consumer) + if last_consumer in consumer_order: + start = (consumer_order.index(last_consumer) + 1) % len(consumer_order) + consumer_order = consumer_order[start:] + consumer_order[:start] + + selected: list[sqlite3.Row] = [] + selected_indices: set[int] = set() + last_selected = last_consumer + while consumer_order and len(selected) < limit: + made_progress = False + for consumer_id in consumer_order: + pending = by_consumer[consumer_id] + while pending: + row = pending.pop(0) + source_index = int(row["source_index"]) + if source_index not in selected_indices: + selected.append(row) + selected_indices.add(source_index) + last_selected = consumer_id + made_progress = True + break + if len(selected) >= limit: + break + if not made_progress: + break + for row in prefetch_rows: + if len(selected) >= limit: + break + source_index = int(row["source_index"]) + if source_index not in selected_indices: + selected.append(row) + selected_indices.add(source_index) + + now = self._clock() + requests: list[CaptureRequest] = [] + for row in selected: + source_index = int(row["source_index"]) + priority = CapturePriority(int(row["priority"])) + updated = conn.execute( + "UPDATE entries SET state=?,captured_at=?,reserved_bytes=? " + "WHERE source_index=? AND state=?", + ( + CaptureState.CAPTURING.value, + now, + self.capture_reservation_bytes, + source_index, + CaptureState.QUEUED.value, + ), + ) + if updated.rowcount != 1: + raise RuntimeError( + f"capture claim lost transactional ownership for {source_index}" + ) + requests.append( + CaptureRequest( + key=self._key(conn, source_index), + source_index=source_index, + generation=int(row["generation"]), + priority=priority, + demand_consumers=owners.get(source_index, ()), + reserved_bytes=self.capture_reservation_bytes, + ) + ) + if last_selected: + conn.execute( + "UPDATE registry_meta SET value=? WHERE key='scheduler_cursor'", + (last_selected,), + ) + self._record_peaks_locked(conn) + return tuple(requests) + + def _validated_ref_json( + self, request: CaptureRequest, ref: SampleRef + ) -> tuple[str, int]: + assert_no_tensors(ref) + # Storage backends own ``generation`` as the physical object locator. + # Recapture scheduling is a separate namespace: runtime adapters attach + # ``window_generation`` without rewriting the store's generation. The + # fallback preserves refs produced directly against the registry. + generation = ref.metadata.get( + "window_generation", ref.metadata.get("generation") + ) + if ( + isinstance(generation, bool) + or not isinstance(generation, int) + or generation != request.generation + ): + raise ValueError( + f"capture ref generation {generation!r} does not match " + f"claimed generation {request.generation}" + ) + if ref.source_task_id != request.key.source_sample_id: + raise ValueError( + f"capture ref source_task_id={ref.source_task_id!r} does not match " + f"{request.key.source_sample_id!r}" + ) + with self._lock: + run = self._conn.execute( + "SELECT run_id FROM run_config WHERE singleton=1" + ).fetchone() + if run is None or ref.run_id != run["run_id"]: + raise ValueError( + f"capture ref run_id={ref.run_id!r} does not match registry run " + f"{None if run is None else run['run_id']!r}" + ) + estimated = ref.estimated_bytes + if ( + isinstance(estimated, bool) + or not isinstance(estimated, int) + or estimated < 0 + ): + raise ValueError("SampleRef.estimated_bytes must be a non-negative integer") + if self.max_live_bytes is not None and estimated > request.reserved_bytes: + raise ValueError( + f"capture estimated_bytes={estimated} exceeds reserved upper bound " + f"{request.reserved_bytes}; reclaim the payload and fail the request" + ) + return json.dumps(ref_to_dict(ref), separators=(",", ":")), estimated + + def mark_committing(self, request: CaptureRequest, ref: SampleRef) -> None: + """Persist payload identity before exposing the READY transition.""" + ref_json, estimated = self._validated_ref_json(request, ref) + with self._transaction() as conn: + row = conn.execute( + "SELECT * FROM entries WHERE source_index=?", (request.source_index,) + ).fetchone() + if ( + row is None + or CaptureState(row["state"]) != CaptureState.CAPTURING + or int(row["generation"]) != request.generation + or row["source_sample_id"] != request.key.source_sample_id + ): + raise RuntimeError( + f"capture commit transition mismatch for {request.key}" + ) + conn.execute( + "UPDATE entries SET state=?,estimated_bytes=?,ref_json=? " + "WHERE source_index=?", + ( + CaptureState.COMMITTING.value, + estimated, + ref_json, + request.source_index, + ), + ) + + def complete_capture(self, request: CaptureRequest, ref: SampleRef) -> None: + """Expose a previously persisted COMMITTING payload to consumers.""" + ref_json, estimated = self._validated_ref_json(request, ref) + with self._transaction() as conn: + row = conn.execute( + "SELECT * FROM entries WHERE source_index=?", (request.source_index,) + ).fetchone() + if ( + row is None + or CaptureState(row["state"]) != CaptureState.COMMITTING + or int(row["generation"]) != request.generation + ): + raise RuntimeError( + f"capture READY transition mismatch for {request.key}" + ) + if row["ref_json"] != ref_json: + raise RuntimeError( + f"capture payload identity changed while committing {request.key}" + ) + recapture = int(row["materializations"]) > 0 + conn.execute( + "UPDATE entries SET state=?,materializations=materializations+1," + "priority=NULL,queued_at=NULL,reserved_bytes=0,estimated_bytes=?," + "retry_count=0,last_error=NULL WHERE source_index=?", + ( + CaptureState.READY.value, + estimated, + request.source_index, + ), + ) + self._set_meta_int( + conn, "capture_count", self._meta_int(conn, "capture_count") + 1 + ) + if recapture: + self._set_meta_int( + conn, + "recapture_count", + self._meta_int(conn, "recapture_count") + 1, + ) + self._record_peaks_locked(conn) + interested = conn.execute( + "SELECT DISTINCT consumer_id FROM interests WHERE source_index=?", + (request.source_index,), + ).fetchall() + for consumer in interested: + self._top_up_prefetch_locked(conn, consumer["consumer_id"]) + + def fail_capture( + self, + request: CaptureRequest, + error: BaseException | str, + *, + retryable: bool, + max_retries: int, + ) -> bool: + """Fail one generation and return whether a retry was queued.""" + if ( + isinstance(max_retries, bool) + or not isinstance(max_retries, int) + or max_retries < 0 + ): + raise ValueError("max_retries must be a non-negative integer") + message = str(error) + with self._transaction() as conn: + row = conn.execute( + "SELECT * FROM entries WHERE source_index=?", (request.source_index,) + ).fetchone() + if ( + row is None + or CaptureState(row["state"]) + not in (CaptureState.CAPTURING, CaptureState.COMMITTING) + or int(row["generation"]) != request.generation + ): + raise RuntimeError( + f"capture failure transition mismatch for {request.key}" + ) + failures = int(row["retry_count"]) + 1 + demand = conn.execute( + "SELECT 1 FROM interests WHERE source_index=? AND kind='demand' " + "LIMIT 1", + (request.source_index,), + ).fetchone() + interested = conn.execute( + "SELECT 1 FROM interests WHERE source_index=? LIMIT 1", + (request.source_index,), + ).fetchone() + should_retry = bool( + retryable and failures <= max_retries and interested is not None + ) + if should_retry: + priority = ( + CapturePriority.DEMAND if demand else CapturePriority.PREFETCH + ) + conn.execute( + "UPDATE entries SET state=?,generation=generation+1,priority=?," + "queued_at=?,captured_at=NULL,reserved_bytes=0,retry_count=?," + "last_error=? WHERE source_index=?", + ( + CaptureState.QUEUED.value, + int(priority), + self._clock(), + failures, + message, + request.source_index, + ), + ) + else: + conn.execute( + "UPDATE entries SET state=?,priority=NULL,queued_at=NULL," + "captured_at=NULL,reserved_bytes=0,retry_count=?,last_error=? " + "WHERE source_index=?", + ( + CaptureState.FAILED.value, + failures, + message, + request.source_index, + ), + ) + self._prune_orphans_locked(conn) + return should_retry + + def release_and_advance( + self, consumer_id: str, leases: Sequence[CaptureReadLease] + ) -> int: + """Release read protection and advance only the contiguous ACK prefix.""" + if not leases: + return self.consumer_cursor(consumer_id) + if len({lease.token for lease in leases}) != len(leases): + raise ValueError("leases contains duplicate tokens") + with self._transaction() as conn: + consumer = conn.execute( + "SELECT * FROM consumers WHERE consumer_id=?", (consumer_id,) + ).fetchone() + if consumer is None: + raise KeyError(f"unknown consumer {consumer_id!r}") + if consumer["state"] not in self._ACTIVE_CONSUMER_STATES: + raise RuntimeError(f"consumer {consumer_id!r} is {consumer['state']!r}") + for lease in leases: + row = conn.execute( + "SELECT * FROM read_leases WHERE token=?", (lease.token,) + ).fetchone() + observed = None + if row is not None: + observed = ( + row["consumer_id"], + int(row["source_index"]), + int(row["generation"]), + ) + expected = (consumer_id, lease.source_index, lease.generation) + if observed != expected: + raise RuntimeError( + f"unknown, released, or mismatched lease {lease.token}" + ) + conn.execute("DELETE FROM read_leases WHERE token=?", (lease.token,)) + conn.execute( + "INSERT OR IGNORE INTO completed_samples VALUES(?,?)", + (consumer_id, lease.source_index), + ) + self._drop_demand_if_idle_locked(conn, consumer_id, lease.source_index) + + cursor = int(consumer["cursor"]) + total = int(consumer["total_samples"]) + while cursor < total: + completed = conn.execute( + "SELECT 1 FROM completed_samples WHERE consumer_id=? AND " + "source_index=?", + (consumer_id, cursor), + ).fetchone() + if completed is None: + break + conn.execute( + "DELETE FROM completed_samples WHERE consumer_id=? AND " + "source_index=?", + (consumer_id, cursor), + ) + cursor += 1 + state = "eof" if cursor == total else consumer["state"] + now = self._clock() + conn.execute( + "UPDATE consumers SET cursor=?,state=?,updated_at=? " + "WHERE consumer_id=?", + (cursor, state, now, consumer_id), + ) + self._refresh_window_locked(conn, consumer_id) + return cursor + + def abandon_leases( + self, consumer_id: str, leases: Sequence[CaptureReadLease] + ) -> None: + """Release read protection without recording consumer progress.""" + if not leases: + return + with self._transaction() as conn: + for lease in leases: + row = conn.execute( + "SELECT * FROM read_leases WHERE token=?", (lease.token,) + ).fetchone() + if row is None: + continue + if ( + row["consumer_id"] != consumer_id + or int(row["source_index"]) != lease.source_index + or int(row["generation"]) != lease.generation + ): + raise RuntimeError( + f"read lease identity mismatch while abandoning {lease.token}" + ) + conn.execute("DELETE FROM read_leases WHERE token=?", (lease.token,)) + self._drop_demand_if_idle_locked(conn, consumer_id, lease.source_index) + self._prune_orphans_locked(conn) + + def begin_evictions( + self, *, limit: int = 64, pressure: bool = False + ) -> tuple[EvictionCandidate, ...]: + """Atomically claim reclaimable captures. + + Normal eviction requires no interest. Pressure eviction may drop soft + window cache entries, but never a demand waiter or read lease. + """ + if isinstance(limit, bool) or not isinstance(limit, int) or limit < 1: + raise ValueError("limit must be a positive integer") + with self._transaction() as conn: + if pressure: + interest_clause = ( + "NOT EXISTS(SELECT 1 FROM interests i WHERE i.source_index=" + "e.source_index AND i.kind='demand')" + ) + else: + interest_clause = ( + "NOT EXISTS(SELECT 1 FROM interests i WHERE i.source_index=" + "e.source_index)" + ) + rows = conn.execute( + f"SELECT e.* FROM entries e WHERE e.state=? AND {interest_clause} " + "AND NOT EXISTS(SELECT 1 FROM waiters w WHERE w.source_index=" + "e.source_index) AND NOT EXISTS(SELECT 1 FROM read_leases l " + "WHERE l.source_index=e.source_index) ORDER BY e.source_index " + "LIMIT ?", + (CaptureState.READY.value, limit), + ).fetchall() + candidates: list[EvictionCandidate] = [] + for row in rows: + if not row["ref_json"]: + raise RuntimeError("READY capture selected for eviction has no ref") + conn.execute( + "UPDATE entries SET state=?,prefetch_suppressed=? " + "WHERE source_index=?", + ( + CaptureState.EVICTING.value, + int(pressure), + row["source_index"], + ), + ) + candidates.append( + EvictionCandidate( + key=self._key(conn, int(row["source_index"])), + source_index=int(row["source_index"]), + generation=int(row["generation"]), + ref=ref_from_dict(json.loads(row["ref_json"])), + ) + ) + return tuple(candidates) + + def finish_evictions(self, candidates: Sequence[EvictionCandidate]) -> None: + if not candidates: + return + with self._transaction() as conn: + affected: set[str] = set() + for candidate in candidates: + row = conn.execute( + "SELECT * FROM entries WHERE source_index=?", + (candidate.source_index,), + ).fetchone() + if ( + row is None + or CaptureState(row["state"]) != CaptureState.EVICTING + or int(row["generation"]) != candidate.generation + or row["source_sample_id"] != candidate.key.source_sample_id + ): + raise RuntimeError( + f"eviction completion mismatch for {candidate.key}" + ) + conn.execute( + "UPDATE entries SET state=?,priority=NULL,queued_at=NULL," + "captured_at=NULL,reserved_bytes=0,estimated_bytes=0," + "ref_json=NULL WHERE source_index=?", + (CaptureState.ABSENT.value, candidate.source_index), + ) + interests = conn.execute( + "SELECT consumer_id,kind FROM interests WHERE source_index=?", + (candidate.source_index,), + ).fetchall() + if any(item["kind"] == "demand" for item in interests): + updated = conn.execute( + "SELECT * FROM entries WHERE source_index=?", + (candidate.source_index,), + ).fetchone() + self._queue_entry(conn, updated, CapturePriority.DEMAND) + affected.update(item["consumer_id"] for item in interests) + for consumer_id in affected: + self._top_up_prefetch_locked(conn, consumer_id) + + def cancel_evictions( + self, candidates: Sequence[EvictionCandidate], error: BaseException | str + ) -> None: + if not candidates: + return + with self._transaction() as conn: + for candidate in candidates: + row = conn.execute( + "SELECT state,generation FROM entries WHERE source_index=?", + (candidate.source_index,), + ).fetchone() + if ( + row is not None + and CaptureState(row["state"]) == CaptureState.EVICTING + and int(row["generation"]) == candidate.generation + ): + conn.execute( + "UPDATE entries SET state=?,last_error=? WHERE source_index=?", + ( + CaptureState.READY.value, + str(error), + candidate.source_index, + ), + ) + + def reclaim( + self, + store: "FeatureStore", + *, + limit: int = 64, + pressure: bool = False, + reason: str = "window-expired", + ) -> int: + """Reclaim eligible payloads and commit metadata eviction afterward.""" + candidates = self.begin_evictions(limit=limit, pressure=pressure) + completed: list[EvictionCandidate] = [] + try: + for candidate in candidates: + store.reclaim(candidate.ref, reason=reason) + completed.append(candidate) + except BaseException as error: + cleanup_candidates = candidates[len(completed) :] + if completed: + try: + self.finish_evictions(completed) + except BaseException as finish_error: + cleanup_candidates = candidates + error.add_note( + f"failed to finish completed evictions: {finish_error!r}" + ) + try: + self.cancel_evictions(cleanup_candidates, error) + except BaseException as cancel_error: + error.add_note(f"failed to cancel pending evictions: {cancel_error!r}") + raise + self.finish_evictions(completed) + return len(completed) + + def cancel_queued_prefetch(self, limit: Optional[int] = None) -> int: + if limit is not None and ( + isinstance(limit, bool) or not isinstance(limit, int) or limit < 0 + ): + raise ValueError("limit must be a non-negative integer or None") + with self._transaction() as conn: + sql = ( + "SELECT source_index FROM entries WHERE state=? AND priority=? " + "ORDER BY queued_at DESC" + ) + params: list[Any] = [ + CaptureState.QUEUED.value, + int(CapturePriority.PREFETCH), + ] + if limit is not None: + sql += " LIMIT ?" + params.append(limit) + rows = conn.execute(sql, params).fetchall() + for row in rows: + conn.execute( + "UPDATE entries SET state=?,priority=NULL,queued_at=NULL," + "prefetch_suppressed=1 WHERE source_index=?", + (CaptureState.ABSENT.value, row["source_index"]), + ) + return len(rows) + + def resume_prefetch(self) -> None: + with self._transaction() as conn: + conn.execute("UPDATE entries SET prefetch_suppressed=0") + consumers = conn.execute( + "SELECT consumer_id FROM consumers WHERE state IN (?,?,?)", + self._ACTIVE_CONSUMER_STATES, + ).fetchall() + for row in consumers: + self._top_up_prefetch_locked(conn, row["consumer_id"]) + + def _release_consumer_transients_locked( + self, conn: sqlite3.Connection, consumer_id: str + ) -> None: + conn.execute("DELETE FROM waiters WHERE consumer_id=?", (consumer_id,)) + conn.execute("DELETE FROM read_leases WHERE consumer_id=?", (consumer_id,)) + conn.execute( + "DELETE FROM interests WHERE consumer_id=? AND kind='demand'", + (consumer_id,), + ) + self._prune_orphans_locked(conn) + + def _release_consumer_locked( + self, conn: sqlite3.Connection, consumer_id: str + ) -> None: + self._release_consumer_transients_locked(conn, consumer_id) + conn.execute("DELETE FROM interests WHERE consumer_id=?", (consumer_id,)) + conn.execute( + "DELETE FROM completed_samples WHERE consumer_id=?", (consumer_id,) + ) + self._prune_orphans_locked(conn) + + def complete_consumer( + self, consumer_id: str, *, allow_partial: bool = False + ) -> None: + """Mark one consumer complete and release all of its cache interests. + + The default remains strict and requires the canonical stream to drain. + ``allow_partial`` is reserved for an explicit launcher step budget: the + consumer completed its configured work, even though other consumers may + continue farther through the shared source stream. + """ + if not isinstance(allow_partial, bool): + raise TypeError("allow_partial must be a bool") + with self._transaction() as conn: + consumer = conn.execute( + "SELECT * FROM consumers WHERE consumer_id=?", (consumer_id,) + ).fetchone() + if consumer is None: + raise KeyError(f"unknown consumer {consumer_id!r}") + if consumer["state"] == "completed": + return + if consumer["state"] == "failed": + raise RuntimeError(f"cannot complete failed consumer {consumer_id!r}") + outstanding = conn.execute( + "SELECT (SELECT COUNT(*) FROM waiters WHERE consumer_id=?) + " + "(SELECT COUNT(*) FROM read_leases WHERE consumer_id=?)", + (consumer_id, consumer_id), + ).fetchone()[0] + if not allow_partial and int(consumer["cursor"]) != int( + consumer["total_samples"] + ): + raise RuntimeError( + f"consumer {consumer_id!r} completed at cursor " + f"{consumer['cursor']}/{consumer['total_samples']}" + ) + if outstanding and not allow_partial: + raise RuntimeError( + f"consumer {consumer_id!r} completed with {outstanding} " + "outstanding acquisitions" + ) + self._release_consumer_locked(conn, consumer_id) + now = self._clock() + conn.execute( + "UPDATE consumers SET state='completed',heartbeat_at=?,updated_at=? " + "WHERE consumer_id=?", + (now, now, consumer_id), + ) + + def fail_consumer(self, consumer_id: str, error: BaseException | str) -> None: + with self._transaction() as conn: + consumer = conn.execute( + "SELECT * FROM consumers WHERE consumer_id=?", (consumer_id,) + ).fetchone() + if consumer is None: + raise KeyError(f"unknown consumer {consumer_id!r}") + if consumer["state"] == "failed": + return + if consumer["state"] == "completed": + raise RuntimeError(f"cannot fail completed consumer {consumer_id!r}") + self._release_consumer_locked(conn, consumer_id) + now = self._clock() + conn.execute( + "UPDATE consumers SET state='failed',failure=?,heartbeat_at=?," + "updated_at=? WHERE consumer_id=?", + (str(error), now, now, consumer_id), + ) + + def expire_consumers(self, timeout_s: float) -> tuple[str, ...]: + if timeout_s <= 0: + raise ValueError("timeout_s must be > 0") + cutoff = self._clock() - timeout_s + with self._transaction() as conn: + rows = conn.execute( + "SELECT consumer_id FROM consumers WHERE state IN (?,?,?) AND " + "heartbeat_at < ? ORDER BY consumer_id", + (*self._ACTIVE_CONSUMER_STATES, cutoff), + ).fetchall() + expired = tuple(row["consumer_id"] for row in rows) + for consumer_id in expired: + self._release_consumer_locked(conn, consumer_id) + conn.execute( + "UPDATE consumers SET state='failed',failure=?,updated_at=? " + "WHERE consumer_id=?", + ("heartbeat expired", self._clock(), consumer_id), + ) + return expired + + def wait_for_consumers(self, timeout_s: float) -> bool: + if timeout_s <= 0: + raise ValueError("timeout_s must be > 0") + deadline = time.monotonic() + timeout_s + while True: + snapshot = self.snapshot() + if set(snapshot["consumers"]) == set(snapshot["expected_consumers"]): + return True + if time.monotonic() >= deadline: + return False + time.sleep(min(self.poll_s, max(0.0, deadline - time.monotonic()))) + + def finalize_run(self) -> str: + """Persist terminal status only after consumers and payloads drain.""" + with self._transaction() as conn: + run = self._run(conn) + consumers = conn.execute( + "SELECT consumer_id,state FROM consumers ORDER BY consumer_id" + ).fetchall() + expected = set(json.loads(run["expected_consumers_json"])) + observed = {row["consumer_id"] for row in consumers} + if observed != expected: + raise RuntimeError( + f"cannot finalize before all consumers register: " + f"missing={sorted(expected - observed)}" + ) + nonterminal = [ + row["consumer_id"] + for row in consumers + if row["state"] not in ("completed", "failed") + ] + if nonterminal: + raise RuntimeError(f"nonterminal consumers remain: {nonterminal}") + inventory = conn.execute( + "SELECT state,COUNT(*) AS count FROM entries WHERE state != ? " + "GROUP BY state", + (CaptureState.ABSENT.value,), + ).fetchall() + outstanding = conn.execute( + "SELECT (SELECT COUNT(*) FROM waiters) + " + "(SELECT COUNT(*) FROM read_leases) + " + "(SELECT COUNT(*) FROM interests)" + ).fetchone()[0] + if inventory or outstanding: + detail = {row["state"]: int(row["count"]) for row in inventory} + raise RuntimeError( + f"registry inventory did not drain: entries={detail}, " + f"outstanding={outstanding}" + ) + status = ( + "completed_with_failures" + if any(row["state"] == "failed" for row in consumers) + else "completed" + ) + conn.execute("UPDATE run_config SET status=? WHERE singleton=1", (status,)) + return status + + def snapshot(self) -> dict[str, Any]: + with self._lock: + run = self._conn.execute( + "SELECT * FROM run_config WHERE singleton=1" + ).fetchone() + if run is None: + raise RuntimeError("windowed capture run is not initialized") + consumers = self._conn.execute( + "SELECT * FROM consumers ORDER BY consumer_id" + ).fetchall() + entry_counts = self._conn.execute( + "SELECT state,COUNT(*) AS count FROM entries GROUP BY state" + ).fetchall() + queued_counts = self._conn.execute( + "SELECT priority,COUNT(*) AS count FROM entries WHERE state=? " + "GROUP BY priority", + (CaptureState.QUEUED.value,), + ).fetchall() + live_refs, live_bytes = self._live_usage_locked(self._conn) + waiters = int( + self._conn.execute("SELECT COUNT(*) FROM waiters").fetchone()[0] + ) + leases = int( + self._conn.execute("SELECT COUNT(*) FROM read_leases").fetchone()[0] + ) + interests = int( + self._conn.execute("SELECT COUNT(*) FROM interests").fetchone()[0] + ) + capture_count = self._meta_int(self._conn, "capture_count") + recapture_count = self._meta_int(self._conn, "recapture_count") + peak_live_refs = self._meta_int(self._conn, "peak_live_refs") + peak_live_bytes = self._meta_int(self._conn, "peak_live_bytes") + priority_names = { + int(CapturePriority.DEMAND): "demand", + int(CapturePriority.PREFETCH): "prefetch", + } + return { + "run_id": run["run_id"], + "status": run["status"], + "total_samples": int(run["total_samples"]), + "expected_consumers": tuple(json.loads(run["expected_consumers_json"])), + "consumers": { + row["consumer_id"]: { + "cursor": int(row["cursor"]), + "state": row["state"], + "lookbehind": int(row["lookbehind"]), + "lookahead": int(row["lookahead"]), + "prefetch_depth": int(row["prefetch_depth"]), + "max_outstanding": int(row["max_outstanding"]), + "failure": row["failure"], + } + for row in consumers + }, + "entries": {row["state"]: int(row["count"]) for row in entry_counts}, + "queued": { + priority_names[int(row["priority"])]: int(row["count"]) + for row in queued_counts + }, + "waiters": waiters, + "leases": leases, + "interests": interests, + "live_refs": live_refs, + "live_bytes": live_bytes, + "max_live_refs": self.max_live_refs, + "max_live_bytes": self.max_live_bytes, + "capture_count": capture_count, + "recapture_count": recapture_count, + "peak_live_refs": peak_live_refs, + "peak_live_bytes": peak_live_bytes, + } + + def close(self) -> None: + with self._lock: + if self._closed: + return + self._conn.close() + self._closed = True + + +__all__ = ["SQLiteWindowedCaptureRegistry"] diff --git a/specforge/runtime/data_plane/windowed_capture_runtime.py b/specforge/runtime/data_plane/windowed_capture_runtime.py new file mode 100644 index 000000000..e555dbdfe --- /dev/null +++ b/specforge/runtime/data_plane/windowed_capture_runtime.py @@ -0,0 +1,432 @@ +# coding=utf-8 +# Copyright 2024 The SpecForge team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +"""Process-safe runtime loops for consumer-driven capture windows.""" + +from __future__ import annotations + +import dataclasses +import sqlite3 +import threading +import time +from dataclasses import dataclass, field +from typing import Any, Callable, Mapping, Optional, Sequence + +from specforge.inference.capture import CaptureConfig +from specforge.runtime.contracts import PromptTask, SampleRef, assert_no_tensors +from specforge.runtime.data_plane.windowed_capture import ( + CaptureRequest, + SQLiteWindowedCaptureRegistry, +) + + +@dataclass +class WindowedConsumerControl: + """Keep one logical consumer live across model initialization and training.""" + + registry: SQLiteWindowedCaptureRegistry + consumer_id: str + heartbeat_interval_s: float + _stop: threading.Event = field(default_factory=threading.Event, init=False) + _ready: threading.Event = field(default_factory=threading.Event, init=False) + _thread: Optional[threading.Thread] = field(default=None, init=False) + _error: Optional[BaseException] = field(default=None, init=False) + _started: bool = field(default=False, init=False) + _terminal: bool = field(default=False, init=False) + + def __post_init__(self) -> None: + if self.heartbeat_interval_s <= 0: + raise ValueError("heartbeat_interval_s must be > 0") + + def _heartbeat_loop(self) -> None: + consecutive_transient_failures = 0 + while not self._stop.wait(self.heartbeat_interval_s): + try: + self.registry.heartbeat(self.consumer_id, ready=self._ready.is_set()) + consecutive_transient_failures = 0 + except sqlite3.OperationalError as exc: + consecutive_transient_failures += 1 + if consecutive_transient_failures < 3: + continue + self._error = exc + return + except BaseException as exc: # exposed synchronously by ensure_healthy + # WindowedCaptureQueue may observe EOF and complete the consumer + # between heartbeat intervals. That terminal race is success. + try: + state = self.registry.snapshot()["consumers"][self.consumer_id][ + "state" + ] + except BaseException: + state = None + if state == "completed": + return + self._error = exc + return + + def start(self) -> "WindowedConsumerControl": + if self._started: + return self + if self._terminal: + raise RuntimeError("cannot start terminal windowed consumer control") + self.registry.heartbeat(self.consumer_id, ready=False) + self._thread = threading.Thread( + target=self._heartbeat_loop, + name=f"windowed-heartbeat-{self.consumer_id}", + daemon=True, + ) + self._started = True + self._thread.start() + return self + + def mark_ready(self) -> None: + self.start() + self.ensure_healthy() + self._ready.set() + self.registry.heartbeat(self.consumer_id, ready=True) + + def _stop_heartbeat(self) -> None: + self._stop.set() + if self._thread is None or not self._thread.is_alive(): + return + self._thread.join(timeout=max(1.0, self.heartbeat_interval_s * 2)) + if self._thread.is_alive(): + raise TimeoutError(f"consumer {self.consumer_id!r} heartbeat did not stop") + + def ensure_healthy(self) -> None: + if self._error is not None: + raise RuntimeError( + f"consumer {self.consumer_id!r} registry heartbeat failed" + ) from self._error + + def complete(self, *, allow_partial: bool = False) -> None: + if self._terminal: + return + self._stop_heartbeat() + self.ensure_healthy() + self.registry.complete_consumer(self.consumer_id, allow_partial=allow_partial) + self._terminal = True + + def fail(self, error: BaseException | str) -> None: + if self._terminal: + return + try: + self._stop_heartbeat() + except BaseException as stop_error: + if isinstance(error, BaseException): + error.add_note(f"failed to stop registry heartbeat: {stop_error!r}") + self.registry.fail_consumer(self.consumer_id, error) + self._terminal = True + + def close(self) -> None: + self._stop_heartbeat() + + +def start_windowed_consumer_control( + registry: SQLiteWindowedCaptureRegistry, + consumer_id: str, + *, + lookbehind: int, + lookahead: int, + prefetch_depth: int, + max_outstanding: int, + heartbeat_interval_s: float, + durable_cursor: Optional[int] = None, +) -> WindowedConsumerControl: + """Register a new consumer or reconcile an interrupted active consumer.""" + snapshot = registry.snapshot() + existing = snapshot["consumers"].get(consumer_id) + cursor = 0 if durable_cursor is None else durable_cursor + if existing is None: + registry.register_consumer( + consumer_id, + lookbehind=lookbehind, + lookahead=lookahead, + prefetch_depth=prefetch_depth, + max_outstanding=max_outstanding, + cursor=cursor, + ) + else: + expected = (lookbehind, lookahead, prefetch_depth, max_outstanding) + observed = tuple( + int(existing[name]) + for name in ( + "lookbehind", + "lookahead", + "prefetch_depth", + "max_outstanding", + ) + ) + if observed != expected: + raise RuntimeError( + f"consumer {consumer_id!r} window mismatch: " + f"expected={expected}, observed={observed}" + ) + if existing["state"] in ("completed", "failed"): + raise RuntimeError( + f"cannot resume terminal consumer {consumer_id!r}: " + f"{existing['state']}" + ) + registry.resume_consumer( + consumer_id, + durable_cursor=( + int(existing["cursor"]) if durable_cursor is None else cursor + ), + ) + return WindowedConsumerControl( + registry=registry, + consumer_id=consumer_id, + heartbeat_interval_s=heartbeat_interval_s, + ).start() + + +class WindowedCaptureService: + """Demand-driven capture producer and sole payload-reclamation owner.""" + + def __init__( + self, + registry: SQLiteWindowedCaptureRegistry, + *, + prompts: Sequence[PromptTask], + feature_source: Any, + capture: CaptureConfig, + owner_store: Any, + capture_batch_size: int = 8, + batch_wait_s: float = 0.002, + max_capture_retries: int = 2, + retry_backoff_s: float = 0.05, + consumer_registration_timeout_s: float = 600.0, + consumer_heartbeat_timeout_s: float = 120.0, + poll_s: float = 0.01, + reclaim_batch_size: int = 64, + ) -> None: + for name, value in ( + ("capture_batch_size", capture_batch_size), + ("reclaim_batch_size", reclaim_batch_size), + ): + if isinstance(value, bool) or not isinstance(value, int) or value < 1: + raise ValueError(f"{name} must be a positive integer") + if max_capture_retries < 0: + raise ValueError("max_capture_retries must be >= 0") + for name, value in ( + ("batch_wait_s", batch_wait_s), + ("retry_backoff_s", retry_backoff_s), + ("consumer_registration_timeout_s", consumer_registration_timeout_s), + ("consumer_heartbeat_timeout_s", consumer_heartbeat_timeout_s), + ("poll_s", poll_s), + ): + if value < 0 or ( + name not in {"batch_wait_s", "retry_backoff_s"} and value == 0 + ): + raise ValueError(f"{name} has invalid duration {value}") + if not callable(getattr(feature_source, "produce_refs", None)): + raise TypeError( + "feature_source must expose produce_refs(tasks, capture=...)" + ) + if not callable(getattr(owner_store, "reclaim", None)): + raise TypeError("owner_store must expose reclaim(ref, reason=...)") + if hasattr(owner_store, "lifetime_owner") and not owner_store.lifetime_owner: + raise ValueError( + "windowed producer must use the feature-store lifetime owner" + ) + + prompt_by_id: dict[str, PromptTask] = {} + for prompt in prompts: + if not isinstance(prompt, PromptTask): + raise TypeError("prompts must contain PromptTask values") + assert_no_tensors(prompt) + if prompt.task_id in prompt_by_id: + raise ValueError(f"duplicate prompt task_id {prompt.task_id!r}") + prompt_by_id[prompt.task_id] = prompt + if not prompt_by_id: + raise ValueError("prompts must not be empty") + + self.registry = registry + self.prompts = prompt_by_id + self.feature_source = feature_source + self.capture = capture + self.owner_store = owner_store + self.capture_batch_size = capture_batch_size + self.batch_wait_s = batch_wait_s + self.max_capture_retries = max_capture_retries + self.retry_backoff_s = retry_backoff_s + self.consumer_registration_timeout_s = consumer_registration_timeout_s + self.consumer_heartbeat_timeout_s = consumer_heartbeat_timeout_s + self.poll_s = poll_s + self.reclaim_batch_size = reclaim_batch_size + self.capture_calls = 0 + self.captured_refs = 0 + self.capture_failures = 0 + + def _task(self, request: CaptureRequest) -> PromptTask: + try: + prompt = self.prompts[request.key.source_sample_id] + except KeyError as exc: + raise KeyError( + f"capture registry requested unknown source " + f"{request.key.source_sample_id!r}" + ) from exc + return dataclasses.replace( + prompt, + attempt=request.generation - 1, + metadata={ + **prompt.metadata, + "capture_generation": request.generation, + }, + ) + + def _fail_requests( + self, + requests: Sequence[CaptureRequest], + error: BaseException | str, + *, + retryable: bool, + ) -> bool: + retried = False + for request in requests: + retried |= self.registry.fail_capture( + request, + error, + retryable=retryable, + max_retries=self.max_capture_retries, + ) + self.capture_failures += 1 + return retried + + def _capture_requests(self, requests: Sequence[CaptureRequest]) -> None: + tasks = [self._task(request) for request in requests] + self.capture_calls += 1 + try: + results = self.feature_source.produce_refs(tasks, capture=self.capture) + except BaseException as exc: + retried = self._fail_requests(requests, exc, retryable=True) + if retried and self.retry_backoff_s: + time.sleep(self.retry_backoff_s) + return + if len(results) != len(requests): + self._fail_requests( + requests, + RuntimeError( + f"produce_refs returned {len(results)} results for " + f"{len(requests)} requests" + ), + retryable=False, + ) + return + + retried = False + for request, result in zip(requests, results): + if not isinstance(result, SampleRef): + result_task_id = getattr( + result, "task_id", request.key.source_sample_id + ) + reason = getattr(result, "reason", repr(result)) + aligned = result_task_id == request.key.source_sample_id + retried |= self._fail_requests( + (request,), + ( + reason + if aligned + else f"misaligned capture result: {result_task_id!r}" + ), + retryable=aligned and bool(getattr(result, "retryable", False)), + ) + continue + result = dataclasses.replace( + result, + metadata={ + **result.metadata, + "window_generation": request.generation, + }, + ) + try: + self.registry.mark_committing(request, result) + self.registry.complete_capture(request, result) + self.captured_refs += 1 + except BaseException as exc: + try: + self.owner_store.reclaim( + result, reason="windowed-registry-commit-failed" + ) + except BaseException as cleanup_error: + exc.add_note( + f"failed to reclaim rejected capture: {cleanup_error!r}" + ) + retried |= self._fail_requests((request,), exc, retryable=False) + if retried and self.retry_backoff_s: + time.sleep(self.retry_backoff_s) + + @staticmethod + def _all_consumers_terminal(snapshot: Mapping[str, Any]) -> bool: + consumers = snapshot["consumers"] + return set(consumers) == set(snapshot["expected_consumers"]) and all( + value["state"] in ("completed", "failed") for value in consumers.values() + ) + + def _reclaim(self, *, pressure: bool = False) -> int: + reclaimed = self.registry.reclaim( + self.owner_store, + limit=self.reclaim_batch_size, + pressure=pressure, + reason="window-pressure" if pressure else "window-expired", + ) + gc = getattr(self.owner_store, "gc", None) + if callable(gc): + gc() + return reclaimed + + def drive( + self, + *, + should_stop: Optional[Callable[[], bool]] = None, + max_rounds: int = 10_000_000, + ) -> int: + """Serve capture demand until every configured consumer is terminal.""" + if max_rounds < 1: + raise ValueError("max_rounds must be >= 1") + if not self.registry.wait_for_consumers(self.consumer_registration_timeout_s): + raise TimeoutError("windowed consumers did not register before deadline") + + for _ in range(max_rounds): + if should_stop is not None and should_stop(): + raise RuntimeError("windowed capture service was stopped") + self.registry.expire_consumers(self.consumer_heartbeat_timeout_s) + progressed = bool(self._reclaim()) + snapshot = self.registry.snapshot() + if self._all_consumers_terminal(snapshot): + while self._reclaim(): + pass + self.registry.finalize_run() + return self.captured_refs + + if snapshot["queued"] and self.batch_wait_s: + time.sleep(self.batch_wait_s) + requests = self.registry.claim_batch(self.capture_batch_size) + if requests: + self._capture_requests(requests) + progressed = True + elif snapshot["queued"]: + progressed = bool(self._reclaim(pressure=True)) or progressed + if not progressed: + time.sleep(self.poll_s) + raise RuntimeError(f"windowed capture exceeded max_rounds={max_rounds}") + + def snapshot(self) -> dict[str, Any]: + return { + "registry": self.registry.snapshot(), + "capture_calls": self.capture_calls, + "captured_refs": self.captured_refs, + "capture_failures": self.capture_failures, + } + + +__all__ = [ + "WindowedCaptureService", + "WindowedConsumerControl", + "start_windowed_consumer_control", +] diff --git a/specforge/training/disaggregated.py b/specforge/training/disaggregated.py index 8eebf9cf2..8d0207bb0 100644 --- a/specforge/training/disaggregated.py +++ b/specforge/training/disaggregated.py @@ -17,9 +17,11 @@ from __future__ import annotations +import hashlib import json import os import time +from collections.abc import Mapping from typing import Callable, Optional, Sequence from specforge.algorithms.registry import AlgorithmRegistration @@ -39,6 +41,59 @@ _OFFLINE_CONTROL_SUFFIXES = (".done", ".consumed", ".failed", ".consumer_failed") +def _stabilize_windowed_prompts(prompts): + """Add deterministic task ids at the fixed windowed-inventory boundary.""" + from specforge.runtime.contracts import PromptTask + + stabilized = [] + seen_ids = set() + for index, prompt in enumerate(prompts): + if isinstance(prompt, PromptTask): + task_id = prompt.task_id + stabilized_prompt = prompt + elif isinstance(prompt, Mapping): + task_id = prompt.get("task_id") + stabilized_prompt = prompt + if task_id is None: + identity = { + key: value for key, value in prompt.items() if key != "task_id" + } + try: + encoded = json.dumps( + identity, + allow_nan=False, + ensure_ascii=True, + separators=(",", ":"), + sort_keys=True, + ).encode() + except (TypeError, ValueError) as exc: + raise ValueError( + "windowed fanout cannot derive a stable task_id from " + f"prompt {index}; non-JSON prompts must provide one explicitly" + ) from exc + digest = hashlib.sha256(encoded).hexdigest() + task_id = f"canonical-prompt-{index:08d}-{digest}" + stabilized_prompt = dict(prompt) + stabilized_prompt["task_id"] = task_id + else: + raise TypeError( + "windowed fanout prompts must be PromptTask instances or mappings; " + f"got {type(prompt).__name__} at index {index}" + ) + + if not isinstance(task_id, str) or not task_id: + raise ValueError( + f"windowed fanout prompt {index} has an invalid explicit task_id" + ) + if task_id in seen_ids: + raise ValueError( + f"windowed fanout prompt task_id {task_id!r} is duplicated" + ) + seen_ids.add(task_id) + stabilized.append(stabilized_prompt) + return stabilized + + def _write_control(path: str, value: str = "") -> None: """Atomically publish one small filesystem control record.""" os.makedirs(os.path.dirname(os.path.abspath(path)), exist_ok=True) @@ -125,7 +180,12 @@ def _env(name: str) -> str: return value -def _mooncake_store(cfg: Config, *, retain_on_release: bool = False): +def _mooncake_store( + cfg: Config, + *, + retain_on_release: bool = False, + lifetime_owner: bool = True, +): from specforge.runtime.data_plane.disaggregated import AuthPolicy from specforge.runtime.data_plane.mooncake_store import MooncakeFeatureStore @@ -145,12 +205,24 @@ def _mooncake_store(cfg: Config, *, retain_on_release: bool = False): ): if os.environ.get(env_name): setup_kwargs[key] = int(os.environ[env_name]) + deployment = cfg.deployment.disaggregated + fanout = deployment.windowed_fanout if deployment is not None else None + lifecycle_db_path = ( + os.environ.get("DISAGG_CAPTURE_LIFECYCLE_DB") + if fanout is not None and lifetime_owner + else None + ) return MooncakeFeatureStore( store_id=os.environ.get("DISAGG_STORE_ID", cfg.run_id), setup_kwargs=setup_kwargs, auth=AuthPolicy(token), credential=token, retain_on_release=retain_on_release, + lifetime_owner=lifetime_owner, + lifecycle_db_path=lifecycle_db_path, + max_resident_bytes=( + fanout.max_live_bytes if fanout is not None and lifetime_owner else None + ), ) @@ -501,6 +573,277 @@ def _producer_capture_metadata(cfg: Config, algorithm: AlgorithmRegistration): ) +def _build_windowed_online( + cfg: Config, + *, + algorithm: AlgorithmRegistration, + build_model_bundle: Callable, + prepare_prompts: Callable, + optimizer_factory: Callable, + logger: Callable, +): + """Assemble one role in a bounded independent-consumer fanout.""" + from specforge.training.assembly import TrainingRun, _load_input_tools + + deployment = cfg.deployment.disaggregated + assert deployment is not None and deployment.windowed_fanout is not None + fanout = deployment.windowed_fanout + registry_db_path = _env("DISAGG_WINDOW_REGISTRY") + modality = cfg.model.input_modality + streaming = algorithm.providers.server_streaming_for(modality) + layers, hidden_size, target_vocab, draft_vocab = _producer_capture_metadata( + cfg, algorithm + ) + from specforge.launch import build_disagg_windowed_capture_contract + + capture, contract_digest = build_disagg_windowed_capture_contract( + strategy=algorithm, + modality=modality, + target_hidden_size=hidden_size, + target_model_version=cfg.model.target_model_path, + tokenizer_version=cfg.model.target_model_path, + target_vocab_size=target_vocab, + draft_vocab_size=draft_vocab, + target_repr=streaming.target_representation, + aux_hidden_state_layer_ids=layers, + vocab_map_version=cfg.model.vocab_mapping_path or None, + ) + + if cfg.training.role == "producer": + from specforge.inference.adapters.server_capture import ( + ServerCaptureSchema, + SGLangServerCaptureAdapter, + ) + from specforge.launch import build_disagg_online_windowed_producer + from specforge.runtime.data_plane.feature_store import ( + drain_feature_store_removals, + ) + from specforge.training.model_loading import resolve_draft_config + + input_adapter = streaming.create_input_adapter(cfg) + input_tools = _load_input_tools(cfg, algorithm, input_adapter=input_adapter) + draft_config = resolve_draft_config( + cfg, provider=algorithm.providers.model.draft_config + ) + if input_adapter is None: + prompts = prepare_prompts(cfg, input_tools, draft_config=draft_config) + else: + prompts = input_adapter.prepare_prompts( + cfg, input_tools, draft_config=draft_config + ) + if len(prompts) != cfg.data.max_prompts: + raise ValueError( + "windowed_fanout prepared prompt count does not match " + f"data.max_prompts: {len(prompts)} != {cfg.data.max_prompts}" + ) + prompts = _stabilize_windowed_prompts(prompts) + urls = _server_urls(cfg) + if len(urls) != 1: + raise ValueError( + "windowed_fanout currently requires exactly one capture server" + ) + store = _mooncake_store(cfg, lifetime_owner=True) + layout = streaming.layout + adapter = SGLangServerCaptureAdapter( + urls[0], + store, + run_id=cfg.run_id, + algorithm=algorithm.name, + schema=ServerCaptureSchema( + aux_feature=layout.aux_feature, + last_hidden_feature=layout.last_hidden_feature, + passthrough=layout.passthrough, + attention_mask_feature=layout.attention_mask_feature, + ), + request_input_adapter=input_adapter, + target_model_version=cfg.model.target_model_path, + ) + runtime = build_disagg_online_windowed_producer( + prompts=prompts, + feature_store=store, + feature_source=adapter, + run_id=cfg.run_id, + consumer_ids=tuple(consumer.consumer_id for consumer in fanout.consumers), + registry_db_path=registry_db_path, + max_live_refs=fanout.max_live_refs, + max_live_bytes=fanout.max_live_bytes, + capture_reservation_bytes=fanout.capture_reservation_bytes, + target_hidden_size=hidden_size, + target_model_version=cfg.model.target_model_path, + tokenizer_version=cfg.model.target_model_path, + strategy=algorithm, + modality=modality, + target_vocab_size=target_vocab, + draft_vocab_size=draft_vocab, + target_repr=streaming.target_representation, + aux_hidden_state_layer_ids=layers, + vocab_map_version=cfg.model.vocab_mapping_path or None, + capture_batch_size=fanout.capture_batch_size, + capture_batch_wait_s=fanout.capture_batch_wait_s, + max_capture_retries=fanout.max_capture_retries, + retry_backoff_s=fanout.capture_retry_backoff_s, + consumer_registration_timeout_s=(fanout.consumer_registration_timeout_s), + consumer_heartbeat_timeout_s=fanout.consumer_heartbeat_timeout_s, + registry_poll_s=fanout.registry_poll_s, + ) + + def produce() -> int: + produced = 0 + primary_error: Optional[BaseException] = None + try: + produced = runtime.drive() + except BaseException as exc: + primary_error = exc + cleanup_errors = [] + try: + runtime.close() + except Exception as exc: + cleanup_errors.append(f"registry close: {type(exc).__name__}: {exc}") + try: + store.abort_all( + reason=( + "windowed-attempt-failed" + if primary_error is not None + else "windowed-attempt-finished" + ), + force=True, + ) + drain_feature_store_removals(store) + except Exception as exc: + cleanup_errors.append( + f"Mooncake owner cleanup: {type(exc).__name__}: {exc}" + ) + if primary_error is not None and cleanup_errors: + raise RuntimeError( + f"windowed producer failed ({type(primary_error).__name__}: " + f"{primary_error}) and cleanup also failed: {cleanup_errors}" + ) from primary_error + if primary_error is not None: + raise primary_error + if cleanup_errors: + raise RuntimeError( + f"windowed producer cleanup failed: {cleanup_errors}" + ) + return produced + + return TrainingRun(execute=produce) + + consumer_id = _env("SPECFORGE_FANOUT_CONSUMER_ID") + consumer = fanout.consumer(consumer_id) + lookbehind, lookahead, prefetch_depth = fanout.window_for(consumer_id) + from specforge.runtime.data_plane.windowed_capture import ( + SQLiteWindowedCaptureRegistry, + ) + from specforge.runtime.data_plane.windowed_capture_runtime import ( + start_windowed_consumer_control, + ) + + registry = SQLiteWindowedCaptureRegistry( + registry_db_path, + max_live_refs=fanout.max_live_refs, + max_live_bytes=fanout.max_live_bytes, + capture_reservation_bytes=fanout.capture_reservation_bytes, + poll_s=fanout.registry_poll_s, + ) + control = None + try: + initialized = registry.wait_initialized(fanout.consumer_registration_timeout_s) + expected = (cfg.run_id, contract_digest, cfg.data.max_prompts) + observed = ( + initialized["run_id"], + initialized["contract_digest"], + initialized["total_samples"], + ) + if observed != expected: + raise RuntimeError( + "windowed fanout registry identity mismatch: " + f"expected={expected!r}, observed={observed!r}" + ) + control = start_windowed_consumer_control( + registry, + consumer_id, + lookbehind=lookbehind, + lookahead=lookahead, + prefetch_depth=prefetch_depth, + max_outstanding=fanout.max_outstanding_per_consumer, + heartbeat_interval_s=fanout.consumer_heartbeat_interval_s, + ) + bundle = build_model_bundle(cfg) + total_steps = cfg.data.max_prompts // ( + cfg.training.batch_size * cfg.training.accumulation_steps + ) + from specforge.launch import build_disagg_online_windowed_consumer + + runtime = build_disagg_online_windowed_consumer( + consumer_id=consumer_id, + registry_db_path=registry_db_path, + max_live_refs=fanout.max_live_refs, + max_live_bytes=fanout.max_live_bytes, + capture_reservation_bytes=fanout.capture_reservation_bytes, + contract_digest=contract_digest, + total_samples=cfg.data.max_prompts, + feature_store=_mooncake_store(cfg, lifetime_owner=False), + draft_model=bundle.model, + optimizer_factory=optimizer_factory(cfg), + run_id=cfg.run_id, + output_dir=cfg.output_dir, + metadata_db_path=_env("DISAGG_DB"), + lookbehind=lookbehind, + lookahead=lookahead, + prefetch_depth=prefetch_depth, + max_outstanding=fanout.max_outstanding_per_consumer, + strategy=algorithm, + modality=modality, + batch_size=cfg.training.batch_size, + accumulation_steps=cfg.training.accumulation_steps, + num_epochs=1, + max_steps=cfg.training.max_steps, + total_steps=cfg.training.total_steps or total_steps, + save_interval=cfg.training.save_interval, + eval_interval=0, + idle_timeout_s=fanout.consumer_idle_timeout_s, + logger=logger, + log_interval=cfg.training.log_interval, + strategy_kwargs=dict(bundle.strategy_kwargs), + resume=consumer.resume_from is not None, + resume_from=consumer.resume_from, + max_checkpoints=cfg.training.max_checkpoints, + heartbeat_interval_s=fanout.consumer_heartbeat_interval_s, + initialization_timeout_s=fanout.consumer_registration_timeout_s, + registry_poll_s=fanout.registry_poll_s, + loader_prefetch_batches=fanout.consumer_prefetch_batches, + consumer_control=control, + ) + except BaseException as exc: + if control is not None: + try: + control.fail(exc) + except BaseException as cleanup_error: + exc.add_note( + f"failed to report fanout consumer setup failure: " + f"{cleanup_error!r}" + ) + try: + control.close() + except BaseException as cleanup_error: + exc.add_note( + f"failed to close fanout consumer control: {cleanup_error!r}" + ) + try: + registry.close() + except BaseException as cleanup_error: + exc.add_note(f"failed to close fanout registry: {cleanup_error!r}") + raise + + def consume() -> int: + try: + return runtime.run() + finally: + runtime.close() + + return TrainingRun(execute=consume) + + def _build_online( cfg: Config, *, @@ -510,6 +853,16 @@ def _build_online( optimizer_factory: Callable, logger: Callable, ): + deployment = cfg.deployment.disaggregated + if deployment is not None and deployment.windowed_fanout is not None: + return _build_windowed_online( + cfg, + algorithm=algorithm, + build_model_bundle=build_model_bundle, + prepare_prompts=prepare_prompts, + optimizer_factory=optimizer_factory, + logger=logger, + ) from specforge.training.assembly import ( TrainingRun, _dataloader_num_workers, diff --git a/specforge/training/trainer.py b/specforge/training/trainer.py index b1279a708..0cb57d37f 100644 --- a/specforge/training/trainer.py +++ b/specforge/training/trainer.py @@ -508,6 +508,10 @@ def micro_step(self) -> int: def last_checkpoint_step(self) -> Optional[int]: return self._controller.last_checkpoint_step + @property + def loader(self): + return self._loader + def fit(self) -> int: """Run training and configured evaluation through one lifecycle.""" loader_close_attempted = False diff --git a/tests/test_config/test_launch_topology.py b/tests/test_config/test_launch_topology.py index c703ae8c0..069b96061 100644 --- a/tests/test_config/test_launch_topology.py +++ b/tests/test_config/test_launch_topology.py @@ -40,6 +40,7 @@ "qwen3-4b-eagle3-online.yaml": 1, "qwen3-8b-dflash-disaggregated.yaml": 4, "qwen3-8b-dflash-1server-dp7-disaggregated.yaml": 7, + "qwen3-8b-dflash-windowed-fanout.yaml": 1, "qwen3-8b-dflash-online.yaml": 8, "qwen3-8b-domino-1server-dp7-disaggregated.yaml": 7, "qwen3-8b-domino-disaggregated.yaml": 4, @@ -94,6 +95,75 @@ "server_urls": ["http://127.0.0.1:30000"], **LOCAL_MOONCAKE_ENDPOINTS, }, + "qwen3-8b-dflash-windowed-fanout.yaml": { + "control_dir": "./outputs/qwen3-8b-dflash-windowed-fanout/control", + "backend": "mooncake", + "client_buffer_size": 1073741824, + "windowed_fanout": { + "window_lookbehind": 2, + "window_lookahead": 16, + "max_prefetch_per_consumer": 8, + "max_outstanding_per_consumer": 8, + "max_live_refs": 48, + "max_live_bytes": 25769803776, + "capture_reservation_bytes": 536870912, + "capture_max_sample_bytes": 536870912, + "capture_batch_size": 8, + "consumer_prefetch_batches": 1, + "consumers": [ + { + "consumer_id": "dflash-b4", + "seed": 42, + "loss_type": "dflash", + "loss_decay_gamma": 7.0, + "dpace_alpha": 0.5, + "draft_block_size": 4, + "num_anchors": 64, + "learning_rate": 0.0006, + "warmup_ratio": 0.04, + }, + { + "consumer_id": "dflash-b8", + "seed": 43, + "loss_type": "dflash", + "loss_decay_gamma": 7.0, + "dpace_alpha": 0.5, + "draft_block_size": 8, + "num_anchors": 128, + "learning_rate": 0.0006, + "warmup_ratio": 0.04, + }, + { + "consumer_id": "dflash-b16", + "seed": 44, + "loss_type": "dflash", + "loss_decay_gamma": 7.0, + "dpace_alpha": 0.5, + "draft_block_size": 16, + "num_anchors": 256, + "learning_rate": 0.0006, + "warmup_ratio": 0.04, + }, + ], + }, + "managed_local": { + "trainer_cuda_visible_devices": ["1", "2", "3"], + "shutdown_grace_s": 120, + "mooncake": { + "protocol": "tcp", + "global_segment_size_bytes": 34359738368, + "local_buffer_size_bytes": 1073741824, + }, + "capture_servers": [ + { + "port": 30000, + "cuda_visible_devices": ["0"], + "tp_size": 1, + "mem_fraction_static": 0.5, + } + ], + }, + }, "qwen3-8b-dflash-1server-dp7-disaggregated.yaml": { "control_dir": ("outputs/qwen3-8b-dflash-1server-dp7-disaggregated/control"), "backend": "mooncake", @@ -267,7 +337,7 @@ def _recipes() -> dict[str, Path]: class ExampleLaunchTopologyTest(unittest.TestCase): def test_every_recipe_has_the_explicit_golden_topology(self): recipes = _recipes() - self.assertEqual(len(EXPECTED_NPROC_PER_NODE), 56) + self.assertEqual(len(EXPECTED_NPROC_PER_NODE), 57) self.assertEqual(set(recipes), set(EXPECTED_NPROC_PER_NODE)) for filename, nproc_per_node in EXPECTED_NPROC_PER_NODE.items(): diff --git a/tests/test_config/test_schema.py b/tests/test_config/test_schema.py index 07e7ea6fc..6ee266675 100644 --- a/tests/test_config/test_schema.py +++ b/tests/test_config/test_schema.py @@ -63,6 +63,61 @@ def _managed_local_payload(*, ep_size: int) -> dict: return payload +def _windowed_fanout_payload(*, managed: bool = True) -> dict: + payload = _online_payload("dflash") + payload["data"]["max_prompts"] = 8 + payload["training"].update( + {"batch_size": 2, "accumulation_steps": 1, "num_epochs": 1} + ) + consumers = [ + { + "consumer_id": "block-4", + "seed": 42, + "loss_type": "dflash", + "loss_decay_gamma": 7.0, + "dpace_alpha": 0.5, + "draft_block_size": 4, + "num_anchors": 128, + "learning_rate": 0.0006, + "warmup_ratio": 0.04, + }, + { + "consumer_id": "block-16", + "seed": 43, + "loss_type": "dpace", + "loss_decay_gamma": None, + "dpace_alpha": 0.5, + "draft_block_size": 16, + "num_anchors": 256, + "learning_rate": 0.0006, + "warmup_ratio": 0.04, + }, + ] + disaggregated = payload["deployment"]["disaggregated"] + disaggregated["windowed_fanout"] = { + "consumers": consumers, + "max_live_bytes": 8 << 30, + "max_outstanding_per_consumer": 4, + } + payload["deployment"]["trainer"] = {"nnodes": 1, "nproc_per_node": 1} + if managed: + disaggregated.pop("server_urls") + disaggregated["managed_local"] = { + "trainer_cuda_visible_devices": ["1", "2"], + "capture_servers": [ + { + "port": 30000, + "cuda_visible_devices": ["0"], + "tp_size": 1, + } + ], + } + else: + consumers[0]["cuda_visible_device"] = "1" + consumers[1]["cuda_visible_device"] = "2" + return payload + + def _write(payload: dict, suffix: str) -> str: fd, path = tempfile.mkstemp(suffix=suffix) with os.fdopen(fd, "w") as f: @@ -76,6 +131,58 @@ def _write(payload: dict, suffix: str) -> str: class ConfigSchemaTest(unittest.TestCase): + def test_windowed_fanout_is_typed_and_bounded(self): + config = Config.model_validate(_windowed_fanout_payload()) + fanout = config.deployment.disaggregated.windowed_fanout + + self.assertEqual( + [consumer.consumer_id for consumer in fanout.consumers], + ["block-4", "block-16"], + ) + self.assertEqual(fanout.window_for("block-4"), (2, 40, 8)) + self.assertEqual(fanout.max_live_bytes, 8 << 30) + + def test_windowed_fanout_rejects_noncanonical_topologies(self): + invalid = _windowed_fanout_payload() + invalid["deployment"]["trainer"]["nproc_per_node"] = 2 + with self.assertRaisesRegex(ValidationError, "independent single-process"): + Config.model_validate(invalid) + + invalid = _windowed_fanout_payload(managed=False) + del invalid["deployment"]["disaggregated"]["windowed_fanout"]["consumers"][1][ + "cuda_visible_device" + ] + with self.assertRaisesRegex(ValidationError, "require cuda_visible_device"): + Config.model_validate(invalid) + + invalid = _windowed_fanout_payload() + invalid["deployment"]["disaggregated"]["managed_local"][ + "trainer_cuda_visible_devices" + ] = ["1"] + with self.assertRaisesRegex(ValidationError, "fanout consumer count"): + Config.model_validate(invalid) + + invalid = _windowed_fanout_payload() + invalid["deployment"]["disaggregated"]["windowed_fanout"].update( + capture_reservation_bytes=128 << 20, + capture_max_sample_bytes=256 << 20, + ) + with self.assertRaisesRegex(ValidationError, "must not exceed"): + Config.model_validate(invalid) + + invalid = _windowed_fanout_payload() + invalid["deployment"]["disaggregated"]["managed_local"][ + "capture_servers" + ].append( + { + "port": 30001, + "cuda_visible_devices": ["3"], + "tp_size": 1, + } + ) + with self.assertRaisesRegex(ValidationError, "exactly one capture server"): + Config.model_validate(invalid) + def test_liger_kernel_flag_is_typed_and_defaults_off(self): default = Config.model_validate(copy.deepcopy(MINIMAL)) self.assertFalse(default.model.use_liger_kernel) diff --git a/tests/test_config/test_unified_feature_reachability.py b/tests/test_config/test_unified_feature_reachability.py index dd0fb2b49..cdfb8db50 100644 --- a/tests/test_config/test_unified_feature_reachability.py +++ b/tests/test_config/test_unified_feature_reachability.py @@ -151,7 +151,7 @@ def test_all_example_configs_validate_through_the_typed_entry(self): for path in EXAMPLE_CONFIG_DIR.glob("*.yaml") if not path.name.startswith(".") ) - self.assertEqual(len(paths), 56) + self.assertEqual(len(paths), 57) resolved_runs = { path.name: resolve_run(Config.from_file(str(path))) for path in paths diff --git a/tests/test_runtime/test_controller_no_tensor.py b/tests/test_runtime/test_controller_no_tensor.py index cc08898bf..3162e0443 100644 --- a/tests/test_runtime/test_controller_no_tensor.py +++ b/tests/test_runtime/test_controller_no_tensor.py @@ -106,6 +106,29 @@ def test_commit_samples_idempotent(self): self.assertEqual(ctrl.status()["samples_committed"], 1) self.assertEqual(ctrl.sample_queue.depth(), 1) + def test_record_external_refs_ledgers_without_enqueueing(self): + ctrl = DataFlowController("run1") + ref = _ref(0) + + self.assertEqual(ctrl.record_external_refs([ref]), 1) + self.assertEqual(ctrl.record_external_refs([ref]), 0) + self.assertEqual(ctrl.store.all_committed_ids(), [ref.sample_id]) + self.assertEqual(ctrl.sample_queue.depth(), 0) + + bad = SampleRef( + sample_id="bad", + run_id="r", + source_task_id=None, + feature_store_uri="mem://x/bad", + feature_keys={}, + feature_specs={}, + strategy="eagle3", + metadata={"sneaky": torch.zeros(1)}, + ) + with self.assertRaises(TypeError): + ctrl.record_external_refs([bad]) + self.assertEqual(ctrl.status()["samples_committed"], 1) + def test_prompt_lease_and_commit_clears_lease(self): ctrl = DataFlowController("run1") ids = ctrl.ingest_prompts([{"payload": {"text": "hi"}, "max_length": 16}]) diff --git a/tests/test_runtime/test_disagg_multiserver.py b/tests/test_runtime/test_disagg_multiserver.py index 444d78d7b..0b2bd57f3 100644 --- a/tests/test_runtime/test_disagg_multiserver.py +++ b/tests/test_runtime/test_disagg_multiserver.py @@ -66,6 +66,7 @@ def _adapter(store, post_fn, url="http://server:30000"): store, run_id="run0", algorithm=ALGORITHM.name, + capture_token="unit-capture-token", schema=ServerCaptureSchema( aux_feature=layout.aux_feature, last_hidden_feature=layout.last_hidden_feature, diff --git a/tests/test_runtime/test_disaggregated.py b/tests/test_runtime/test_disaggregated.py index f00640104..b341f50ea 100644 --- a/tests/test_runtime/test_disaggregated.py +++ b/tests/test_runtime/test_disaggregated.py @@ -58,6 +58,24 @@ def test_consumer_release_of_stale_handle_does_not_free_fresh_reput(self): out, _ = consumer.get(ref2) # gen2 must still be intact self.assertEqual(out["x"].sum().item(), 4.0) # ones(1,4) intact + def test_fanout_reader_release_defers_free_to_lifetime_owner(self): + producer = SharedDirFeatureStore(self.root, store_id="st", lifetime_owner=True) + reader = SharedDirFeatureStore(self.root, store_id="st", lifetime_owner=False) + ref = producer.put({"x": torch.ones(1, 4)}, sample_id="s0", metadata={}) + + out, handle = reader.get(ref) + self.assertEqual(out["x"].sum().item(), 4.0) + reader.release(handle) + out_again, handle_again = reader.get(ref) + reader.release(handle_again) + self.assertEqual(out_again["x"].sum().item(), 4.0) + + with self.assertRaisesRegex(RuntimeError, "lifetime owner"): + reader.reclaim(ref) + producer.reclaim(ref) + with self.assertRaises(KeyError): + reader.get(ref) + def test_use_after_free_get_raises(self): store = SharedDirFeatureStore(self.root) ref = store.put({"x": torch.randn(1, 4)}, sample_id="s0", metadata={}) diff --git a/tests/test_runtime/test_launch_plan.py b/tests/test_runtime/test_launch_plan.py index a819dd966..007991f3a 100644 --- a/tests/test_runtime/test_launch_plan.py +++ b/tests/test_runtime/test_launch_plan.py @@ -106,6 +106,58 @@ def _managed_config(control_dir, *, nproc=2, servers=None): return Config.model_validate(raw) +def _windowed_fanout_config(control_dir): + raw = _config(mode="disaggregated").model_dump() + raw["data"]["max_prompts"] = 8 + raw["training"].update({"batch_size": 2, "accumulation_steps": 1, "num_epochs": 1}) + raw["deployment"]["trainer"] = {"nnodes": 1, "nproc_per_node": 1} + raw["deployment"]["disaggregated"].update( + { + "control_dir": control_dir, + "server_urls": [], + "windowed_fanout": { + "max_live_bytes": 8 << 30, + "max_outstanding_per_consumer": 4, + "consumers": [ + { + "consumer_id": "block-4", + "seed": 42, + "loss_type": "dflash", + "loss_decay_gamma": 7.0, + "dpace_alpha": 0.5, + "draft_block_size": 4, + "num_anchors": 128, + "learning_rate": 0.0006, + "warmup_ratio": 0.04, + }, + { + "consumer_id": "block-16", + "seed": 43, + "loss_type": "dpace", + "loss_decay_gamma": None, + "dpace_alpha": 0.5, + "draft_block_size": 16, + "num_anchors": 256, + "learning_rate": 0.0007, + "warmup_ratio": 0.05, + }, + ], + }, + "managed_local": { + "trainer_cuda_visible_devices": ["1", "2"], + "capture_servers": [ + { + "port": 30000, + "cuda_visible_devices": ["0"], + "tp_size": 1, + } + ], + }, + } + ) + return Config.model_validate(raw) + + CAPTURE_CONTRACT = ServerCaptureContract( method="dflash", aux_layer_ids=(1, 9, 17, 25, 33), @@ -657,6 +709,141 @@ def test_managed_local_plan_owns_mooncake_and_multiple_capture_servers(self): rendered = json.loads(plan.render()) self.assertEqual(len(rendered["services"]), 3) + def test_windowed_fanout_uses_canonical_independent_consumer_commands(self): + with tempfile.TemporaryDirectory() as root: + control_dir = os.path.join(root, "attempt") + cfg = _windowed_fanout_config(control_dir) + with mock.patch( + "specforge.training.capture_contract.resolve_server_capture_contract", + return_value=CAPTURE_CONTRACT, + ): + plan = build_launch_plan( + cfg, + config_path="run.yaml", + worker_prefix=("specforge",), + env={}, + ) + + self.assertEqual((plan.kind, plan.role), ("managed_supervisor", "both")) + self.assertEqual( + [command.label for command in plan.commands], + ["block-4", "block-16", "producer"], + ) + for command, consumer_id, device in zip( + plan.commands[:2], ("block-4", "block-16"), ("1", "2") + ): + self.assertEqual( + command.env["SPECFORGE_FANOUT_CONSUMER_ID"], consumer_id + ) + self.assertEqual(command.env["CUDA_VISIBLE_DEVICES"], device) + self.assertTrue( + command.env["DISAGG_DB"].endswith( + f"consumers/{consumer_id}/consumer.sqlite" + ) + ) + self.assertEqual( + command.argv[command.argv.index("--consumer-id") + 1], + consumer_id, + ) + producer = plan.commands[-1] + self.assertEqual(producer.env["CUDA_VISIBLE_DEVICES"], "") + capture_service = plan.services[1] + argv = capture_service.command.argv + self.assertEqual( + argv[argv.index("--spec-capture-store-id") + 1], cfg.run_id + ) + self.assertIn("--spec-capture-inventory-db", argv) + self.assertIn("--spec-capture-lifecycle-db", argv) + self.assertEqual( + capture_service.command.env["SGLANG_SPEC_CAPTURE_TOKEN"], + producer.env["SGLANG_SPEC_CAPTURE_TOKEN"], + ) + rendered = json.loads(plan.render()) + self.assertEqual( + rendered["commands"][-1]["env"]["SGLANG_SPEC_CAPTURE_TOKEN"], + "", + ) + + projected = _config_for_role(cfg, "consumer", "block-16") + self.assertEqual(projected.model.draft_block_size, 16) + self.assertEqual(projected.training.num_anchors, 256) + self.assertEqual(projected.training.loss_type, "dpace") + self.assertEqual(projected.training.seed, 43) + self.assertTrue(projected.output_dir.endswith("block-16")) + self.assertIsNone(projected.deployment.disaggregated.managed_local) + + os.makedirs(os.path.join(control_dir, "logs")) + child = build_launch_plan( + cfg, + config_path="run.yaml", + requested_role="consumer", + consumer_id="block-4", + env=plan.commands[0].env, + ) + self.assertEqual((child.kind, child.role), ("worker", "consumer")) + + def test_windowed_fanout_requires_explicit_consumer_selection(self): + with tempfile.TemporaryDirectory() as root: + control_dir = os.path.join(root, "attempt") + cfg = _windowed_fanout_config(control_dir) + os.makedirs(os.path.join(control_dir, "logs")) + with self.assertRaisesRegex(ValueError, "requires --consumer-id"): + build_launch_plan( + cfg, + config_path="run.yaml", + requested_role="consumer", + env={"SPECFORGE_MANAGED_LOCAL_CHILD": "1"}, + ) + + def test_windowed_fanout_projects_external_consumer_resume_canonically(self): + with tempfile.TemporaryDirectory() as root: + raw = _windowed_fanout_config(root).model_dump() + disaggregated = raw["deployment"]["disaggregated"] + disaggregated["managed_local"] = None + disaggregated["server_urls"] = ["http://capture:30000"] + for consumer, device in zip( + disaggregated["windowed_fanout"]["consumers"], ("1", "2") + ): + consumer["cuda_visible_device"] = device + checkpoint = os.path.join(root, "block-4-checkpoint") + disaggregated["windowed_fanout"]["consumers"][0]["resume_from"] = checkpoint + cfg = Config.model_validate(raw) + + projected = _config_for_role(cfg, "consumer", "block-4") + + self.assertIsNone(projected.training.resume_from) + self.assertEqual( + projected.deployment.disaggregated.windowed_fanout.consumer( + "block-4" + ).resume_from, + checkpoint, + ) + + def test_windowed_fanout_rejects_multiple_environment_capture_servers(self): + with tempfile.TemporaryDirectory() as root: + raw = _windowed_fanout_config(root).model_dump() + disaggregated = raw["deployment"]["disaggregated"] + disaggregated["managed_local"] = None + disaggregated["server_urls"] = [] + for consumer, device in zip( + disaggregated["windowed_fanout"]["consumers"], ("1", "2") + ): + consumer["cuda_visible_device"] = device + cfg = Config.model_validate(raw) + + with self.assertRaisesRegex(ValueError, "exactly one capture server"): + build_launch_plan( + cfg, + config_path="run.yaml", + requested_role="producer", + env={ + "DISAGG_SERVER_URLS": "http://capture-a:30000," + "http://capture-b:30000", + "MOONCAKE_METADATA_SERVER": "http://metadata:8080", + "MOONCAKE_MASTER_SERVER_ADDR": "master:50051", + }, + ) + def test_multiserver_example_yaml_builds_the_managed_plan(self): path = ( Path(__file__).resolve().parents[2] @@ -1305,6 +1492,18 @@ def test_mooncake_readiness_accepts_missing_key_but_rejects_server_errors(self): ): self.assertEqual(_http_ready(readiness), expected) + def test_http_readiness_allows_slow_model_health_checks(self): + response = mock.MagicMock() + response.__enter__.return_value.status = 200 + readiness = ReadinessSpec("http", "http://127.0.0.1:30000/health", 1800) + + with mock.patch( + "specforge.launch_plan.urllib_request.urlopen", return_value=response + ) as open_url: + self.assertTrue(_http_ready(readiness)) + + open_url.assert_called_once_with(readiness.url, timeout=5.0) + def test_supervisor_terminates_a_sibling_after_child_failure(self): producer = _FakeProcess([None]) consumer = _FakeProcess([7]) diff --git a/tests/test_runtime/test_materializing_ref_source.py b/tests/test_runtime/test_materializing_ref_source.py new file mode 100644 index 000000000..af1f6e09e --- /dev/null +++ b/tests/test_runtime/test_materializing_ref_source.py @@ -0,0 +1,187 @@ +# coding=utf-8 +"""Local feature generation -> persisted ref adapter tests.""" + +import tempfile +import unittest + +import torch + +from specforge.inference.adapters.materializing import ( + MaterializationFailure, + MaterializingRefSource, +) +from specforge.inference.capture import CaptureConfig +from specforge.runtime.contracts import PromptTask, SampleRef +from specforge.runtime.data_plane.disaggregated import SharedDirFeatureStore + + +def _capture(): + return CaptureConfig.from_strategy( + required_features={"input_ids", "hidden_states", "loss_mask"}, + aux_hidden_state_layer_ids=(2, 18, 33), + target_repr="hidden_state", + target_hidden_size=8, + ) + + +def _task(index: int) -> PromptTask: + return PromptTask( + task_id=f"task-{index}", + run_id="run", + source_id="fixture", + payload={"input_ids": [index, index + 1], "loss_mask": [1, 1]}, + metadata={"num_tokens": 2}, + max_length=8, + ) + + +class _FeatureSource: + def __init__( + self, *, missing_loss_mask: int = -1, include_aux_metadata: bool = True + ) -> None: + self.missing_loss_mask = missing_loss_mask + self.include_aux_metadata = include_aux_metadata + self.generated = [] + + def generate_features(self, tasks, *, capture): + del capture + self.generated = [] + for index, _task_value in enumerate(tasks): + features = { + "input_ids": torch.tensor([[index, index + 1]]), + "hidden_states": torch.full((1, 2, 24), float(index)), + "loss_mask": torch.ones(1, 2, dtype=torch.long), + } + if self.include_aux_metadata: + features["__aux_layer_ids__"] = (2, 18, 33) + if index == self.missing_loss_mask: + features.pop("loss_mask") + self.generated.append(features) + return self.generated + + +class _FailingPutStore: + def __init__(self) -> None: + self.aborted = [] + + def put(self, tensors, *, sample_id, metadata): + del tensors, sample_id, metadata + raise OSError("store full") + + def abort(self, sample_id, *, reason): + self.aborted.append((sample_id, reason)) + + +class _ShortFeatureSource: + def generate_features(self, tasks, *, capture): + del tasks, capture + return [] + + +class TestMaterializingRefSource(unittest.TestCase): + def test_persists_ordered_refs_with_stable_identity_and_metadata(self): + with tempfile.TemporaryDirectory() as root: + store = SharedDirFeatureStore(root, store_id="features") + generated = _FeatureSource() + source = MaterializingRefSource( + generated, + store, + run_id="run", + strategy="dflash", + target_model_version="target-v1", + tokenizer_version="tokenizer-v1", + ) + + results = source.produce_refs([_task(0), _task(1)], capture=_capture()) + + self.assertTrue(all(isinstance(result, SampleRef) for result in results)) + self.assertEqual( + [result.sample_id for result in results], + ["run:task-0", "run:task-1"], + ) + self.assertEqual( + [result.source_task_id for result in results], + ["task-0", "task-1"], + ) + self.assertEqual(results[0].strategy, "dflash") + self.assertEqual(results[0].target_model_version, "target-v1") + self.assertEqual(results[0].tokenizer_version, "tokenizer-v1") + self.assertEqual(results[0].metadata["generation"], 1) + tensors, handle = store.get(results[1]) + self.assertEqual(tensors["hidden_states"].sum().item(), 48.0) + store.release(handle) + self.assertIn("loss_mask", generated.generated[0]) + + def test_capture_mismatch_is_non_retryable_and_does_not_write(self): + with tempfile.TemporaryDirectory() as root: + store = SharedDirFeatureStore(root, store_id="features") + source = MaterializingRefSource( + _FeatureSource(missing_loss_mask=1), + store, + run_id="run", + strategy="dflash", + ) + + first, second = source.produce_refs( + [_task(0), _task(1)], capture=_capture() + ) + + self.assertIsInstance(first, SampleRef) + self.assertIsInstance(second, MaterializationFailure) + self.assertEqual(second.task_id, "task-1") + self.assertFalse(second.retryable) + self.assertEqual(store.health()["resident_samples"], 1) + + def test_omitted_aux_layer_ids_are_rejected_before_write(self): + with tempfile.TemporaryDirectory() as root: + store = SharedDirFeatureStore(root, store_id="features") + source = MaterializingRefSource( + _FeatureSource(include_aux_metadata=False), + store, + run_id="run", + strategy="dflash", + ) + + (result,) = source.produce_refs([_task(0)], capture=_capture()) + + self.assertIsInstance(result, MaterializationFailure) + self.assertIn("omitted aux-layer ids", result.reason) + self.assertEqual(store.health()["resident_samples"], 0) + + def test_put_failure_is_aligned_retryable_and_aborted(self): + store = _FailingPutStore() + source = MaterializingRefSource( + _FeatureSource(), store, run_id="run", strategy="dflash" + ) + + (result,) = source.produce_refs([_task(0)], capture=_capture()) + + self.assertEqual( + result, + MaterializationFailure( + task_id="task-0", reason="put_failed:store full", retryable=True + ), + ) + self.assertEqual(store.aborted, [("run:task-0", "put_failed:store full")]) + + def test_wrong_feature_count_fails_every_task_without_retry(self): + store = _FailingPutStore() + source = MaterializingRefSource( + _ShortFeatureSource(), store, run_id="run", strategy="dflash" + ) + + results = source.produce_refs([_task(0), _task(1)], capture=_capture()) + + self.assertEqual([result.task_id for result in results], ["task-0", "task-1"]) + self.assertTrue(all(not result.retryable for result in results)) + self.assertTrue( + all( + "returned 0 feature records for 2 tasks" in result.reason + for result in results + ) + ) + self.assertEqual(store.aborted, []) + + +if __name__ == "__main__": + unittest.main(verbosity=2) diff --git a/tests/test_runtime/test_mooncake_store.py b/tests/test_runtime/test_mooncake_store.py index 3c1590b9f..d903443ec 100644 --- a/tests/test_runtime/test_mooncake_store.py +++ b/tests/test_runtime/test_mooncake_store.py @@ -9,7 +9,10 @@ """ import ctypes +import dataclasses import importlib.util +import os +import tempfile import unittest import torch @@ -37,6 +40,7 @@ def __init__(self) -> None: self.fail_remove = False self.lease_defer = False # remove() returns ok but keeps bytes (Mooncake lease) self.put_calls = 0 + self.fail_put_call = None self.remove_calls = 0 def is_exist(self, key): @@ -51,8 +55,10 @@ def unregister_buffer(self, ptr): def put_from(self, key, ptr, size, config=None): self.last_config = config - self._d[key] = ctypes.string_at(ptr, size) # DMA-equivalent read of src self.put_calls += 1 + if self.put_calls == self.fail_put_call: + return -1 + self._d[key] = ctypes.string_at(ptr, size) # DMA-equivalent read of src return 0 def get_into(self, key, ptr, size): @@ -63,15 +69,18 @@ def get_into(self, key, ptr, size): ctypes.memmove(ptr, data, n) # DMA-equivalent write into dst return n - def remove(self, key): + def remove(self, key, force=False): self.remove_calls += 1 if self.fail_remove: return -1 - if self.lease_defer: + if self.lease_defer and not force: return 0 # report success but keep the object (lease-deferred free) self._d.pop(key, None) return 0 + def get_size(self, key): + return len(self._d.get(key, b"")) + def _phys_resident(fake, sid="s0", store_id="run0"): """Do any per-tensor objects for the sample remain in the fake?""" @@ -177,6 +186,51 @@ def test_retain_on_release_keeps_data(self): out, _ = fs.get(ref) # still available for the next epoch self.assertIn("hidden_state", out) + def test_fanout_owner_reclaims_after_retained_consumer_leases(self): + fs = _store(retain_on_release=True) + ref = fs.put(_tensors(), sample_id="s0", metadata=_meta()) + _, first = fs.get(ref) + _, second = fs.get(ref) + fs.release(first) + fs.release(second) + self.assertEqual(fs.health()["resident_samples"], 1) + + fs.reclaim(ref, reason="fanout-globally-consumed") + self.assertEqual(fs.health()["resident_samples"], 0) + with self.assertRaisesRegex(KeyError, "untracked sample"): + fs.reclaim(ref, reason="duplicate-global-ack") + + def test_stale_fanout_reclaim_cannot_delete_new_generation(self): + fs = _store(retain_on_release=True) + stale = fs.put(_tensors(), sample_id="s0", metadata=_meta()) + current = fs.put(_tensors(), sample_id="s0", metadata=_meta()) + + with self.assertRaisesRegex(KeyError, "stale sample"): + fs.reclaim(stale, reason="delayed-consumer-ack") + out, handle = fs.get(current) + self.assertIn("hidden_state", out) + fs.release(handle) + fs.reclaim(current, reason="fanout-globally-consumed") + self.assertEqual(fs.health()["resident_samples"], 0) + + def test_fanout_consumer_drops_local_state_without_remote_delete(self): + fake = _FakeMooncakeStore() + owner = MooncakeFeatureStore(store=fake, store_id="run0") + consumer = MooncakeFeatureStore( + store=fake, store_id="run0", lifetime_owner=False + ) + ref = owner.put(_tensors(), sample_id="s0", metadata=_meta()) + + _, handle = consumer.get(ref) + consumer.release(handle) + self.assertTrue(_phys_resident(fake)) + self.assertEqual(consumer.health()["resident_samples"], 0) + with self.assertRaisesRegex(PermissionError, "non-owner"): + consumer.put(_tensors(), sample_id="s1", metadata=_meta()) + + owner.reclaim(ref, reason="fanout-globally-consumed") + self.assertFalse(_phys_resident(fake)) + def test_consume_once_free_on_last_lease(self): fs = _store() ref = fs.put(_tensors(), sample_id="s0", metadata=_meta()) @@ -206,6 +260,46 @@ def test_max_resident_bytes_raises_when_behind(self): with self.assertRaises(MemoryError): fs.put(_tensors(), sample_id="s0", metadata=_meta()) + def test_adopt_recovers_remote_size_before_enforcing_budget(self): + fake = _FakeMooncakeStore() + producer = MooncakeFeatureStore(store=fake, store_id="run0") + ref = producer.put(_tensors(), sample_id="s0", metadata=_meta()) + unknown_size_ref = dataclasses.replace(ref, estimated_bytes=0) + owner = MooncakeFeatureStore( + store=fake, + store_id="run0", + max_resident_bytes=ref.estimated_bytes - 1, + ) + + with self.assertRaises(MemoryError): + owner.adopt(unknown_size_ref) + + def test_partial_raw_put_removes_written_keys(self): + fake = _FakeMooncakeStore() + fake.fail_put_call = 2 + fs = MooncakeFeatureStore(store=fake, store_id="run0") + + with self.assertRaises(RuntimeError): + fs.put(_tensors(), sample_id="s0", metadata=_meta()) + + self.assertEqual(fake._d, {}) + self.assertEqual(fs.health()["release_pending"], 0) + + def test_partial_raw_put_tracks_failed_cleanup_until_drain(self): + fake = _FakeMooncakeStore() + fake.fail_put_call = 2 + fake.fail_remove = True + fs = MooncakeFeatureStore(store=fake, store_id="run0") + + with self.assertRaises(RuntimeError): + fs.put(_tensors(), sample_id="s0", metadata=_meta()) + self.assertEqual(fs.health()["release_pending"], 1) + + fake.fail_remove = False + report = fs.drain_pending_removals(max_attempts=1, retry_interval_s=0) + self.assertEqual(report["release_pending"], 0) + self.assertEqual(fake._d, {}) + def test_gc_force_frees_past_max_hold(self): clock = _FakeClock() fs = _store(max_hold_age_s=10.0, clock=clock) @@ -338,6 +432,50 @@ def test_lifecycle_drain_is_bounded_and_never_hides_remote_leak(self): self.assertEqual(fs.health()["force_freed_total"], 0) self.assertTrue(_phys_resident(fake)) + def test_drain_includes_superseded_generation_removals(self): + fake = _FakeMooncakeStore() + with tempfile.TemporaryDirectory() as tempdir: + fs = MooncakeFeatureStore( + store=fake, + store_id="run0", + lifecycle_db_path=os.path.join(tempdir, "lifecycle.db"), + ) + fs.put(_tensors(), sample_id="s0", metadata=_meta()) + fake.fail_remove = True + fs.put(_tensors(), sample_id="s0", metadata=_meta()) + fs.gc() + self.assertEqual(fs.health()["release_pending"], 1) + + fake.fail_remove = False + report = fs.drain_pending_removals(max_attempts=1, retry_interval_s=0) + + self.assertEqual(report["release_pending"], 0) + self.assertFalse(any("/g1/" in key for key in fake._d)) + self.assertTrue(any("/g2/" in key for key in fake._d)) + + def test_external_gc_retry_budget_is_capped_and_drain_stays_authoritative(self): + fake = _FakeMooncakeStore() + with tempfile.TemporaryDirectory() as tempdir: + fs = MooncakeFeatureStore( + store=fake, + store_id="run0", + max_release_attempts=1, + lifecycle_db_path=os.path.join(tempdir, "lifecycle.db"), + ) + fs.put(_tensors(), sample_id="s0", metadata=_meta()) + fake.fail_remove = True + fs.put(_tensors(), sample_id="s0", metadata=_meta()) + + for _ in range(5): + report = fs.gc() + self.assertEqual(report["release_pending"], 1) + self.assertEqual(set(fs._external_release_pending.values()), {1}) + with self.assertRaisesRegex(RuntimeError, "could not drain 1 pending"): + fs.drain_pending_removals( + max_attempts=2, + retry_interval_s=0, + ) + def test_restart_authority_adopts_acked_remote_ref_before_removing(self): fake = _FakeMooncakeStore() producer = MooncakeFeatureStore(store=fake, store_id="run0") @@ -423,6 +561,18 @@ class _ObjectOnly(_FakeMooncakeStore): self.assertIn("Upgrade", message) self.assertIn("serialized put/get transport is not supported", message) + @unittest.skipUnless(torch.cuda.is_available(), "CUDA is required for pinning") + def test_selected_get_can_fill_pinned_destination_directly(self): + fs = _store() + src = _tensors() + ref = fs.put(src, sample_id="s0", metadata=_meta()) + + out, _ = fs.get(ref, names=["hidden_state"], pin_memory=True) + + self.assertEqual(set(out), {"hidden_state"}) + self.assertTrue(out["hidden_state"].is_pinned()) + torch.testing.assert_close(out["hidden_state"], src["hidden_state"]) + def test_one_object_per_tensor_with_generation_in_key(self): fake = _FakeMooncakeStore() fs = MooncakeFeatureStore(store=fake, store_id="run0") @@ -438,7 +588,7 @@ def test_wire_bytes_are_exactly_the_raw_tensor(self): # The stored bytes are exactly the raw tensor buffer, not an archive. fake = _FakeMooncakeStore() fs = MooncakeFeatureStore(store=fake, store_id="run0") - ref = fs.put(_tensors(), sample_id="s0", metadata=_meta()) + fs.put(_tensors(), sample_id="s0", metadata=_meta()) wire_bytes = fake._d["run0/s0/g1/hidden_state"] self.assertEqual(len(wire_bytes), 4 * 8 * 4) self.assertFalse(wire_bytes[:2] == b"PK") @@ -511,7 +661,7 @@ def test_cross_process_stale_generation_rejected(self): def test_cross_process_abort_blocks_consumer_get(self): # With a normal (immediate) remove, producer.abort physically deletes the - # objects, so a separate consumer's get() raises (B5 holds cross-process via + # blob, so a separate consumer's get() raises (B5 holds cross-process via # physical removal, not the per-process tombstone). fake, producer, consumer = _shared_pair(retain_on_release=True) ref = producer.put(_tensors(), sample_id="s0", metadata=_meta()) @@ -521,8 +671,8 @@ def test_cross_process_abort_blocks_consumer_get(self): consumer.get(ref) def test_cross_process_consume_once_free_by_consumer(self): - # Consume-once consumer frees the shared tensor objects on release; the - # producer can then no longer resolve the ref. + # Consume-once consumer (retain_on_release=False) frees the shared tensor + # objects on release; the producer can then no longer resolve the ref. fake, producer, consumer = _shared_pair() # consumer frees on release ref = producer.put(_tensors(), sample_id="s0", metadata=_meta()) _, handle = consumer.get(ref) @@ -551,6 +701,77 @@ def test_cross_process_abort_under_lease_defer_is_known_gap(self): consumer.get(ref) # SHOULD raise; currently returns stale -> xfail +class TestMooncakeDurableLifecycle(unittest.TestCase): + def test_tombstones_stay_out_of_process_memory(self): + fake = _FakeMooncakeStore() + with tempfile.TemporaryDirectory() as workdir: + lifecycle = os.path.join(workdir, "lifecycle.db") + owner = MooncakeFeatureStore( + store=fake, store_id="run0", lifecycle_db_path=lifecycle + ) + first_ref = owner.put(_tensors(), sample_id="s0", metadata=_meta()) + owner.reclaim(first_ref) + for index in range(1, 50): + ref = owner.put(_tensors(), sample_id=f"s{index}", metadata=_meta()) + owner.reclaim(ref) + + self.assertEqual(owner.health()["local_tombstones"], 0) + self.assertEqual(owner.health()["resident_samples"], 0) + self.assertEqual(owner._lifecycle.pending(), ()) + with self.assertRaisesRegex(KeyError, "lifecycle state"): + owner.get(first_ref) + + def test_restarted_owner_cleans_resident_inventory(self): + fake = _FakeMooncakeStore() + with tempfile.TemporaryDirectory() as workdir: + lifecycle = os.path.join(workdir, "lifecycle.db") + original = MooncakeFeatureStore( + store=fake, store_id="run0", lifecycle_db_path=lifecycle + ) + ref = original.put(_tensors(), sample_id="s0", metadata=_meta()) + + recovered = MooncakeFeatureStore( + store=fake, store_id="run0", lifecycle_db_path=lifecycle + ) + self.assertEqual(recovered.health()["resident_samples"], 1) + self.assertEqual(recovered.abort_all(reason="restart-cleanup"), 1) + self.assertEqual(recovered.health()["resident_samples"], 0) + self.assertFalse(_phys_resident(fake)) + + reader = MooncakeFeatureStore( + store=fake, + store_id="run0", + lifetime_owner=False, + lifecycle_db_path=lifecycle, + ) + with self.assertRaises(KeyError): + reader.get(ref) + + def test_restarted_owner_retries_partial_write_cleanup(self): + fake = _FakeMooncakeStore() + fake.fail_put_call = 2 + fake.fail_remove = True + with tempfile.TemporaryDirectory() as workdir: + lifecycle = os.path.join(workdir, "lifecycle.db") + original = MooncakeFeatureStore( + store=fake, store_id="run0", lifecycle_db_path=lifecycle + ) + with self.assertRaisesRegex(RuntimeError, "put_from failed"): + original.put(_tensors(), sample_id="s0", metadata=_meta()) + + self.assertEqual(original._lifecycle.state("s0", 1), "tombstoned") + self.assertTrue(_phys_resident(fake)) + fake.fail_put_call = None + fake.fail_remove = False + + recovered = MooncakeFeatureStore( + store=fake, store_id="run0", lifecycle_db_path=lifecycle + ) + self.assertEqual(recovered.abort_all(reason="restart-cleanup"), 1) + self.assertEqual(recovered._lifecycle.state("s0", 1), "cleaned") + self.assertFalse(_phys_resident(fake)) + + @unittest.skipUnless( importlib.util.find_spec("mooncake") is not None, "mooncake package not installed; real end-to-end store test skipped", diff --git a/tests/test_runtime/test_package_architecture.py b/tests/test_runtime/test_package_architecture.py index 3ee02f4ac..66b0cf46e 100644 --- a/tests/test_runtime/test_package_architecture.py +++ b/tests/test_runtime/test_package_architecture.py @@ -247,11 +247,14 @@ "build_disagg_offline_runtime", "build_disagg_online_consumer", "build_disagg_online_producer", + "build_disagg_online_windowed_consumer", + "build_disagg_online_windowed_producer", "build_offline_runtime", } DRAFT_MODEL_BUILDERS = CANONICAL_LAUNCH_EXPORTS - { "build_disagg_online_producer", + "build_disagg_online_windowed_producer", } CANONICAL_DRAFT_CONFIGS = { @@ -519,7 +522,7 @@ def test_public_training_lifecycle_has_one_surface(self): self.assertIsInstance(assembler_returns[0].value, ast.Name) self.assertEqual("trainer", assembler_returns[0].value.id) - trainer_builders = CANONICAL_LAUNCH_EXPORTS - {"build_disagg_online_producer"} + trainer_builders = DRAFT_MODEL_BUILDERS for name in trainer_builders: returns = [ node diff --git a/tests/test_runtime/test_server_capture.py b/tests/test_runtime/test_server_capture.py index 68e6ad62c..5ee6f2a28 100644 --- a/tests/test_runtime/test_server_capture.py +++ b/tests/test_runtime/test_server_capture.py @@ -14,10 +14,12 @@ """ import ctypes +import dataclasses import os import tempfile import unittest from typing import Any, Dict, List +from unittest import mock import torch @@ -87,12 +89,14 @@ def __init__( self, backend: _FakeMooncakeStore, *, + capture_token: str = "unit-capture-token", hidden: int = HIDDEN, aux_width: int = len(AUX_LAYERS) * HIDDEN, aux_layer_ids=AUX_LAYERS, error_sample_ids=(), ) -> None: self.backend = backend + self.capture_token = capture_token self.hidden = hidden self.aux_width = aux_width self.aux_layer_ids = None if aux_layer_ids is None else list(aux_layer_ids) @@ -107,6 +111,7 @@ def __call__(self, url: str, json_body: Dict[str, Any], timeout: float): assert url.endswith("/generate") rows: List[Dict[str, Any]] = [] for input_ids, spec in zip(json_body["input_ids"], json_body["spec_capture"]): + assert spec["auth_token"] == self.capture_token sid, gen = spec["sample_id"], int(spec["gen"]) if sid in self.error_sample_ids: rows.append( @@ -282,11 +287,49 @@ def _mk( schema=_capture_schema(algorithm), request_input_adapter=request_input_adapter, post_fn=server, + capture_token="unit-capture-token", ) return backend, server, store, adapter class TestServerCaptureAdapter(unittest.TestCase): + def test_capture_capability_and_namespace_are_client_guarded(self): + backend = _FakeMooncakeStore() + store = MooncakeFeatureStore(store=backend, store_id="run0") + schema = _capture_schema("eagle3") + with mock.patch.dict(os.environ, {"SGLANG_SPEC_CAPTURE_TOKEN": ""}): + with self.assertRaisesRegex(ValueError, "requires capture_token"): + SGLangServerCaptureAdapter( + "http://server:30000", + store, + run_id="run0", + algorithm="eagle3", + schema=schema, + ) + with self.assertRaisesRegex(ValueError, "run_id == store.store_id"): + SGLangServerCaptureAdapter( + "http://server:30000", + store, + run_id="other", + algorithm="eagle3", + schema=schema, + capture_token="unit-capture-token", + ) + + def test_request_sends_capability_and_bound_store_namespace(self): + backend = _FakeMooncakeStore() + server = _StubCaptureServer(backend) + captured = [] + + def inspect_request(url, json_body, timeout): + captured.extend(json_body["spec_capture"]) + return server(url, json_body, timeout) + + _, _, _, adapter = _mk(server=inspect_request, backend=backend) + adapter.produce_refs([_task(0, 4)], capture=_eagle3_contract()) + self.assertEqual(captured[0]["auth_token"], "unit-capture-token") + self.assertEqual(captured[0]["store_id"], "run0") + def test_generic_adapter_inputs_merge_with_runtime_owned_request_fields(self): backend = _FakeMooncakeStore() server = _StubCaptureServer(backend) @@ -514,6 +557,71 @@ def test_per_task_server_error_becomes_failure_marker(self): self.assertIn("injected sink error", results[0].reason) self.assertIsInstance(results[1], SampleRef) + def test_retryable_rejection_does_not_block_the_replacement_attempt(self): + class BusyRemovalStore(_FakeMooncakeStore): + def is_exist(self, key): + return 1 + + def remove(self, key): + return -1 + + backend = BusyRemovalStore() + server = _StubCaptureServer(backend, error_sample_ids={"run0:t0"}) + _, _, store, adapter = _mk(server=server, backend=backend) + + (failure,) = adapter.produce_refs([_task(0, 4)], capture=_eagle3_contract()) + self.assertIsInstance(failure, ServerCaptureFailure) + self.assertTrue(failure.retryable) + server.error_sample_ids.clear() + (replacement,) = adapter.produce_refs([_task(0, 4)], capture=_eagle3_contract()) + + self.assertIsInstance(replacement, SampleRef) + self.assertEqual(store.health()["release_pending"], 0) + + def test_response_feature_set_mismatch_reclaims_server_keys(self): + backend = _FakeMooncakeStore() + server = _StubCaptureServer(backend) + + def omit_feature(url, json_body, timeout): + rows = server(url, json_body, timeout) + rows[0]["meta_info"]["spec_capture"]["features"].pop("loss_mask") + return rows + + _, _, _, adapter = _mk(server=omit_feature, backend=backend) + (result,) = adapter.produce_refs([_task(0, 4)], capture=_eagle3_contract()) + self.assertIsInstance(result, ServerCaptureFailure) + self.assertFalse(result.retryable) + self.assertIn("feature set mismatch", result.reason) + self.assertFalse(backend._d) + + def test_malformed_metadata_reclaims_only_the_bad_row(self): + backend = _FakeMooncakeStore() + server = _StubCaptureServer(backend) + + def corrupt_first_row(url, json_body, timeout): + rows = server(url, json_body, timeout) + rows[0]["meta_info"]["spec_capture"]["features"]["hidden_state"].pop( + "shape" + ) + return rows + + _, _, _, adapter = _mk(server=corrupt_first_row, backend=backend) + results = adapter.produce_refs( + [_task(0, 4), _task(1, 5)], capture=_eagle3_contract() + ) + self.assertIsInstance(results[0], ServerCaptureFailure) + self.assertIsInstance(results[1], SampleRef) + self.assertFalse(any("/run0:t0/g1/" in key for key in backend._d)) + self.assertTrue(any("/run0:t1/g1/" in key for key in backend._d)) + + def test_health_reports_rpc_batch_profile(self): + _, _, _, adapter = _mk(algorithm="dflash") + adapter.produce_refs([_task(0, 4), _task(1, 5)], capture=_dflash_contract()) + health = adapter.health() + self.assertEqual(health["rpc_calls"], 1) + self.assertEqual(health["rpc_tasks"], 2) + self.assertEqual(health["mean_batch_size"], 2.0) + def test_rollout_worker_ref_path(self): from specforge.inference.rollout_worker import RolloutWorker @@ -588,6 +696,29 @@ def lose_first_response(url, json_body, timeout): self.assertTrue(all("/g1/" in key for key in backend._d)) self.assertFalse(any("/g2/" in key for key in backend._d)) + def test_windowed_recapture_uses_the_registry_generation(self): + backend, server, store, adapter = _mk() + requests = [] + original_post = adapter.post_fn + + def record_request(url, json_body, timeout): + requests.append(json_body["spec_capture"][0]) + return original_post(url, json_body, timeout) + + adapter.post_fn = record_request + task = dataclasses.replace( + _task(0, 4), + attempt=6, + metadata={"capture_generation": 7}, + ) + + (ref,) = adapter.produce_refs([task], capture=_eagle3_contract()) + + self.assertEqual(ref.metadata["generation"], 7) + self.assertEqual(requests[0]["gen"], 7) + self.assertTrue(requests[0]["replace"]) + self.assertTrue(all("/g7/" in key for key in backend._d)) + def test_terminal_lost_response_is_reclaimed_without_a_ref(self): backend, server, store, adapter = _mk() @@ -704,6 +835,7 @@ def test_transport_requires_an_injected_capture_schema(self): run_id="run0", algorithm="eagle3", schema=object(), + capture_token="unit-capture-token", ) @@ -784,6 +916,7 @@ def test_producer_streams_refs_via_feature_source(self): algorithm="dflash", schema=_capture_schema("dflash"), post_fn=server, + capture_token="unit-capture-token", ) prompts = [ { diff --git a/tests/test_runtime/test_server_capture_gate.py b/tests/test_runtime/test_server_capture_gate.py index 81598adf2..7598e540b 100644 --- a/tests/test_runtime/test_server_capture_gate.py +++ b/tests/test_runtime/test_server_capture_gate.py @@ -26,6 +26,7 @@ import importlib.util import os import shutil +import signal import subprocess import sys import tempfile @@ -37,6 +38,8 @@ CUDA = torch.cuda.is_available() ENABLED = os.environ.get("SPECFORGE_RUN_SERVER_CAPTURE_TESTS") == "1" PORT = 30989 +STORE_ID = "server-capture-gate" +CAPTURE_TOKEN = "server-capture-gate-token" AUX_LAYER_IDS = [1, 3, 4] H, TOL = 64, 2e-2 # fixture hidden size; documented bf16 tolerance @@ -87,6 +90,8 @@ def setUpClass(cls): from transformers import LlamaConfig, LlamaForCausalLM cls.workdir = tempfile.mkdtemp(prefix="spec_capture_gate_") + cls.inventory_db_path = os.path.join(cls.workdir, "capture_inventory.sqlite3") + cls.lifecycle_db_path = os.path.join(cls.workdir, "mooncake_lifecycle.sqlite3") cfg = LlamaConfig( hidden_size=H, intermediate_size=128, @@ -127,7 +132,10 @@ def setUpClass(cls): MOONCAKE_LOCAL_HOSTNAME=os.environ.get( "MOONCAKE_LOCAL_HOSTNAME", "127.0.0.1" ), + SGLANG_SPEC_CAPTURE_TOKEN=CAPTURE_TOKEN, ) + cls.server_log = open(os.path.join(cls.workdir, "server.log"), "w") + cls.addClassCleanup(cls.server_log.close) cls.server = subprocess.Popen( [ sys.executable, @@ -144,15 +152,23 @@ def setUpClass(cls): "-1", "--disable-radix-cache", "--enable-spec-capture", + "--spec-capture-store-id", + STORE_ID, + "--spec-capture-inventory-db", + cls.inventory_db_path, + "--spec-capture-lifecycle-db", + cls.lifecycle_db_path, "--spec-capture-aux-layer-ids", *[str(i) for i in AUX_LAYER_IDS], "--port", str(PORT), ], - stdout=open(os.path.join(cls.workdir, "server.log"), "w"), + stdout=cls.server_log, stderr=subprocess.STDOUT, env=env, + start_new_session=True, ) + cls.addClassCleanup(cls._terminate_process_group, cls.server) import requests deadline = time.time() + 300 @@ -180,11 +196,15 @@ def _ensure_mooncake_master(cls): "no MOONCAKE_MASTER_SERVER_ADDR in env and no mooncake_master " "binary on PATH" ) + cls.master_log = open(os.path.join(cls.workdir, "mooncake_master.log"), "w") + cls.addClassCleanup(cls.master_log.close) cls.master = subprocess.Popen( [binary, "--enable-http-metadata-server=true"], - stdout=open(os.path.join(cls.workdir, "mooncake_master.log"), "w"), + stdout=cls.master_log, stderr=subprocess.STDOUT, + start_new_session=True, ) + cls.addClassCleanup(cls._terminate_process_group, cls.master) time.sleep(3) if cls.master.poll() is not None: raise unittest.SkipTest( @@ -198,22 +218,25 @@ def _ensure_mooncake_master(cls): os.environ.setdefault("MOONCAKE_LOCAL_HOSTNAME", "127.0.0.1") os.environ.setdefault("MOONCAKE_PROTOCOL", "tcp") - @classmethod - def tearDownClass(cls): - for proc in (cls.server, cls.master): - if proc is not None and proc.poll() is None: - proc.terminate() - try: - proc.wait(timeout=30) - except subprocess.TimeoutExpired: - proc.kill() + @staticmethod + def _terminate_process_group(proc): + try: + os.killpg(proc.pid, signal.SIGTERM) + except ProcessLookupError: + return + try: + proc.wait(timeout=30) + except subprocess.TimeoutExpired: + os.killpg(proc.pid, signal.SIGKILL) + proc.wait(timeout=5) # -- helpers --------------------------------------------------------------- - def _store(self, store_id): + def _store(self): from specforge.runtime.data_plane.mooncake_store import MooncakeFeatureStore return MooncakeFeatureStore( - store_id=store_id, + store_id=STORE_ID, + lifecycle_db_path=self.lifecycle_db_path, setup_kwargs={ "local_hostname": os.environ["MOONCAKE_LOCAL_HOSTNAME"], "metadata_server": os.environ["MOONCAKE_METADATA_SERVER"], @@ -225,13 +248,13 @@ def _store(self, store_id): }, ) - def _tasks(self, rows): + def _tasks(self, rows, *, prefix): from specforge.runtime.contracts import PromptTask return [ PromptTask( - task_id=f"t{i}", - run_id="gate0", + task_id=f"{prefix}-t{i}", + run_id=STORE_ID, source_id="gate", payload={ "input_ids": list(r), @@ -279,13 +302,14 @@ def test_eagle3_zero_copy_end_to_end(self): from specforge.runtime.contracts import SampleRef rows = [[5, 6, 7, 8, 9, 10], [11, 12, 13, 14]] - store = self._store("gate-eagle3") + store = self._store() adapter = SGLangServerCaptureAdapter( f"http://localhost:{PORT}", store, - run_id="gate0", + run_id=STORE_ID, algorithm="eagle3", schema=_capture_schema("eagle3"), + capture_token=CAPTURE_TOKEN, ) contract = CaptureConfig.from_strategy( required_features={ @@ -299,7 +323,9 @@ def test_eagle3_zero_copy_end_to_end(self): target_repr="hidden_state", target_hidden_size=H, ) - refs = adapter.produce_refs(self._tasks(rows), capture=contract) + refs = adapter.produce_refs( + self._tasks(rows, prefix="eagle3"), capture=contract + ) for ref in refs: self.assertIsInstance(ref, SampleRef, f"expected a ref, got failure: {ref}") @@ -392,13 +418,14 @@ def test_dflash_capture_same_server(self): from specforge.runtime.contracts import SampleRef rows = [[3, 1, 4, 1, 5]] - store = self._store("gate-dflash") + store = self._store() adapter = SGLangServerCaptureAdapter( f"http://localhost:{PORT}", store, - run_id="gate1", + run_id=STORE_ID, algorithm="dflash", schema=_capture_schema("dflash"), + capture_token=CAPTURE_TOKEN, ) contract = CaptureConfig.from_strategy( required_features={"input_ids", "hidden_states", "loss_mask"}, @@ -406,7 +433,9 @@ def test_dflash_capture_same_server(self): target_repr="hidden_state", target_hidden_size=H, ) - (ref,) = adapter.produce_refs(self._tasks(rows), capture=contract) + (ref,) = adapter.produce_refs( + self._tasks(rows, prefix="dflash"), capture=contract + ) self.assertIsInstance(ref, SampleRef, f"expected a ref, got: {ref}") out, handle = store.get(ref) self.assertEqual(sorted(out), ["hidden_states", "input_ids", "loss_mask"]) diff --git a/tests/test_runtime/test_sglang_capture_patch.py b/tests/test_runtime/test_sglang_capture_patch.py new file mode 100644 index 000000000..9f9d80bd2 --- /dev/null +++ b/tests/test_runtime/test_sglang_capture_patch.py @@ -0,0 +1,619 @@ +# coding=utf-8 +"""Executable contract tests for the SGLang spec-capture patch artifact.""" + +import hashlib +import importlib.metadata +import importlib.util +import json +import shutil +import sqlite3 +import subprocess +import sys +import tempfile +import unittest +import venv +from pathlib import Path + +import torch + +_ROOT = Path(__file__).resolve().parents[2] +_PATCH = _ROOT / "patches" / "sglang" / "v0.5.14" / "spec-capture.patch" +_BASE_TO_CURRENT_PATCH = ( + _ROOT / "patches" / "sglang" / "v0.5.14" / "spec-capture-base-to-current.patch" +) +_INSTALLER = _ROOT / "scripts" / "apply_sglang_spec_capture_patch.sh" + + +def _extract_sink() -> Path: + lines = _PATCH.read_text().splitlines() + target = "+++ b/python/sglang/srt/spec_capture_sink.py" + start = lines.index(target) + 1 + while start < len(lines) and not lines[start].startswith("@@"): + start += 1 + hunk_header = lines[start] + declared_lines = int(hunk_header.split("+1,", 1)[1].split(" ", 1)[0]) + start += 1 + source = [] + for line in lines[start:]: + if line.startswith("diff --git "): + break + if line.startswith("+"): + source.append(line[1:]) + elif line == "\\ No newline at end of file": + continue + else: + raise AssertionError(f"unexpected non-addition in new sink hunk: {line}") + if len(source) != declared_lines: + raise AssertionError( + f"sink hunk declares {declared_lines} lines but contains {len(source)}" + ) + root = Path(tempfile.mkdtemp(prefix="sglang-capture-patch-")) + path = root / "spec_capture_sink.py" + path.write_text("\n".join(source) + "\n") + return path + + +def _load_sink_module(): + path = _extract_sink() + spec = importlib.util.spec_from_file_location("_patched_spec_capture_sink", path) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module + + +class TestSGLangCapturePatch(unittest.TestCase): + def test_patch_is_well_formed(self): + subprocess.run( + ["git", "apply", "--numstat", str(_PATCH)], + cwd=_ROOT, + check=True, + capture_output=True, + text=True, + ) + + def test_patch_captures_final_pre_norm_hidden_state(self): + text = _PATCH.read_text(encoding="utf-8") + self.assertIn("diff --git a/python/sglang/srt/models/qwen2.py", text) + self.assertIn("if self.end_layer in self.layers_to_capture:", text) + self.assertLess( + text.index("if self.end_layer in self.layers_to_capture:"), + text.index( + "if hidden_states.shape[0] != 0:", text.index("models/qwen2.py") + ), + ) + + def test_patch_copies_only_requested_capture_artifacts(self): + text = _PATCH.read_text(encoding="utf-8") + start = text.index("+ def _append_spec_capture_states(") + end = text.index("+ def _sink_spec_capture(", start) + body = text[start:end] + + self.assertIn( + '+ if "aux" in requested_artifacts:\n' + "+ req.spec_capture_aux.append(\n" + "+ logits_output.hidden_states[start:end].cpu().clone()", + body, + ) + self.assertIn( + "+ if (\n" + '+ "last_hidden" in requested_artifacts\n' + "+ and logits_output.last_hidden_states is not None\n" + "+ ):\n" + "+ req.spec_capture_last_hidden.append(\n" + "+ logits_output.last_hidden_states[start:end].cpu().clone()", + body, + ) + + def test_base_to_current_migration_is_well_formed_and_scoped(self): + result = subprocess.run( + ["git", "apply", "--numstat", str(_BASE_TO_CURRENT_PATCH)], + cwd=_ROOT, + check=True, + capture_output=True, + text=True, + ) + targets = {line.split("\t", 2)[-1] for line in result.stdout.splitlines()} + self.assertEqual( + targets, + { + "python/sglang/srt/managers/schedule_batch.py", + "python/sglang/srt/managers/scheduler_components/" + "batch_result_processor.py", + "python/sglang/srt/models/qwen2.py", + "python/sglang/srt/server_args.py", + "python/sglang/srt/spec_capture_sink.py", + }, + ) + + def test_installer_pins_patch_and_migration_digests(self): + installer = _INSTALLER.read_text(encoding="utf-8") + patch_digest = hashlib.sha256(_PATCH.read_bytes()).hexdigest() + migration_digest = hashlib.sha256( + _BASE_TO_CURRENT_PATCH.read_bytes() + ).hexdigest() + self.assertIn(f'EXPECTED_PATCH_SHA256="{patch_digest}"', installer) + self.assertIn( + f'EXPECTED_BASE_TO_CURRENT_SHA256="{migration_digest}"', installer + ) + + def test_installer_upgrades_the_published_base_patch(self): + try: + installed_version = importlib.metadata.version("sglang") + except importlib.metadata.PackageNotFoundError: + self.skipTest("sglang is not installed") + if installed_version != "0.5.14": + self.skipTest(f"requires sglang==0.5.14, found {installed_version}") + + package_spec = importlib.util.find_spec("sglang") + self.assertIsNotNone(package_spec) + self.assertIsNotNone(package_spec.origin) + package_root = Path(package_spec.origin).resolve().parent + + with tempfile.TemporaryDirectory(prefix="sglang-legacy-installer-") as root: + env_root = Path(root) / "env" + venv.EnvBuilder(with_pip=False).create(env_root) + python = env_root / "bin/python" + site_packages = Path( + subprocess.run( + [ + str(python), + "-c", + "import sysconfig; print(sysconfig.get_paths()['purelib'])", + ], + check=True, + capture_output=True, + text=True, + ).stdout.strip() + ) + parent_paths = sorted( + { + str(Path(entry).resolve()) + for entry in sys.path + if entry and Path(entry).is_dir() + } + ) + (site_packages / "specforge-parent-runtime.pth").write_text( + "\n".join(parent_paths) + "\n", + encoding="utf-8", + ) + shutil.copytree(package_root, site_packages / "sglang") + dist_info = site_packages / "sglang-0.5.14.dist-info" + dist_info.mkdir() + (dist_info / "METADATA").write_text( + "Metadata-Version: 2.1\nName: sglang\nVersion: 0.5.14\n", + encoding="utf-8", + ) + + def run_installer(mode: str, *, check: bool = True): + return subprocess.run( + [ + "bash", + str(_INSTALLER), + "--python", + str(python), + mode, + ], + check=check, + capture_output=True, + text=True, + timeout=180, + ) + + # Normalize either a clean or already-patched source installation. + run_installer("--apply") + subprocess.run( + [ + "patch", + "--batch", + "--reverse", + "-p2", + "-d", + str(site_packages), + ], + input=_BASE_TO_CURRENT_PATCH.read_bytes(), + check=True, + capture_output=True, + ) + + base_check = run_installer("--check", check=False) + self.assertEqual(base_check.returncode, 2) + self.assertIn("published base spec-capture patch", base_check.stderr) + upgraded = run_installer("--apply") + self.assertIn("applied and verified spec-capture patch", upgraded.stdout) + + def test_remove_exact_uses_force_without_granting_a_read_lease(self): + module = _load_sink_module() + root = Path(tempfile.mkdtemp(prefix="capture-remove-")) + sink = module.SpecCaptureSink( + store_id="run0", + auth_token="secret", + max_sample_bytes=1 << 20, + inventory_db_path=str(root / "inventory.db"), + lifecycle_db_path=str(root / "lifecycle.db"), + aux_layer_ids=[1, 2], + ) + + class Store: + status = -704 + + def __init__(self): + self.calls = [] + + def remove(self, key, force=False): + self.calls.append((key, force)) + return self.status + + def is_exist(self, key): + raise AssertionError("remove cleanup must not grant a read lease") + + store = Store() + sink._connect = lambda: store + sink._remove_exact("run0/s0/g1/hidden_states", force=True) + self.assertEqual(store.calls, [("run0/s0/g1/hidden_states", True)]) + + store.status = -706 + with self.assertRaisesRegex(RuntimeError, "status -706"): + sink._remove_exact("run0/s0/g1/hidden_states", force=False) + self.assertEqual(store.calls[-1], ("run0/s0/g1/hidden_states", False)) + + def _sink(self, *, max_bytes=1 << 20, inventory_path=None): + module = _load_sink_module() + inventory_path = inventory_path or str( + Path(tempfile.mkdtemp(prefix="capture-inventory-")) / "inventory.db" + ) + sink = module.SpecCaptureSink( + store_id="run0", + auth_token="secret", + max_sample_bytes=max_bytes, + inventory_db_path=inventory_path, + lifecycle_db_path=inventory_path + ".lifecycle", + aux_layer_ids=[1, 2], + ) + writes = [] + sink._remove_exact = lambda key, *, force: None + sink._put_tensor = lambda key, tensor, **_kwargs: writes.append( + (key, tensor.clone()) + ) + return sink, writes + + @staticmethod + def _request(**overrides): + request = { + "auth_token": "secret", + "store_id": "run0", + "sample_id": "run0:s0", + "gen": 1, + "replace": False, + "features": {"aux": "hidden_states"}, + "passthrough": [ + { + "name": "input_ids", + "data": [1, 2, 3], + "shape": [1, 3], + "dtype": "int64", + } + ], + } + request.update(overrides) + return request + + def test_valid_authenticated_request_uses_server_namespace(self): + sink, writes = self._sink() + result = sink.put_sample( + self._request(), + aux=torch.ones(3, 4, dtype=torch.bfloat16), + last_hidden=None, + ) + self.assertEqual(result["store_id"], "run0") + self.assertEqual( + [key for key, _ in writes], + ["run0/run0:s0/g1/hidden_states", "run0/run0:s0/g1/input_ids"], + ) + + def test_lifecycle_is_planned_before_first_write_and_then_resident(self): + sink, writes = self._sink() + states_at_write = [] + + def record_write(key, tensor, **_kwargs): + row = sink._lifecycle.execute( + "SELECT state, feature_names_json, estimated_bytes FROM " + "mooncake_objects WHERE store_id=? AND sample_id=? AND generation=?", + ("run0", "run0:s0", 1), + ).fetchone() + states_at_write.append(row) + writes.append((key, tensor.clone())) + + sink._put_tensor = record_write + sink.put_sample( + self._request(), + aux=torch.ones(3, 4, dtype=torch.bfloat16), + last_hidden=None, + ) + + self.assertEqual([row[0] for row in states_at_write], ["planned", "planned"]) + self.assertEqual(states_at_write[0][1], '["hidden_states","input_ids"]') + self.assertEqual(states_at_write[0][2], 48) + state = sink._lifecycle.execute( + "SELECT state FROM mooncake_objects WHERE store_id=? AND sample_id=? " + "AND generation=?", + ("run0", "run0:s0", 1), + ).fetchone()[0] + self.assertEqual(state, "resident") + + def test_lifecycle_bytes_exclude_unrequested_artifacts(self): + sink, writes = self._sink() + sink.put_sample( + self._request(), + aux=torch.ones(3, 4, dtype=torch.bfloat16), + # SGLang produces this tensor for DFlash capture too, but the request + # does not persist it and the owner must not account for it. + last_hidden=torch.ones(3, 8, dtype=torch.bfloat16), + ) + + row = sink._lifecycle.execute( + "SELECT feature_names_json, estimated_bytes FROM mooncake_objects " + "WHERE store_id=? AND sample_id=? AND generation=?", + ("run0", "run0:s0", 1), + ).fetchone() + self.assertEqual(row, ('["hidden_states","input_ids"]', 48)) + self.assertEqual( + [key for key, _ in writes], + ["run0/run0:s0/g1/hidden_states", "run0/run0:s0/g1/input_ids"], + ) + + def test_failed_write_cleans_planned_lifecycle(self): + sink, _ = self._sink() + removed = [] + calls = 0 + + def fail_second_write(key, tensor, **_kwargs): + nonlocal calls + calls += 1 + if calls == 2: + raise RuntimeError("injected put failure") + + sink._put_tensor = fail_second_write + sink._remove_exact = lambda key, *, force: removed.append((key, force)) + + with self.assertRaisesRegex(RuntimeError, "injected put failure"): + sink.put_sample( + self._request(), + aux=torch.ones(3, 4, dtype=torch.bfloat16), + last_hidden=None, + ) + + self.assertEqual( + removed, + [ + ("run0/run0:s0/g1/hidden_states", True), + ("run0/run0:s0/g1/input_ids", True), + ], + ) + state = sink._lifecycle.execute( + "SELECT state FROM mooncake_objects WHERE store_id=? AND sample_id=? " + "AND generation=?", + ("run0", "run0:s0", 1), + ).fetchone()[0] + self.assertEqual(state, "cleaned") + + def test_failed_cleanup_leaves_planned_row_for_owner_takeover(self): + sink, _ = self._sink() + + def fail_write(key, tensor, **_kwargs): + raise RuntimeError("injected put failure") + + def fail_remove(key, *, force): + raise RuntimeError("injected remove failure") + + sink._put_tensor = fail_write + sink._remove_exact = fail_remove + + with self.assertRaisesRegex(RuntimeError, "injected remove failure"): + sink.put_sample( + self._request(), + aux=torch.ones(3, 4, dtype=torch.bfloat16), + last_hidden=None, + ) + + with sqlite3.connect(sink.lifecycle_db_path) as lifecycle: + state = lifecycle.execute( + "SELECT state FROM mooncake_objects WHERE store_id=? AND " + "sample_id=? AND generation=?", + ("run0", "run0:s0", 1), + ).fetchone()[0] + self.assertEqual(state, "planned") + + def test_capability_namespace_schema_and_quota_are_enforced(self): + cases = ( + ({"auth_token": "wrong"}, "capability"), + ({"store_id": "attacker"}, "server-owned namespace"), + ({"sample_id": "other:s0"}, "prefixed"), + ({"features": {"aux": "../../escape"}}, "not allowed"), + ( + { + "passthrough": [ + { + "name": "input_ids", + "data": [1], + "shape": [1, 2], + "dtype": "int64", + } + ] + }, + "requires 2 values", + ), + ) + for override, error in cases: + with self.subTest(override=override): + sink, writes = self._sink() + with self.assertRaisesRegex((ValueError, PermissionError), error): + sink.put_sample( + self._request(**override), + aux=torch.ones(3, 4, dtype=torch.bfloat16), + last_hidden=None, + ) + self.assertEqual(writes, []) + + sink, writes = self._sink(max_bytes=1) + with self.assertRaisesRegex(ValueError, "above the"): + sink.put_sample( + self._request(), + aux=torch.ones(3, 4, dtype=torch.bfloat16), + last_hidden=None, + ) + self.assertEqual(writes, []) + + def test_response_loss_retry_is_idempotent_and_reclaims_prior_generation(self): + inventory = str( + Path(tempfile.mkdtemp(prefix="capture-response-loss-")) / "inventory.db" + ) + first, first_writes = self._sink(inventory_path=inventory) + result = first.put_sample( + self._request(gen=1), + aux=torch.ones(3, 4, dtype=torch.bfloat16), + last_hidden=None, + ) + write_count = len(first_writes) + + repeated = first.put_sample( + self._request(gen=1), + aux=torch.zeros(3, 4, dtype=torch.bfloat16), + last_hidden=None, + ) + self.assertEqual(repeated, result) + self.assertEqual(len(first_writes), write_count) + + restarted, second_writes = self._sink(inventory_path=inventory) + removed = [] + restarted._remove_exact = lambda key, *, force: removed.append((key, force)) + retried = restarted.put_sample( + self._request(gen=2), + aux=torch.ones(3, 4, dtype=torch.bfloat16), + last_hidden=None, + ) + self.assertEqual(retried["gen"], 2) + self.assertEqual( + removed, + [ + ("run0/run0:s0/g1/hidden_states", False), + ("run0/run0:s0/g1/input_ids", False), + ], + ) + self.assertTrue(all("/g2/" in key for key, _ in second_writes)) + + def test_replacement_journal_recovers_at_every_prior_key_delete(self): + for crash_after in (0, 1, 2): + with self.subTest(crash_after=crash_after): + inventory = str( + Path(tempfile.mkdtemp(prefix="capture-replacement-journal-")) + / "inventory.db" + ) + keys = set() + + def attach(sink): + sink._put_tensor = lambda key, tensor, **_kwargs: keys.add(key) + sink._remove_exact = lambda key, *, force: keys.discard(key) + + first, _ = self._sink(inventory_path=inventory) + attach(first) + first.put_sample( + self._request(gen=1), + aux=torch.ones(3, 4, dtype=torch.bfloat16), + last_hidden=None, + ) + old_keys = { + "run0/run0:s0/g1/hidden_states", + "run0/run0:s0/g1/input_ids", + } + self.assertEqual(keys, old_keys) + + interrupted, _ = self._sink(inventory_path=inventory) + remove_calls = 0 + + def crash_during_remove(key, *, force): + nonlocal remove_calls + self.assertFalse(force) + remove_calls += 1 + if crash_after == 0 and remove_calls == 1: + raise RuntimeError("injected replacement crash") + keys.discard(key) + if remove_calls == crash_after: + raise RuntimeError("injected replacement crash") + + interrupted._put_tensor = lambda key, tensor, **_kwargs: keys.add(key) + interrupted._remove_exact = crash_during_remove + with self.assertRaisesRegex(RuntimeError, "replacement crash"): + interrupted.put_sample( + self._request(gen=2), + aux=torch.ones(3, 4, dtype=torch.bfloat16), + last_hidden=None, + ) + + row = interrupted._inventory.execute( + "SELECT generation, state, prior_generation, prior_keys_json " + "FROM captures WHERE sample_id=?", + ("run0:s0",), + ).fetchone() + self.assertEqual(row[:3], (2, "replacing", 1)) + self.assertEqual(set(json.loads(row[3])), old_keys) + + delayed, delayed_writes = self._sink(inventory_path=inventory) + with self.assertRaisesRegex(ValueError, "stale"): + delayed.put_sample( + self._request(gen=1), + aux=torch.ones(3, 4, dtype=torch.bfloat16), + last_hidden=None, + ) + self.assertEqual(delayed_writes, []) + + resumed, _ = self._sink(inventory_path=inventory) + attach(resumed) + result = resumed.put_sample( + self._request(gen=2), + aux=torch.ones(3, 4, dtype=torch.bfloat16), + last_hidden=None, + ) + self.assertEqual(result["gen"], 2) + self.assertEqual( + keys, + { + "run0/run0:s0/g2/hidden_states", + "run0/run0:s0/g2/input_ids", + }, + ) + final = resumed._inventory.execute( + "SELECT generation, state, prior_generation, prior_keys_json " + "FROM captures WHERE sample_id=?", + ("run0:s0",), + ).fetchone() + self.assertEqual(final, (2, "committed", None, None)) + states = dict( + resumed._lifecycle.execute( + "SELECT generation, state FROM mooncake_objects WHERE " + "store_id=? AND sample_id=? ORDER BY generation", + ("run0", "run0:s0"), + ).fetchall() + ) + self.assertEqual(states, {1: "cleaned", 2: "resident"}) + + def test_inventory_schema_migrates_existing_capture_database(self): + inventory = str( + Path(tempfile.mkdtemp(prefix="capture-inventory-migration-")) + / "inventory.db" + ) + with sqlite3.connect(inventory) as connection: + connection.execute( + "CREATE TABLE captures (sample_id TEXT PRIMARY KEY, " + "generation INTEGER NOT NULL, keys_json TEXT NOT NULL, " + "result_json TEXT, state TEXT NOT NULL)" + ) + sink, _ = self._sink(inventory_path=inventory) + columns = { + row[1] for row in sink._inventory.execute("PRAGMA table_info(captures)") + } + self.assertIn("prior_generation", columns) + self.assertIn("prior_keys_json", columns) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_runtime/test_windowed_capture.py b/tests/test_runtime/test_windowed_capture.py new file mode 100644 index 000000000..61917edb4 --- /dev/null +++ b/tests/test_runtime/test_windowed_capture.py @@ -0,0 +1,692 @@ +# coding=utf-8 +"""Deterministic CPU tests for consumer-driven capture windows.""" + +from __future__ import annotations + +import tempfile +import threading +import unittest +from unittest import mock + +from specforge.runtime.contracts import FeatureSpec, SampleRef +from specforge.runtime.data_plane.windowed_capture import ( + CaptureFailedError, + ConsumerFailedError, + SQLiteWindowedCaptureRegistry, + WindowedCaptureQueue, + capture_contract_digest, +) + + +class _Clock: + def __init__(self) -> None: + self.now = 1000.0 + + def __call__(self) -> float: + return self.now + + +class _ReclaimStore: + def __init__(self) -> None: + self.reclaimed: list[tuple[str, int, str]] = [] + + def reclaim(self, ref: SampleRef, *, reason: str = "consumed") -> None: + self.reclaimed.append((ref.sample_id, int(ref.metadata["generation"]), reason)) + + +class _AlreadyReclaimedStore(_ReclaimStore): + def reclaim(self, ref: SampleRef, *, reason: str = "consumed") -> None: + super().reclaim(ref, reason=reason) + raise KeyError("cannot reclaim untracked sample") + + +class WindowedRegistryFixture(unittest.TestCase): + def setUp(self) -> None: + self.tempdir = tempfile.TemporaryDirectory() + self.addCleanup(self.tempdir.cleanup) + self.clock = _Clock() + self.digest = capture_contract_digest( + { + "strategy": "dflash", + "target_model_version": "qwen3-8b-fixture", + "features": ["input_ids", "hidden_states", "loss_mask"], + } + ) + self.registry = self.make_registry(max_live_refs=8) + + def make_registry( + self, + *, + path: str | None = None, + max_live_refs: int, + max_live_bytes: int | None = None, + reservation_bytes: int | None = None, + ) -> SQLiteWindowedCaptureRegistry: + registry = SQLiteWindowedCaptureRegistry( + path or f"{self.tempdir.name}/registry.db", + max_live_refs=max_live_refs, + max_live_bytes=max_live_bytes, + capture_reservation_bytes=reservation_bytes, + clock=self.clock, + poll_s=0.001, + ) + self.addCleanup(registry.close) + return registry + + def initialize( + self, *, samples: int = 8, consumers: tuple[str, ...] = ("a", "b", "c") + ) -> None: + self.registry.initialize_run( + run_id="capture-run", + contract_digest=self.digest, + source_sample_ids=[f"source-{index}" for index in range(samples)], + expected_consumers=consumers, + ) + + def register( + self, + consumer_id: str, + *, + lookbehind: int = 0, + lookahead: int = 0, + prefetch_depth: int = 0, + max_outstanding: int = 1, + ) -> None: + self.registry.register_consumer( + consumer_id, + lookbehind=lookbehind, + lookahead=lookahead, + prefetch_depth=prefetch_depth, + max_outstanding=max_outstanding, + ) + + @staticmethod + def ref(request, *, estimated_bytes: int = 80) -> SampleRef: + sample_id = f"capture-run:{request.key.source_sample_id}" + return SampleRef( + sample_id=sample_id, + run_id="capture-run", + source_task_id=request.key.source_sample_id, + feature_store_uri=f"fixture://capture-run/{sample_id}", + feature_keys={"input_ids": f"{sample_id}/input_ids"}, + feature_specs={ + "input_ids": FeatureSpec(name="input_ids", shape=(1, 4), dtype="int64") + }, + strategy="dflash", + estimated_bytes=estimated_bytes, + metadata={"generation": request.generation, "target_repr": None}, + ) + + def complete(self, request, *, estimated_bytes: int = 80) -> SampleRef: + ref = self.ref(request, estimated_bytes=estimated_bytes) + self.registry.mark_committing(request, ref) + self.registry.complete_capture(request, ref) + return ref + + +class TestWindowedCaptureIdentity(WindowedRegistryFixture): + def test_contract_digest_is_canonical_and_contract_sensitive(self): + left = capture_contract_digest({"layers": (1, 2), "features": {"a", "b"}}) + right = capture_contract_digest({"features": {"b", "a"}, "layers": [1, 2]}) + other = capture_contract_digest({"features": {"a"}, "layers": [1, 2]}) + + self.assertEqual(left, right) + self.assertNotEqual(left, other) + + def test_existing_run_rejects_source_or_capacity_identity_change(self): + self.initialize(samples=2, consumers=("a",)) + with self.assertRaisesRegex(RuntimeError, "identity mismatch"): + self.registry.initialize_run( + run_id="capture-run", + contract_digest=self.digest, + source_sample_ids=["source-1", "source-0"], + expected_consumers=("a",), + ) + + other = self.make_registry( + path=f"{self.tempdir.name}/registry.db", max_live_refs=7 + ) + with self.assertRaisesRegex(RuntimeError, "identity mismatch"): + other.initialize_run( + run_id="capture-run", + contract_digest=self.digest, + source_sample_ids=["source-0", "source-1"], + expected_consumers=("a",), + ) + + +class TestWindowedCaptureStateMachine(WindowedRegistryFixture): + def test_three_misses_singleflight_and_share_one_generation(self): + self.initialize(samples=2) + for consumer_id in ("a", "b", "c"): + self.register(consumer_id) + tickets = [ + self.registry.request_acquire(consumer_id, 0) + for consumer_id in ("a", "b", "c") + ] + + requests = self.registry.claim_batch(8) + + self.assertEqual(len(requests), 1) + self.assertEqual(requests[0].demand_consumers, ("a", "b", "c")) + self.complete(requests[0]) + leases = [self.registry.wait_ready(ticket, timeout_s=1.0) for ticket in tickets] + self.assertEqual({lease.generation for lease in leases}, {1}) + for lease in leases: + self.registry.release_and_advance(lease.consumer_id, [lease]) + candidates = self.registry.begin_evictions() + self.assertEqual([candidate.source_index for candidate in candidates], [0]) + + def test_duplicate_request_is_rejected_without_duplicate_capture(self): + self.initialize(samples=1, consumers=("a", "b")) + self.register("a") + self.register("b") + first = self.registry.request_acquire("a", 0) + with self.assertRaisesRegex(RuntimeError, "already has"): + self.registry.request_acquire("a", 0) + second = self.registry.request_acquire("b", 0) + + request = self.registry.claim_batch(2)[0] + self.complete(request) + self.assertEqual( + self.registry.wait_ready(first, timeout_s=1.0).generation, + self.registry.wait_ready(second, timeout_s=1.0).generation, + ) + self.assertEqual(self.registry.snapshot()["capture_count"], 1) + + def test_heterogeneous_windows_replenish_independent_frontiers(self): + self.registry.close() + self.registry = self.make_registry(max_live_refs=5) + self.initialize(samples=7, consumers=("fast", "slow")) + self.register("fast", lookahead=3, prefetch_depth=4, max_outstanding=2) + self.register("slow", lookahead=0, prefetch_depth=1) + initial = self.registry.claim_batch(8) + self.assertEqual([request.source_index for request in initial], [0, 1, 2, 3]) + for request in initial: + self.complete(request) + + ticket = self.registry.request_acquire("fast", 0) + lease = self.registry.wait_ready(ticket, timeout_s=1.0) + self.assertTrue(lease.ready_at_request) + self.registry.release_and_advance("fast", [lease]) + frontier = self.registry.claim_batch(8) + + self.assertEqual([request.source_index for request in frontier], [4]) + snapshot = self.registry.snapshot() + self.assertEqual(snapshot["consumers"]["fast"]["cursor"], 1) + self.assertEqual(snapshot["consumers"]["slow"]["cursor"], 0) + + def test_prefetch_slots_roll_forward_through_the_legal_window(self): + self.initialize(samples=8, consumers=("fast",)) + self.register("fast", lookahead=5, prefetch_depth=2) + + for expected in ((0, 1), (2, 3), (4, 5)): + requests = self.registry.claim_batch(8) + self.assertEqual( + tuple(request.source_index for request in requests), + expected, + ) + for request in requests: + self.complete(request) + + self.assertEqual(self.registry.claim_batch(8), ()) + snapshot = self.registry.snapshot() + self.assertEqual(snapshot["capture_count"], 6) + self.assertEqual(snapshot["consumers"]["fast"]["cursor"], 0) + + def test_round_robin_demand_fairness(self): + self.initialize(samples=6) + for consumer_id in ("a", "b", "c"): + self.register(consumer_id, lookahead=5, max_outstanding=3) + for consumer_id, indices in ( + ("a", (0, 1, 2)), + ("b", (3,)), + ("c", (4,)), + ): + for source_index in indices: + self.registry.request_acquire(consumer_id, source_index) + + first = self.registry.claim_batch(3) + + self.assertEqual({request.source_index for request in first}, {0, 3, 4}) + + def test_retry_generation_is_unique_and_terminal_failure_wakes_waiter(self): + self.initialize(samples=1, consumers=("a",)) + self.register("a") + ticket = self.registry.request_acquire("a", 0) + first = self.registry.claim_batch(1)[0] + + self.assertTrue( + self.registry.fail_capture( + first, "transient", retryable=True, max_retries=1 + ) + ) + second = self.registry.claim_batch(1)[0] + self.assertEqual((first.generation, second.generation), (1, 2)) + self.assertFalse( + self.registry.fail_capture( + second, "still broken", retryable=True, max_retries=1 + ) + ) + with self.assertRaisesRegex(CaptureFailedError, "still broken"): + self.registry.wait_ready(ticket, timeout_s=1.0) + + def test_terminal_failure_is_observed_by_all_waiters_before_recapture(self): + self.initialize(samples=1) + for consumer_id in ("a", "b", "c"): + self.register(consumer_id) + first = self.registry.request_acquire("a", 0) + second = self.registry.request_acquire("b", 0) + failed = self.registry.claim_batch(1)[0] + self.assertFalse( + self.registry.fail_capture( + failed, "terminal generation", retryable=False, max_retries=0 + ) + ) + + with self.assertRaisesRegex(CaptureFailedError, "terminal generation"): + self.registry.wait_ready(first, timeout_s=1.0) + late = self.registry.request_acquire("c", 0) + self.assertEqual(self.registry.claim_batch(1), ()) + for ticket in (second, late): + with self.assertRaisesRegex(CaptureFailedError, "terminal generation"): + self.registry.wait_ready(ticket, timeout_s=1.0) + + retry = self.registry.request_acquire("a", 0) + replacement = self.registry.claim_batch(1)[0] + self.assertEqual(replacement.generation, failed.generation + 1) + self.complete(replacement) + self.assertEqual( + self.registry.wait_ready(retry, timeout_s=1.0).generation, + replacement.generation, + ) + + def test_fresh_demand_resets_exhausted_capture_retry_budget(self): + self.initialize(samples=1, consumers=("a",)) + self.register("a") + failed_ticket = self.registry.request_acquire("a", 0) + failed = self.registry.claim_batch(1)[0] + self.assertFalse( + self.registry.fail_capture( + failed, "first demand exhausted", retryable=True, max_retries=0 + ) + ) + with self.assertRaisesRegex(CaptureFailedError, "first demand exhausted"): + self.registry.wait_ready(failed_ticket, timeout_s=1.0) + + retry_ticket = self.registry.request_acquire("a", 0) + fresh = self.registry.claim_batch(1)[0] + self.assertTrue( + self.registry.fail_capture( + fresh, "fresh transient", retryable=True, max_retries=1 + ) + ) + replacement = self.registry.claim_batch(1)[0] + self.complete(replacement) + self.assertEqual( + self.registry.wait_ready(retry_ticket, timeout_s=1.0).generation, + replacement.generation, + ) + + def test_failed_consumer_drops_only_its_waiter_and_window(self): + self.initialize(samples=2, consumers=("failed", "healthy")) + self.register("failed") + self.register("healthy") + failed_ticket = self.registry.request_acquire("failed", 0) + healthy_ticket = self.registry.request_acquire("healthy", 0) + request = self.registry.claim_batch(1)[0] + + self.registry.fail_consumer("failed", "trainer exited") + self.complete(request) + with self.assertRaisesRegex(ConsumerFailedError, "trainer exited"): + self.registry.wait_ready(failed_ticket, timeout_s=1.0) + healthy_lease = self.registry.wait_ready(healthy_ticket, timeout_s=1.0) + self.assertEqual(healthy_lease.source_index, 0) + self.registry.release_and_advance("healthy", [healthy_lease]) + self.assertEqual(self.registry.snapshot()["consumers"]["healthy"]["cursor"], 1) + + def test_read_lease_blocks_pressure_eviction_until_release(self): + self.initialize(samples=1, consumers=("a",)) + self.register("a") + ticket = self.registry.request_acquire("a", 0) + request = self.registry.claim_batch(1)[0] + self.complete(request) + lease = self.registry.wait_ready(ticket, timeout_s=1.0) + + self.assertEqual(self.registry.begin_evictions(pressure=True), ()) + self.registry.release_and_advance("a", [lease]) + candidates = self.registry.begin_evictions(pressure=True) + + self.assertEqual([item.source_index for item in candidates], [0]) + + def test_expired_consumer_releases_its_lease_without_stopping_peer(self): + self.initialize(samples=1, consumers=("expired", "healthy")) + self.register("expired") + self.register("healthy") + ticket = self.registry.request_acquire("expired", 0) + request = self.registry.claim_batch(1)[0] + self.complete(request) + self.registry.wait_ready(ticket, timeout_s=1.0) + self.registry.heartbeat("healthy") + self.clock.now += 11.0 + self.registry.heartbeat("healthy") + + self.assertEqual(self.registry.expire_consumers(10.0), ("expired",)) + snapshot = self.registry.snapshot() + + self.assertEqual(snapshot["leases"], 0) + self.assertNotEqual(snapshot["consumers"]["healthy"]["state"], "failed") + + def test_pressure_reclaim_unblocks_demand_without_evicting_hard_interest(self): + self.registry.close() + self.registry = self.make_registry(max_live_refs=2) + self.initialize(samples=3, consumers=("fast",)) + self.register("fast", lookahead=2, prefetch_depth=2) + for request in self.registry.claim_batch(2): + self.complete(request) + ticket = self.registry.request_acquire("fast", 2) + self.assertEqual(self.registry.claim_batch(1), ()) + + store = _ReclaimStore() + self.assertEqual(self.registry.reclaim(store, limit=1, pressure=True), 1) + demand = self.registry.claim_batch(1) + + self.assertEqual([request.source_index for request in demand], [2]) + self.assertNotIn(ticket.key.source_sample_id, store.reclaimed[0][0]) + self.assertLessEqual(self.registry.snapshot()["live_refs"], 2) + + def test_byte_reservations_are_hard_and_oversize_completion_is_rejected(self): + self.registry.close() + self.registry = self.make_registry( + max_live_refs=4, + max_live_bytes=200, + reservation_bytes=100, + ) + self.initialize(samples=3, consumers=("a",)) + self.register("a", lookahead=2, prefetch_depth=3) + requests = self.registry.claim_batch(4) + self.assertEqual(len(requests), 2) + self.complete(requests[0], estimated_bytes=80) + with self.assertRaisesRegex(ValueError, "exceeds reserved"): + self.registry.mark_committing( + requests[1], self.ref(requests[1], estimated_bytes=101) + ) + self.registry.fail_capture( + requests[1], "oversize", retryable=False, max_retries=0 + ) + snapshot = self.registry.snapshot() + self.assertLessEqual(snapshot["peak_live_bytes"], 200) + self.assertLessEqual(snapshot["peak_live_refs"], 2) + + +class TestWindowedCaptureResume(WindowedRegistryFixture): + def test_owner_recovery_fences_old_generation_and_durable_cursor_rewinds(self): + self.initialize(samples=2, consumers=("a",)) + self.register("a", lookahead=1, max_outstanding=2) + ticket = self.registry.request_acquire("a", 0) + interrupted = self.registry.claim_batch(1)[0] + path = self.registry.path + self.registry.close() + + resumed = self.make_registry(path=path, max_live_refs=8) + self.registry = resumed + with self.assertRaisesRegex(RuntimeError, "recover_inflight"): + resumed.initialize_run( + run_id="capture-run", + contract_digest=self.digest, + source_sample_ids=["source-0", "source-1"], + expected_consumers=("a",), + ) + resumed.initialize_run( + run_id="capture-run", + contract_digest=self.digest, + source_sample_ids=["source-0", "source-1"], + expected_consumers=("a",), + recover_inflight=True, + ) + replacement = resumed.claim_batch(1)[0] + self.assertEqual((interrupted.generation, replacement.generation), (1, 2)) + self.complete(replacement) + lease = resumed.wait_ready(ticket, timeout_s=1.0) + resumed.release_and_advance("a", [lease]) + self.assertEqual(resumed.consumer_cursor("a"), 1) + + resumed.resume_consumer("a", durable_cursor=0) + replay = resumed.request_acquire("a", 0) + replay_lease = resumed.wait_ready(replay, timeout_s=1.0) + self.assertTrue(replay_lease.ready_at_request) + self.assertEqual(replay_lease.generation, 2) + + def test_committing_recovery_reclaims_payload_before_requeue(self): + self.initialize(samples=1, consumers=("a",)) + self.register("a") + ticket = self.registry.request_acquire("a", 0) + interrupted = self.registry.claim_batch(1)[0] + ref = self.ref(interrupted) + self.registry.mark_committing(interrupted, ref) + path = self.registry.path + self.registry.close() + + resumed = self.make_registry(path=path, max_live_refs=8) + self.registry = resumed + with self.assertRaisesRegex(RuntimeError, "recovery_store"): + resumed.initialize_run( + run_id="capture-run", + contract_digest=self.digest, + source_sample_ids=["source-0"], + expected_consumers=("a",), + recover_inflight=True, + ) + store = _ReclaimStore() + resumed.initialize_run( + run_id="capture-run", + contract_digest=self.digest, + source_sample_ids=["source-0"], + expected_consumers=("a",), + recover_inflight=True, + recovery_store=store, + ) + + self.assertEqual( + store.reclaimed, [(ref.sample_id, 1, "interrupted-capture-recovery")] + ) + replacement = resumed.claim_batch(1)[0] + self.assertEqual(replacement.generation, 2) + self.complete(replacement) + lease = resumed.wait_ready(ticket, timeout_s=1.0) + self.assertEqual(lease.generation, 2) + + def test_committing_recovery_treats_already_reclaimed_payload_as_complete(self): + self.initialize(samples=1, consumers=("a",)) + self.register("a") + ticket = self.registry.request_acquire("a", 0) + interrupted = self.registry.claim_batch(1)[0] + self.registry.mark_committing(interrupted, self.ref(interrupted)) + path = self.registry.path + self.registry.close() + + resumed = self.make_registry(path=path, max_live_refs=8) + self.registry = resumed + store = _AlreadyReclaimedStore() + resumed.initialize_run( + run_id="capture-run", + contract_digest=self.digest, + source_sample_ids=["source-0"], + expected_consumers=("a",), + recover_inflight=True, + recovery_store=store, + ) + + replacement = resumed.claim_batch(1)[0] + self.assertEqual(replacement.generation, interrupted.generation + 1) + self.complete(replacement) + self.assertEqual(resumed.wait_ready(ticket, timeout_s=1.0).generation, 2) + + +class TestWindowedCaptureQueueAndSoak(WindowedRegistryFixture): + def test_1p1c_queue_preserves_order_and_requires_prefix_ack(self): + self.registry.close() + self.registry = self.make_registry(max_live_refs=3) + self.initialize(samples=3, consumers=("trainer",)) + self.register("trainer", lookahead=2, prefetch_depth=3, max_outstanding=2) + for request in self.registry.claim_batch(3): + self.complete(request) + queue = WindowedCaptureQueue(self.registry, "trainer", idle_timeout_s=1.0) + + first = queue.get(2) + self.assertEqual( + [ref.source_task_id for ref in first], ["source-0", "source-1"] + ) + with self.assertRaisesRegex(RuntimeError, "leased prefix"): + queue.ack_ids([first[1].sample_id]) + queue.ack_ids([ref.sample_id for ref in first]) + with self.assertRaisesRegex(RuntimeError, "leased prefix"): + queue.ack(first) + last = queue.get(1) + queue.ack(last) + self.assertEqual(queue.get(1), []) + metrics = queue.metrics() + self.assertEqual(metrics["refs"], 3) + self.assertEqual(metrics["ready_at_request_refs"], 3) + self.assertEqual(metrics["ready_at_request_ratio"], 1.0) + self.assertEqual(metrics["next_fetch"], 3) + self.assertEqual(metrics["in_flight"], 0) + self.assertEqual( + self.registry.snapshot()["consumers"]["trainer"]["state"], + "completed", + ) + + def test_concurrent_retry_rewind_is_not_overwritten_by_get(self): + self.initialize(samples=4, consumers=("trainer",)) + self.register("trainer", lookahead=3, prefetch_depth=4, max_outstanding=4) + for request in self.registry.claim_batch(4): + self.complete(request) + queue = WindowedCaptureQueue(self.registry, "trainer", idle_timeout_s=1.0) + first = queue.get(2) + original_request_many = self.registry.request_many + entered = threading.Event() + release = threading.Event() + + def blocked_request_many(consumer_id, source_indices): + indices = tuple(source_indices) + tickets = original_request_many(consumer_id, indices) + if indices == (2, 3): + entered.set() + self.assertTrue(release.wait(timeout=1.0)) + return tickets + + result: list[list[SampleRef]] = [] + with mock.patch.object( + self.registry, "request_many", side_effect=blocked_request_many + ): + thread = threading.Thread(target=lambda: result.append(queue.get(2))) + thread.start() + self.assertTrue(entered.wait(timeout=1.0)) + queue.fail(first, "retry", retryable=True) + release.set() + thread.join(timeout=2.0) + self.assertFalse(thread.is_alive()) + + self.assertEqual( + [ref.source_task_id for ref in result[0]], ["source-0", "source-1"] + ) + queue.ack(result[0]) + remaining = queue.get(2) + self.assertEqual( + [ref.source_task_id for ref in remaining], ["source-2", "source-3"] + ) + queue.ack(remaining) + + def test_reclaim_preserves_store_error_when_metadata_cleanup_also_fails(self): + self.initialize(samples=2, consumers=("trainer",)) + self.register("trainer", lookahead=1, prefetch_depth=2, max_outstanding=2) + for request in self.registry.claim_batch(2): + self.complete(request) + queue = WindowedCaptureQueue(self.registry, "trainer", idle_timeout_s=1.0) + refs = queue.get(2) + queue.ack(refs) + + store = _ReclaimStore() + original_reclaim = store.reclaim + + def fail_second(ref, *, reason="consumed"): + if store.reclaimed: + raise OSError("physical reclaim failed") + original_reclaim(ref, reason=reason) + + store.reclaim = fail_second + with ( + mock.patch.object( + self.registry, + "finish_evictions", + side_effect=RuntimeError("metadata finish failed"), + ), + self.assertRaisesRegex(OSError, "physical reclaim failed") as raised, + ): + self.registry.reclaim(store) + + self.assertTrue( + any("metadata finish failed" in note for note in raised.exception.__notes__) + ) + self.assertEqual(len(self.registry.begin_evictions()), 2) + + def test_skewed_1p3c_soak_stays_within_slots_and_bytes(self): + self.registry.close() + self.registry = self.make_registry( + max_live_refs=9, + max_live_bytes=900, + reservation_bytes=100, + ) + total = 200 + consumers = ("fast", "medium", "slow") + self.initialize(samples=total, consumers=consumers) + self.register("fast", lookahead=4, prefetch_depth=5) + self.register("medium", lookahead=2, prefetch_depth=3) + self.register("slow", lookahead=0, prefetch_depth=1) + store = _ReclaimStore() + delivered = {consumer_id: [] for consumer_id in consumers} + + def consume_one(consumer_id: str) -> None: + cursor = self.registry.consumer_cursor(consumer_id) + if cursor >= total: + return + ticket = self.registry.request_acquire(consumer_id, cursor) + requests = self.registry.claim_batch(16) + if not requests: + reclaimed = self.registry.reclaim( + store, limit=1, pressure=True, reason="capacity" + ) + self.assertEqual(reclaimed, 1) + requests = self.registry.claim_batch(16) + for request in requests: + self.complete(request) + lease = self.registry.wait_ready(ticket, timeout_s=1.0) + delivered[consumer_id].append(lease.source_index) + self.registry.release_and_advance(consumer_id, [lease]) + self.registry.reclaim(store, limit=32) + snapshot = self.registry.snapshot() + self.assertLessEqual(snapshot["live_refs"], 9) + self.assertLessEqual(snapshot["live_bytes"], 900) + + while any(self.registry.consumer_cursor(item) < total for item in consumers): + for _ in range(5): + consume_one("fast") + for _ in range(2): + consume_one("medium") + consume_one("slow") + + for consumer_id in consumers: + self.registry.complete_consumer(consumer_id) + self.registry.reclaim(store, limit=32) + snapshot = self.registry.snapshot() + + for consumer_id in consumers: + self.assertEqual(delivered[consumer_id], list(range(total))) + self.assertLessEqual(snapshot["peak_live_refs"], 9) + self.assertLessEqual(snapshot["peak_live_bytes"], 900) + self.assertEqual(snapshot["live_refs"], 0) + self.assertEqual(self.registry.finalize_run(), "completed") + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_runtime/test_windowed_capture_runtime.py b/tests/test_runtime/test_windowed_capture_runtime.py new file mode 100644 index 000000000..662db440b --- /dev/null +++ b/tests/test_runtime/test_windowed_capture_runtime.py @@ -0,0 +1,763 @@ +# coding=utf-8 +"""Deterministic tests for process-facing windowed capture runtime loops.""" + +from __future__ import annotations + +import dataclasses +import inspect +import os +import sqlite3 +import tempfile +import threading +import time +import unittest +from dataclasses import dataclass +from unittest import mock + +import torch + +from specforge.algorithms.builtin import builtin_algorithm_registry +from specforge.inference.capture import CaptureConfig +from specforge.launch import ( + build_disagg_online_windowed_consumer, + build_disagg_online_windowed_producer, + build_disagg_windowed_capture_contract, +) +from specforge.runtime.contracts import FeatureSpec, PromptTask, SampleRef +from specforge.runtime.control_plane.metadata_store import SQLiteMetadataStore +from specforge.runtime.data_plane.windowed_capture import ( + CaptureFailedError, + SQLiteWindowedCaptureRegistry, + WindowedCaptureQueue, + capture_contract_digest, +) +from specforge.runtime.data_plane.windowed_capture_runtime import ( + WindowedCaptureService, + WindowedConsumerControl, + start_windowed_consumer_control, +) +from specforge.training.checkpoint import STATE_FILE + + +class _OwnerStore: + lifetime_owner = True + + def __init__(self) -> None: + self.live: dict[str, int] = {} + self.reclaimed: list[tuple[str, int, str]] = [] + + def adopt(self, ref: SampleRef) -> None: + self.live[ref.sample_id] = int(ref.metadata["generation"]) + + def reclaim(self, ref: SampleRef, *, reason: str) -> None: + generation = int(ref.metadata["generation"]) + self.reclaimed.append((ref.sample_id, generation, reason)) + if self.live.get(ref.sample_id) == generation: + self.live.pop(ref.sample_id) + + def gc(self): + return {} + + +class _ReaderStore: + lifetime_owner = False + + +@dataclass(frozen=True) +class _Failure: + task_id: str + reason: str + retryable: bool + + +class _RefSource: + def __init__(self, store: _OwnerStore, *, fail_once: bool = False) -> None: + self.store = store + self.fail_once = fail_once + self.calls = 0 + self.generations: dict[str, int] = {} + + def produce_refs(self, tasks, *, capture): + del capture + self.calls += 1 + out = [] + for task in tasks: + if self.fail_once and self.calls == 1: + out.append(_Failure(task.task_id, "transient", True)) + continue + generation = self.generations.get(task.task_id, 0) + 1 + self.generations[task.task_id] = generation + sample_id = f"run:{task.task_id}" + ref = SampleRef( + sample_id=sample_id, + run_id="run", + source_task_id=task.task_id, + feature_store_uri=f"fixture://run/{sample_id}?generation={generation}", + feature_keys={"input_ids": f"{sample_id}/input_ids"}, + feature_specs={ + "input_ids": FeatureSpec( + name="input_ids", shape=(1, 4), dtype="int64" + ) + }, + strategy="dflash", + estimated_bytes=64, + metadata={"generation": generation}, + ) + self.store.adopt(ref) + out.append(ref) + return out + + +def _capture() -> CaptureConfig: + return CaptureConfig.from_strategy( + required_features={"input_ids"}, + aux_hidden_state_layer_ids=(), + target_repr="hidden_state", + target_hidden_size=8, + ) + + +def _prompts(total: int) -> list[PromptTask]: + return [ + PromptTask( + task_id=f"source-{index}", + run_id="run", + source_id="fixture", + payload={"input_ids": [index, index + 1]}, + max_length=8, + ) + for index in range(total) + ] + + +class TestWindowedCaptureService(unittest.TestCase): + def setUp(self) -> None: + self.tempdir = tempfile.TemporaryDirectory() + self.addCleanup(self.tempdir.cleanup) + + def registry(self, *, max_live_refs=4, consumers=("a",)): + registry = SQLiteWindowedCaptureRegistry( + os.path.join(self.tempdir.name, "window.db"), + max_live_refs=max_live_refs, + max_live_bytes=max_live_refs * 100, + capture_reservation_bytes=100, + poll_s=0.001, + ) + registry.initialize_run( + run_id="run", + contract_digest=capture_contract_digest(_capture()), + source_sample_ids=[task.task_id for task in _prompts(12)], + expected_consumers=consumers, + ) + self.addCleanup(registry.close) + return registry + + def test_rejects_negative_batch_wait(self): + registry = self.registry() + store = _OwnerStore() + with self.assertRaisesRegex(ValueError, "batch_wait_s"): + WindowedCaptureService( + registry, + prompts=_prompts(12), + feature_source=_RefSource(store), + capture=_capture(), + owner_store=store, + batch_wait_s=-0.001, + ) + + def test_capture_task_carries_the_registry_generation(self): + registry = self.registry(max_live_refs=1) + registry.register_consumer("a") + registry.request_acquire("a", 0) + request = registry.claim_batch(1)[0] + store = _OwnerStore() + service = WindowedCaptureService( + registry, + prompts=_prompts(12), + feature_source=_RefSource(store), + capture=_capture(), + owner_store=store, + ) + + task = service._task(request) + + self.assertEqual(task.metadata["capture_generation"], request.generation) + self.assertEqual(task.attempt, request.generation - 1) + + def test_physical_generation_is_not_rewritten_by_window_generation(self): + registry = self.registry(max_live_refs=1) + registry.register_consumer("a") + ticket = registry.request_acquire("a", 0) + request = registry.claim_batch(1)[0] + ref = SampleRef( + sample_id="run:source-0", + run_id="run", + source_task_id="source-0", + feature_store_uri="fixture://run/source-0?generation=41", + feature_keys={}, + feature_specs={}, + strategy="dflash", + estimated_bytes=1, + metadata={"generation": 41, "window_generation": request.generation}, + ) + registry.mark_committing(request, ref) + registry.complete_capture(request, ref) + + lease = registry.wait_ready(ticket, timeout_s=1.0) + self.assertEqual(lease.ref.metadata["generation"], 41) + self.assertEqual(lease.ref.metadata["window_generation"], request.generation) + + def test_registry_rejects_capture_ref_from_another_run(self): + registry = self.registry(max_live_refs=1) + registry.register_consumer("a") + registry.request_acquire("a", 0) + request = registry.claim_batch(1)[0] + ref = _RefSource(_OwnerStore()).produce_refs( + [_prompts(1)[0]], capture=_capture() + )[0] + ref = dataclasses.replace( + ref, + run_id="another-run", + metadata={**ref.metadata, "window_generation": request.generation}, + ) + + with self.assertRaisesRegex(ValueError, "does not match registry run"): + registry.mark_committing(request, ref) + + def test_heartbeat_treats_concurrent_completion_as_terminal_success(self): + registry = self.registry(max_live_refs=1) + registry.register_consumer("a", cursor=12) + control = start_windowed_consumer_control( + registry, + "a", + lookbehind=0, + lookahead=0, + prefetch_depth=0, + max_outstanding=1, + heartbeat_interval_s=0.01, + ) + registry.complete_consumer("a") + time.sleep(0.03) + + control.ensure_healthy() + control.close() + + def test_heartbeat_recovers_from_a_transient_sqlite_failure(self): + registry = mock.Mock() + registry.heartbeat.side_effect = [ + None, + sqlite3.OperationalError("database is locked"), + None, + ] + control = WindowedConsumerControl( + registry=registry, + consumer_id="a", + heartbeat_interval_s=0.001, + ).start() + deadline = time.monotonic() + 1.0 + while registry.heartbeat.call_count < 3 and time.monotonic() < deadline: + time.sleep(0.001) + + control.ensure_healthy() + control.close() + self.assertGreaterEqual(registry.heartbeat.call_count, 3) + + def test_external_ledger_failure_abandons_queue_lease(self): + registry = self.registry(max_live_refs=1) + registry.register_consumer("a") + ticket = registry.request_acquire("a", 0) + request = registry.claim_batch(1)[0] + source = _RefSource(_OwnerStore()) + ref = source.produce_refs([_prompts(1)[0]], capture=_capture())[0] + ref = SampleRef( + **{ + **ref.__dict__, + "metadata": {**ref.metadata, "window_generation": request.generation}, + } + ) + registry.mark_committing(request, ref) + registry.complete_capture(request, ref) + registry.cancel_acquire(ticket) + queue = WindowedCaptureQueue( + registry, + "a", + idle_timeout_s=1.0, + record_refs=lambda _refs: (_ for _ in ()).throw(RuntimeError("ledger")), + ) + + with self.assertRaisesRegex(RuntimeError, "ledger"): + queue.get(1) + self.assertEqual(registry.snapshot()["leases"], 0) + + def test_skewed_1p3c_and_partial_completion_stay_bounded(self): + consumers = ("fast", "medium", "short") + registry = self.registry(max_live_refs=4, consumers=consumers) + registry.register_consumer( + "fast", lookahead=3, prefetch_depth=4, max_outstanding=1 + ) + registry.register_consumer( + "medium", lookahead=1, prefetch_depth=2, max_outstanding=1 + ) + registry.register_consumer("short", max_outstanding=1) + store = _OwnerStore() + source = _RefSource(store) + service = WindowedCaptureService( + registry, + prompts=_prompts(12), + feature_source=source, + capture=_capture(), + owner_store=store, + capture_batch_size=3, + consumer_registration_timeout_s=1.0, + consumer_heartbeat_timeout_s=30.0, + poll_s=0.001, + ) + delivered = {consumer: [] for consumer in consumers} + errors = [] + + def consume(consumer: str, count: int, delay: float) -> None: + try: + queue = WindowedCaptureQueue(registry, consumer, idle_timeout_s=2.0) + while len(delivered[consumer]) < count: + refs = queue.get(1) + if not refs: + break + delivered[consumer].append(refs[0].source_task_id) + queue.ack(refs) + if delay: + time.sleep(delay) + if count < 12: + queue.complete(allow_partial=True) + else: + self.assertEqual(queue.get(1), []) + except BaseException as exc: + errors.append(exc) + + service_result = [] + + def produce() -> None: + try: + service_result.append(service.drive(max_rounds=100_000)) + except BaseException as exc: + errors.append(exc) + + threads = [ + threading.Thread(target=produce), + threading.Thread(target=consume, args=("fast", 12, 0.0)), + threading.Thread(target=consume, args=("medium", 12, 0.001)), + threading.Thread(target=consume, args=("short", 3, 0.002)), + ] + for thread in threads: + thread.start() + for thread in threads: + thread.join(timeout=10.0) + + self.assertFalse(any(thread.is_alive() for thread in threads)) + self.assertEqual(errors, []) + self.assertEqual(delivered["fast"], [f"source-{i}" for i in range(12)]) + self.assertEqual(delivered["medium"], [f"source-{i}" for i in range(12)]) + self.assertEqual(delivered["short"], [f"source-{i}" for i in range(3)]) + snapshot = registry.snapshot() + self.assertLessEqual(snapshot["peak_live_refs"], 4) + self.assertLessEqual(snapshot["peak_live_bytes"], 400) + self.assertEqual(snapshot["live_refs"], 0) + self.assertEqual(snapshot["status"], "completed") + self.assertTrue(service_result) + self.assertEqual(store.live, {}) + + def test_retryable_source_result_is_retried(self): + registry = self.registry(max_live_refs=2) + registry.register_consumer("a") + store = _OwnerStore() + source = _RefSource(store, fail_once=True) + service = WindowedCaptureService( + registry, + prompts=_prompts(12), + feature_source=source, + capture=_capture(), + owner_store=store, + max_capture_retries=1, + retry_backoff_s=0, + consumer_registration_timeout_s=1.0, + consumer_heartbeat_timeout_s=30.0, + poll_s=0.001, + ) + queue = WindowedCaptureQueue(registry, "a", idle_timeout_s=2.0) + errors = [] + producer = threading.Thread( + target=lambda: self._drive_or_record(service, errors) + ) + producer.start() + received = [] + while True: + refs = queue.get(1) + if not refs: + break + received.extend(ref.source_task_id for ref in refs) + queue.ack(refs) + producer.join(timeout=5.0) + + self.assertFalse(producer.is_alive()) + self.assertEqual(errors, []) + self.assertEqual(received, [f"source-{i}" for i in range(12)]) + self.assertGreaterEqual(service.capture_failures, 1) + + def test_non_retryable_capture_failure_terminates_and_reclaims(self): + registry = self.registry(max_live_refs=2) + registry.register_consumer("a") + store = _OwnerStore() + service = WindowedCaptureService( + registry, + prompts=_prompts(12), + feature_source=_RefSource(store, fail_once=True), + capture=_capture(), + owner_store=store, + max_capture_retries=0, + retry_backoff_s=0, + consumer_registration_timeout_s=1.0, + consumer_heartbeat_timeout_s=30.0, + poll_s=0.001, + ) + queue = WindowedCaptureQueue(registry, "a", idle_timeout_s=2.0) + errors = [] + producer = threading.Thread( + target=lambda: self._drive_or_record(service, errors) + ) + producer.start() + + with self.assertRaisesRegex(CaptureFailedError, "transient"): + queue.get(1) + queue.close("capture failed") + producer.join(timeout=5.0) + + self.assertFalse(producer.is_alive()) + self.assertEqual(errors, []) + self.assertEqual(service.capture_failures, 1) + self.assertEqual(registry.snapshot()["status"], "completed_with_failures") + self.assertEqual(store.live, {}) + + @staticmethod + def _drive_or_record(service, errors): + try: + service.drive(max_rounds=100_000) + except BaseException as exc: + errors.append(exc) + + +class TestWindowedLaunchBuilders(unittest.TestCase): + def test_dflash_window_contract_does_not_default_to_logits(self): + capture, _digest = build_disagg_windowed_capture_contract( + strategy="dflash", + target_hidden_size=8, + target_model_version="fixture", + tokenizer_version="fixture", + ) + self.assertIsNone(capture.target_repr) + + def test_shared_assembler_forwards_loader_prefetch_depth(self): + from specforge.launch import _assemble_trainer + + assembled = mock.Mock() + assembled.controller = mock.sentinel.trainer + assembled.loader = mock.sentinel.loader + with mock.patch( + "specforge.training.Trainer", return_value=assembled + ) as trainer: + result = _assemble_trainer( + algorithm=builtin_algorithm_registry().resolve("dflash"), + controller=mock.sentinel.controller, + store=mock.sentinel.store, + ref_source={"queue": mock.sentinel.queue}, + model=mock.sentinel.model, + target_head=None, + optimizer_factory=mock.sentinel.optimizer, + run_id="run", + output_dir="output", + batch_size=1, + accumulation_steps=1, + num_epochs=1, + max_steps=1, + save_interval=0, + eval_interval=0, + tp_size=1, + sp_ulysses_size=1, + sp_ring_size=1, + logger=None, + log_interval=1, + collate_fn=mock.sentinel.collate, + dataloader_num_workers=3, + ) + + self.assertIs(result, assembled) + self.assertEqual(trainer.call_args.kwargs["dataloader_num_workers"], 3) + + def test_producer_requires_stable_task_ids(self): + with tempfile.TemporaryDirectory() as root: + with self.assertRaisesRegex(ValueError, "stable task_id"): + build_disagg_online_windowed_producer( + prompts=[{"payload": {"input_ids": [1]}}], + feature_store=_OwnerStore(), + feature_source=_RefSource(_OwnerStore()), + run_id="run", + consumer_ids=("a",), + registry_db_path=os.path.join(root, "window.db"), + max_live_refs=2, + target_hidden_size=8, + target_model_version="fixture", + tokenizer_version="fixture", + strategy="dflash", + target_repr="hidden_state", + ) + + def test_consumer_builder_ledgers_window_refs_and_validates_capacity(self): + from specforge.launch import _assemble_trainer + + with tempfile.TemporaryDirectory() as root: + owner = _OwnerStore() + producer = build_disagg_online_windowed_producer( + prompts=[ + { + "task_id": "source-0", + "payload": {"input_ids": [1, 2]}, + } + ], + feature_store=owner, + feature_source=_RefSource(owner), + run_id="run", + consumer_ids=("a",), + registry_db_path=os.path.join(root, "window.db"), + max_live_refs=2, + target_hidden_size=8, + target_model_version="fixture", + tokenizer_version="fixture", + strategy="dflash", + target_repr="hidden_state", + ) + fake_trainer = mock.Mock() + fake_loader = mock.Mock() + fake_trainer.loader = fake_loader + with mock.patch( + "specforge.launch._assemble_trainer", + return_value=fake_trainer, + ) as assemble: + runtime = build_disagg_online_windowed_consumer( + consumer_id="a", + registry_db_path=os.path.join(root, "window.db"), + max_live_refs=2, + contract_digest=producer.contract_digest, + total_samples=1, + feature_store=_ReaderStore(), + draft_model=object(), + optimizer_factory=mock.Mock(), + run_id="run", + output_dir=os.path.join(root, "output"), + metadata_db_path=os.path.join(root, "consumer.db"), + strategy="dflash", + max_outstanding=2, + batch_size=1, + accumulation_steps=2, + loader_prefetch_batches=3, + initialization_timeout_s=1.0, + heartbeat_interval_s=0.1, + ) + try: + self.assertIsNone(runtime.controller.sample_queue) + self.assertTrue(assemble.call_args.kwargs["durable_ack"]) + self.assertEqual(assemble.call_args.kwargs["dataloader_num_workers"], 3) + unexpected = set(assemble.call_args.kwargs) - set( + inspect.signature(_assemble_trainer).parameters + ) + self.assertEqual(unexpected, set()) + self.assertEqual(runtime.accounting_snapshot()["queue"]["refs"], 0) + finally: + runtime.control.fail("test cleanup") + runtime.close() + producer.close() + + def test_consumer_resume_requires_checkpoint_for_durable_prefix(self): + with tempfile.TemporaryDirectory() as root: + registry_path = os.path.join(root, "window.db") + registry = SQLiteWindowedCaptureRegistry( + registry_path, + max_live_refs=2, + poll_s=0.001, + ) + digest = capture_contract_digest(_capture()) + registry.initialize_run( + run_id="run", + contract_digest=digest, + source_sample_ids=("source-0",), + expected_consumers=("a",), + ) + registry.close() + metadata_path = os.path.join(root, "consumer.db") + metadata = SQLiteMetadataStore(metadata_path) + ref = _RefSource(_OwnerStore()).produce_refs( + [_prompts(1)[0]], capture=_capture() + )[0] + metadata.commit_sample(ref) + metadata.record_train_ack( + [ref.sample_id], global_step=1, optimizer_durable=True + ) + metadata.close() + + with self.assertRaisesRegex(ValueError, "no resume_from checkpoint"): + build_disagg_online_windowed_consumer( + consumer_id="a", + registry_db_path=registry_path, + max_live_refs=2, + contract_digest=digest, + total_samples=1, + feature_store=_ReaderStore(), + draft_model=object(), + optimizer_factory=mock.Mock(), + run_id="run", + output_dir=os.path.join(root, "output"), + metadata_db_path=metadata_path, + strategy="dflash", + resume=True, + initialization_timeout_s=1.0, + ) + + def test_consumer_resume_restores_durable_cursor_and_checkpoint(self): + with tempfile.TemporaryDirectory() as root: + registry_path = os.path.join(root, "window.db") + registry = SQLiteWindowedCaptureRegistry( + registry_path, + max_live_refs=2, + poll_s=0.001, + ) + digest = capture_contract_digest(_capture()) + registry.initialize_run( + run_id="run", + contract_digest=digest, + source_sample_ids=("source-0", "source-1"), + expected_consumers=("a",), + ) + registry.close() + + metadata_path = os.path.join(root, "consumer.db") + metadata = SQLiteMetadataStore(metadata_path) + ref = _RefSource(_OwnerStore()).produce_refs( + [_prompts(1)[0]], capture=_capture() + )[0] + metadata.commit_sample(ref) + metadata.record_train_ack( + [ref.sample_id], global_step=1, optimizer_durable=True + ) + metadata.close() + + checkpoint = os.path.join(root, "checkpoint") + os.makedirs(checkpoint) + torch.save({"global_step": 1}, os.path.join(checkpoint, STATE_FILE)) + fake_trainer = mock.Mock() + fake_trainer.loader = mock.Mock() + + with mock.patch( + "specforge.launch._assemble_trainer", + return_value=fake_trainer, + ) as assemble: + runtime = build_disagg_online_windowed_consumer( + consumer_id="a", + registry_db_path=registry_path, + max_live_refs=2, + contract_digest=digest, + total_samples=2, + feature_store=_ReaderStore(), + draft_model=object(), + optimizer_factory=mock.Mock(), + run_id="run", + output_dir=os.path.join(root, "output"), + metadata_db_path=metadata_path, + strategy="dflash", + resume=True, + resume_from=checkpoint, + initialization_timeout_s=1.0, + heartbeat_interval_s=0.1, + ) + try: + self.assertEqual( + runtime.registry.snapshot()["consumers"]["a"]["cursor"], 1 + ) + self.assertEqual(assemble.call_args.kwargs["resume_from"], checkpoint) + finally: + runtime.control.fail("test cleanup") + runtime.close() + + def test_producer_recovery_reclaims_committing_generation_and_replays(self): + with tempfile.TemporaryDirectory() as root: + path = os.path.join(root, "window.db") + prompt = {"task_id": "source-0", "payload": {"input_ids": [1, 2]}} + store = _OwnerStore() + source = _RefSource(store) + first = build_disagg_online_windowed_producer( + prompts=[prompt], + feature_store=store, + feature_source=source, + run_id="run", + consumer_ids=("a",), + registry_db_path=path, + max_live_refs=2, + target_hidden_size=8, + target_model_version="fixture", + tokenizer_version="fixture", + strategy="dflash", + target_repr="hidden_state", + ) + first.registry.register_consumer("a") + first.registry.request_acquire("a", 0) + request = first.registry.claim_batch(1)[0] + ref = source.produce_refs([_prompts(1)[0]], capture=_capture())[0] + ref = dataclasses.replace( + ref, + metadata={ + **ref.metadata, + "window_generation": request.generation, + }, + ) + first.registry.mark_committing(request, ref) + first.close() + + resumed = build_disagg_online_windowed_producer( + prompts=[prompt], + feature_store=store, + feature_source=source, + run_id="run", + consumer_ids=("a",), + registry_db_path=path, + max_live_refs=2, + target_hidden_size=8, + target_model_version="fixture", + tokenizer_version="fixture", + strategy="dflash", + target_repr="hidden_state", + recover=True, + consumer_registration_timeout_s=1.0, + consumer_heartbeat_timeout_s=30.0, + registry_poll_s=0.001, + ) + resumed.registry.resume_consumer("a", durable_cursor=0) + queue = WindowedCaptureQueue(resumed.registry, "a", idle_timeout_s=2.0) + errors = [] + thread = threading.Thread( + target=lambda: TestWindowedCaptureService._drive_or_record( + resumed.service, errors + ) + ) + thread.start() + refs = queue.get(1) + queue.ack(refs) + self.assertEqual(queue.get(1), []) + thread.join(timeout=5.0) + try: + self.assertFalse(thread.is_alive()) + self.assertEqual(errors, []) + self.assertEqual([item[1] for item in store.reclaimed], [1, 2]) + self.assertEqual(refs[0].metadata["generation"], 2) + self.assertEqual(resumed.registry.snapshot()["status"], "completed") + finally: + resumed.close() + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_runtime/test_windowed_fanout_assembly.py b/tests/test_runtime/test_windowed_fanout_assembly.py new file mode 100644 index 000000000..0d4e3ba57 --- /dev/null +++ b/tests/test_runtime/test_windowed_fanout_assembly.py @@ -0,0 +1,381 @@ +# coding=utf-8 +"""Canonical ``specforge train`` assembly for windowed fanout roles.""" + +from __future__ import annotations + +import os +import tempfile +import types +import unittest +from unittest import mock + +from specforge.config import Config +from specforge.training.disaggregated import ( + _build_windowed_online, + _stabilize_windowed_prompts, +) + + +def _config(root: str, *, role: str, resume_from: str | None = None) -> Config: + return Config.model_validate( + { + "run_id": "fanout-run", + "output_dir": os.path.join(root, "output"), + "model": { + "target_model_path": "target", + "draft_model_config": "draft", + "target_backend": "sglang", + }, + "data": { + "train_data_path": "/prompts.jsonl", + "max_prompts": 4, + }, + "training": { + "strategy": "dflash", + "role": role, + "batch_size": 2, + "accumulation_steps": 1, + "num_epochs": 1, + "max_steps": 2, + }, + "deployment": { + "mode": "disaggregated", + "trainer": {"nnodes": 1, "nproc_per_node": 1}, + "disaggregated": { + "control_dir": root, + "backend": "mooncake", + "server_urls": ["http://capture:30000"], + "windowed_fanout": { + "max_live_bytes": 1 << 30, + "max_outstanding_per_consumer": 4, + "consumers": [ + { + "consumer_id": "block-4", + "seed": 42, + "loss_type": "dflash", + "loss_decay_gamma": 7.0, + "dpace_alpha": 0.5, + "draft_block_size": 4, + "num_anchors": 128, + "learning_rate": 0.0006, + "warmup_ratio": 0.04, + "cuda_visible_device": "1", + "resume_from": resume_from, + } + ], + }, + }, + }, + } + ) + + +def _algorithm(): + streaming = types.SimpleNamespace( + target_representation=None, + layout=types.SimpleNamespace( + aux_feature="aux", + last_hidden_feature="last_hidden", + passthrough=(), + attention_mask_feature=None, + ), + create_input_adapter=mock.Mock(return_value=None), + ) + providers = types.SimpleNamespace( + server_streaming_for=mock.Mock(return_value=streaming), + model=types.SimpleNamespace(draft_config=object()), + ) + return types.SimpleNamespace(name="dflash", providers=providers), streaming + + +def _environment(root: str, **extra: str) -> dict[str, str]: + return { + "DISAGG_WINDOW_REGISTRY": os.path.join(root, "window.db"), + "DISAGG_DB": os.path.join(root, "consumer.db"), + **extra, + } + + +class WindowedFanoutAssemblyTest(unittest.TestCase): + def test_producer_uses_canonical_prompts_and_cleans_owner_on_failure(self): + with tempfile.TemporaryDirectory() as root: + cfg = _config(root, role="producer") + algorithm, _ = _algorithm() + runtime = mock.Mock() + failure = RuntimeError("capture failed") + runtime.drive.side_effect = failure + store = mock.Mock() + prompts = [ + {"payload": {"input_ids": [index], "loss_mask": [1]}} + for index in range(4) + ] + prepare_prompts = mock.Mock(return_value=prompts) + + with ( + mock.patch.dict(os.environ, _environment(root), clear=False), + mock.patch( + "specforge.training.disaggregated._producer_capture_metadata", + return_value=([1, 2], 16, 32, 24), + ), + mock.patch( + "specforge.launch.build_disagg_windowed_capture_contract", + return_value=(object(), "digest"), + ), + mock.patch( + "specforge.training.assembly._load_input_tools", + return_value=object(), + ), + mock.patch( + "specforge.training.model_loading.resolve_draft_config", + return_value=object(), + ), + mock.patch( + "specforge.inference.adapters.server_capture." + "SGLangServerCaptureAdapter" + ) as adapter, + mock.patch( + "specforge.launch.build_disagg_online_windowed_producer", + autospec=True, + return_value=runtime, + ) as build_producer, + mock.patch( + "specforge.training.disaggregated._mooncake_store", + return_value=store, + ) as build_store, + mock.patch( + "specforge.runtime.data_plane.feature_store." + "drain_feature_store_removals" + ) as drain, + ): + run = _build_windowed_online( + cfg, + algorithm=algorithm, + build_model_bundle=mock.Mock(), + prepare_prompts=prepare_prompts, + optimizer_factory=mock.Mock(), + logger=mock.Mock(), + ) + with self.assertRaises(RuntimeError) as raised: + run.run() + + self.assertIs(raised.exception, failure) + prepare_prompts.assert_called_once() + adapter.assert_called_once() + build_store.assert_called_once_with(cfg, lifetime_owner=True) + prepared = build_producer.call_args.kwargs["prompts"] + self.assertEqual( + [prompt["payload"] for prompt in prepared], + [prompt["payload"] for prompt in prompts], + ) + self.assertEqual( + [prompt["task_id"] for prompt in prepared], + [prompt["task_id"] for prompt in _stabilize_windowed_prompts(prompts)], + ) + self.assertTrue(all("task_id" not in prompt for prompt in prompts)) + self.assertEqual( + build_producer.call_args.kwargs["consumer_ids"], ("block-4",) + ) + self.assertEqual(build_producer.call_args.kwargs["modality"], "text") + runtime.close.assert_called_once_with() + store.abort_all.assert_called_once_with( + reason="windowed-attempt-failed", force=True + ) + drain.assert_called_once_with(store) + + def test_producer_rejects_incomplete_fixed_prompt_inventory(self): + with tempfile.TemporaryDirectory() as root: + cfg = _config(root, role="producer") + algorithm, _ = _algorithm() + with ( + mock.patch.dict(os.environ, _environment(root), clear=False), + mock.patch( + "specforge.training.disaggregated._producer_capture_metadata", + return_value=([], 16, 32, 24), + ), + mock.patch( + "specforge.launch.build_disagg_windowed_capture_contract", + return_value=(object(), "digest"), + ), + mock.patch( + "specforge.training.assembly._load_input_tools", + return_value=object(), + ), + mock.patch( + "specforge.training.model_loading.resolve_draft_config", + return_value=object(), + ), + mock.patch( + "specforge.launch.build_disagg_online_windowed_producer" + ) as build_producer, + ): + with self.assertRaisesRegex(ValueError, "prepared prompt count"): + _build_windowed_online( + cfg, + algorithm=algorithm, + build_model_bundle=mock.Mock(), + prepare_prompts=mock.Mock(return_value=[{"task_id": "0"}]), + optimizer_factory=mock.Mock(), + logger=mock.Mock(), + ) + build_producer.assert_not_called() + + def test_prompt_ids_are_deterministic_and_conflicts_are_rejected(self): + prompts = [ + {"payload": {"input_ids": [1, 2], "loss_mask": [0, 1]}}, + {"payload": {"input_ids": [1, 2], "loss_mask": [0, 1]}}, + ] + first = _stabilize_windowed_prompts(prompts) + second = _stabilize_windowed_prompts(prompts) + + self.assertEqual( + [prompt["task_id"] for prompt in first], + [prompt["task_id"] for prompt in second], + ) + self.assertNotEqual(first[0]["task_id"], first[1]["task_id"]) + with self.assertRaisesRegex(ValueError, "duplicated"): + _stabilize_windowed_prompts([{"task_id": "same"}, {"task_id": "same"}]) + with self.assertRaisesRegex(ValueError, "must provide one explicitly"): + _stabilize_windowed_prompts([{"payload": {"opaque": object()}}]) + + def test_consumer_registers_before_model_build_and_forwards_resume(self): + with tempfile.TemporaryDirectory() as root: + checkpoint = os.path.join(root, "checkpoint-1") + cfg = _config(root, role="consumer", resume_from=checkpoint) + algorithm, _ = _algorithm() + events: list[str] = [] + registry = mock.Mock() + registry.wait_initialized.return_value = { + "run_id": cfg.run_id, + "contract_digest": "digest", + "total_samples": cfg.data.max_prompts, + } + control = mock.Mock() + runtime = mock.Mock() + runtime.run.return_value = 2 + bundle = types.SimpleNamespace(model=object(), strategy_kwargs={"x": 1}) + + def register(*_args, **_kwargs): + events.append("register") + return control + + def build_model(_cfg): + events.append("model") + return bundle + + store = mock.Mock(lifetime_owner=False) + optimizer = mock.Mock(return_value=object()) + with ( + mock.patch.dict( + os.environ, + _environment(root, SPECFORGE_FANOUT_CONSUMER_ID="block-4"), + clear=False, + ), + mock.patch( + "specforge.training.disaggregated._producer_capture_metadata", + return_value=([], 16, 32, 24), + ), + mock.patch( + "specforge.launch.build_disagg_windowed_capture_contract", + return_value=(object(), "digest"), + ), + mock.patch( + "specforge.runtime.data_plane.windowed_capture." + "SQLiteWindowedCaptureRegistry", + return_value=registry, + ), + mock.patch( + "specforge.runtime.data_plane.windowed_capture_runtime." + "start_windowed_consumer_control", + side_effect=register, + ), + mock.patch( + "specforge.training.disaggregated._mooncake_store", + return_value=store, + ) as build_store, + mock.patch( + "specforge.launch.build_disagg_online_windowed_consumer", + return_value=runtime, + ) as build_consumer, + ): + run = _build_windowed_online( + cfg, + algorithm=algorithm, + build_model_bundle=build_model, + prepare_prompts=mock.Mock(), + optimizer_factory=optimizer, + logger=mock.Mock(), + ) + self.assertEqual(run.run(), 2) + + self.assertEqual(events, ["register", "model"]) + build_store.assert_called_once_with(cfg, lifetime_owner=False) + kwargs = build_consumer.call_args.kwargs + self.assertEqual(kwargs["consumer_id"], "block-4") + self.assertEqual( + kwargs["metadata_db_path"], _environment(root)["DISAGG_DB"] + ) + self.assertTrue(kwargs["resume"]) + self.assertEqual(kwargs["resume_from"], checkpoint) + self.assertEqual(kwargs["strategy_kwargs"], {"x": 1}) + self.assertIs(kwargs["consumer_control"], control) + runtime.close.assert_called_once_with() + + def test_model_build_failure_is_reported_and_registry_is_closed(self): + with tempfile.TemporaryDirectory() as root: + cfg = _config(root, role="consumer") + algorithm, _ = _algorithm() + registry = mock.Mock() + registry.wait_initialized.return_value = { + "run_id": cfg.run_id, + "contract_digest": "digest", + "total_samples": cfg.data.max_prompts, + } + control = mock.Mock() + failure = RuntimeError("model build failed") + with ( + mock.patch.dict( + os.environ, + _environment(root, SPECFORGE_FANOUT_CONSUMER_ID="block-4"), + clear=False, + ), + mock.patch( + "specforge.training.disaggregated._producer_capture_metadata", + return_value=([], 16, 32, 24), + ), + mock.patch( + "specforge.launch.build_disagg_windowed_capture_contract", + return_value=(object(), "digest"), + ), + mock.patch( + "specforge.runtime.data_plane.windowed_capture." + "SQLiteWindowedCaptureRegistry", + return_value=registry, + ), + mock.patch( + "specforge.runtime.data_plane.windowed_capture_runtime." + "start_windowed_consumer_control", + return_value=control, + ), + mock.patch( + "specforge.launch.build_disagg_online_windowed_consumer" + ) as build_consumer, + ): + with self.assertRaises(RuntimeError) as raised: + _build_windowed_online( + cfg, + algorithm=algorithm, + build_model_bundle=mock.Mock(side_effect=failure), + prepare_prompts=mock.Mock(), + optimizer_factory=mock.Mock(), + logger=mock.Mock(), + ) + + self.assertIs(raised.exception, failure) + control.fail.assert_called_once_with(failure) + control.close.assert_called_once_with() + registry.close.assert_called_once_with() + build_consumer.assert_not_called() + + +if __name__ == "__main__": + unittest.main()