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
13 changes: 12 additions & 1 deletion nemo_gym/base_responses_api_agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@
from collections.abc import Mapping
from functools import wraps
from typing import Any, Optional
from warnings import warn

from fastapi import Body, FastAPI, Request

Expand Down Expand Up @@ -43,7 +44,8 @@


class BaseResponsesAPIAgentConfig(BaseRunServerInstanceConfig):
pass
skip_verification: bool = False
skip_verification_reward: float = 0.0


class BaseResponsesAPIAgent(BaseServer):
Expand Down Expand Up @@ -145,6 +147,15 @@ async def run(self, body: BaseRunRequest = Body()) -> BaseVerifyResponse:

async def aggregate_metrics(self, body: AggregateMetricsRequest = Body()) -> AggregateMetrics:
"""Default: same RewardProfiler aggregation as resources server. Override to proxy."""
if self.config.skip_verification:
warn(
"Skipping aggregate metrics because skip_verification=True; "
"use disable_aggregation=True to avoid writing aggregate metric files.",
RuntimeWarning,
stacklevel=2,
)
return AggregateMetrics()

return compute_aggregate_metrics(
body.verify_responses,
compute_metrics_fn=self.compute_metrics,
Expand Down
17 changes: 17 additions & 0 deletions nemo_gym/global_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -96,6 +96,8 @@
OBSERVABILITY_ENABLED_KEY_NAME = "observability_enabled"
MODEL_CALL_CAPTURE_DIR_KEY_NAME = "model_call_capture_dir"
COMPONENT_NAME_KEY_NAME = "component_name"
SKIP_VERIFICATION_KEY_NAME = "skip_verification"
SKIP_VERIFICATION_REWARD_KEY_NAME = "skip_verification_reward"
NEMO_GYM_RESERVED_TOP_LEVEL_KEYS = [
CONFIG_PATHS_KEY_NAME,
ENTRYPOINT_KEY_NAME,
Expand Down Expand Up @@ -126,6 +128,8 @@
OBSERVABILITY_ENABLED_KEY_NAME,
MODEL_CALL_CAPTURE_DIR_KEY_NAME,
COMPONENT_NAME_KEY_NAME,
SKIP_VERIFICATION_KEY_NAME,
SKIP_VERIFICATION_REWARD_KEY_NAME,
]

# Data keys
Expand Down Expand Up @@ -418,6 +422,8 @@ def validate_and_populate_defaults(
port_range_low: int,
port_range_high: int,
initial_disallowed_ports: Optional[List[int]] = None,
skip_verification: Optional[bool] = None,
skip_verification_reward: Optional[float] = None,
probe_ports: bool = True,
) -> List[int]:
server_refs = [c.get_server_ref() for c in server_instance_configs]
Expand Down Expand Up @@ -467,6 +473,15 @@ def validate_and_populate_defaults(
# Port already exists, add it to the disallowed list.
disallowed_ports.append(run_server_config_dict["port"])

if server_instance_config.SERVER_TYPE == "responses_api_agents":
if skip_verification is not None and SKIP_VERIFICATION_KEY_NAME not in run_server_config_dict:
run_server_config_dict[SKIP_VERIFICATION_KEY_NAME] = skip_verification
if (
skip_verification_reward is not None
and SKIP_VERIFICATION_REWARD_KEY_NAME not in run_server_config_dict
):
run_server_config_dict[SKIP_VERIFICATION_REWARD_KEY_NAME] = skip_verification_reward

return disallowed_ports

def collect_missing_value_paths(self, config: DictConfig) -> List[str]:
Expand Down Expand Up @@ -775,6 +790,8 @@ def parse(self, parse_config: Optional[GlobalConfigDictParserConfig] = None) ->
initial_disallowed_ports=initial_disallowed_ports,
port_range_low=port_range_low,
port_range_high=port_range_high,
skip_verification=global_config_dict.get(SKIP_VERIFICATION_KEY_NAME),
skip_verification_reward=global_config_dict.get(SKIP_VERIFICATION_REWARD_KEY_NAME),
probe_ports=not parse_config.offline,
)

Expand Down
31 changes: 21 additions & 10 deletions responses_api_agents/simple_agent/app.py
Original file line number Diff line number Diff line change
Expand Up @@ -324,16 +324,24 @@ async def run(self, request: Request, body: SimpleAgentRunRequest) -> SimpleAgen
}
)

