From 5f4ddab376628e2d330ff2fb4a782bae936c76ef Mon Sep 17 00:00:00 2001 From: Laura Dang Date: Thu, 10 Sep 2026 15:16:06 -0400 Subject: [PATCH 01/12] feat(token-capture): add Megatron worker capture backend Add the Megatron worker capture adapter and consolidate backend lifecycle handling. Preserve deferred-call custody and multi-turn lineage, keep capture metadata nested, strip compact prompt token echoes, and cover the shared backend lifecycle with parameterized tests. Co-authored-by: Claude Fable 5.1 Signed-off-by: Laura Dang --- nemo_gym/token_id_capture/config.py | 7 +- nemo_gym/token_id_capture/external_capture.py | 313 ++++++++++++++++++ nemo_gym/token_id_capture/staging/capture.py | 10 +- responses_api_models/vllm_model/app.py | 198 +---------- .../vllm_model/tests/test_app.py | 141 +++++++- .../test_external_capture_handlers.py | 248 ++++++++++++++ .../test_token_capture_staging_worker.py | 6 + tests/unit_tests/test_token_id_capture.py | 11 + 8 files changed, 744 insertions(+), 190 deletions(-) create mode 100644 nemo_gym/token_id_capture/external_capture.py create mode 100644 tests/unit_tests/test_external_capture_handlers.py diff --git a/nemo_gym/token_id_capture/config.py b/nemo_gym/token_id_capture/config.py index 7ce486cfb4..7ffe0837d7 100644 --- a/nemo_gym/token_id_capture/config.py +++ b/nemo_gym/token_id_capture/config.py @@ -71,7 +71,7 @@ from collections.abc import Mapping from importlib import import_module from pathlib import Path -from typing import Any +from typing import Any, Literal from pydantic import BaseModel, ConfigDict, Field, model_validator @@ -86,6 +86,7 @@ logger = logging.getLogger(__name__) TOKEN_ID_CAPTURE_BLOCK = "token_id_capture" +ExternalStagingBackend = Literal["vllm_worker", "megatron_worker"] class TokenIdCaptureSettings(BaseModel): @@ -133,6 +134,8 @@ class TokenIdCaptureSettings(BaseModel): # It also makes committed parents visible to every serving worker. # No additional in-memory coordinator is used. external_staging: bool = False + # Both backends stage a canonical delta before returning coordinates. + external_staging_backend: ExternalStagingBackend = "vllm_worker" # Name of the environment variable containing the manifest-route bearer token. # The serving process reads the token without adding it to serialized configuration. control_auth_token_env: str = Field( @@ -170,6 +173,8 @@ def _validate(self) -> "TokenIdCaptureConfig": "token_id_capture.external_staging requires rebuild_response=false because the " "framework owns staged-record finalization" ) + if block.external_staging_backend != "vllm_worker" and not block.external_staging: + raise ValueError("token_id_capture.external_staging_backend requires external_staging=true") if not block.enabled: # Keep inactive settings for templated configurations. # A run may toggle only ``enabled``. diff --git a/nemo_gym/token_id_capture/external_capture.py b/nemo_gym/token_id_capture/external_capture.py new file mode 100644 index 0000000000..fcca36754f --- /dev/null +++ b/nemo_gym/token_id_capture/external_capture.py @@ -0,0 +1,313 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Backend strategies for framework-owned token capture.""" + +from __future__ import annotations + +import logging +from abc import ABC, abstractmethod +from typing import Any, Protocol + +from nemo_gym.token_id_capture.config import ExternalStagingBackend +from nemo_gym.token_id_capture.fingerprint import FINGERPRINT_VERSION, assistant_fingerprint +from nemo_gym.token_id_capture.protocols import CaptureLedger +from nemo_gym.token_id_capture.records import ( + TOKEN_FIELDS, + response_to_output_items, + strip_token_fields, +) +from nemo_gym.token_id_capture.sink import ( + NG_CAPTURE_FIELD, + NG_COMMIT_COORDS_FIELD, + CaptureContext, + current_capture_context, + mark_external_staging_committed, +) +from nemo_gym.token_id_capture.staging.records import ( + INVALID_COMMIT_COORDS_REASON, + WORKER_CAPTURE_FAILED_REASON, + WORKER_MISSING_COMMIT_COORDS_REASON, + CallRecord, + CaptureAdmission, + CaptureLedgerCommit, + CommitCoords, +) + + +LOGGER = logging.getLogger(__name__) + +# Megatron's ``return_tokenized_data`` echoes the exact prompt token form it +# uses for lossless multi-turn prefix stitching alongside ``TOKEN_FIELDS``. +_MEGATRON_TRANSPORT_FIELDS = ("compact_prompt_token_ids",) + + +class ExternalCaptureHandler(Protocol): + """Prepare and finalize one backend-specific external capture call.""" + + def prepare_request(self, request_payload: dict[str, Any]) -> dict[str, Any]: + """Attach capture instructions to an admitted engine request.""" + ... + + async def finalize_response(self, response_payload: dict[str, Any]) -> None: + """Commit lineage and remove capture-only response fields.""" + ... + + +def _strip_capture_transport_fields(payload: dict[str, Any]) -> None: + """Keep token IDs, logprobs, routes, and coordinates off the agent hop.""" + payload.pop(NG_COMMIT_COORDS_FIELD, None) + payload.pop("prompt_token_ids", None) + for choice in payload.get("choices") or []: + if not isinstance(choice, dict): + continue + choice.pop("logprobs", None) + choice.pop("token_ids", None) + message = choice.get("message") + if isinstance(message, dict): + for field_name in (*TOKEN_FIELDS, *_MEGATRON_TRANSPORT_FIELDS): + message.pop(field_name, None) + + +class _BaseExternalCaptureHandler(ABC): + """Own the lifecycle shared by external capture backends.""" + + _INVALID_CAPTURE_REASON: str + _CAPTURE_ERROR_MESSAGE: str + _POISON_ERROR_MESSAGE: str + + def prepare_request(self, request_payload: dict[str, Any]) -> dict[str, Any]: + """Attach capture instructions to an engine-bound request. + + An unadmitted call (``UNRESOLVED`` — already poisoned in the ledger) + is forwarded as plain traffic: the backend captures nothing and the + completion still serves the agent. + """ + context = current_capture_context() + if context is None or not context.external_staging: + return request_payload + admission = context.capture_admission + if admission is None: + return request_payload + return self._prepare_admitted_request(request_payload, admission) + + @abstractmethod + def _prepare_admitted_request( + self, + request_payload: dict[str, Any], + admission: CaptureAdmission, + ) -> dict[str, Any]: + """Attach backend-specific fields after shared admission checks.""" + + async def finalize_response(self, response_payload: dict[str, Any]) -> None: + context = current_capture_context() + if context is None or not context.external_staging or context.lineage_store is None: + return + ledger = context.lineage_store + if not isinstance(ledger, CaptureLedger): + raise ValueError("external staging requires a CaptureLedger on the capture context") + try: + admission = context.capture_admission + if admission is None: + # UNRESOLVED — the ledger already carries this call's poison row. + return + try: + await self._finalize_admitted_response( + response_payload, + context=context, + ledger=ledger, + admission=admission, + ) + except Exception: + # Backend/framework payloads are an external integrity boundary. + # Poison capture without turning a valid model completion into a + # harness failure. + LOGGER.exception( + self._CAPTURE_ERROR_MESSAGE, + context.rollout_id, + context.model_call_id, + ) + try: + await ledger.record_failure( + context.rollout_id, + context.model_call_id, + self._INVALID_CAPTURE_REASON, + ) + except Exception: + LOGGER.exception( + self._POISON_ERROR_MESSAGE, + context.rollout_id, + context.model_call_id, + ) + finally: + _strip_capture_transport_fields(response_payload) + + @abstractmethod + async def _finalize_admitted_response( + self, + response_payload: dict[str, Any], + *, + context: CaptureContext, + ledger: CaptureLedger, + admission: CaptureAdmission, + ) -> None: + """Publish backend-specific custody for an admitted response.""" + + +class VLLMWorkerCaptureHandler(_BaseExternalCaptureHandler): + """Commit lineage after a vLLM worker durably stages the token delta.""" + + _INVALID_CAPTURE_REASON = INVALID_COMMIT_COORDS_REASON + _CAPTURE_ERROR_MESSAGE = "Worker capture acknowledgement failed for rollout %s call %s" + _POISON_ERROR_MESSAGE = "Could not poison rollout %s call %s after a failed acknowledgement" + + def _prepare_admitted_request( + self, + request_payload: dict[str, Any], + admission: CaptureAdmission, + ) -> dict[str, Any]: + request_payload[NG_CAPTURE_FIELD] = admission.model_dump(mode="json") + request_payload.update( + logprobs=True, + top_logprobs=0, + return_tokens_as_token_ids=True, + ) + if admission.mode == "token_in": + request_payload["required_prefix_token_ids"] = list(admission.required_prefix_token_ids) + return request_payload + + async def _finalize_admitted_response( + self, + response_payload: dict[str, Any], + *, + context: CaptureContext, + ledger: CaptureLedger, + admission: CaptureAdmission, + ) -> None: + """Publish the worker's coordinates as a ledger row. + + The ordering invariant the external sink requires — a call must not + become a lineage parent until its staged record is durable — holds + structurally: the worker stages before acknowledging, so the ledger + row (which is what makes the call resolvable) is written only after + the coordinates arrive. The shared lifecycle strips custody fields + after this method returns. + """ + coords_payload = response_payload.pop(NG_COMMIT_COORDS_FIELD, None) + if coords_payload is None: + await ledger.record_failure( + context.rollout_id, + context.model_call_id, + WORKER_MISSING_COMMIT_COORDS_REASON, + ) + return + coords = CommitCoords.model_validate(coords_payload) + if coords.rollout_id != context.rollout_id or coords.model_call_id != context.model_call_id: + raise ValueError( + f"coordinates for {coords.rollout_id}/{coords.model_call_id} do not match the " + f"active capture context {context.rollout_id}/{context.model_call_id}" + ) + if coords.disposition == "capture_failed": + await ledger.record_failure( + context.rollout_id, + context.model_call_id, + WORKER_CAPTURE_FAILED_REASON, + ) + return + if coords.parent_call_id != admission.parent_call_id or coords.prev_len != admission.prev_len: + raise ValueError(f"coordinates for {coords.model_call_id} diverge from admission") + # The served envelope id is the terminal-attribution join key: the + # agent proves which response it kept by possessing it. Observe the + # payload's own id; never mint one. A served completion without an + # id is a stamping bug and fails closed (poisons the call below). + response_id = str(response_payload.get("id") or "") + if not response_id: + raise ValueError(f"served response for {coords.model_call_id} carries no envelope id") + child_staging_chain = list(context.parent_staging_chain) + [str(coords.staging_key)] + response_items, _ = strip_token_fields(response_to_output_items(response_payload)) + # Content-witness keys, hashed while the response is still + # server-side: this call's own output, and request + output (the + # cumulative reading). Unfingerprintable content abstains (None) + # rather than poisoning a valid completion. + try: + output_fingerprint = assistant_fingerprint(list(response_items)) or None + continuation_fingerprint = ( + assistant_fingerprint(list(context.request_items or []) + list(response_items)) or None + ) + except (TypeError, ValueError): + output_fingerprint = None + continuation_fingerprint = None + # The lineage row omits token arrays because the worker stores token + # deltas separately. ``CallRecord`` re-validates its wire invariants. + record = CallRecord( + model_call_id=coords.model_call_id, + parent_call_id=coords.parent_call_id, + prev_len=coords.prev_len, + delta_len=coords.delta_len, + cum_len=coords.cum_len, + weight_version=coords.weight_version, + digest=coords.digest, + extras_digest=coords.extras_digest, + staging_key=coords.staging_key, + mode=admission.mode, + admitted_at=context.admitted_at, + chain_hash=coords.chain_hash, + cumulative_hash=coords.cumulative_hash, + response_id=response_id, + output_fingerprint=output_fingerprint, + continuation_fingerprint=continuation_fingerprint, + fingerprint_version=FINGERPRINT_VERSION, + ) + commit = CaptureLedgerCommit( + rollout_id=context.rollout_id, + record=record, + staging_chain=tuple(child_staging_chain), + request_items=list(context.request_items or []), + response_items=response_items, + ) + await ledger.record(commit) + mark_external_staging_committed( + rollout_id=coords.rollout_id, + model_call_id=coords.model_call_id, + ) + + +class MegatronWorkerCaptureHandler(VLLMWorkerCaptureHandler): + """Commit lineage after an MInf worker durably stages a canonical delta.""" + + _INVALID_CAPTURE_REASON = "invalid_megatron_commit_coordinates" + _CAPTURE_ERROR_MESSAGE = "Megatron capture acknowledgement failed for rollout %s call %s" + _POISON_ERROR_MESSAGE = "Could not poison rollout %s call %s after a failed MInf acknowledgement" + + def _prepare_admitted_request( + self, + request_payload: dict[str, Any], + admission: CaptureAdmission, + ) -> dict[str, Any]: + choice_count = request_payload.get("n") + if choice_count is not None and choice_count != 1: + raise ValueError("Megatron token capture requires n=1") + request_metadata = request_payload.get("request_metadata") + if request_metadata is None: + request_metadata = {} + request_payload["request_metadata"] = request_metadata + if not isinstance(request_metadata, dict): + raise ValueError("Megatron request_metadata must be an object") + request_metadata[NG_CAPTURE_FIELD] = admission.model_dump(mode="json") + request_payload.update( + logprobs=True, + top_logprobs=0, + return_tokenized_data=True, + ) + if admission.mode == "token_in": + request_payload["required_prefix_token_ids"] = list(admission.required_prefix_token_ids) + return request_payload + + +def make_external_capture_handler(backend: ExternalStagingBackend) -> ExternalCaptureHandler: + """Create the external capture strategy selected by typed configuration.""" + if backend == "vllm_worker": + return VLLMWorkerCaptureHandler() + if backend == "megatron_worker": + return MegatronWorkerCaptureHandler() + raise ValueError(f"Unsupported external staging backend: {backend}") diff --git a/nemo_gym/token_id_capture/staging/capture.py b/nemo_gym/token_id_capture/staging/capture.py index 9cacda7797..d821eb78b9 100644 --- a/nemo_gym/token_id_capture/staging/capture.py +++ b/nemo_gym/token_id_capture/staging/capture.py @@ -93,6 +93,7 @@ def begin_call( *, prefix_token_ids: list[int] | None = None, stream: bool = False, + weight_version: int | None = None, ) -> ActiveCall: """Admit a typed gate contract and stamp its generation weight version. @@ -103,6 +104,10 @@ def begin_call( When the admission carries the prefix inline the argument may be omitted; if given it must match. A text root accepts no prefix. Violations are caller bugs and raise ``CaptureError``; they never poison the call. + + ``weight_version`` is supplied when engine-finished metadata owns the + authoritative per-call epoch. Other worker capture paths leave it unset + and read the serving worker's current version provider instead. """ if not isinstance(admission, CaptureAdmission): raise TypeError("admission must be a CaptureAdmission") @@ -112,9 +117,10 @@ def begin_call( "token capture does not support streaming responses" ) resolved_prefix = self._resolve_prefix(admission, prefix_token_ids) - weight_version = self._weight_version_fn() + if weight_version is None: + weight_version = self._weight_version_fn() if type(weight_version) is not int or weight_version < 0: - raise CaptureError(f"weight_version_fn must return a non-negative int, got {weight_version!r}") + raise CaptureError(f"weight_version must be a non-negative int, got {weight_version!r}") return ActiveCall(admission=admission, weight_version=weight_version, prefix_token_ids=resolved_prefix) @staticmethod diff --git a/responses_api_models/vllm_model/app.py b/responses_api_models/vllm_model/app.py index 749272d88a..bdc417c15c 100644 --- a/responses_api_models/vllm_model/app.py +++ b/responses_api_models/vllm_model/app.py @@ -51,26 +51,12 @@ ) from nemo_gym.server_utils import SESSION_ID_KEY, is_nemo_gym_fastapi_entrypoint from nemo_gym.token_id_capture import ( - NG_CAPTURE_FIELD, - NG_COMMIT_COORDS_FIELD, current_capture_context, - mark_external_staging_committed, ) from nemo_gym.token_id_capture.config import token_id_capture_config -from nemo_gym.token_id_capture.fingerprint import FINGERPRINT_VERSION, assistant_fingerprint -from nemo_gym.token_id_capture.protocols import CaptureLedger -from nemo_gym.token_id_capture.records import ( - TOKEN_FIELDS, - response_to_output_items, - strip_token_fields, -) -from nemo_gym.token_id_capture.staging.records import ( - INVALID_COMMIT_COORDS_REASON, - WORKER_CAPTURE_FAILED_REASON, - WORKER_MISSING_COMMIT_COORDS_REASON, - CallRecord, - CaptureLedgerCommit, - CommitCoords, +from nemo_gym.token_id_capture.external_capture import ( + ExternalCaptureHandler, + make_external_capture_handler, ) @@ -301,7 +287,7 @@ class VLLMModel(SimpleResponsesAPIModel): "mm_processor_kwargs", "required_prefix_token_ids", ) - _external_capture_enabled: bool = PrivateAttr(default=False) + _external_capture_handler: ExternalCaptureHandler | None = PrivateAttr(default=None) def setup_exception_middleware(self, app) -> None: @app.middleware("http") @@ -363,10 +349,8 @@ def _post_init(self) -> None: global_config = getattr(self.server_client, "global_config_dict", None) capture_config = token_id_capture_config(global_config) if global_config is not None else None - self._external_capture_enabled = bool( - capture_config is not None and capture_config.token_id_capture.external_staging - ) - if self._external_capture_enabled: + self._external_capture_handler = None + if capture_config is not None and capture_config.token_id_capture.external_staging: if self.config.use_completions_api: raise ValueError("token_id_capture.external_staging does not support use_completions_api=true") if self.config.is_responses_native: @@ -376,6 +360,9 @@ def _post_init(self) -> None: "token_id_capture.external_staging requires return_token_id_information=false; " "worker custody replaces the token echo" ) + self._external_capture_handler = make_external_capture_handler( + capture_config.token_id_capture.external_staging_backend + ) self._chat_template_tokenizer = None if self.config.use_completions_api and self.config.render_chat_template: @@ -731,8 +718,8 @@ def _preprocess_chat_completion_create_params(self, request: Request, body_dict: self._apply_sampling_overrides(body_dict) self._validate_single_choice_token_request(body_dict) - if self._external_capture_enabled: - body_dict = self._apply_external_capture(body_dict) + if self._external_capture_handler is not None: + body_dict = self._external_capture_handler.prepare_request(body_dict) else: body_dict = self._apply_prefix_supply(body_dict) @@ -743,29 +730,6 @@ def _preserve_envelope_id(self) -> bool: context = current_capture_context() return context is not None and context.external_staging - def _apply_external_capture(self, body_dict: Dict[str, Any]) -> Dict[str, Any]: - """Add worker capture metadata to an admitted chat request. - - If parent resolution did not admit the call, forward the request without capture metadata. - The lineage store has already recorded that capture failure. - The worker returns the completion without staging token data. - """ - context = current_capture_context() - if context is None or not context.external_staging: - return body_dict - admission = context.capture_admission - if admission is None: - return body_dict - body_dict[NG_CAPTURE_FIELD] = admission.model_dump(mode="json") - body_dict.update( - logprobs=True, - top_logprobs=0, - return_tokens_as_token_ids=True, - ) - if admission.mode == "token_in": - body_dict["required_prefix_token_ids"] = list(admission.required_prefix_token_ids) - return body_dict - # Protect the ``[supplied, eligible, total]`` diagnostic counts. # Eligible calls have a resolved parent. _prefix_supply_counts: List[int] = PrivateAttr(default_factory=lambda: [0, 0, 0]) @@ -1010,8 +974,8 @@ async def chat_completions( f"NeMo Gym server `{self.config.name}` config has explicitly been set to not use a reasoning parser i.e. `uses_reasoning_parser: false`. Please do not use a reasoning parser in your vLLM endpoint, or fix the `{self.config.name}` server config!" ) - if self._external_capture_enabled: - await self._finalize_external_capture(chat_completion_dict) + if self._external_capture_handler is not None: + await self._external_capture_handler.finalize_response(chat_completion_dict) if self.config.return_token_id_information: message_dict = choice_dict["message"] @@ -1079,142 +1043,6 @@ async def chat_completions( return NeMoGymChatCompletion.model_validate(chat_completion_dict) - async def _finalize_external_capture(self, payload: Dict[str, Any]) -> None: - """Validate and record a response staged by the inference worker. - - The worker returns commit coordinates only after ``StagingSink.stage`` succeeds. - This method validates those coordinates against the active call. - It then records the call in the lineage store. - Finally, it removes token data and commit coordinates from the served response. - """ - context = current_capture_context() - if context is None or not context.external_staging or context.lineage_store is None: - return - ledger = context.lineage_store - if not isinstance(ledger, CaptureLedger): - raise ValueError("external staging requires a CaptureLedger on the capture context") - coords_payload = payload.pop(NG_COMMIT_COORDS_FIELD, None) - admission = context.capture_admission - if admission is None: - # UNRESOLVED — the ledger already carries this call's poison row. - self._strip_capture_transport_fields(payload) - return - try: - if coords_payload is None: - await ledger.record_failure( - context.rollout_id, - context.model_call_id, - WORKER_MISSING_COMMIT_COORDS_REASON, - ) - return - coords = CommitCoords.model_validate(coords_payload) - if coords.rollout_id != context.rollout_id or coords.model_call_id != context.model_call_id: - raise ValueError( - f"coordinates for {coords.rollout_id}/{coords.model_call_id} do not match the " - f"active capture context {context.rollout_id}/{context.model_call_id}" - ) - if coords.disposition == "capture_failed": - await ledger.record_failure( - context.rollout_id, - context.model_call_id, - WORKER_CAPTURE_FAILED_REASON, - ) - return - if coords.parent_call_id != admission.parent_call_id or coords.prev_len != admission.prev_len: - raise ValueError(f"coordinates for {coords.model_call_id} diverge from admission") - # Store the response ID returned to the agent with the corresponding lineage row. - # A missing ID makes that association impossible. - # Treat a missing ID as a capture failure. - response_id = str(payload.get("id") or "") - if not response_id: - raise ValueError(f"served response for {coords.model_call_id} carries no envelope id") - child_staging_chain = list(context.parent_staging_chain) + [str(coords.staging_key)] - response_items, _ = strip_token_fields(response_to_output_items(payload)) - # Compute one fingerprint for the response items. - # Compute another for the request and response items together. - # If either input cannot be fingerprinted, store no fingerprints and continue recording the call. - try: - output_fingerprint = assistant_fingerprint(list(response_items)) or None - continuation_fingerprint = ( - assistant_fingerprint(list(context.request_items or []) + list(response_items)) or None - ) - except (TypeError, ValueError): - output_fingerprint = None - continuation_fingerprint = None - # The lineage row omits token arrays because the worker stores token deltas separately. - # Finalization verifies both hashes against those staged token deltas. - # ``CallRecord`` re-validates the manifest-row invariants (contiguous - # lengths, root/child mode); a ValidationError poisons the call below. - record = CallRecord( - model_call_id=coords.model_call_id, - parent_call_id=coords.parent_call_id, - prev_len=coords.prev_len, - delta_len=coords.delta_len, - cum_len=coords.cum_len, - weight_version=coords.weight_version, - digest=coords.digest, - extras_digest=coords.extras_digest, - staging_key=coords.staging_key, - mode=admission.mode, - chain_hash=coords.chain_hash, - cumulative_hash=coords.cumulative_hash, - response_id=response_id, - admitted_at=context.admitted_at, - output_fingerprint=output_fingerprint, - continuation_fingerprint=continuation_fingerprint, - fingerprint_version=FINGERPRINT_VERSION, - ) - commit = CaptureLedgerCommit( - rollout_id=context.rollout_id, - record=record, - staging_chain=tuple(child_staging_chain), - request_items=list(context.request_items or []), - response_items=response_items, - ) - await ledger.record(commit) - mark_external_staging_committed( - rollout_id=coords.rollout_id, - model_call_id=coords.model_call_id, - ) - except Exception: - # Worker/framework payloads are an external integrity boundary. - # Poison capture without turning a valid model completion into a - # harness failure. - LOG.exception( - "Worker capture acknowledgement failed for rollout %s call %s", - context.rollout_id, - context.model_call_id, - ) - try: - await ledger.record_failure( - context.rollout_id, - context.model_call_id, - INVALID_COMMIT_COORDS_REASON, - ) - except Exception: - LOG.exception( - "Could not poison rollout %s call %s after a failed acknowledgement", - context.rollout_id, - context.model_call_id, - ) - finally: - self._strip_capture_transport_fields(payload) - - @staticmethod - def _strip_capture_transport_fields(payload: Dict[str, Any]) -> None: - """Keep token IDs, logprobs, routes, and coordinates off the agent hop.""" - payload.pop(NG_COMMIT_COORDS_FIELD, None) - payload.pop("prompt_token_ids", None) - for choice in payload.get("choices") or []: - if not isinstance(choice, dict): - continue - choice.pop("logprobs", None) - choice.pop("token_ids", None) - message = choice.get("message") - if isinstance(message, dict): - for field_name in TOKEN_FIELDS: - message.pop(field_name, None) - @staticmethod def _require_token_id_list(value: Any, field_name: str) -> List[Any]: """Check the container without scanning or copying token IDs.""" diff --git a/responses_api_models/vllm_model/tests/test_app.py b/responses_api_models/vllm_model/tests/test_app.py index a7032eda63..eb4458e0ac 100644 --- a/responses_api_models/vllm_model/tests/test_app.py +++ b/responses_api_models/vllm_model/tests/test_app.py @@ -64,6 +64,7 @@ resolve_parent, set_token_sink, ) +from nemo_gym.token_id_capture.staging.records import CaptureAdmission from responses_api_models.vllm_model.app import ( VLLMConverter, VLLMModel, @@ -762,7 +763,13 @@ class FakeUUID: class TestApp: - def _setup_server(self, monkeypatch: MonkeyPatch, *, propagate_context_overflow_errors: bool = False): + def _setup_server( + self, + monkeypatch: MonkeyPatch, + *, + propagate_context_overflow_errors: bool = False, + external_staging_backend: str | None = None, + ): config = VLLMModelConfig( host="0.0.0.0", port=8081, @@ -780,7 +787,20 @@ def _setup_server(self, monkeypatch: MonkeyPatch, *, propagate_context_overflow_ get_global_config_dict_mock.return_value = dict() monkeypatch.setattr(nemo_gym.server_utils, "get_global_config_dict", get_global_config_dict_mock) - return VLLMModel(config=config, server_client=MagicMock(spec=ServerClient, global_config_dict={})) + global_config = {} + if external_staging_backend is not None: + global_config = { + "token_id_capture": { + "enabled": True, + "external_staging": True, + "external_staging_backend": external_staging_backend, + "rebuild_response": False, + } + } + return VLLMModel( + config=config, + server_client=MagicMock(spec=ServerClient, global_config_dict=global_config), + ) async def test_sanity(self, monkeypatch: MonkeyPatch) -> None: assert not self._setup_server(monkeypatch).config.propagate_context_overflow_errors @@ -809,6 +829,123 @@ def test_context_overflow_propagation_flag(self, monkeypatch: MonkeyPatch, propa assert response.status_code == 200 assert '"finish_reason": "length"' in response.text + def test_megatron_capture_handler_prepares_an_admitted_child_request(self, monkeypatch: MonkeyPatch) -> None: + server = self._setup_server(monkeypatch, external_staging_backend="megatron_worker") + context = CaptureContext( + rollout_id="rollout-1", + model_call_id="c2", + token_sink=None, + external_staging=True, + capture_admission=CaptureAdmission( + rollout_id="rollout-1", + model_call_id="c2", + parent_call_id="c1", + prev_len=3, + mode="token_in", + required_prefix_token_ids=[10, 11, 12], + parent_chain_hash="0" * 64, + ), + ) + token = set_token_sink(context) + try: + outbound = server._preprocess_chat_completion_create_params( + MagicMock(), + {"messages": [{"role": "user", "content": "continue"}]}, + ) + finally: + reset_token_sink(token) + + assert outbound["return_tokenized_data"] is True + assert outbound["required_prefix_token_ids"] == [10, 11, 12] + assert outbound["logprobs"] is True + assert outbound["top_logprobs"] == 0 + assert outbound["request_metadata"]["ng_capture"] == context.capture_admission.model_dump(mode="json") + assert "ng_capture" not in outbound + assert "return_tokens_as_token_ids" not in outbound + + async def test_megatron_capture_handler_finalizes_through_chat_completions(self, monkeypatch: MonkeyPatch) -> None: + server = self._setup_server(monkeypatch, external_staging_backend="megatron_worker") + client = MagicMock(spec=NeMoGymAsyncOpenAI) + client.create_chat_completion = AsyncMock( + return_value={ + "id": "minf-17", + "object": "chat.completion", + "created": FIXED_TIME, + "model": "dummy_model", + "ng_commit_coords": { + "schema_version": 2, + "digest_version": 2, + "extras_digest_version": 1, + "rollout_id": "rollout-1", + "model_call_id": "c1", + "parent_call_id": None, + "prev_len": 0, + "delta_len": 3, + "cum_len": 3, + "weight_version": 7, + "disposition": "staged", + "digest": "0" * 64, + "extras_digest": "1" * 64, + "staging_key": "rollout-1/c1", + "chain_hash": "2" * 64, + "cumulative_hash": "3" * 64, + }, + "choices": [ + { + "index": 0, + "finish_reason": "stop", + "message": { + "role": "assistant", + "content": "done", + }, + } + ], + } + ) + server._clients = [client] + request = MagicMock() + request.session = {SESSION_ID_KEY: "session-1"} + request.headers = {} + lineage_store = InMemoryLineageStore() + context = CaptureContext( + rollout_id="rollout-1", + model_call_id="c1", + token_sink=None, + lineage_store=lineage_store, + external_staging=True, + admitted_at=1.5, + request_items=[{"role": "user", "content": "go"}], + capture_admission=CaptureAdmission( + rollout_id="rollout-1", + model_call_id="c1", + mode="text", + ), + ) + token = set_token_sink(context) + try: + response = await server.chat_completions( + request, + NeMoGymChatCompletionCreateParamsNonStreaming( + messages=[{"role": "user", "content": "go"}], + ), + ) + finally: + reset_token_sink(token) + + outbound = client.create_chat_completion.await_args.kwargs + assert outbound["return_tokenized_data"] is True + assert outbound["request_metadata"]["ng_capture"] == context.capture_admission.model_dump(mode="json") + assert "ng_capture" not in outbound + assert context.committed is True + manifest = await lineage_store.manifest("rollout-1") + assert manifest["records"][0]["staging_key"] == "rollout-1/c1" + assert manifest["records"][0]["weight_version"] == 7 + assert manifest["records"][0]["response_id"] == "minf-17" + response_payload = response.model_dump() + assert "prompt_token_ids" not in response_payload + assert "prompt_token_ids" not in response_payload["choices"][0]["message"] + assert "generation_token_ids" not in response_payload["choices"][0]["message"] + def test_session_client_routing_is_stable_across_workers(self, monkeypatch: MonkeyPatch) -> None: workers = [self._setup_server(monkeypatch) for _ in range(2)] for worker in workers: diff --git a/tests/unit_tests/test_external_capture_handlers.py b/tests/unit_tests/test_external_capture_handlers.py new file mode 100644 index 0000000000..030c875a82 --- /dev/null +++ b/tests/unit_tests/test_external_capture_handlers.py @@ -0,0 +1,248 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""External capture strategy lifecycle tests.""" + +from typing import Any + +import pytest + +from nemo_gym.token_id_capture.external_capture import ( + MegatronWorkerCaptureHandler, + VLLMWorkerCaptureHandler, + make_external_capture_handler, +) +from nemo_gym.token_id_capture.lineage import InMemoryLineageStore +from nemo_gym.token_id_capture.sink import CaptureContext, reset_token_sink, set_token_sink +from nemo_gym.token_id_capture.staging.records import CaptureAdmission, CommitCoords + + +def _root_context(store: InMemoryLineageStore) -> CaptureContext: + return CaptureContext( + rollout_id="rollout-1", + model_call_id="c1", + token_sink=None, + lineage_store=store, + external_staging=True, + request_items=[{"role": "user", "content": "go"}], + capture_admission=CaptureAdmission( + rollout_id="rollout-1", + model_call_id="c1", + mode="text", + ), + ) + + +def _transport_payload(**fields: Any) -> dict[str, Any]: + message = { + "role": "assistant", + "content": "done", + "prompt_token_ids": [10, 11], + "generation_token_ids": [12], + "generation_log_probs": [-0.2], + "routed_experts": {"data": "unused"}, + # Megatron ``return_tokenized_data`` echo, absent from vLLM payloads. + "compact_prompt_token_ids": [10, 11], + } + message.update(fields) + return { + "id": "request-1", + "prompt_token_ids": [10, 11], + "choices": [ + { + "token_ids": [12], + "logprobs": {"content": []}, + "message": message, + } + ], + } + + +def _assert_transport_fields_stripped(payload: dict[str, Any]) -> None: + assert "ng_commit_coords" not in payload + assert "prompt_token_ids" not in payload + choice = payload["choices"][0] + assert "token_ids" not in choice + assert "logprobs" not in choice + message = choice["message"] + assert "prompt_token_ids" not in message + assert "generation_token_ids" not in message + assert "generation_log_probs" not in message + assert "routed_experts" not in message + assert "compact_prompt_token_ids" not in message + + +@pytest.mark.parametrize( + ("handler", "request_payload", "metadata_field", "token_return_field"), + [ + (VLLMWorkerCaptureHandler(), {}, None, "return_tokens_as_token_ids"), + (MegatronWorkerCaptureHandler(), {}, "request_metadata", "return_tokenized_data"), + ( + MegatronWorkerCaptureHandler(), + {"request_metadata": {"caller_metadata": "preserved"}}, + "request_metadata", + "return_tokenized_data", + ), + ], + ids=["vllm", "megatron", "megatron-existing-metadata"], +) +def test_handler_prepares_worker_staged_request( + handler, request_payload: dict[str, Any], metadata_field, token_return_field +) -> None: + store = InMemoryLineageStore() + context = _root_context(store) + token = set_token_sink(context) + try: + payload = handler.prepare_request(request_payload) + finally: + reset_token_sink(token) + + capture_container = payload if metadata_field is None else payload[metadata_field] + assert capture_container["ng_capture"] == context.capture_admission.model_dump(mode="json") + if metadata_field is not None: + assert capture_container == { + **request_payload.get(metadata_field, {}), + "ng_capture": context.capture_admission.model_dump(mode="json"), + } + assert "ng_capture" not in payload + assert payload["logprobs"] is True + assert payload["top_logprobs"] == 0 + assert payload[token_return_field] is True + other_token_return_field = ( + "return_tokenized_data" if token_return_field == "return_tokens_as_token_ids" else "return_tokens_as_token_ids" + ) + assert other_token_return_field not in payload + + +@pytest.mark.parametrize( + ("request_payload", "error"), + [ + ({"n": 2}, "requires n=1"), + ({"request_metadata": []}, "request_metadata must be an object"), + ], +) +def test_megatron_handler_rejects_invalid_request_contract(request_payload: dict[str, Any], error: str) -> None: + store = InMemoryLineageStore() + context = _root_context(store) + token = set_token_sink(context) + try: + with pytest.raises(ValueError, match=error): + MegatronWorkerCaptureHandler().prepare_request(request_payload) + finally: + reset_token_sink(token) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("handler", "coords_kwargs", "drop_response_id", "expected_failure"), + [ + ( + VLLMWorkerCaptureHandler(), + { + "delta_len": 3, + "cum_len": 3, + "digest": "0" * 64, + "extras_digest": "1" * 64, + "staging_key": "r0/c1", + "chain_hash": "2" * 64, + "cumulative_hash": "3" * 64, + }, + False, + None, + ), + ( + VLLMWorkerCaptureHandler(), + {"delta_len": 0, "cum_len": 0, "disposition": "capture_failed"}, + False, + "worker_capture_failed", + ), + ( + MegatronWorkerCaptureHandler(), + None, + True, + "worker_response_missing_commit_coordinates", + ), + ], + ids=["vllm-staged", "vllm-capture-failed", "megatron-missing-coordinates"], +) +async def test_handler_finalization_updates_lineage_and_cleans_transport( + handler, coords_kwargs, drop_response_id, expected_failure +) -> None: + store = InMemoryLineageStore() + context = _root_context(store) + payload = _transport_payload() + if coords_kwargs is not None: + payload["ng_commit_coords"] = CommitCoords( + rollout_id="rollout-1", + model_call_id="c1", + prev_len=0, + weight_version=7, + **coords_kwargs, + ).model_dump(mode="json") + if drop_response_id: + payload.pop("id") + token = set_token_sink(context) + try: + await handler.finalize_response(payload) + finally: + reset_token_sink(token) + + manifest = await store.manifest("rollout-1") + assert context.committed is (expected_failure is None) + if expected_failure is None: + record = manifest["records"][0] + assert record["staging_key"] == "r0/c1" + assert record["weight_version"] == 7 + assert record["chain_hash"] == "2" * 64 + assert record["cumulative_hash"] == "3" * 64 + assert record["response_id"] == "request-1" + else: + assert manifest["failures"] == [ + { + "schema_version": 2, + "model_call_id": "c1", + "reason": expected_failure, + } + ] + _assert_transport_fields_stripped(payload) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "handler", + [VLLMWorkerCaptureHandler(), MegatronWorkerCaptureHandler()], + ids=["vllm", "megatron"], +) +async def test_handlers_strip_unadmitted_capture_responses(handler) -> None: + store = InMemoryLineageStore() + context = _root_context(store) + context.capture_admission = None + payload = _transport_payload() + payload["ng_commit_coords"] = {"unused": True} + token = set_token_sink(context) + try: + await handler.finalize_response(payload) + finally: + reset_token_sink(token) + _assert_transport_fields_stripped(payload) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "handler", + [VLLMWorkerCaptureHandler(), MegatronWorkerCaptureHandler()], + ids=["vllm", "megatron"], +) +async def test_handlers_leave_uncorrelated_traffic_untouched(handler) -> None: + payload = _transport_payload() + await handler.finalize_response(payload) + assert payload["prompt_token_ids"] == [10, 11] + assert payload["choices"][0]["message"]["generation_token_ids"] == [12] + + +@pytest.mark.parametrize( + ("backend", "handler_type"), + [("vllm_worker", VLLMWorkerCaptureHandler), ("megatron_worker", MegatronWorkerCaptureHandler)], +) +def test_factory_selects_the_typed_backend_strategy(backend, handler_type) -> None: + assert isinstance(make_external_capture_handler(backend), handler_type) diff --git a/tests/unit_tests/test_token_capture_staging_worker.py b/tests/unit_tests/test_token_capture_staging_worker.py index 58d63b83c3..52b2269673 100644 --- a/tests/unit_tests/test_token_capture_staging_worker.py +++ b/tests/unit_tests/test_token_capture_staging_worker.py @@ -178,6 +178,12 @@ def test_weight_version_is_stamped_at_admission() -> None: assert (first.weight_version, second.weight_version) == (3, 9) +def test_explicit_worker_weight_version_overrides_provider() -> None: + capture, _ = _capture(weight_version=1) + call = capture.begin_call(_root(), weight_version=9) + assert call.weight_version == 9 + + @pytest.mark.parametrize("bad_version", [-1, 1.5, True]) def test_weight_version_must_be_a_non_negative_int(bad_version: Any) -> None: capture, _ = _capture() diff --git a/tests/unit_tests/test_token_id_capture.py b/tests/unit_tests/test_token_id_capture.py index 48a332a4c4..297da1a0d0 100644 --- a/tests/unit_tests/test_token_id_capture.py +++ b/tests/unit_tests/test_token_id_capture.py @@ -482,6 +482,17 @@ def test_external_staging_requires_framework_owned_rebuild_and_active_capture(): assert config.token_id_capture.external_staging is True +def test_megatron_worker_backend_requires_external_staging(): + with pytest.raises(ValueError, match="requires external_staging=true"): + TokenIdCaptureConfig.model_validate( + { + "token_id_capture": { + "external_staging_backend": "megatron_worker", + } + } + ) + + def test_agent_capture_selection_uses_static_agent_config_or_all_agents(): config = { "token_id_capture": {"enabled": True, "rebuild_response": False, "allow_unresolved_continuations": True}, From 7e89d2fe334e6b11283bf3a2d4047b0c2561690f Mon Sep 17 00:00:00 2001 From: Laura Dang Date: Wed, 16 Sep 2026 01:21:27 -0700 Subject: [PATCH 02/12] feat(token-capture): add Gym-owned Megatron extraction adapter Add MegatronCaptureAdapter mirroring VLLMCaptureAdapter so MInf payload extraction runs inside RolloutTokenCapture.complete_call_from_response. Extraction errors now poison the call with capture_failed coordinates (surfacing as worker_capture_failed) instead of leaving Gym without coordinates, matching the vLLM path. Widen CaptureAdapter payload types to Any since MInf hands an offloaded payload object, not a dict. Co-Authored-By: Claude Fable 5.1 Signed-off-by: Laura Dang --- .../token_id_capture/adapters/megatron.py | 65 +++++++++++++++++++ nemo_gym/token_id_capture/staging/capture.py | 2 +- .../token_id_capture/staging/protocols.py | 12 ++-- .../test_token_capture_staging_worker.py | 53 +++++++++++++++ 4 files changed, 127 insertions(+), 5 deletions(-) create mode 100644 nemo_gym/token_id_capture/adapters/megatron.py diff --git a/nemo_gym/token_id_capture/adapters/megatron.py b/nemo_gym/token_id_capture/adapters/megatron.py new file mode 100644 index 0000000000..6dca704db3 --- /dev/null +++ b/nemo_gym/token_id_capture/adapters/megatron.py @@ -0,0 +1,65 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Dependency-light extraction adapter for Megatron Inference offloaded payloads.""" + +from __future__ import annotations + +from collections.abc import Mapping, Sequence +from typing import Any + + +PREFIX_IDS_FIELD = "required_prefix_token_ids" +PROMPT_IDS_FIELD = "prompt_token_ids" +GENERATED_IDS_FIELD = "generated_token_ids" +GENERATED_LOGPROBS_FIELD = "generated_log_probs" + + +def _field(payload: Any, name: str) -> Any: + """Read one field from an MInf payload object or an equivalent mapping.""" + if isinstance(payload, Mapping): + return payload.get(name) + return getattr(payload, name, None) + + +def _sequence(payload: Any, name: str) -> Sequence[Any]: + value = _field(payload, name) + if value is None: + raise ValueError(f"Megatron offloaded payload carries no {name}") + if isinstance(value, (str, bytes)) or not isinstance(value, Sequence): + raise ValueError(f"Megatron offloaded payload field {name} must be a token-id sequence") + return value + + +class MegatronCaptureAdapter: + """Translate MInf request/response material at the framework boundary. + + MInf offloads exact prompt ids, generated ids, and selected-token log + probabilities as attributes on a payload object rather than as a chat + completion dict. The adapter reads either shape so extraction failures + flow through ``RolloutTokenCapture.complete_call_from_response`` and + poison the call with ``capture_failed`` coordinates, the same outcome the + vLLM adapter produces. + """ + + def enter_prefix(self, request_payload: dict[str, Any], prefix_ids: list[int]) -> dict[str, Any]: + request_payload[PREFIX_IDS_FIELD] = list(prefix_ids) + return request_payload + + def extract_prompt_ids(self, response_payload: Any) -> list[int]: + return [int(token_id) for token_id in _sequence(response_payload, PROMPT_IDS_FIELD)] + + def extract_generation(self, response_payload: Any) -> tuple[list[int], list[float]]: + token_ids = [int(token_id) for token_id in _sequence(response_payload, GENERATED_IDS_FIELD)] + log_probs = [float(value) for value in _sequence(response_payload, GENERATED_LOGPROBS_FIELD)] + if len(token_ids) != len(log_probs): + raise ValueError( + f"Megatron generated token and log-probability lengths differ: {len(token_ids)} != {len(log_probs)}" + ) + return token_ids, log_probs + + def extract_extras(self, response_payload: Any) -> dict[str, Any] | None: + # MInf routed-experts rows are total_tokens - 1 long while the staging + # contract is delta-token aligned. Until that shift has a first-class + # representation the adapter stages no extras. + return None diff --git a/nemo_gym/token_id_capture/staging/capture.py b/nemo_gym/token_id_capture/staging/capture.py index d821eb78b9..69da856b5a 100644 --- a/nemo_gym/token_id_capture/staging/capture.py +++ b/nemo_gym/token_id_capture/staging/capture.py @@ -268,7 +268,7 @@ def complete_call( def complete_call_from_response( self, call: ActiveCall, - response_payload: dict[str, Any], + response_payload: Any, ) -> CommitCoords: """Extract engine-native material and stage it as one atomic lifecycle step.""" if self._adapter is None: diff --git a/nemo_gym/token_id_capture/staging/protocols.py b/nemo_gym/token_id_capture/staging/protocols.py index 1d8b97485b..e235b254de 100644 --- a/nemo_gym/token_id_capture/staging/protocols.py +++ b/nemo_gym/token_id_capture/staging/protocols.py @@ -79,15 +79,19 @@ def enter_prefix(self, request_payload: dict[str, Any], prefix_ids: list[int]) - """ ... - def extract_prompt_ids(self, response_payload: dict[str, Any]) -> list[int]: - """Return the exact prompt token IDs used for generation.""" + def extract_prompt_ids(self, response_payload: Any) -> list[int]: + """Return the exact prompt token IDs used for generation. + + ``response_payload`` is engine-native: a chat completion dict for + vLLM, an offloaded payload object for Megatron Inference. + """ ... - def extract_generation(self, response_payload: dict[str, Any]) -> tuple[list[int], list[float]]: + def extract_generation(self, response_payload: Any) -> tuple[list[int], list[float]]: """Return exact generated token IDs and selected-token log probabilities.""" ... - def extract_extras(self, response_payload: dict[str, Any]) -> dict[str, Any] | None: + def extract_extras(self, response_payload: Any) -> dict[str, Any] | None: """Return optional versioned engine-native per-token material.""" ... diff --git a/tests/unit_tests/test_token_capture_staging_worker.py b/tests/unit_tests/test_token_capture_staging_worker.py index 52b2269673..d6dd52cf9a 100644 --- a/tests/unit_tests/test_token_capture_staging_worker.py +++ b/tests/unit_tests/test_token_capture_staging_worker.py @@ -5,10 +5,12 @@ import asyncio import threading +from types import SimpleNamespace from typing import Any import pytest +from nemo_gym.token_id_capture.adapters.megatron import MegatronCaptureAdapter from nemo_gym.token_id_capture.adapters.vllm import ( VLLMCaptureAdapter, extract_generation_token_info, @@ -438,6 +440,57 @@ def test_vllm_extraction_failure_returns_poisoned_coords() -> None: assert sink.events == [] +def _minf_payload(**overrides: Any) -> SimpleNamespace: + fields: dict[str, Any] = { + "prompt_token_ids": [10, 11], + "generated_token_ids": [12, 13], + "generated_log_probs": [-0.2, -0.3], + } + fields.update(overrides) + return SimpleNamespace(**{name: value for name, value in fields.items() if value is not ...}) + + +def test_megatron_adapter_stages_offloaded_payload_objects() -> None: + capture, sink = _capture(adapter=MegatronCaptureAdapter()) + coords = capture.complete_call_from_response(capture.begin_call(_root()), _minf_payload()) + assert coords.disposition == "staged" + assert sink.records[0].token_ids_delta == [10, 11, 12, 13] + assert sink.records[0].generation_log_probs_delta == [0.0, 0.0, -0.2, -0.3] + assert sink.records[0].extras is None + + +def test_megatron_adapter_reads_mapping_payloads_and_casts_scalars() -> None: + adapter = MegatronCaptureAdapter() + payload = {"prompt_token_ids": (1, 2), "generated_token_ids": [3], "generated_log_probs": [-1]} + assert adapter.extract_prompt_ids(payload) == [1, 2] + assert adapter.extract_generation(payload) == ([3], [-1.0]) + assert adapter.extract_extras(payload) is None + + +@pytest.mark.parametrize("missing", ["prompt_token_ids", "generated_token_ids", "generated_log_probs"]) +def test_megatron_adapter_missing_field_poisons_capture(missing: str) -> None: + capture, sink = _capture(adapter=MegatronCaptureAdapter()) + coords = capture.complete_call_from_response( + capture.begin_call(_root()), + _minf_payload(**{missing: ...}), + ) + assert coords.disposition == "capture_failed" + assert sink.events == [] + + +def test_megatron_adapter_rejects_malformed_fields() -> None: + adapter = MegatronCaptureAdapter() + with pytest.raises(ValueError, match="prompt_token_ids must be a token-id sequence"): + adapter.extract_prompt_ids(_minf_payload(prompt_token_ids=5)) + with pytest.raises(ValueError, match="lengths differ"): + adapter.extract_generation(_minf_payload(generated_log_probs=[-0.1])) + + +def test_megatron_adapter_enter_prefix_writes_the_required_prefix_field() -> None: + request = MegatronCaptureAdapter().enter_prefix({"n": 1}, [1, 2]) + assert request == {"n": 1, "required_prefix_token_ids": [1, 2]} + + def test_install_capture_uses_the_worker_host_seam() -> None: host = CaptureHost() sink = _MemorySink() From ba47b88e0a035719497aed2c71a775787cddb532 Mon Sep 17 00:00:00 2001 From: Laura Dang Date: Wed, 16 Sep 2026 01:30:48 -0700 Subject: [PATCH 03/12] refactor(token-capture): share acknowledgement commit path across backends Move _finalize_admitted_response into _BaseExternalCaptureHandler so the vLLM and Megatron handlers are siblings that differ only in request preparation. The commit path operates solely on shared Gym contracts. Replace per-backend log strings with a single backend label. Co-Authored-By: Claude Fable 5.1 Signed-off-by: Laura Dang --- nemo_gym/token_id_capture/external_capture.py | 73 ++++++++----------- 1 file changed, 31 insertions(+), 42 deletions(-) diff --git a/nemo_gym/token_id_capture/external_capture.py b/nemo_gym/token_id_capture/external_capture.py index 36abfd40d8..759341075d 100644 --- a/nemo_gym/token_id_capture/external_capture.py +++ b/nemo_gym/token_id_capture/external_capture.py @@ -77,8 +77,7 @@ class _BaseExternalCaptureHandler(ABC): """Own the lifecycle shared by external capture backends.""" _INVALID_CAPTURE_REASON: str - _CAPTURE_ERROR_MESSAGE: str - _POISON_ERROR_MESSAGE: str + _BACKEND_LABEL: str def prepare_request(self, request_payload: dict[str, Any]) -> dict[str, Any]: """Attach capture instructions to an engine-bound request. @@ -147,7 +146,7 @@ async def finalize_response(self, served_payload: dict[str, Any]) -> None: # Poison capture without turning a valid model completion into a # harness failure. LOGGER.exception( - self._CAPTURE_ERROR_MESSAGE, + f"{self._BACKEND_LABEL} worker capture acknowledgement failed for rollout %s call %s", context.rollout_id, context.model_call_id, ) @@ -159,12 +158,11 @@ async def finalize_response(self, served_payload: dict[str, Any]) -> None: ) except Exception: LOGGER.exception( - self._POISON_ERROR_MESSAGE, + f"Could not poison rollout %s call %s after a failed {self._BACKEND_LABEL} worker acknowledgement", context.rollout_id, context.model_call_id, ) - @abstractmethod async def _finalize_admitted_response( self, served_payload: dict[str, Any], @@ -174,41 +172,11 @@ async def _finalize_admitted_response( ledger: CaptureLedger, admission: CaptureAdmission, ) -> None: - """Publish backend-specific custody for an admitted response.""" - - -class VLLMWorkerCaptureHandler(_BaseExternalCaptureHandler): - """Commit lineage after a vLLM worker durably stages the token delta.""" - - _INVALID_CAPTURE_REASON = INVALID_COMMIT_COORDS_REASON - _CAPTURE_ERROR_MESSAGE = "Worker capture acknowledgement failed for rollout %s call %s" - _POISON_ERROR_MESSAGE = "Could not poison rollout %s call %s after a failed acknowledgement" - - def _prepare_admitted_request( - self, - request_payload: dict[str, Any], - admission: CaptureAdmission, - ) -> dict[str, Any]: - request_payload[NG_CAPTURE_FIELD] = admission.model_dump(mode="json") - request_payload.update( - logprobs=True, - top_logprobs=0, - return_tokens_as_token_ids=True, - ) - if admission.mode == "token_in": - request_payload["required_prefix_token_ids"] = list(admission.required_prefix_token_ids) - return request_payload + """Validate the worker acknowledgement and commit lineage for an admitted response. - async def _finalize_admitted_response( - self, - served_payload: dict[str, Any], - *, - coords_payload: dict[str, Any] | None, - context: CaptureContext, - ledger: CaptureLedger, - admission: CaptureAdmission, - ) -> None: - """Publish the worker's coordinates as a ledger row. + This path operates only on shared Gym contracts (``CommitCoords``, + ``CallRecord``, ``CaptureLedgerCommit``); backends differ only in how + ``_prepare_admitted_request`` asks the engine to stage tokens. The ordering invariant the external sink requires — a call must not become a lineage parent until its staged record is durable — holds @@ -295,12 +263,33 @@ async def _finalize_admitted_response( ) -class MegatronWorkerCaptureHandler(VLLMWorkerCaptureHandler): +class VLLMWorkerCaptureHandler(_BaseExternalCaptureHandler): + """Commit lineage after a vLLM worker durably stages the token delta.""" + + _INVALID_CAPTURE_REASON = INVALID_COMMIT_COORDS_REASON + _BACKEND_LABEL = "vLLM" + + def _prepare_admitted_request( + self, + request_payload: dict[str, Any], + admission: CaptureAdmission, + ) -> dict[str, Any]: + request_payload[NG_CAPTURE_FIELD] = admission.model_dump(mode="json") + request_payload.update( + logprobs=True, + top_logprobs=0, + return_tokens_as_token_ids=True, + ) + if admission.mode == "token_in": + request_payload["required_prefix_token_ids"] = list(admission.required_prefix_token_ids) + return request_payload + + +class MegatronWorkerCaptureHandler(_BaseExternalCaptureHandler): """Commit lineage after an MInf worker durably stages a canonical delta.""" _INVALID_CAPTURE_REASON = "invalid_megatron_commit_coordinates" - _CAPTURE_ERROR_MESSAGE = "Megatron capture acknowledgement failed for rollout %s call %s" - _POISON_ERROR_MESSAGE = "Could not poison rollout %s call %s after a failed MInf acknowledgement" + _BACKEND_LABEL = "Megatron" def _prepare_admitted_request( self, From 5f46ece4a134eb62348617b70374e3039b2d3071 Mon Sep 17 00:00:00 2001 From: Laura Dang Date: Wed, 16 Sep 2026 01:38:50 -0700 Subject: [PATCH 04/12] fix(token-capture): share the invalid-coordinates failure reason across backends Both handlers validate the same CommitCoords contract, so the Megatron handler now reuses INVALID_COMMIT_COORDS_REASON from the centralized wire vocabulary instead of a backend-specific string. The default lives on the base handler; backends stay identifiable through log labels. Co-Authored-By: Claude Fable 5.1 Signed-off-by: Laura Dang --- nemo_gym/token_id_capture/external_capture.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/nemo_gym/token_id_capture/external_capture.py b/nemo_gym/token_id_capture/external_capture.py index 759341075d..5ea353a36b 100644 --- a/nemo_gym/token_id_capture/external_capture.py +++ b/nemo_gym/token_id_capture/external_capture.py @@ -76,7 +76,7 @@ def _strip_capture_transport_fields(payload: dict[str, Any]) -> None: class _BaseExternalCaptureHandler(ABC): """Own the lifecycle shared by external capture backends.""" - _INVALID_CAPTURE_REASON: str + _INVALID_CAPTURE_REASON = INVALID_COMMIT_COORDS_REASON _BACKEND_LABEL: str def prepare_request(self, request_payload: dict[str, Any]) -> dict[str, Any]: @@ -266,7 +266,6 @@ async def _finalize_admitted_response( class VLLMWorkerCaptureHandler(_BaseExternalCaptureHandler): """Commit lineage after a vLLM worker durably stages the token delta.""" - _INVALID_CAPTURE_REASON = INVALID_COMMIT_COORDS_REASON _BACKEND_LABEL = "vLLM" def _prepare_admitted_request( @@ -288,7 +287,6 @@ def _prepare_admitted_request( class MegatronWorkerCaptureHandler(_BaseExternalCaptureHandler): """Commit lineage after an MInf worker durably stages a canonical delta.""" - _INVALID_CAPTURE_REASON = "invalid_megatron_commit_coordinates" _BACKEND_LABEL = "Megatron" def _prepare_admitted_request( From 91250ecea585f053e468b9c094ed0a87cfedd087 Mon Sep 17 00:00:00 2001 From: Laura Dang Date: Wed, 16 Sep 2026 01:38:51 -0700 Subject: [PATCH 05/12] test(token-capture): exercise the acknowledgement contract on both handlers Run every finalization case against the vLLM and Megatron handlers: staged, missing coordinates, capture_failed, malformed coordinates, mismatched identity, mismatched admission, and missing response id. Add ledger-write failure cases covering the poison path and the nested log-and-continue branch when poisoning itself fails. Populate the chat_completions mock with every transport field the strip step removes so the absence assertions are load-bearing. Co-Authored-By: Claude Fable 5.1 Signed-off-by: Laura Dang --- .../vllm_model/tests/test_app.py | 35 ++- .../test_external_capture_handlers.py | 257 ++++++++++++++---- 2 files changed, 234 insertions(+), 58 deletions(-) diff --git a/responses_api_models/vllm_model/tests/test_app.py b/responses_api_models/vllm_model/tests/test_app.py index 302d04426d..2403dcc7a0 100644 --- a/responses_api_models/vllm_model/tests/test_app.py +++ b/responses_api_models/vllm_model/tests/test_app.py @@ -891,13 +891,29 @@ async def test_megatron_capture_handler_finalizes_through_chat_completions(self, "chain_hash": "2" * 64, "cumulative_hash": "3" * 64, }, + # Transport-only token data at every location the capture handler + # must scrub before the completion leaves the model server. + "prompt_token_ids": [10, 11], "choices": [ { "index": 0, "finish_reason": "stop", + "token_ids": [12, 13, 14], + "logprobs": { + "content": [ + {"token": "token_id:12", "logprob": -0.1, "bytes": None, "top_logprobs": []}, + {"token": "token_id:13", "logprob": -0.2, "bytes": None, "top_logprobs": []}, + {"token": "token_id:14", "logprob": -0.3, "bytes": None, "top_logprobs": []}, + ] + }, "message": { "role": "assistant", "content": "done", + "prompt_token_ids": [10, 11], + # Megatron ``return_tokenized_data`` echo of the exact prompt form. + "compact_prompt_token_ids": [10, 11], + "generation_token_ids": [12, 13, 14], + "generation_log_probs": [-0.1, -0.2, -0.3], }, } ], @@ -948,10 +964,25 @@ async def test_megatron_capture_handler_finalizes_through_chat_completions(self, assert manifest["records"][0]["staging_key"] == "rollout-1/c1" assert manifest["records"][0]["weight_version"] == 7 assert manifest["records"][0]["response_id"] == "minf-17" + # ``NeMoGymChatCompletion`` inherits the OpenAI SDK's ``extra="allow"``, so + # any transport field left on the dict would be re-admitted verbatim into + # the served response. Every injected location must therefore be gone. response_payload = response.model_dump() assert "prompt_token_ids" not in response_payload - assert "prompt_token_ids" not in response_payload["choices"][0]["message"] - assert "generation_token_ids" not in response_payload["choices"][0]["message"] + assert "ng_commit_coords" not in response_payload + served_choice = response_payload["choices"][0] + assert "token_ids" not in served_choice + # ``logprobs`` is a declared ``Choice`` field: stripping the dict entry + # leaves the pydantic default rather than removing the key. + assert served_choice["logprobs"] is None + served_message = served_choice["message"] + for field_name in ( + "prompt_token_ids", + "compact_prompt_token_ids", + "generation_token_ids", + "generation_log_probs", + ): + assert field_name not in served_message, field_name def test_session_client_routing_is_stable_across_workers(self, monkeypatch: MonkeyPatch) -> None: workers = [self._setup_server(monkeypatch) for _ in range(2)] diff --git a/tests/unit_tests/test_external_capture_handlers.py b/tests/unit_tests/test_external_capture_handlers.py index 231873d833..bedeb60977 100644 --- a/tests/unit_tests/test_external_capture_handlers.py +++ b/tests/unit_tests/test_external_capture_handlers.py @@ -3,6 +3,8 @@ """External capture strategy lifecycle tests.""" +import logging +from dataclasses import dataclass from typing import Any import pytest @@ -14,7 +16,21 @@ ) from nemo_gym.token_id_capture.lineage import InMemoryLineageStore from nemo_gym.token_id_capture.sink import CaptureContext, reset_token_sink, set_token_sink -from nemo_gym.token_id_capture.staging.records import CaptureAdmission, CommitCoords +from nemo_gym.token_id_capture.staging.records import ( + INVALID_COMMIT_COORDS_REASON, + WORKER_CAPTURE_FAILED_REASON, + WORKER_MISSING_COMMIT_COORDS_REASON, + CaptureAdmission, + CaptureLedgerCommit, + CommitCoords, +) + + +HANDLER_CLASSES = pytest.mark.parametrize( + "handler_cls", + [VLLMWorkerCaptureHandler, MegatronWorkerCaptureHandler], + ids=["vllm", "megatron"], +) def _root_context(store: InMemoryLineageStore) -> CaptureContext: @@ -132,82 +148,211 @@ def test_megatron_handler_rejects_invalid_request_contract(request_payload: dict reset_token_sink(token) -@pytest.mark.asyncio -@pytest.mark.parametrize( - ("handler", "coords_kwargs", "drop_response_id", "expected_failure"), - [ - ( - VLLMWorkerCaptureHandler(), - { - "delta_len": 3, - "cum_len": 3, - "digest": "0" * 64, - "extras_digest": "1" * 64, - "staging_key": "r0/c1", - "chain_hash": "2" * 64, - "cumulative_hash": "3" * 64, - }, - False, - None, - ), - ( - VLLMWorkerCaptureHandler(), - {"delta_len": 0, "cum_len": 0, "disposition": "capture_failed"}, - False, - "worker_capture_failed", - ), - ( - MegatronWorkerCaptureHandler(), - None, - True, - "worker_response_missing_commit_coordinates", - ), - ], - ids=["vllm-staged", "vllm-capture-failed", "megatron-missing-coordinates"], -) -async def test_handler_finalization_updates_lineage_and_cleans_transport( - handler, coords_kwargs, drop_response_id, expected_failure -) -> None: - store = InMemoryLineageStore() - context = _root_context(store) - payload = _transport_payload() - if coords_kwargs is not None: - payload["ng_commit_coords"] = CommitCoords( +def _staged_coords(**overrides: Any) -> dict[str, Any]: + """Return a valid ``staged`` acknowledgement for the ``_root_context`` call, with overrides.""" + kwargs: dict[str, Any] = { + "rollout_id": "rollout-1", + "model_call_id": "c1", + "prev_len": 0, + "weight_version": 7, + "delta_len": 3, + "cum_len": 3, + "digest": "0" * 64, + "extras_digest": "1" * 64, + "staging_key": "r0/c1", + "chain_hash": "2" * 64, + "cumulative_hash": "3" * 64, + } + kwargs.update(overrides) + return CommitCoords(**kwargs).model_dump(mode="json") + + +@dataclass(frozen=True) +class _FinalizeCase: + """One worker-acknowledgement scenario, run against every handler backend. + + ``coords_payload`` is placed verbatim on ``ng_commit_coords`` (``None`` + means the worker sent no acknowledgement at all). ``expected_failure`` is + the poison reason the ledger must carry, or ``None`` for a committed call. + """ + + coords_payload: Any + expected_failure: str | None + drop_response_id: bool = False + + +_FINALIZE_CASES = { + "staged": _FinalizeCase(coords_payload=_staged_coords(), expected_failure=None), + # ``drop_response_id`` is deliberately False: finalization returns at the + # missing-coordinates branch before the envelope id is ever inspected. + "missing-coordinates": _FinalizeCase( + coords_payload=None, + expected_failure=WORKER_MISSING_COMMIT_COORDS_REASON, + ), + "capture-failed": _FinalizeCase( + coords_payload=CommitCoords( rollout_id="rollout-1", model_call_id="c1", prev_len=0, weight_version=7, - **coords_kwargs, - ).model_dump(mode="json") - if drop_response_id: - payload.pop("id") + delta_len=0, + cum_len=0, + disposition="capture_failed", + ).model_dump(mode="json"), + expected_failure=WORKER_CAPTURE_FAILED_REASON, + ), + "malformed-not-a-dict": _FinalizeCase( + coords_payload=["not", "a", "mapping"], + expected_failure=INVALID_COMMIT_COORDS_REASON, + ), + "malformed-missing-fields": _FinalizeCase( + coords_payload={"rollout_id": "rollout-1", "model_call_id": "c1"}, + expected_failure=INVALID_COMMIT_COORDS_REASON, + ), + "mismatched-rollout-id": _FinalizeCase( + coords_payload=_staged_coords(rollout_id="rollout-other"), + expected_failure=INVALID_COMMIT_COORDS_REASON, + ), + "mismatched-model-call-id": _FinalizeCase( + coords_payload=_staged_coords(model_call_id="c9"), + expected_failure=INVALID_COMMIT_COORDS_REASON, + ), + # Self-consistent child coordinates (parent set, prev_len > 0) that + # diverge from the parentless text admission on the context. + "mismatched-parent-and-prev-len": _FinalizeCase( + coords_payload=_staged_coords(parent_call_id="c0", prev_len=2, cum_len=5), + expected_failure=INVALID_COMMIT_COORDS_REASON, + ), + "missing-response-id": _FinalizeCase( + coords_payload=_staged_coords(), + expected_failure=INVALID_COMMIT_COORDS_REASON, + drop_response_id=True, + ), +} + + +def _assert_poisoned( + manifest: dict[str, Any], + context: CaptureContext, + payload: dict[str, Any], + reason: str, +) -> None: + """Assert a call failed closed: not committed, exactly one poison row, transport scrubbed.""" + assert context.committed is False + assert manifest["records"] == [] + assert manifest["failures"] == [ + { + "schema_version": 2, + "model_call_id": "c1", + "reason": reason, + } + ] + _assert_transport_fields_stripped(payload) + + +async def _prepare_and_finalize(handler, context: CaptureContext, payload: dict[str, Any]) -> None: token = set_token_sink(context) try: handler.prepare_response(payload) _assert_transport_fields_stripped(payload) - assert "ng_commit_coords" not in payload await handler.finalize_response(payload) finally: reset_token_sink(token) + +@pytest.mark.asyncio +@HANDLER_CLASSES +@pytest.mark.parametrize("case", list(_FINALIZE_CASES.values()), ids=list(_FINALIZE_CASES)) +async def test_handler_finalization_updates_lineage_and_cleans_transport(handler_cls, case: _FinalizeCase) -> None: + handler = handler_cls() + store = InMemoryLineageStore() + context = _root_context(store) + payload = _transport_payload() + if case.coords_payload is not None: + payload["ng_commit_coords"] = case.coords_payload + if case.drop_response_id: + payload.pop("id") + + await _prepare_and_finalize(handler, context, payload) + manifest = await store.manifest("rollout-1") - assert context.committed is (expected_failure is None) - if expected_failure is None: + if case.expected_failure is None: + assert context.committed is True + assert manifest["failures"] == [] record = manifest["records"][0] assert record["staging_key"] == "r0/c1" assert record["weight_version"] == 7 assert record["chain_hash"] == "2" * 64 assert record["cumulative_hash"] == "3" * 64 assert record["response_id"] == "request-1" + _assert_transport_fields_stripped(payload) else: - assert manifest["failures"] == [ - { - "schema_version": 2, - "model_call_id": "c1", - "reason": expected_failure, - } - ] + _assert_poisoned(manifest, context, payload, case.expected_failure) + + +class _FaultyLedger(InMemoryLineageStore): + """Ledger whose writes can be made to raise, to exercise the poison fallback paths.""" + + def __init__(self, *, record_fails: bool, record_failure_fails: bool) -> None: + super().__init__() + self._record_fails = record_fails + self._record_failure_fails = record_failure_fails + self.record_failure_calls: list[tuple[str, str, str]] = [] + + async def record(self, commit: CaptureLedgerCommit) -> None: + if self._record_fails: + raise RuntimeError("ledger write failed") + await super().record(commit) + + async def record_failure(self, rollout_id: str, model_call_id: str, reason: str) -> None: + self.record_failure_calls.append((rollout_id, model_call_id, reason)) + if self._record_failure_fails: + raise RuntimeError("ledger poison write failed") + await super().record_failure(rollout_id, model_call_id, reason) + + +@pytest.mark.asyncio +@HANDLER_CLASSES +async def test_handler_poisons_call_when_ledger_record_raises(handler_cls) -> None: + handler = handler_cls() + store = _FaultyLedger(record_fails=True, record_failure_fails=False) + context = _root_context(store) + payload = _transport_payload() + payload["ng_commit_coords"] = _staged_coords() + + await _prepare_and_finalize(handler, context, payload) + + manifest = await store.manifest("rollout-1") + _assert_poisoned(manifest, context, payload, INVALID_COMMIT_COORDS_REASON) + assert store.record_failure_calls == [("rollout-1", "c1", INVALID_COMMIT_COORDS_REASON)] + + +@pytest.mark.asyncio +@HANDLER_CLASSES +async def test_handler_logs_and_continues_when_poisoning_also_fails(handler_cls, caplog) -> None: + handler = handler_cls() + store = _FaultyLedger(record_fails=True, record_failure_fails=True) + context = _root_context(store) + payload = _transport_payload() + payload["ng_commit_coords"] = _staged_coords() + + with caplog.at_level(logging.ERROR, logger="nemo_gym.token_id_capture.external_capture"): + # Must return normally: a valid completion is never turned into a harness failure. + await _prepare_and_finalize(handler, context, payload) + + assert context.committed is False _assert_transport_fields_stripped(payload) + assert store.record_failure_calls == [("rollout-1", "c1", INVALID_COMMIT_COORDS_REASON)] + manifest = await store.manifest("rollout-1") + assert manifest["records"] == [] + assert manifest["failures"] == [] + assert ( + f"{handler._BACKEND_LABEL} worker capture acknowledgement failed for rollout rollout-1 call c1" in caplog.text + ) + assert ( + f"Could not poison rollout rollout-1 call c1 after a failed {handler._BACKEND_LABEL} worker acknowledgement" + in caplog.text + ) @pytest.mark.asyncio From a3cef777968a2ec8de9a7d2b12e599b71733bd42 Mon Sep 17 00:00:00 2001 From: Laura Dang Date: Wed, 16 Sep 2026 22:46:37 -0700 Subject: [PATCH 06/12] fix(token-capture): send offload_params to the Megatron endpoint Megatron Inference renamed the opaque per-request dict that carries the capture admission to offload_params (NVIDIA/Megatron-LM PR #7015); write that body key instead of request_metadata. Stop sending required_prefix_token_ids on token-in calls: Megatron's endpoint no longer accepts it, and the worker-side prompt preparer resolves the prefix from staging_chain anyway, so the field was always an empty list on this path. Co-Authored-By: Claude Fable 5.1 Signed-off-by: Laura Dang --- nemo_gym/token_id_capture/external_capture.py | 19 ++++++++++--------- .../test_external_capture_handlers.py | 8 ++++---- 2 files changed, 14 insertions(+), 13 deletions(-) diff --git a/nemo_gym/token_id_capture/external_capture.py b/nemo_gym/token_id_capture/external_capture.py index 5ea353a36b..6a949717ff 100644 --- a/nemo_gym/token_id_capture/external_capture.py +++ b/nemo_gym/token_id_capture/external_capture.py @@ -297,20 +297,21 @@ def _prepare_admitted_request( choice_count = request_payload.get("n") if choice_count is not None and choice_count != 1: raise ValueError("Megatron token capture requires n=1") - request_metadata = request_payload.get("request_metadata") - if request_metadata is None: - request_metadata = {} - request_payload["request_metadata"] = request_metadata - if not isinstance(request_metadata, dict): - raise ValueError("Megatron request_metadata must be an object") - request_metadata[NG_CAPTURE_FIELD] = admission.model_dump(mode="json") + # Megatron Inference forwards ``offload_params`` opaquely to its prompt preparer and + # payload stager; the admission rides inside it. The prefix itself is resolved on the + # worker from ``staging_chain``, so no prefix token ids travel on the request. + offload_params = request_payload.get("offload_params") + if offload_params is None: + offload_params = {} + request_payload["offload_params"] = offload_params + if not isinstance(offload_params, dict): + raise ValueError("Megatron offload_params must be an object") + offload_params[NG_CAPTURE_FIELD] = admission.model_dump(mode="json") request_payload.update( logprobs=True, top_logprobs=0, return_tokenized_data=True, ) - if admission.mode == "token_in": - request_payload["required_prefix_token_ids"] = list(admission.required_prefix_token_ids) return request_payload diff --git a/tests/unit_tests/test_external_capture_handlers.py b/tests/unit_tests/test_external_capture_handlers.py index bedeb60977..f98e34eb76 100644 --- a/tests/unit_tests/test_external_capture_handlers.py +++ b/tests/unit_tests/test_external_capture_handlers.py @@ -92,11 +92,11 @@ def _assert_transport_fields_stripped(payload: dict[str, Any]) -> None: ("handler", "request_payload", "metadata_field", "token_return_field"), [ (VLLMWorkerCaptureHandler(), {}, None, "return_tokens_as_token_ids"), - (MegatronWorkerCaptureHandler(), {}, "request_metadata", "return_tokenized_data"), + (MegatronWorkerCaptureHandler(), {}, "offload_params", "return_tokenized_data"), ( MegatronWorkerCaptureHandler(), - {"request_metadata": {"caller_metadata": "preserved"}}, - "request_metadata", + {"offload_params": {"caller_metadata": "preserved"}}, + "offload_params", "return_tokenized_data", ), ], @@ -134,7 +134,7 @@ def test_handler_prepares_worker_staged_request( ("request_payload", "error"), [ ({"n": 2}, "requires n=1"), - ({"request_metadata": []}, "request_metadata must be an object"), + ({"offload_params": []}, "offload_params must be an object"), ], ) def test_megatron_handler_rejects_invalid_request_contract(request_payload: dict[str, Any], error: str) -> None: From 37dc751f45d6c666ee5ae760a0c36882b50b3936 Mon Sep 17 00:00:00 2001 From: Laura Dang Date: Wed, 16 Sep 2026 23:03:51 -0700 Subject: [PATCH 07/12] fix(token-capture): allow RolloutReceipt.terminal_selection to be unset None means no attribution stage ran (for example the manifest failed to parse), so consumers can keep such receipts out of the per-method buckets. Signed-off-by: Laura Dang Co-Authored-By: Claude Fable 5.1 --- nemo_gym/token_id_capture/staging/records.py | 6 ++++-- tests/unit_tests/test_token_capture_staging_core.py | 6 ++++++ 2 files changed, 10 insertions(+), 2 deletions(-) diff --git a/nemo_gym/token_id_capture/staging/records.py b/nemo_gym/token_id_capture/staging/records.py index dae68325d9..6e4733d993 100644 --- a/nemo_gym/token_id_capture/staging/records.py +++ b/nemo_gym/token_id_capture/staging/records.py @@ -369,8 +369,10 @@ class RolloutReceipt(_DigestWireModel): # named the response it kept, ``response_id``/``content`` when a witness # joined the scored response to a manifest row, ``heuristic`` when the # parent-link walk inferred it (also stamped on failed selections — it - # names the last stage attempted). - terminal_selection: Literal["declared", "response_id", "content", "heuristic"] + # names the last stage attempted). ``None`` when no attribution stage ran + # at all (for example the manifest failed to parse), so such receipts + # stay out of the per-method buckets. + terminal_selection: Literal["declared", "response_id", "content", "heuristic"] | None = None # The witness abstention/corroboration trail from attribution, kept on # success and failure alike so per-method metrics stay diagnosable. terminal_attribution_reason: str | None = None diff --git a/tests/unit_tests/test_token_capture_staging_core.py b/tests/unit_tests/test_token_capture_staging_core.py index fb1e1d7254..e77eec124d 100644 --- a/tests/unit_tests/test_token_capture_staging_core.py +++ b/tests/unit_tests/test_token_capture_staging_core.py @@ -333,6 +333,12 @@ def test_stage_result_and_receipt_validate_identity() -> None: manifest=[manifest], terminal_selection="declared", ) + # No attribution stage ran (e.g. the manifest failed to parse): the + # selection stays unset rather than being stamped with a method. + unattributed = RolloutReceipt(rollout_id="rollout-1", manifest=[manifest]) + assert unattributed.terminal_selection is None + with pytest.raises(ValidationError): + RolloutReceipt(rollout_id="rollout-1", manifest=[manifest], terminal_selection="guess") def test_staging_namespace_has_no_serving_or_framework_dependencies() -> None: From 04f035455cdf16431f56027a914f0f90c344a853 Mon Sep 17 00:00:00 2001 From: Laura Dang Date: Wed, 16 Sep 2026 23:04:34 -0700 Subject: [PATCH 08/12] refactor(token-capture): stop requesting return_tokenized_data from Megatron The worker stages the token delta itself and Gym strips every token transport field before the completion leaves the model server, so the HTTP echo was never read. Drop the flag and keep the compact_prompt_token_ids strip as a defensive measure. Also repoint the vLLM model server Megatron tests at offload_params; they still asserted the old request_metadata key and the dropped required_prefix_token_ids field after a3cef7779. Co-Authored-By: Claude Fable 5.1 Signed-off-by: Laura Dang --- nemo_gym/token_id_capture/external_capture.py | 13 ++++++------- .../vllm_model/tests/test_app.py | 13 ++++++++----- .../test_external_capture_handlers.py | 18 ++++++++++-------- 3 files changed, 24 insertions(+), 20 deletions(-) diff --git a/nemo_gym/token_id_capture/external_capture.py b/nemo_gym/token_id_capture/external_capture.py index 6a949717ff..2e249c175a 100644 --- a/nemo_gym/token_id_capture/external_capture.py +++ b/nemo_gym/token_id_capture/external_capture.py @@ -37,8 +37,9 @@ LOGGER = logging.getLogger(__name__) -# Megatron's ``return_tokenized_data`` echoes the exact prompt token form it -# uses for lossless multi-turn prefix stitching alongside ``TOKEN_FIELDS``. +# Megatron may echo the exact prompt token form it uses for lossless +# multi-turn prefix stitching. Gym does not request or read it, but strips it +# defensively alongside ``TOKEN_FIELDS`` so it never reaches the agent hop. _MEGATRON_TRANSPORT_FIELDS = ("compact_prompt_token_ids",) @@ -307,11 +308,9 @@ def _prepare_admitted_request( if not isinstance(offload_params, dict): raise ValueError("Megatron offload_params must be an object") offload_params[NG_CAPTURE_FIELD] = admission.model_dump(mode="json") - request_payload.update( - logprobs=True, - top_logprobs=0, - return_tokenized_data=True, - ) + # The worker stages the token delta itself, so Gym requests no token + # echo (``return_tokenized_data``) on the HTTP path. + request_payload.update(logprobs=True, top_logprobs=0) return request_payload diff --git a/responses_api_models/vllm_model/tests/test_app.py b/responses_api_models/vllm_model/tests/test_app.py index 2403dcc7a0..a3775dec2b 100644 --- a/responses_api_models/vllm_model/tests/test_app.py +++ b/responses_api_models/vllm_model/tests/test_app.py @@ -856,11 +856,14 @@ def test_megatron_capture_handler_prepares_an_admitted_child_request(self, monke finally: reset_token_sink(token) - assert outbound["return_tokenized_data"] is True - assert outbound["required_prefix_token_ids"] == [10, 11, 12] + # The prefix contract travels only inside the admission in offload_params; + # the worker resolves it from staging_chain, so the body carries no prefix + # copy and requests no token echo. + assert "return_tokenized_data" not in outbound + assert "required_prefix_token_ids" not in outbound assert outbound["logprobs"] is True assert outbound["top_logprobs"] == 0 - assert outbound["request_metadata"]["ng_capture"] == context.capture_admission.model_dump(mode="json") + assert outbound["offload_params"]["ng_capture"] == context.capture_admission.model_dump(mode="json") assert "ng_capture" not in outbound assert "return_tokens_as_token_ids" not in outbound @@ -956,8 +959,8 @@ async def test_megatron_capture_handler_finalizes_through_chat_completions(self, reset_token_sink(token) outbound = client.create_chat_completion.await_args.kwargs - assert outbound["return_tokenized_data"] is True - assert outbound["request_metadata"]["ng_capture"] == context.capture_admission.model_dump(mode="json") + assert "return_tokenized_data" not in outbound + assert outbound["offload_params"]["ng_capture"] == context.capture_admission.model_dump(mode="json") assert "ng_capture" not in outbound assert context.committed is True manifest = await lineage_store.manifest("rollout-1") diff --git a/tests/unit_tests/test_external_capture_handlers.py b/tests/unit_tests/test_external_capture_handlers.py index f98e34eb76..0ac7ee9061 100644 --- a/tests/unit_tests/test_external_capture_handlers.py +++ b/tests/unit_tests/test_external_capture_handlers.py @@ -57,7 +57,7 @@ def _transport_payload(**fields: Any) -> dict[str, Any]: "generation_token_ids": [12], "generation_log_probs": [-0.2], "routed_experts": {"data": "unused"}, - # Megatron ``return_tokenized_data`` echo, absent from vLLM payloads. + # Megatron prompt-form echo, absent from vLLM payloads; stripped defensively. "compact_prompt_token_ids": [10, 11], } message.update(fields) @@ -92,12 +92,12 @@ def _assert_transport_fields_stripped(payload: dict[str, Any]) -> None: ("handler", "request_payload", "metadata_field", "token_return_field"), [ (VLLMWorkerCaptureHandler(), {}, None, "return_tokens_as_token_ids"), - (MegatronWorkerCaptureHandler(), {}, "offload_params", "return_tokenized_data"), + (MegatronWorkerCaptureHandler(), {}, "offload_params", None), ( MegatronWorkerCaptureHandler(), {"offload_params": {"caller_metadata": "preserved"}}, "offload_params", - "return_tokenized_data", + None, ), ], ids=["vllm", "megatron", "megatron-existing-metadata"], @@ -123,11 +123,13 @@ def test_handler_prepares_worker_staged_request( assert "ng_capture" not in payload assert payload["logprobs"] is True assert payload["top_logprobs"] == 0 - assert payload[token_return_field] is True - other_token_return_field = ( - "return_tokenized_data" if token_return_field == "return_tokens_as_token_ids" else "return_tokens_as_token_ids" - ) - assert other_token_return_field not in payload + if token_return_field is None: + # Megatron stages the delta on the worker; Gym asks for no token echo. + assert "return_tokenized_data" not in payload + assert "return_tokens_as_token_ids" not in payload + else: + assert payload[token_return_field] is True + assert "return_tokenized_data" not in payload @pytest.mark.parametrize( From 8293c3f7d6d2d9afbadbf953c3ae5494b6016e74 Mon Sep 17 00:00:00 2001 From: Laura Dang Date: Fri, 18 Sep 2026 09:56:28 -0700 Subject: [PATCH 09/12] Update nemo_gym/token_id_capture/adapters/megatron.py Co-authored-by: Ananth Subramaniam Signed-off-by: Laura Dang --- nemo_gym/token_id_capture/adapters/megatron.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/nemo_gym/token_id_capture/adapters/megatron.py b/nemo_gym/token_id_capture/adapters/megatron.py index 6dca704db3..77cec42d2a 100644 --- a/nemo_gym/token_id_capture/adapters/megatron.py +++ b/nemo_gym/token_id_capture/adapters/megatron.py @@ -32,9 +32,9 @@ def _sequence(payload: Any, name: str) -> Sequence[Any]: class MegatronCaptureAdapter: - """Translate MInf request/response material at the framework boundary. + """Translate Megatron inference request/response material at the framework boundary. - MInf offloads exact prompt ids, generated ids, and selected-token log + Megatron inference offloads exact prompt ids, generated ids, and selected-token log probabilities as attributes on a payload object rather than as a chat completion dict. The adapter reads either shape so extraction failures flow through ``RolloutTokenCapture.complete_call_from_response`` and From c04f7ef1a5e734c6ec418c2ce32a287dafed1f74 Mon Sep 17 00:00:00 2001 From: Laura Dang Date: Fri, 18 Sep 2026 14:25:31 -0700 Subject: [PATCH 10/12] fix(token-capture): enforce RolloutReceipt attribution-state invariants An unpoisoned receipt must name both terminal_model_call_id and terminal_selection and carry no failure_reason. An unset terminal_selection is reserved for receipts where no attribution stage ran: such a receipt names no terminal and must be poisoned with a reason. A poisoned receipt may still retain its terminal and method, because attribution can succeed before a later call poisons the rollout. The linearize compatibility wrapper now reports missing_terminal itself instead of constructing a receipt the schema rejects. Signed-off-by: Laura Dang Co-Authored-By: Claude Fable 5.1 Signed-off-by: Laura Dang --- nemo_gym/token_id_capture/staging/rebuild.py | 4 ++ nemo_gym/token_id_capture/staging/records.py | 22 +++++++ .../test_token_capture_staging_core.py | 61 ++++++++++++++++--- 3 files changed, 79 insertions(+), 8 deletions(-) diff --git a/nemo_gym/token_id_capture/staging/rebuild.py b/nemo_gym/token_id_capture/staging/rebuild.py index 4f73b06b59..1ac9b09198 100644 --- a/nemo_gym/token_id_capture/staging/rebuild.py +++ b/nemo_gym/token_id_capture/staging/rebuild.py @@ -369,6 +369,10 @@ def linearize( terminal_hint: str | None = None, ) -> LinearizedRow: """Compatibility wrapper that still executes the production verifier.""" + if terminal_hint is None: + # ``RolloutReceipt`` rejects an unpoisoned receipt without a terminal; + # surface the same rebuild failure the verifier reports for that case. + raise _fail("missing_terminal", "successful receipt has no terminal call") receipt = RolloutReceipt( rollout_id=rollout_id, terminal_model_call_id=terminal_hint, diff --git a/nemo_gym/token_id_capture/staging/records.py b/nemo_gym/token_id_capture/staging/records.py index 6e4733d993..a4b27018eb 100644 --- a/nemo_gym/token_id_capture/staging/records.py +++ b/nemo_gym/token_id_capture/staging/records.py @@ -385,3 +385,25 @@ def _validate_manifest(self) -> Self: if self.terminal_model_call_id is not None and self.terminal_model_call_id not in set(model_call_ids): raise ValueError("terminal_model_call_id is absent from the receipt manifest") return self + + @model_validator(mode="after") + def _validate_attribution_state(self) -> Self: + """Keep the terminal fields in one of the three documented states. + + A trainable receipt names its terminal and the method that chose it. An + unset ``terminal_selection`` is reserved for receipts where no attribution + stage ran, which cannot name a terminal and must be poisoned with a reason. + A poisoned receipt may still carry a terminal and method, because + attribution can succeed before another call poisons the rollout. + """ + if not self.capture_poisoned: + if self.failure_reason is not None: + raise ValueError("an unpoisoned receipt cannot carry a failure_reason") + if self.terminal_model_call_id is None or self.terminal_selection is None: + raise ValueError("an unpoisoned receipt must name terminal_model_call_id and terminal_selection") + if self.terminal_selection is None: + if self.terminal_model_call_id is not None: + raise ValueError("terminal_model_call_id requires a terminal_selection method") + if self.failure_reason is None: + raise ValueError("a receipt with no attribution stage must be poisoned with a failure_reason") + return self diff --git a/tests/unit_tests/test_token_capture_staging_core.py b/tests/unit_tests/test_token_capture_staging_core.py index e77eec124d..2d11187989 100644 --- a/tests/unit_tests/test_token_capture_staging_core.py +++ b/tests/unit_tests/test_token_capture_staging_core.py @@ -298,13 +298,9 @@ def test_coords_enforce_staged_and_failed_shapes() -> None: ) -def test_stage_result_and_receipt_validate_identity() -> None: - assert StageResult(ok=True, staging_key="backend/key").staging_key == "backend/key" - with pytest.raises(ValidationError, match="requires an error"): - StageResult(ok=False) - +def _manifest_row() -> CallRecord: payload = _record_payload() - manifest = CallRecord( + return CallRecord( model_call_id=payload["model_call_id"], parent_call_id=payload["parent_call_id"], prev_len=payload["prev_len"], @@ -319,6 +315,14 @@ def test_stage_result_and_receipt_validate_identity() -> None: cumulative_hash=payload["cumulative_hash"], response_id="chatcmpl-call-1", ) + + +def test_stage_result_and_receipt_validate_identity() -> None: + assert StageResult(ok=True, staging_key="backend/key").staging_key == "backend/key" + with pytest.raises(ValidationError, match="requires an error"): + StageResult(ok=False) + + manifest = _manifest_row() receipt = RolloutReceipt( rollout_id="rollout-1", terminal_model_call_id="call-1", @@ -334,11 +338,52 @@ def test_stage_result_and_receipt_validate_identity() -> None: terminal_selection="declared", ) # No attribution stage ran (e.g. the manifest failed to parse): the - # selection stays unset rather than being stamped with a method. - unattributed = RolloutReceipt(rollout_id="rollout-1", manifest=[manifest]) + # selection stays unset rather than being stamped with a method, and the + # receipt must be poisoned so the row is not trainable. + unattributed = RolloutReceipt( + rollout_id="rollout-1", + manifest=[manifest], + capture_poisoned=True, + failure_reason="invalid_manifest_row", + ) assert unattributed.terminal_selection is None with pytest.raises(ValidationError): RolloutReceipt(rollout_id="rollout-1", manifest=[manifest], terminal_selection="guess") + # Attribution succeeded before a later call poisoned the rollout: the + # terminal and its method stay on the poisoned receipt. + poisoned_after_attribution = RolloutReceipt( + rollout_id="rollout-1", + terminal_model_call_id="call-1", + manifest=[manifest], + capture_poisoned=True, + failure_reason="worker_capture_failed", + terminal_selection="heuristic", + ) + assert poisoned_after_attribution.terminal_model_call_id == "call-1" + + +@pytest.mark.parametrize( + ("fields", "match"), + [ + # Unpoisoned receipts must be fully attributed. + ({"terminal_model_call_id": None, "terminal_selection": "declared"}, "unpoisoned receipt must name"), + ({"terminal_model_call_id": "call-1", "terminal_selection": None}, "unpoisoned receipt must name"), + ({}, "unpoisoned receipt must name"), + ( + {"terminal_model_call_id": "call-1", "terminal_selection": "declared", "failure_reason": "x"}, + "cannot carry a failure_reason", + ), + # An unset selection is reserved for "no attribution stage ran". + ( + {"terminal_model_call_id": "call-1", "capture_poisoned": True, "failure_reason": "x"}, + "requires a terminal_selection", + ), + ({"capture_poisoned": True}, "must be poisoned with a failure_reason"), + ], +) +def test_rollout_receipt_rejects_inconsistent_attribution_state(fields: dict, match: str) -> None: + with pytest.raises(ValidationError, match=match): + RolloutReceipt(rollout_id="rollout-1", manifest=[_manifest_row()], **fields) def test_staging_namespace_has_no_serving_or_framework_dependencies() -> None: From f11130e28a57e59c5c31848c008735a0b71ed425 Mon Sep 17 00:00:00 2001 From: Laura Dang Date: Fri, 18 Sep 2026 14:25:31 -0700 Subject: [PATCH 11/12] fix(token-capture): fail closed on malformed or multimodal Megatron capture input Validate element types in the Megatron adapter before conversion: token id fields must contain only plain integers and log-probability fields only numbers, so a float, string, or bool element poisons the call instead of being silently coerced into a plausible id. Reject megatron_worker capture requests whose messages carry non-text content parts. The adapter stages no media geometry yet, so a multimodal prompt would commit rows misaligned with the expanded engine prompt. Spell out Megatron Inference instead of the MInf shorthand. Signed-off-by: Laura Dang Co-Authored-By: Claude Fable 5.1 Signed-off-by: Laura Dang --- .../token_id_capture/adapters/megatron.py | 31 ++++++++++++++++--- nemo_gym/token_id_capture/external_capture.py | 25 ++++++++++++++- .../test_external_capture_handlers.py | 12 +++++++ .../test_token_capture_staging_worker.py | 28 +++++++++++++++++ 4 files changed, 90 insertions(+), 6 deletions(-) diff --git a/nemo_gym/token_id_capture/adapters/megatron.py b/nemo_gym/token_id_capture/adapters/megatron.py index 77cec42d2a..260f129805 100644 --- a/nemo_gym/token_id_capture/adapters/megatron.py +++ b/nemo_gym/token_id_capture/adapters/megatron.py @@ -16,7 +16,7 @@ def _field(payload: Any, name: str) -> Any: - """Read one field from an MInf payload object or an equivalent mapping.""" + """Read one field from a Megatron Inference payload object or an equivalent mapping.""" if isinstance(payload, Mapping): return payload.get(name) return getattr(payload, name, None) @@ -31,6 +31,27 @@ def _sequence(payload: Any, name: str) -> Sequence[Any]: return value +def _token_ids(payload: Any, name: str) -> list[int]: + """Read a token-id field, rejecting anything but plain integers. + + Megatron Inference hands host-side ``list[int]`` values. A float, string, or + bool element means the payload is malformed; ``int()`` would silently + truncate or coerce it into a plausible-looking id. + """ + values = _sequence(payload, name) + if any(type(value) is not int for value in values): + raise ValueError(f"Megatron offloaded payload field {name} must contain only integer token ids") + return list(values) + + +def _log_probs(payload: Any, name: str) -> list[float]: + """Read a log-probability field, rejecting non-numeric elements.""" + values = _sequence(payload, name) + if any(isinstance(value, bool) or not isinstance(value, (int, float)) for value in values): + raise ValueError(f"Megatron offloaded payload field {name} must contain only numeric log probabilities") + return [float(value) for value in values] + + class MegatronCaptureAdapter: """Translate Megatron inference request/response material at the framework boundary. @@ -47,11 +68,11 @@ def enter_prefix(self, request_payload: dict[str, Any], prefix_ids: list[int]) - return request_payload def extract_prompt_ids(self, response_payload: Any) -> list[int]: - return [int(token_id) for token_id in _sequence(response_payload, PROMPT_IDS_FIELD)] + return _token_ids(response_payload, PROMPT_IDS_FIELD) def extract_generation(self, response_payload: Any) -> tuple[list[int], list[float]]: - token_ids = [int(token_id) for token_id in _sequence(response_payload, GENERATED_IDS_FIELD)] - log_probs = [float(value) for value in _sequence(response_payload, GENERATED_LOGPROBS_FIELD)] + token_ids = _token_ids(response_payload, GENERATED_IDS_FIELD) + log_probs = _log_probs(response_payload, GENERATED_LOGPROBS_FIELD) if len(token_ids) != len(log_probs): raise ValueError( f"Megatron generated token and log-probability lengths differ: {len(token_ids)} != {len(log_probs)}" @@ -59,7 +80,7 @@ def extract_generation(self, response_payload: Any) -> tuple[list[int], list[flo return token_ids, log_probs def extract_extras(self, response_payload: Any) -> dict[str, Any] | None: - # MInf routed-experts rows are total_tokens - 1 long while the staging + # Megatron Inference routed-experts rows are total_tokens - 1 long while the staging # contract is delta-token aligned. Until that shift has a first-class # representation the adapter stages no extras. return None diff --git a/nemo_gym/token_id_capture/external_capture.py b/nemo_gym/token_id_capture/external_capture.py index f99c24eb1d..ee8fed83f2 100644 --- a/nemo_gym/token_id_capture/external_capture.py +++ b/nemo_gym/token_id_capture/external_capture.py @@ -290,7 +290,7 @@ def _prepare_admitted_request( class MegatronWorkerCaptureHandler(_BaseExternalCaptureHandler): - """Commit lineage after an MInf worker durably stages a canonical delta.""" + """Commit lineage after a Megatron Inference worker durably stages a canonical delta.""" _BACKEND_LABEL = "Megatron" @@ -302,6 +302,7 @@ def _prepare_admitted_request( choice_count = request_payload.get("n") if choice_count is not None and choice_count != 1: raise ValueError("Megatron token capture requires n=1") + _reject_multimodal_content(request_payload) # Megatron Inference forwards ``offload_params`` opaquely to its prompt preparer and # payload stager; the admission rides inside it. The prefix itself is resolved on the # worker from ``staging_chain``, so no prefix token ids travel on the request. @@ -318,6 +319,28 @@ def _prepare_admitted_request( return request_payload +def _reject_multimodal_content(request_payload: dict[str, Any]) -> None: + """Fail closed when a Megatron capture request carries media or audio parts. + + The Megatron adapter stages no media geometry, so a multimodal prompt would + commit token rows whose lengths disagree with the expanded engine prompt. + Until multimodal staging lands, refuse the request rather than train on it. + """ + messages = request_payload.get("messages") + if not isinstance(messages, list): + return + for message in messages: + content = message.get("content") if isinstance(message, dict) else None + if not isinstance(content, list): + continue + for part in content: + part_type = part.get("type") if isinstance(part, dict) else None + if part_type != "text": + raise ValueError( + f"Megatron token capture does not support multimodal content (got part type {part_type!r})" + ) + + def make_external_capture_handler(backend: ExternalStagingBackend) -> ExternalCaptureHandler: """Create the external capture strategy selected by typed configuration.""" if backend == "vllm_worker": diff --git a/tests/unit_tests/test_external_capture_handlers.py b/tests/unit_tests/test_external_capture_handlers.py index 0ac7ee9061..253b3715eb 100644 --- a/tests/unit_tests/test_external_capture_handlers.py +++ b/tests/unit_tests/test_external_capture_handlers.py @@ -137,6 +137,18 @@ def test_handler_prepares_worker_staged_request( [ ({"n": 2}, "requires n=1"), ({"offload_params": []}, "offload_params must be an object"), + ( + {"messages": [{"role": "user", "content": [{"type": "image_url", "image_url": {"url": "data:,"}}]}]}, + "does not support multimodal content", + ), + ( + {"messages": [{"role": "user", "content": [{"type": "audio_url", "audio_url": {"url": "data:,"}}]}]}, + "does not support multimodal content", + ), + ( + {"messages": [{"role": "user", "content": [{"type": "input_audio", "input_audio": {"data": ""}}]}]}, + "does not support multimodal content", + ), ], ) def test_megatron_handler_rejects_invalid_request_contract(request_payload: dict[str, Any], error: str) -> None: diff --git a/tests/unit_tests/test_token_capture_staging_worker.py b/tests/unit_tests/test_token_capture_staging_worker.py index d6dd52cf9a..60bd9f3afd 100644 --- a/tests/unit_tests/test_token_capture_staging_worker.py +++ b/tests/unit_tests/test_token_capture_staging_worker.py @@ -486,6 +486,34 @@ def test_megatron_adapter_rejects_malformed_fields() -> None: adapter.extract_generation(_minf_payload(generated_log_probs=[-0.1])) +@pytest.mark.parametrize("bad_ids", [[1.0, 2], ["1", 2], [True, 2], [None]]) +def test_megatron_adapter_rejects_non_integer_token_ids(bad_ids: list) -> None: + adapter = MegatronCaptureAdapter() + with pytest.raises(ValueError, match="prompt_token_ids must contain only integer token ids"): + adapter.extract_prompt_ids(_minf_payload(prompt_token_ids=bad_ids)) + with pytest.raises(ValueError, match="generated_token_ids must contain only integer token ids"): + adapter.extract_generation( + _minf_payload(generated_token_ids=bad_ids, generated_log_probs=[-0.1] * len(bad_ids)) + ) + + +@pytest.mark.parametrize("bad_log_probs", [["-0.1", -0.2], [True, -0.2], [None, -0.2]]) +def test_megatron_adapter_rejects_non_numeric_log_probs(bad_log_probs: list) -> None: + adapter = MegatronCaptureAdapter() + with pytest.raises(ValueError, match="generated_log_probs must contain only numeric log probabilities"): + adapter.extract_generation(_minf_payload(generated_log_probs=bad_log_probs)) + + +def test_megatron_adapter_malformed_element_poisons_capture() -> None: + capture, sink = _capture(adapter=MegatronCaptureAdapter()) + coords = capture.complete_call_from_response( + capture.begin_call(_root()), + _minf_payload(generated_token_ids=[12.0, 13.0]), + ) + assert coords.disposition == "capture_failed" + assert sink.events == [] + + def test_megatron_adapter_enter_prefix_writes_the_required_prefix_field() -> None: request = MegatronCaptureAdapter().enter_prefix({"n": 1}, [1, 2]) assert request == {"n": 1, "required_prefix_token_ids": [1, 2]} From b601670aaa82973fbc88155a5fd99379cbd62260 Mon Sep 17 00:00:00 2001 From: Laura Dang Date: Fri, 18 Sep 2026 14:25:31 -0700 Subject: [PATCH 12/12] refactor(token-capture): name the backend in the external_staging config check Signed-off-by: Laura Dang Co-Authored-By: Claude Fable 5.1 Signed-off-by: Laura Dang --- nemo_gym/token_id_capture/config.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/nemo_gym/token_id_capture/config.py b/nemo_gym/token_id_capture/config.py index 7ffe0837d7..8c3bfd9abd 100644 --- a/nemo_gym/token_id_capture/config.py +++ b/nemo_gym/token_id_capture/config.py @@ -173,7 +173,7 @@ def _validate(self) -> "TokenIdCaptureConfig": "token_id_capture.external_staging requires rebuild_response=false because the " "framework owns staged-record finalization" ) - if block.external_staging_backend != "vllm_worker" and not block.external_staging: + if block.external_staging_backend == "megatron_worker" and not block.external_staging: raise ValueError("token_id_capture.external_staging_backend requires external_staging=true") if not block.enabled: # Keep inactive settings for templated configurations.