diff --git a/nemo_gym/token_id_capture/external_capture.py b/nemo_gym/token_id_capture/external_capture.py index ee8fed83f2..d3a0e52b4f 100644 --- a/nemo_gym/token_id_capture/external_capture.py +++ b/nemo_gym/token_id_capture/external_capture.py @@ -51,7 +51,18 @@ def prepare_request(self, request_payload: dict[str, Any]) -> dict[str, Any]: ... def prepare_response(self, response_payload: dict[str, Any]) -> None: - """Retain the worker acknowledgement and remove capture-only response fields.""" + """Retain the worker acknowledgement and remove capture-only response fields. + + Model servers must call this for every completion the worker returns, + even one whose acknowledgement is missing. It marks the request as + having received a worker completion, and ``finalize_response`` commits + or poisons the call only when that mark is present. Skipping it for a + real completion would leave the call merely uncommitted instead of + failing closed with ``worker_response_missing_commit_coordinates``. + Completions the model server synthesizes itself (the sequential + reasoning guard, or a backend context-limit error converted into an + empty completion) never pass through here, so they stay uncommitted. + """ ... async def finalize_response(self, served_payload: dict[str, Any]) -> None: @@ -119,6 +130,7 @@ def prepare_response(self, response_payload: dict[str, Any]) -> None: if context is None or not context.external_staging: return context.external_commit_coords = response_payload.pop(NG_COMMIT_COORDS_FIELD, None) + context.external_worker_response_seen = True _strip_capture_transport_fields(response_payload) async def finalize_response(self, served_payload: dict[str, Any]) -> None: @@ -138,6 +150,15 @@ async def finalize_response(self, served_payload: dict[str, Any]) -> None: if admission is None: # UNRESOLVED — the ledger already carries this call's poison row. return + if not context.external_worker_response_seen: + # No worker completion reached ``prepare_response``: the model + # server built this response itself. The reasoning guard never + # calls the worker, and on context overflow the worker returns an + # HTTP 400 instead of a completion. Leave the call uncommitted for + # the middleware to record. A completion the worker did return + # without ``ng_commit_coords`` sets the flag and still fails + # closed below with ``worker_response_missing_commit_coordinates``. + return try: await self._finalize_admitted_response( served_payload, diff --git a/nemo_gym/token_id_capture/sink.py b/nemo_gym/token_id_capture/sink.py index 1cd7fa3c72..910fe870c2 100644 --- a/nemo_gym/token_id_capture/sink.py +++ b/nemo_gym/token_id_capture/sink.py @@ -107,6 +107,9 @@ class CaptureContext: request_items: list[dict] | None = None # Retain the worker acknowledgement privately until API conversion finishes. external_commit_coords: dict[str, Any] | None = None + # A normal worker completion was received, even if its acknowledgement is + # missing. Synthetic guard/overflow completions leave this false. + external_worker_response_seen: bool = False @property def parent_call_id(self) -> str | None: diff --git a/responses_api_models/vllm_model/tests/test_external_capture_streaming.py b/responses_api_models/vllm_model/tests/test_external_capture_streaming.py index 51883134db..558c8b9589 100644 --- a/responses_api_models/vllm_model/tests/test_external_capture_streaming.py +++ b/responses_api_models/vllm_model/tests/test_external_capture_streaming.py @@ -9,6 +9,7 @@ from unittest.mock import AsyncMock, MagicMock import pytest +from aiohttp import ClientResponseError from fastapi.testclient import TestClient from nemo_gym import chat_streaming, responses_streaming @@ -19,10 +20,11 @@ from nemo_gym.token_id_capture.lineage import FileLineageStore from nemo_gym.token_id_capture.records import UNCOMMITTED_CALL_REASON from nemo_gym.token_id_capture.sink import current_capture_context -from nemo_gym.token_id_capture.staging import resolve_terminal +from nemo_gym.token_id_capture.staging import resolve_terminal, select_terminal_call from nemo_gym.token_id_capture.staging.capture import RolloutTokenCapture from nemo_gym.token_id_capture.staging.rebuild import verify_and_linearize from nemo_gym.token_id_capture.staging.records import ( + WORKER_MISSING_COMMIT_COORDS_REASON, CaptureAdmission, RolloutManifest, RolloutReceipt, @@ -44,6 +46,7 @@ def __init__(self): self.context = None self.tool_call = False self.reasoning = False + self.reasoning_only = False self.refusal = False self.capture = RolloutTokenCapture(sink=self, weight_version_fn=lambda: 7, adapter=VLLMCaptureAdapter()) @@ -80,6 +83,8 @@ async def create_chat_completion(self, **body): } if self.reasoning: payload["choices"][0]["message"]["reasoning_content"] = "Check the requested calculation." + if self.reasoning_only: + payload["choices"][0]["message"]["content"] = None if self.refusal: payload["choices"][0]["message"].update(content=None, refusal="I cannot help with that.") if self.tool_call and turn == 1: @@ -97,9 +102,10 @@ async def create_chat_completion(self, **body): } ], ) - if "ng_capture" not in body: + capture_params = body.get("ng_capture") or (body.get("offload_params") or {}).get("ng_capture") + if capture_params is None: return payload - admission = CaptureAdmission.model_validate(body["ng_capture"]) + admission = CaptureAdmission.model_validate(capture_params) prefix = [token for key in admission.staging_chain for token in self.records[key].token_ids_delta] call = self.capture.begin_call(admission, prefix_token_ids=prefix, stream=body["stream"]) prompt = prefix + [turn * 10] @@ -116,12 +122,13 @@ async def create_chat_completion(self, **body): @pytest.fixture def make_harness(tmp_path, monkeypatch): - def make(dialect="responses", evaluation=False, reasoning=False, **overrides): + def make(dialect="responses", evaluation=False, reasoning=False, backend="vllm_worker", **overrides): root = tmp_path / dialect.replace("/", "-") global_config = { "token_id_capture": { "enabled": True, "external_staging": True, + "external_staging_backend": backend, "rebuild_response": False, "lineage_store": "nemo_gym.token_id_capture.lineage:FileLineageStore", "lineage_store_kwargs": {"root": str(root)}, @@ -264,6 +271,111 @@ async def check_send(message): assert events[-1]["type"] == "response.completed" +@pytest.mark.parametrize("backend", ["vllm_worker", "megatron_worker"]) +@pytest.mark.parametrize("stream", [False, True]) +@pytest.mark.parametrize( + "dialect,trigger,prior_call", + [(dialect, "overflow", prior) for dialect in DIALECTS for prior in (False, True)] + + [(dialect, "reasoning_guard", True) for dialect in ("responses", "compaction")], +) +async def test_synthetic_completion_leaves_call_uncommitted( + make_harness, monkeypatch, backend, stream, dialect, trigger, prior_call +): + h = make_harness(dialect, reasoning=True, backend=backend, sequential_reasoning_allowed=False) + h.worker.reasoning_only = trigger == "reasoning_guard" + worker_call = AsyncMock(wraps=h.worker.create_chat_completion) + monkeypatch.setattr(h.worker, "create_chat_completion", worker_call) + body = _body(dialect, stream=False) + if prior_call: + first_messages = await _request(h.app, _path(dialect), body) + assert first_messages[0]["status"] == 200 + first = json.loads(b"".join(message.get("body", b"") for message in first_messages)) + if dialect in ("responses", "compaction"): + body["input"].extend(first["output"]) + if trigger == "reasoning_guard": + assert [item["type"] for item in first["output"]] == ["reasoning"] + elif dialect == "chat/completions": + body["messages"].append( + {key: value for key, value in first["choices"][0]["message"].items() if value is not None} + ) + else: + body["messages"].append({"role": "assistant", "content": first["content"]}) + if trigger == "overflow": + body["input" if dialect in ("responses", "compaction") else "messages"].append( + {"role": "user", "content": "continue"} + ) + + if trigger == "overflow": + error = ClientResponseError(MagicMock(real_url="http://worker/v1/chat/completions"), (), status=400) + error.response_content = b'{"error":{"message":"maximum context length","code":400}}' + worker_call.side_effect = error + body["stream"] = stream + messages = await _request(h.app, _path(dialect), body) + assert messages[0]["status"] == 200, messages + assert worker_call.await_count == int(prior_call) + int(trigger == "overflow") + assert h.finalize.await_count == int(prior_call) + 1 + raw = b"".join(message.get("body", b"") for message in messages) + for internal in ("ng_commit_coords", "prompt_token_ids", "generation_token_ids", "generation_log_probs"): + assert internal.encode() not in raw + wire_dialect = {"chat/completions": "chat_completions", "compaction": "responses"}.get(dialect, dialect) + served = _reconstruct_streamed_response(raw, wire_dialect) if stream else json.loads(raw) + assert served is not None + + manifest = RolloutManifest.model_validate(await h.ledger.manifest("r1")) + assert len(manifest.records) == int(prior_call) + assert len(h.worker.records) == int(prior_call) + assert [failure.reason for failure in manifest.failures] == [UNCOMMITTED_CALL_REASON] + selection = select_terminal_call(manifest.records) + if prior_call: + record = manifest.records[0] + assert record.response_id == first["id"] + assert manifest.failures[0].model_call_id != record.model_call_id + assert selection.terminal_model_call_id == record.model_call_id + receipt = RolloutReceipt( + rollout_id="r1", + manifest=manifest.records, + terminal_model_call_id=selection.terminal_model_call_id, + terminal_selection="heuristic", + ) + row = verify_and_linearize(receipt, h.worker.fetch([record.staging_key])) + assert row.token_ids == [10, 11] + assert row.token_mask == [0.0, 1.0] + assert row.logprobs == [0.0, -0.25] + else: + assert selection.terminal_model_call_id is None + assert selection.reason == "no_records" + # A harness that explicitly selects the synthetic response cannot attribute + # it to an earlier generated turn. + attribution = resolve_terminal(manifest.records, served, declared_response_id=served["id"]) + assert not attribution.attributed + assert "declared_terminal_not_captured" in attribution.reason + + +@pytest.mark.parametrize("backend", ["vllm_worker", "megatron_worker"]) +@pytest.mark.parametrize("stream", [False, True]) +async def test_worker_completion_without_coordinates_still_fails_closed(make_harness, monkeypatch, backend, stream): + # A real worker completion that is missing its acknowledgement must still poison the call. + # It must not be treated like a synthetic completion, which is left uncommitted. + h = make_harness("responses", backend=backend) + original = h.worker.create_chat_completion + + async def drop_coords(**body): + payload = await original(**body) + payload.pop("ng_commit_coords") + return payload + + monkeypatch.setattr(h.worker, "create_chat_completion", drop_coords) + messages = await _request(h.app, _path("responses"), _body("responses", stream)) + assert messages[0]["status"] == 200 + assert h.finalize.await_count == 1 + manifest = RolloutManifest.model_validate(await h.ledger.manifest("r1")) + assert manifest.records == [] + assert [failure.reason for failure in manifest.failures] == [ + WORKER_MISSING_COMMIT_COORDS_REASON, + UNCOMMITTED_CALL_REASON, + ] + + @pytest.mark.parametrize("override", ["extra_body", "sampling_overrides"]) def test_static_streaming_override_rejected(make_harness, override): with pytest.raises(ValueError, match="non-streaming backend"): diff --git a/tests/unit_tests/test_external_capture_handlers.py b/tests/unit_tests/test_external_capture_handlers.py index 253b3715eb..d1b6018067 100644 --- a/tests/unit_tests/test_external_capture_handlers.py +++ b/tests/unit_tests/test_external_capture_handlers.py @@ -304,6 +304,37 @@ async def test_handler_finalization_updates_lineage_and_cleans_transport(handler _assert_poisoned(manifest, context, payload, case.expected_failure) +@pytest.mark.asyncio +@HANDLER_CLASSES +@pytest.mark.parametrize("previous_worker_response", [False, True]) +async def test_handler_leaves_synthetic_completion_uncommitted(handler_cls, previous_worker_response) -> None: + handler = handler_cls() + if previous_worker_response: + previous_context = _root_context(InMemoryLineageStore()) + payload = _transport_payload() + payload["ng_commit_coords"] = _staged_coords() + await _prepare_and_finalize(handler, previous_context, payload) + assert previous_context.committed + + store = InMemoryLineageStore() + context = _root_context(store) + token = set_token_sink(context) + try: + # Sending a request does not imply that a worker completion arrived: + # context-overflow errors are converted into synthetic completions. + handler.prepare_request({}) + await handler.finalize_response( + {"id": "synthetic", "choices": [{"message": {"role": "assistant", "content": None}}]} + ) + finally: + reset_token_sink(token) + + assert not context.committed + manifest = await store.manifest("rollout-1") + assert manifest["records"] == [] + assert manifest["failures"] == [] + + class _FaultyLedger(InMemoryLineageStore): """Ledger whose writes can be made to raise, to exercise the poison fallback paths."""