Skip to content
Open
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
10 changes: 9 additions & 1 deletion nemo_gym/base_resources_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -101,6 +101,10 @@ class BaseVerifyResponse(BaseVerifyRequest):
# Machine-readable handling belongs to `mask_sample`/`failure_kind`.
failure_reason: Optional[str] = None

# The same environment-side handle reported by ``/seed_session``, returned with the
# score so a rollout record and an environment log line share a join key.
env_session_id: Optional[str] = None


class BaseMultiRewardVerifyResponse(BaseVerifyResponse):
"""Base verify response for environments with multiple reward objectives.
Expand All @@ -123,7 +127,11 @@ class BaseSeedSessionRequest(BaseModel):


class BaseSeedSessionResponse(BaseModel):
pass
# The environment-side handle for this session: a container id, a browser context id,
# a provider session id. Opaque to Gym, and optional - an environment that does not
# report one is unaffected. Returning it lets the training side join its own rollout
# records against environment- and provider-side logs without timestamp guessing.
env_session_id: Optional[str] = None


class MCPServerMetadata(BaseModel):
Expand Down
27 changes: 27 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,
BaseSeedSessionResponse,
BaseVerifyResponse,
ReverifyMode,
SimpleResourcesServer,
Expand Down Expand Up @@ -70,3 +71,29 @@ def test_sanity(self) -> None:

def test_reverify_mode(self) -> None:
assert asyncio.run(_resources_server().get_reverify_mode()) == ReverifyMode.UNKNOWN


class TestEnvironmentSessionId:
"""`env_session_id` gives a rollout record and an environment log a join key."""

def _params(self) -> NeMoGymResponseCreateParamsNonStreaming:
return NeMoGymResponseCreateParamsNonStreaming(input="hi")

def _response(self) -> NeMoGymResponse:
return NeMoGymResponse.model_construct(id="resp-1", output=[])

def test_absent_by_default_on_both_ends(self) -> None:
assert BaseSeedSessionResponse().env_session_id is None
verify = BaseVerifyResponse(responses_create_params=self._params(), response=self._response(), reward=1.0)
assert verify.env_session_id is None

def test_survives_the_seed_to_verify_round_trip(self) -> None:
seeded = BaseSeedSessionResponse(env_session_id="browser-ctx-9c41")
verify = BaseVerifyResponse(
responses_create_params=self._params(),
response=self._response(),
reward=0.0,
env_session_id=seeded.env_session_id,
)
assert seeded.model_dump()["env_session_id"] == "browser-ctx-9c41"
assert verify.model_dump()["env_session_id"] == "browser-ctx-9c41"