Skip to content
Draft
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
13 changes: 13 additions & 0 deletions miles/rollout/generate_hub/agentic_tool_call.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
*,
Expand Down
28 changes: 28 additions & 0 deletions tests/fast/rollout/generate_hub/test_agentic_turn_budget.py
Original file line number Diff line number Diff line change
@@ -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
15 changes: 15 additions & 0 deletions tests/fast/rollout/generate_hub/test_multi_turn.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
Loading