verify_request = SimpleAgentVerifyRequest.model_validate(body.model_dump() | {"response": model_response_json})

verify_response = await self.server_client.post(
server_name=self.config.resources_server.name,
url_path="/verify",
json=verify_request.model_dump(),
cookies=cookies,
)
await raise_for_status(verify_response)
result = await get_response_json(verify_response)
if self.config.skip_verification:
result = body.model_dump() | {
"response": model_response_json,
"reward": float(self.config.skip_verification_reward),
"verification_skipped": True,
}
else:
verify_request = SimpleAgentVerifyRequest.model_validate(
body.model_dump() | {"response": model_response_json}
)
verify_response = await self.server_client.post(
server_name=self.config.resources_server.name,
url_path="/verify",
json=verify_request.model_dump(),
cookies=cookies,
)
await raise_for_status(verify_response)
result = await get_response_json(verify_response)
if trajectory is not None:
resolved = result.get("resolved")
if isinstance(resolved, bool) and trajectory.turns:
Expand All @@ -345,6 +353,9 @@ async def run(self, request: Request, body: SimpleAgentRunRequest) -> SimpleAgen

async def aggregate_metrics(self, body: AggregateMetricsRequest = Body()) -> AggregateMetrics:
"""Proxy aggregate_metrics to the resources server."""
if self.config.skip_verification:
return await super().aggregate_metrics(body)

response = await self.server_client.post(
server_name=self.config.resources_server.name,
url_path="/aggregate_metrics",
Expand Down
62 changes: 62 additions & 0 deletions responses_api_agents/simple_agent/tests/test_app.py
Original file line number Diff line number Diff line change
Expand Up @@ -898,3 +898,65 @@ async def _test_incomplete_details_helper(self, monkeypatch: MonkeyPatch, incomp
"safety_identifier": None,
}
assert expected_responses_dict == actual_responses_dict

async def test_run_skip_verification_uses_configured_reward(self) -> None:
config = SimpleAgentConfig(
host="0.0.0.0",
port=8080,
entrypoint="",
name="simple_agent",
model_server=ModelServerRef(
type="responses_api_models",
name="my model server",
),
resources_server=ResourcesServerRef(
type="resources_servers",
name="my resources server",
),
skip_verification=True,
skip_verification_reward=0.25,
)
server = SimpleAgent(config=config, server_client=MagicMock(spec=ServerClient))
app = server.setup_webserver()
client = TestClient(app)

seed_response = AsyncMock()
seed_response.ok = True
seed_response.cookies = {"session": "seeded"}

model_response_payload = {
"id": "response_id",
"created_at": 1,
"model": "dummy_model",
"object": "response",
"output": [],
"parallel_tool_calls": True,
"tool_choice": "auto",
"tools": [],
}
model_response = AsyncMock()
model_response.ok = True
model_response.cookies = {"session": "model"}
model_response.read.return_value = json.dumps(model_response_payload).encode()

server.server_client.post.side_effect = [seed_response, model_response]

response = client.post(
"/run",
json={"responses_create_params": {"input": [{"role": "user", "content": "hello"}]}},
)

assert response.status_code == 200
response_json = response.json()
assert response_json["reward"] == 0.25
assert response_json["verification_skipped"] is True
assert response_json["response"]["id"] == "response_id"

post_call_kwargs = [post_call.kwargs for post_call in server.server_client.post.call_args_list]
assert [kwargs["url_path"] for kwargs in post_call_kwargs] == [
"/seed_session",
"/v1/responses",
]
assert post_call_kwargs[0]["server_name"] == "my resources server"
assert post_call_kwargs[1]["server_name"] == "simple_agent"
assert post_call_kwargs[1]["cookies"] == {"session": "seeded"}
33 changes: 20 additions & 13 deletions responses_api_agents/tool_simulation_agent/app.py
Original file line number Diff line number Diff line change
Expand Up @@ -73,19 +73,26 @@ async def run(self, body: ToolSimulationAgentRunRequest = Body()) -> ToolSimulat
)
await raise_for_status(response)

verify_request_body = body.model_dump()
verify_request_body["response"] = await response.json()
verify_request = ToolSimulationAgentVerifyRequest.model_validate(verify_request_body)

verify_response = await self.server_client.post(
server_name=config.resources_server.name,
url_path="/verify",
json=verify_request.model_dump(),
)
await raise_for_status(verify_response)
verify_response_json = await verify_response.json()

