Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
23 changes: 22 additions & 1 deletion nemo_gym/token_id_capture/external_capture.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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:
Expand All @@ -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
Comment thread
ananthsub marked this conversation as resolved.
try:
await self._finalize_admitted_response(
served_payload,
Expand Down
3 changes: 3 additions & 0 deletions nemo_gym/token_id_capture/sink.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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,
Expand All @@ -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())

Expand Down Expand Up @@ -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:
Expand All @@ -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]
Expand All @@ -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)},
Expand Down Expand Up @@ -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
Comment thread
ananthsub marked this conversation as resolved.


@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"):
Expand Down
31 changes: 31 additions & 0 deletions tests/unit_tests/test_external_capture_handlers.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""

Expand Down
Loading