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
4 changes: 4 additions & 0 deletions nemo_gym/base_resources_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -96,6 +96,10 @@ class BaseVerifyRequest(BaseRunRequest):
class BaseVerifyResponse(BaseVerifyRequest):
reward: float

# Human-readable diagnosis of why `reward` may not reflect policy quality.
# Machine-readable handling belongs to `mask_sample`/`failure_kind`.
failure_reason: Optional[str] = None


class BaseMultiRewardVerifyResponse(BaseVerifyResponse):
"""Base verify response for environments with multiple reward objectives.
Expand Down
2 changes: 2 additions & 0 deletions resources_servers/math_with_judge/tests/test_app.py
Original file line number Diff line number Diff line change
Expand Up @@ -211,6 +211,7 @@ async def test_verify(self, config: LibraryJudgeMathResourcesServerConfig) -> No
assert sorted(list(not_equal_verify_response.model_dump())) == [
"expected_answer",
"extracted_answer",
"failure_reason",
"judge_evaluations",
"library_reward",
"response",
Expand Down Expand Up @@ -249,6 +250,7 @@ async def test_verify(self, config: LibraryJudgeMathResourcesServerConfig) -> No
assert sorted(list(equal_verify_response.model_dump())) == [
"expected_answer",
"extracted_answer",
"failure_reason",
"judge_evaluations",
"library_reward",
"response",
Expand Down
2 changes: 1 addition & 1 deletion responses_api_agents/tau2/tests/test_data.json

Large diffs are not rendered by default.

Original file line number Diff line number Diff line change
Expand Up @@ -444,6 +444,7 @@ async def test_run(self, agent_config: ToolSimulationAgentConfig) -> None:
},
"response": full_tool_call_response,
"reward": 1,
"failure_reason": None,
}
assert valid_verify_response.json() == expected_valid_verify_response_json
assert server_client_post_mock.call_args_list == expected_invalid_verify_response_calls
Expand Down
16 changes: 16 additions & 0 deletions tests/unit_tests/test_base_resources_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@
from nemo_gym.base_resources_server import (
BaseMultiRewardVerifyResponse,
BaseResourcesServerConfig,
BaseVerifyResponse,
ReverifyMode,
SimpleResourcesServer,
)
Expand All @@ -35,6 +36,21 @@ async def verify(self, body):
return TestSimpleResourcesServer(config=config, server_client=MagicMock(spec=ServerClient))


class TestBaseVerifyResponse:
def test_failure_reason_defaults_none_and_round_trips(self) -> None:
response = BaseVerifyResponse(
responses_create_params=NeMoGymResponseCreateParamsNonStreaming(input="hi"),
response=NeMoGymResponse.model_construct(id="resp-1", output=[]),
reward=0.0,
)
assert response.failure_reason is None
assert response.model_dump()["failure_reason"] is None

rescued = response.model_copy(update={"failure_reason": "judge response unparseable after 3 attempts"})
assert rescued.model_dump()["failure_reason"] == "judge response unparseable after 3 attempts"
assert rescued.reward == 0.0


class TestBaseMultiRewardVerifyResponse:
def test_reward_components_round_trip(self) -> None:
response = BaseMultiRewardVerifyResponse(
Expand Down
Loading