Skip to content
Closed
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
11 changes: 11 additions & 0 deletions responses_api_agents/simple_agent_with_compaction/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,17 @@ through a configured context-compaction policy before each model call.
Context compaction is opt-in through this agent. The existing `simple_agent`
remains unchanged.

## Rollout trace evidence

The `/run` response includes a `rollout_trace_contract`. Canonical prompt,
generation, and log-probability arrays remain in the ordinary response output.
One `model_call_metadata` record accompanies each trainable model call, and its
digest binds the metadata to those canonical arrays. Boundary events describe
intentional history rewrites so training consumers can construct physical
traces without retokenizing the rollout. The same bounded trace representation
is returned whether resource verification runs or is explicitly skipped;
verification status affects reward provenance, not trace encoding.

# Licensing information

Code: Apache 2.0
Expand Down
8 changes: 3 additions & 5 deletions responses_api_agents/simple_agent_with_compaction/app.py
Original file line number Diff line number Diff line change
Expand Up @@ -534,7 +534,7 @@ async def run(
)

context_compacted_response = None
if model_response_json.get("context_compaction_contract") is not None:
if model_response_json.get("rollout_trace_contract") is not None:
context_compacted_response = ContextCompactedResponse.model_validate(model_response_json)
original_input = body.responses_create_params.input
if isinstance(original_input, str):
Expand Down Expand Up @@ -588,10 +588,10 @@ async def run(
if context_compacted_response is None:
return verified

contract = context_compacted_response.context_compaction_contract
contract = context_compacted_response.rollout_trace_contract
context_compacted_response = context_compacted_response.model_copy(
update={
"context_compaction_contract": contract.model_copy(
"rollout_trace_contract": contract.model_copy(
update={
"group_id": body.context_compaction_group_id,
"task_id": body.context_compaction_task_id,
Expand All @@ -601,8 +601,6 @@ async def run(
)
}
)
if self.config.skip_verification:
return verified.model_copy(update={"response": context_compacted_response})
return verified.model_copy(update={"response": build_transport_response(context_compacted_response)})

async def aggregate_metrics(self, body: AggregateMetricsRequest = Body()) -> AggregateMetrics:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -60,11 +60,10 @@
from responses_api_agents.simple_agent_with_compaction.compaction.session import (
ContextCompactedResponse,
ContextCompactedTransportResponse,
ContextCompactionContract,
ContextCompactionSession,
ModelCallMetadata,
PreparedContextCompactionCall,
TransportContextCompactionContract,
RolloutTraceContract,
build_generation_contract,
build_transport_response,
)
Expand All @@ -75,7 +74,6 @@
"CompactionScheduleConfig",
"ContextCompactedResponse",
"ContextCompactedTransportResponse",
"ContextCompactionContract",
"ContextCompactionSession",
"ContextGuardConfig",
"ContextHistoryConfig",
Expand All @@ -100,10 +98,10 @@
"ReasoningRecencyConfig",
"RecencyHistoryPolicy",
"RecencyHistoryPolicyConfig",
"RolloutTraceContract",
"RewriteBoundaryEvent",
"SemanticHistory",
"TransformationLineageDeltaRecord",
"TransportContextCompactionContract",
"TurnChunkedHistoryController",
"UnitLineageRecord",
"build_generation_contract",
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -489,7 +489,6 @@ class GenerationContract(BaseModel):

model_config = ConfigDict(extra="forbid", frozen=True)

schema_version: Literal[1] = 1
model_contract_id: str
tokenizer_contract_id: str
template_contract_id: str
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -51,28 +51,11 @@
LOGGER = logging.getLogger(__name__)


class ContextCompactionContract(BaseModel):
"""Versioned marker that makes exact Gym evidence authoritative."""
class RolloutTraceContract(BaseModel):
"""Bind exact model-call trace evidence to one caller-owned rollout."""

model_config = ConfigDict(extra="forbid", frozen=True)

schema_version: Literal[2] = 2
mode: Literal["exact_trace_authority"] = "exact_trace_authority"
rollout_id: str
group_id: str | None = None
task_id: str | None = None
rollout_index: int | None = Field(default=None, ge=0)
attempt_index: int | None = Field(default=None, ge=0)
generation_contract: GenerationContract


class TransportContextCompactionContract(BaseModel):
"""Post-verification contract for the bounded Gym-to-NeMo-RL envelope."""

model_config = ConfigDict(extra="forbid", frozen=True)

schema_version: Literal[3] = 3
mode: Literal["exact_trace_authority"] = "exact_trace_authority"
rollout_id: str
group_id: str | None = None
task_id: str | None = None
Expand Down Expand Up @@ -139,11 +122,11 @@ class ContextCompactedResponse(NeMoGymResponse):
chunk_records: list[FinalizedChunkRecord] = Field(default_factory=list)
boundary_events: list[RewriteBoundaryEvent] = Field(default_factory=list)
guard_records: list[GuardOutcomeRecord] = Field(default_factory=list)
context_compaction_contract: ContextCompactionContract
rollout_trace_contract: RolloutTraceContract


class ContextCompactedTransportResponse(NeMoGymResponse):
"""Bounded exact-evidence response returned after resource verification."""
"""Bounded exact-evidence response returned by the agent's run endpoint."""

media_assets: dict[str, dict[str, Any]] = Field(default_factory=dict)
model_call_metadata: list[ModelCallMetadata] = Field(default_factory=list)
Expand All @@ -152,13 +135,13 @@ class ContextCompactedTransportResponse(NeMoGymResponse):
chunk_records: list[FinalizedChunkRecord] = Field(default_factory=list)
boundary_events: list[RewriteBoundaryEvent] = Field(default_factory=list)
guard_records: list[GuardOutcomeRecord] = Field(default_factory=list)
context_compaction_contract: TransportContextCompactionContract
rollout_trace_contract: RolloutTraceContract


def build_transport_response(
response: ContextCompactedResponse,
) -> ContextCompactedTransportResponse:
"""Drop validation-only duplication after the resource verifier succeeds."""
"""Drop validation-only duplication before returning an agent response."""

result = response.model_dump(
exclude={
Expand All @@ -182,9 +165,6 @@ def build_transport_response(
result["model_call_metadata"] = [
ModelCallMetadata.from_observed(evidence) for evidence in response.completion_evidence
]
result["context_compaction_contract"] = TransportContextCompactionContract.model_validate(
response.context_compaction_contract.model_dump() | {"schema_version": 3}
)
return ContextCompactedTransportResponse.model_validate(result)


Expand Down Expand Up @@ -581,7 +561,7 @@ def build_response(
),
"boundary_events": list(self.history_controller.boundary_events),
"guard_records": self.guard_records,
"context_compaction_contract": ContextCompactionContract(
"rollout_trace_contract": RolloutTraceContract(
rollout_id=self.rollout_id,
generation_contract=self.generation_contract,
),
Expand Down
57 changes: 41 additions & 16 deletions responses_api_agents/simple_agent_with_compaction/tests/test_app.py
Original file line number Diff line number Diff line change
Expand Up @@ -44,9 +44,11 @@
SimpleAgentWithCompactionRunRequest,
)
from responses_api_agents.simple_agent_with_compaction.compaction import (
ContextCompactedTransportResponse,
ContextCompactionSession,
ContextHistoryConfig,
build_generation_contract,
canonical_digest,
normalize_semantic_items,
)

Expand Down Expand Up @@ -333,7 +335,14 @@ async def test_identity_history_preserves_requests(self) -> None:
assert second_input[1]["type"] == "reasoning"
assert second_input[1]["summary"] == [{"text": "thinking", "type": "summary_text"}]
assert "prompt_token_ids" not in second_input[1]
assert response.json()["context_compaction_contract"]["mode"] == "exact_trace_authority"
assert set(response.json()["rollout_trace_contract"]) == {
"rollout_id",
"group_id",
"task_id",
"rollout_index",
"attempt_index",
"generation_contract",
}

async def test_responses_prefers_caller_owned_rollout_id_cookie(self) -> None:
config = SimpleAgentWithCompactionConfig(
Expand Down Expand Up @@ -381,7 +390,7 @@ async def test_responses_prefers_caller_owned_rollout_id_cookie(self) -> None:
)

assert response.status_code == 200
assert response.json()["context_compaction_contract"]["rollout_id"] == "caller-rollout"
assert response.json()["rollout_trace_contract"]["rollout_id"] == "caller-rollout"

async def test_authority_identity_history_tracks_model_and_tool_outputs(self) -> None:
config = SimpleAgentWithCompactionConfig(
Expand Down Expand Up @@ -624,7 +633,7 @@ def content(item):
assert len(payload["boundary_events"]) == 1

@pytest.mark.parametrize("skip_verification", [False, True])
async def test_run_preserves_authority_contract_with_optional_resource_verification(
async def test_run_returns_one_trace_envelope_with_optional_resource_verification(
self,
skip_verification: bool,
) -> None:
Expand Down Expand Up @@ -713,23 +722,39 @@ async def test_run_preserves_authority_contract_with_optional_resource_verificat

result = await server.run(request, request_body)

assert result.response.context_compaction_contract.rollout_id == "rollout-run"
assert result.response.context_compaction_contract.group_id == "group-run"
assert result.response.context_compaction_contract.task_id == "task-run"
assert result.response.context_compaction_contract.rollout_index == 2
assert result.response.context_compaction_contract.attempt_index == 1
expected_schema_version = 2 if skip_verification else 3
assert result.response.context_compaction_contract.schema_version == expected_schema_version
assert isinstance(result.response, ContextCompactedTransportResponse)
contract_payload = result.response.rollout_trace_contract.model_dump(mode="json")
assert not ({"schema_version", "mode", "format"} & contract_payload.keys())
assert "schema_version" not in contract_payload["generation_contract"]
assert result.response.rollout_trace_contract.rollout_id == "rollout-run"
assert result.response.rollout_trace_contract.group_id == "group-run"
assert result.response.rollout_trace_contract.task_id == "task-run"
assert result.response.rollout_trace_contract.rollout_index == 2
assert result.response.rollout_trace_contract.attempt_index == 1
assert len(result.response.model_call_metadata) == 1
assert result.response.model_call_metadata[0].generation_evidence_digest == canonical_digest(
{
"prompt_token_ids": [10],
"sampled_token_ids": [11],
"sampled_logprobs": [-0.1],
}
)
assert not hasattr(result.response, "completion_evidence")
assert not hasattr(result.response, "agent_input")
assert not hasattr(result.response, "seed_obs")
if skip_verification:
assert result.reward == 0.25
assert result.verification_skipped is True
assert len(result.response.completion_evidence) == 1
assert result.response.agent_input
assert server.server_client.post.call_count == 2
else:
assert len(result.response.model_call_metadata) == 1
assert not hasattr(result.response, "completion_evidence")
assert not hasattr(result.response, "agent_input")
assert not hasattr(result.response, "seed_obs")
verifier_payload = server.server_client.post.call_args_list[2].kwargs["json"]["response"]
assert len(verifier_payload["completion_evidence"]) == 1
observed = verifier_payload["completion_evidence"][0]
assert observed["prompt_token_ids"] == [10]
assert observed["sampled_token_ids"] == [11]
assert observed["sampled_logprobs"] == [-0.1]
assert verifier_payload["agent_input"]
assert "model_call_metadata" not in verifier_payload
inner_responses_call = server.server_client.post.call_args_list[1]
assert inner_responses_call.kwargs["cookies"][_CONTEXT_COMPACTION_ROLLOUT_ID_COOKIE] == "rollout-run"

Expand Down
Loading