From aa800c065317bafd08fde4161274cfe9759583e9 Mon Sep 17 00:00:00 2001 From: Ali Roshan Ghias Date: Sat, 29 Aug 2026 08:49:40 +0200 Subject: [PATCH] fix(agent): unify context compaction trace contract Signed-off-by: Ali Roshan Ghias --- .../simple_agent_with_compaction/README.md | 11 ++++ .../simple_agent_with_compaction/app.py | 8 +-- .../compaction/__init__.py | 6 +- .../compaction/history.py | 1 - .../compaction/session.py | 34 +++-------- .../tests/test_app.py | 57 +++++++++++++------ 6 files changed, 64 insertions(+), 53 deletions(-) diff --git a/responses_api_agents/simple_agent_with_compaction/README.md b/responses_api_agents/simple_agent_with_compaction/README.md index 81fab9c773..04095dbb75 100644 --- a/responses_api_agents/simple_agent_with_compaction/README.md +++ b/responses_api_agents/simple_agent_with_compaction/README.md @@ -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 diff --git a/responses_api_agents/simple_agent_with_compaction/app.py b/responses_api_agents/simple_agent_with_compaction/app.py index 7ddd5f9984..5d7b750c03 100644 --- a/responses_api_agents/simple_agent_with_compaction/app.py +++ b/responses_api_agents/simple_agent_with_compaction/app.py @@ -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): @@ -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, @@ -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: diff --git a/responses_api_agents/simple_agent_with_compaction/compaction/__init__.py b/responses_api_agents/simple_agent_with_compaction/compaction/__init__.py index 65d8d850af..3c77ddddab 100644 --- a/responses_api_agents/simple_agent_with_compaction/compaction/__init__.py +++ b/responses_api_agents/simple_agent_with_compaction/compaction/__init__.py @@ -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, ) @@ -75,7 +74,6 @@ "CompactionScheduleConfig", "ContextCompactedResponse", "ContextCompactedTransportResponse", - "ContextCompactionContract", "ContextCompactionSession", "ContextGuardConfig", "ContextHistoryConfig", @@ -100,10 +98,10 @@ "ReasoningRecencyConfig", "RecencyHistoryPolicy", "RecencyHistoryPolicyConfig", + "RolloutTraceContract", "RewriteBoundaryEvent", "SemanticHistory", "TransformationLineageDeltaRecord", - "TransportContextCompactionContract", "TurnChunkedHistoryController", "UnitLineageRecord", "build_generation_contract", diff --git a/responses_api_agents/simple_agent_with_compaction/compaction/history.py b/responses_api_agents/simple_agent_with_compaction/compaction/history.py index 572c041c32..b2b0a13aab 100644 --- a/responses_api_agents/simple_agent_with_compaction/compaction/history.py +++ b/responses_api_agents/simple_agent_with_compaction/compaction/history.py @@ -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 diff --git a/responses_api_agents/simple_agent_with_compaction/compaction/session.py b/responses_api_agents/simple_agent_with_compaction/compaction/session.py index 95df6399ee..cddf88dcad 100644 --- a/responses_api_agents/simple_agent_with_compaction/compaction/session.py +++ b/responses_api_agents/simple_agent_with_compaction/compaction/session.py @@ -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 @@ -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) @@ -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={ @@ -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) @@ -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, ), diff --git a/responses_api_agents/simple_agent_with_compaction/tests/test_app.py b/responses_api_agents/simple_agent_with_compaction/tests/test_app.py index 76632db7f0..8cb1dae800 100644 --- a/responses_api_agents/simple_agent_with_compaction/tests/test_app.py +++ b/responses_api_agents/simple_agent_with_compaction/tests/test_app.py @@ -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, ) @@ -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( @@ -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( @@ -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: @@ -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"