return ToolSimulationAgentVerifyResponse.model_validate(verify_response_json)
response_json = await response.json()
if config.skip_verification:
result = body.model_dump() | {
"response": response_json,
"reward": float(config.skip_verification_reward),
"verification_skipped": True,
}
else:
verify_request = ToolSimulationAgentVerifyRequest.model_validate(
body.model_dump() | {"response": response_json}
)
verify_response = await self.server_client.post(
server_name=config.resources_server.name,
url_path="/verify",
json=verify_request.model_dump(),
)
await raise_for_status(verify_response)
result = await verify_response.json()

return ToolSimulationAgentVerifyResponse.model_validate(result)


if __name__ == "__main__":
Expand Down
63 changes: 63 additions & 0 deletions responses_api_agents/tool_simulation_agent/tests/test_app.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import json
from typing import Any
from unittest.mock import AsyncMock, MagicMock, call

Expand Down Expand Up @@ -61,7 +62,9 @@ def _set_server_client_post_responses(
post_responses = []
for response in responses:
post_response_mock = AsyncMock()
post_response_mock.ok = True
post_response_mock.json.return_value = response
post_response_mock.read.return_value = json.dumps(response).encode()
post_responses.append(post_response_mock)

if additional_responses_present:
Expand Down Expand Up @@ -444,3 +447,63 @@ async def test_run(self, agent_config: ToolSimulationAgentConfig) -> None:
}
assert valid_verify_response.json() == expected_valid_verify_response_json
assert server_client_post_mock.call_args_list == expected_invalid_verify_response_calls

async def test_run_skip_verification_uses_configured_reward(self, agent_config: ToolSimulationAgentConfig) -> None:
server_client_post_mock = AsyncMock()
server_client_mock = MagicMock(spec=ServerClient)
server_client_mock.post = server_client_post_mock
agent_server = ToolSimulationAgent(
config=agent_config.model_copy(
update={
"skip_verification": True,
"skip_verification_reward": 0.5,
}
),
server_client=server_client_mock,
)
webserver = agent_server.setup_webserver()
test_client = TestClient(webserver)

response_object = {
"id": "chat_response_id",
"created_at": 1,
"model": "response_model",
"object": "response",
"output": [],
"parallel_tool_calls": False,
"tool_choice": "auto",
"tools": [],
}
self._set_server_client_post_responses(server_client_post_mock, response_object)

response = test_client.post(
"/run",
json={
"responses_create_params": {
"input": [
{
"role": "user",
"content": "Please answer directly.",
}
]
}
},
)

assert response.status_code == 200
response_json = response.json()
assert response_json["reward"] == 0.5
assert response_json["verification_skipped"] is True
assert response_json["response"]["id"] == "chat_response_id"
server_client_post_mock.assert_called_once_with(
server_name="tool_agent",
url_path="/v1/responses",
json=NeMoGymResponseCreateParamsNonStreaming(
input=[
NeMoGymEasyInputMessage(
role="user",
content="Please answer directly.",
)
],
),
)
29 changes: 29 additions & 0 deletions tests/unit_tests/test_base_responses_api_agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,9 @@
# limitations under the License.
from unittest.mock import MagicMock

import pytest

from nemo_gym.base_resources_server import AggregateMetricsRequest
from nemo_gym.base_responses_api_agent import (
BaseResponsesAPIAgent,
BaseResponsesAPIAgentConfig,
Expand All @@ -39,3 +42,29 @@ async def run(self, body=...):

agent = TestSimpleResponsesAPIAgent(config=config, server_client=MagicMock(spec=ServerClient))
agent.setup_webserver()

async def test_aggregate_metrics_skip_verification_warns_and_returns_empty_metrics(self) -> None:
config = BaseResponsesAPIAgentConfig(
host="",
port=0,
entrypoint="",
name="",
skip_verification=True,
)

class TestSimpleResponsesAPIAgent(SimpleResponsesAPIAgent):
async def responses(self, body=...):
raise NotImplementedError

async def run(self, body=...):
raise NotImplementedError

agent = TestSimpleResponsesAPIAgent(config=config, server_client=MagicMock(spec=ServerClient))
body = AggregateMetricsRequest(verify_responses=[])

with pytest.warns(RuntimeWarning, match="skip_verification=True"):
result = await agent.aggregate_metrics(body)

assert result.group_level_metrics == []
assert result.agent_metrics == {}
assert result.key_metrics == {}
Loading
Loading