diff --git a/miles/rollout/generate_hub/agentic_tool_call.py b/miles/rollout/generate_hub/agentic_tool_call.py index 46f89bccba6..bd498bba0bb 100644 --- a/miles/rollout/generate_hub/agentic_tool_call.py +++ b/miles/rollout/generate_hub/agentic_tool_call.py @@ -144,9 +144,22 @@ async def generate(input: GenerateFnInput) -> GenerateFnOutput: samples.metadata.update(session_metadata) else: samples[-1].metadata.update(session_metadata) + + _mark_limits_exceeded_truncated(samples, agent_metadata) + return GenerateFnOutput(samples=samples) +def _mark_limits_exceeded_truncated( + samples: Sample | list[Sample], + agent_metadata: dict[str, Any] | None, +) -> None: + """Record turn-budget exhaustion even when the final model turn completed.""" + if (agent_metadata or {}).get("exit_status") == "LimitsExceeded": + final_sample = samples if isinstance(samples, Sample) else samples[-1] + final_sample.status = Sample.Status.TRUNCATED + + def build_agent_function_kwargs( args: argparse.Namespace, *, diff --git a/tests/fast/rollout/generate_hub/test_agentic_turn_budget.py b/tests/fast/rollout/generate_hub/test_agentic_turn_budget.py new file mode 100644 index 00000000000..1ef219826d4 --- /dev/null +++ b/tests/fast/rollout/generate_hub/test_agentic_turn_budget.py @@ -0,0 +1,28 @@ +import pytest + +from miles.rollout.generate_hub.agentic_tool_call import _mark_limits_exceeded_truncated +from miles.utils.types import Sample + + +@pytest.mark.parametrize("as_list", [False, True], ids=["merged", "multi-sample"]) +def test_limits_exceeded_marks_only_final_sample_truncated(as_list: bool) -> None: + samples = [ + Sample(status=Sample.Status.COMPLETED), + Sample(status=Sample.Status.COMPLETED), + ] + value = samples if as_list else samples[-1] + + _mark_limits_exceeded_truncated(value, {"exit_status": "LimitsExceeded"}) + + assert samples[-1].status == Sample.Status.TRUNCATED + if as_list: + assert samples[0].status == Sample.Status.COMPLETED + + +@pytest.mark.parametrize("agent_metadata", [None, {}, {"exit_status": "Submitted"}]) +def test_other_exit_statuses_remain_completed(agent_metadata: dict | None) -> None: + sample = Sample(status=Sample.Status.COMPLETED) + + _mark_limits_exceeded_truncated(sample, agent_metadata) + + assert sample.status == Sample.Status.COMPLETED diff --git a/tests/fast/rollout/generate_hub/test_multi_turn.py b/tests/fast/rollout/generate_hub/test_multi_turn.py index 38145f93d2c..170101930cf 100644 --- a/tests/fast/rollout/generate_hub/test_multi_turn.py +++ b/tests/fast/rollout/generate_hub/test_multi_turn.py @@ -668,6 +668,21 @@ def test_agent_returns_none_metadata_unchanged(self, variant, generation_env): assert s.metadata.get("instance_id") == "test-123" assert "reward" not in s.metadata + @pytest.mark.parametrize( + "generation_env", + [{"args_kwargs": {"agentic_return_metadata": {"exit_status": "LimitsExceeded"}}}], + indirect=True, + ) + def test_limits_exceeded_marks_final_sample_truncated(self, variant, generation_env): + generation_env.mock_server.process_fn = TwoTurnStub.process_fn + + result = _run_generate(variant, generation_env, make_sample(prompt=TwoTurnStub.PROMPT)) + + samples = listify(result.sample) + assert all(sample.status == Sample.Status.COMPLETED for sample in samples[:-1]) + assert samples[-1].status == Sample.Status.TRUNCATED + assert samples[-1].metadata["exit_status"] == "LimitsExceeded" + def test_session_server_identity_forwarded_to_agent_metadata(self, variant, generation_env): from miles.utils.test_utils import mock_tools