diff --git a/3rdparty/Gym-workspace/Gym b/3rdparty/Gym-workspace/Gym index c3bac96314a..cc8ec02281a 160000 --- a/3rdparty/Gym-workspace/Gym +++ b/3rdparty/Gym-workspace/Gym @@ -1 +1 @@ -Subproject commit c3bac96314a59f28b896f597eb9845d175bb0252 +Subproject commit cc8ec02281a590ff9e5111acae8d70fe1dd354ab diff --git a/nemo_rl/environments/nemo_gym.py b/nemo_rl/environments/nemo_gym.py index e702f462042..5dd0d410b6c 100644 --- a/nemo_rl/environments/nemo_gym.py +++ b/nemo_rl/environments/nemo_gym.py @@ -47,6 +47,7 @@ RolloutDataFailure, http_status_is_infra, ) +from nemo_rl.models.generation.dynamo.token_wrapper import DYNAMO_SESSION_ID_HEADER from nemo_rl.models.generation.interfaces import should_use_async_rollouts from nemo_rl.models.policy import PolicyConfig, TokenizerConfig from nemo_rl.utils.routed_experts_codec import decode_routed_experts @@ -1120,6 +1121,15 @@ def setup_nemo_gym_config(config, tokenizer) -> None: generation_config["stop_strings"] = None generation_config["stop_token_ids"] = None + if generation_config["backend"] == "dynamo": + model_config = ( + config.env.setdefault("nemo_gym", {}) + .setdefault("policy_model", {}) + .setdefault("responses_api_models", {}) + .setdefault("vllm_model", {}) + ) + model_config["session_id_header"] = DYNAMO_SESSION_ID_HEADER + # For VLM runs, plumb the tokenizer config into the gym env config so the # NemoGym actor can reconstruct the processor inside itself (needed for # multi-turn multimodal postprocessing). diff --git a/nemo_rl/experience/rollout_manager.py b/nemo_rl/experience/rollout_manager.py index 6717542d642..ffd57d8eab6 100644 --- a/nemo_rl/experience/rollout_manager.py +++ b/nemo_rl/experience/rollout_manager.py @@ -45,8 +45,10 @@ from nemo_rl.experience.metric_utils import calculate_single_metric, pct from nemo_rl.experience.rollouts import ( EffortLevelsConfig, + _add_dynamo_session_id, _apply_effort_shaping, _attach_routed_experts_to_message_log_prefix, + _create_dynamo_session_id, _dummy_routed_experts_for_tokens, _effort_shaping_metrics, _find_routed_experts_template, @@ -439,6 +441,7 @@ async def _run_single_rollout( current_extra_env_info = copy.deepcopy(input_sample["extra_env_info"]) current_stop_strings = input_sample.get("stop_strings", None) task_name = input_sample["task_name"] + session_id = _create_dynamo_session_id(self._policy_generation) total_reward = 0.0 turn_count = 0 @@ -476,6 +479,7 @@ async def _run_single_rollout( ) = await self._generate_response( current_message_log, current_stop_strings, + session_id=session_id, ) except Exception as e: raise _classify_generation_failure( @@ -602,6 +606,8 @@ async def _generate_response( self, message_log: list[dict], stop_strings: list[str] | None, + *, + session_id: str | None = None, ) -> tuple[dict, torch.Tensor, dict[str, Any]]: """Generate a single-turn response for one sample. @@ -618,6 +624,11 @@ async def _generate_response( "stop_strings": [stop_strings], } ) + _add_dynamo_session_id( + generation_input_data, + self._policy_generation, + session_id, + ) # Generate response # TODO: update generate_async to return a single item directly diff --git a/nemo_rl/experience/rollouts.py b/nemo_rl/experience/rollouts.py index 851ac4a660e..b87a9a67836 100644 --- a/nemo_rl/experience/rollouts.py +++ b/nemo_rl/experience/rollouts.py @@ -24,6 +24,7 @@ from collections.abc import AsyncGenerator, Mapping, Sequence from dataclasses import dataclass from typing import Any, Optional +from uuid import uuid4 import ray import torch @@ -77,6 +78,25 @@ TokenizerType = PreTrainedTokenizerBase +def _create_dynamo_session_id( + policy_generation: GenerationInterface, +) -> str | None: + generation_config = getattr(policy_generation, "cfg", {}) + if generation_config.get("backend") == "dynamo": + return str(uuid4()) + return None + + +def _add_dynamo_session_id( + generation_input_data: BatchedDataDict[GenerationDatumSpec], + policy_generation: GenerationInterface, + session_id: str | None, +) -> None: + generation_config = getattr(policy_generation, "cfg", {}) + if session_id is not None and generation_config.get("backend") == "dynamo": + generation_input_data["session_ids"] = [session_id] + + def attach_initial_nemo_gym_image_payloads( batch: BatchedDataDict[DatumSpec], processor: Any, @@ -1127,6 +1147,7 @@ async def async_generate_response_for_sample_turn( max_seq_len: int, greedy: bool = False, *, + session_id: str | None = None, sample_multimodal_data: dict[str, Any] | None = None, deduplicate_multimodal_data: bool = False, ) -> tuple[list[dict], torch.Tensor, torch.Tensor, dict[str, float]]: @@ -1139,6 +1160,8 @@ async def async_generate_response_for_sample_turn( tokenizer: Tokenizer to use max_seq_len: Maximum sequence length greedy: Whether to use greedy decoding + session_id: Stable Dynamo session ID reused across every turn of one + trajectory attempt. Ignored unless the generation backend is Dynamo. sample_multimodal_data: Native vLLM media fields for this sample. deduplicate_multimodal_data: Avoid sending both native and policy-ready media through the async generation boundary. @@ -1165,6 +1188,11 @@ async def async_generate_response_for_sample_turn( "stop_strings": [sample_stop_strings], } ) + _add_dynamo_session_id( + generation_input_data, + policy_generation, + session_id, + ) # Create a dummy batch for generate_responses_async dummy_batch = BatchedDataDict[DatumSpec]( @@ -1237,6 +1265,7 @@ async def run_sample_multi_turn_rollout( current_extra_env_info = copy.deepcopy(initial_sample_state["extra_env_info"]) current_stop_strings = initial_sample_state.get("stop_strings", None) task_name = initial_sample_state["task_name"] + session_id = _create_dynamo_session_id(policy_generation) sample_multimodal_data = { key: initial_sample_state[key] for key in NATIVE_MULTIMODAL_KEYS @@ -1287,6 +1316,7 @@ async def run_sample_multi_turn_rollout( tokenizer, max_seq_len, greedy=greedy, + session_id=session_id, sample_multimodal_data=turn_multimodal_data, deduplicate_multimodal_data=deduplicate_multimodal_data, ) diff --git a/nemo_rl/models/generation/dynamo/dynamo_generation.py b/nemo_rl/models/generation/dynamo/dynamo_generation.py index 1c488855c7d..61f2c37cfaf 100644 --- a/nemo_rl/models/generation/dynamo/dynamo_generation.py +++ b/nemo_rl/models/generation/dynamo/dynamo_generation.py @@ -30,7 +30,10 @@ from nemo_rl.models.generation.dynamo.managed_runtime import ManagedDynamoRuntime from nemo_rl.models.generation.dynamo.metrics import DynamoMetricsSampler from nemo_rl.models.generation.dynamo.refit import DynamoRefitChannel -from nemo_rl.models.generation.dynamo.token_wrapper import DynamoTokenWrapperServer +from nemo_rl.models.generation.dynamo.token_wrapper import ( + DYNAMO_SESSION_ID_HEADER, + DynamoTokenWrapperServer, +) from nemo_rl.models.generation.interfaces import ( CollectiveSenderSpec, GenerationDatumSpec, @@ -441,6 +444,7 @@ async def _post_completion_request( greedy: bool, stop_strings: Optional[list[str]], max_new_tokens: int, + session_id: Optional[str] = None, ) -> tuple[list[int], list[float], bool]: request_url = self._completion_url() payload = self._build_completion_request( @@ -449,12 +453,16 @@ async def _post_completion_request( stop_strings=stop_strings, max_new_tokens=max_new_tokens, ) + request_headers = ( + {DYNAMO_SESSION_ID_HEADER: session_id} if session_id is not None else None + ) response: dict[str, Any] = {} for attempt in range(1, _HTTP_MAX_ATTEMPTS + 1): response = await async_http_post_json( request_url, payload, self._request_timeout_s(), + headers=request_headers, ) if not _is_retryable_http_response(response): break @@ -569,6 +577,16 @@ async def generate_async( "outside this method." ) sample_idx = 0 + session_ids = data.get("session_ids") + session_id = None + if session_ids is not None: + if len(session_ids) != batch_size: + raise ValueError( + "Dynamo session_ids must contain one value for each input sample." + ) + session_id = session_ids[sample_idx] + if not isinstance(session_id, str) or not session_id.strip(): + raise ValueError("Dynamo session IDs must be non-empty strings.") input_length = int(input_lengths_batch[sample_idx].item()) batch_stop_strings = data.get("stop_strings", [[] for _ in range(batch_size)]) per_sample_stop_strings = None @@ -589,6 +607,7 @@ async def generate_async( greedy=greedy, stop_strings=final_stop_strings, max_new_tokens=allowed_new_tokens, + session_id=session_id, ) yield ( diff --git a/nemo_rl/models/generation/dynamo/http_client.py b/nemo_rl/models/generation/dynamo/http_client.py index 8f69bee6b4d..853cced1a88 100644 --- a/nemo_rl/models/generation/dynamo/http_client.py +++ b/nemo_rl/models/generation/dynamo/http_client.py @@ -17,6 +17,7 @@ import json import urllib.error import urllib.request +from collections.abc import Mapping from typing import Any import aiohttp @@ -63,13 +64,17 @@ def http_post_json( async def async_http_post_json( - url: str, payload: dict[str, Any], timeout_s: float + url: str, + payload: dict[str, Any], + timeout_s: float, + *, + headers: Mapping[str, str] | None = None, ) -> dict[str, Any]: """POST JSON without blocking the rollout actor event loop.""" timeout = aiohttp.ClientTimeout(total=timeout_s) try: async with aiohttp.ClientSession(timeout=timeout) as session: - async with session.post(url, json=payload) as response: + async with session.post(url, json=payload, headers=headers) as response: body = await response.read() if response.status >= 400: return { diff --git a/nemo_rl/models/generation/dynamo/token_wrapper.py b/nemo_rl/models/generation/dynamo/token_wrapper.py index eab81fafa35..b62282f4b83 100644 --- a/nemo_rl/models/generation/dynamo/token_wrapper.py +++ b/nemo_rl/models/generation/dynamo/token_wrapper.py @@ -35,6 +35,7 @@ "generation_log_probs", ) _TOOL_ARGUMENT_MAPPING_ERROR = "Can only get item pairs from a mapping." +DYNAMO_SESSION_ID_HEADER = "X-Dynamo-Session-ID" def _coerce_token_id_list(value: Any, field_name: str) -> list[int]: @@ -474,6 +475,7 @@ async def chat_completions(request: Request) -> JSONResponse: status_code, response_body = await self._forward_chat_completion( prepared_body, authorization=request.headers.get("authorization"), + session_id=request.headers.get(DYNAMO_SESSION_ID_HEADER), ) if 200 <= status_code < 300: try: @@ -512,6 +514,7 @@ async def _forward_chat_completion( request_body: dict[str, Any], *, authorization: Optional[str], + session_id: Optional[str] = None, ) -> tuple[int, dict[str, Any]]: import aiohttp @@ -519,6 +522,8 @@ async def _forward_chat_completion( headers = {"Content-Type": "application/json"} if authorization: headers["Authorization"] = authorization + if session_id: + headers[DYNAMO_SESSION_ID_HEADER] = session_id session = self._client_session if session is None: diff --git a/nemo_rl/models/generation/interfaces.py b/nemo_rl/models/generation/interfaces.py index 8691e735321..6cd0479f9eb 100644 --- a/nemo_rl/models/generation/interfaces.py +++ b/nemo_rl/models/generation/interfaces.py @@ -294,6 +294,7 @@ class GenerationDatumSpec(TypedDict): - input_ids: Tensor of token IDs representing the input sequences (right padded) - input_lengths: Tensor containing the actual length of each sequence (without padding) - stop_strings: Optional list of strings to stop generation (per sample) + - session_ids: Optional per-sample stable session IDs; honored only by the Dynamo backend - __extra__: Additional model-specific data fields Example of a batch with 4 entries with different sequence lengths: @@ -319,6 +320,7 @@ class GenerationDatumSpec(TypedDict): input_ids: torch.Tensor input_lengths: torch.Tensor stop_strings: NotRequired[list[str]] + session_ids: NotRequired[list[str]] __extra__: Any diff --git a/tests/unit/environments/test_nemo_gym.py b/tests/unit/environments/test_nemo_gym.py index 8f5a8e510ec..e76285467a1 100644 --- a/tests/unit/environments/test_nemo_gym.py +++ b/tests/unit/environments/test_nemo_gym.py @@ -60,6 +60,7 @@ _reattach_original_multimodal_payloads, attach_static_multimodal_payload, ) +from nemo_rl.models.generation.dynamo.token_wrapper import DYNAMO_SESSION_ID_HEADER from nemo_rl.models.generation.vllm import VllmGeneration # cluster and tokenizer are fixture imports @@ -953,6 +954,45 @@ def test_validate_reward_components_match_scalar(): ) +def test_setup_nemo_gym_config_merges_session_header_for_dynamo() -> None: + config = SimpleNamespace( + policy={"generation": {"backend": "dynamo", "vllm_cfg": {}}}, + env={ + "nemo_gym": { + "num_gpu_nodes": 2, + "policy_model": { + "responses_api_models": { + "vllm_model": {"model_name": "keep-me"}, + "other_model": {"model_name": "untouched"}, + } + }, + } + }, + ) + + setup_nemo_gym_config(config, tokenizer=object()) + + nemo_gym = config.env["nemo_gym"] + responses_api_models = nemo_gym["policy_model"]["responses_api_models"] + assert responses_api_models["vllm_model"] == { + "model_name": "keep-me", + "session_id_header": DYNAMO_SESSION_ID_HEADER, + } + assert responses_api_models["other_model"] == {"model_name": "untouched"} + assert nemo_gym["num_gpu_nodes"] == 2 + + +def test_setup_nemo_gym_config_does_not_set_session_header_for_vllm() -> None: + config = SimpleNamespace( + policy={"generation": {"backend": "vllm", "vllm_cfg": {}}}, + env={}, + ) + + setup_nemo_gym_config(config, tokenizer=object()) + + assert config.env == {} + + @pytest.mark.nemo_gym def test_nemo_gym_stub_module(): from nemo_gym import config_types diff --git a/tests/unit/experience/test_rollout_manager_router_replay.py b/tests/unit/experience/test_rollout_manager_router_replay.py index 6f5390364bb..a6f524654c5 100644 --- a/tests/unit/experience/test_rollout_manager_router_replay.py +++ b/tests/unit/experience/test_rollout_manager_router_replay.py @@ -48,12 +48,20 @@ def __call__(self, text: str, **kwargs) -> SimpleNamespace: class _FakeGeneration: - def __init__(self, outputs: BatchedDataDict | list[BatchedDataDict]) -> None: + def __init__( + self, + outputs: BatchedDataDict | list[BatchedDataDict], + *, + backend: str | None = None, + ) -> None: self._outputs = outputs if isinstance(outputs, list) else [outputs] self._next_output = 0 + self.calls = [] + if backend is not None: + self.cfg = {"backend": backend} async def generate_async(self, data: BatchedDataDict): - del data + self.calls.append(data) output = self._outputs[self._next_output] self._next_output += 1 yield 0, output @@ -63,6 +71,7 @@ def _rollout_impl( output: BatchedDataDict | list[BatchedDataDict], *, max_rollout_turns: int = 1, + backend: str | None = None, ) -> AsyncRolloutImpl: return AsyncRolloutImpl( tokenizer=_FakeTokenizer(), # type: ignore[arg-type] @@ -70,7 +79,7 @@ def _rollout_impl( num_generations_per_prompt=1, max_seq_len=32, max_rollout_turns=max_rollout_turns, - policy_generation=_FakeGeneration(output), # type: ignore[arg-type] + policy_generation=_FakeGeneration(output, backend=backend), # type: ignore[arg-type] ) @@ -211,3 +220,65 @@ def test_second_turn_overwrites_prefix_fallback_routes() -> None: torch.cat((_routes(1, start=107), _fallback_routes(1))), ) assert torch.equal(final_env["routed_experts"], _fallback_routes(2)) + + +def test_dynamo_session_id_is_stable_across_turns_and_distinct_for_siblings() -> None: + second_turn_output = _generation_output( + (10, 11, 12, 20, 21, 31, 32, 40, 41), + route_start=100, + ) + impl = _rollout_impl( + [ + _generation_output(), + second_turn_output, + _generation_output(), + ], + max_rollout_turns=2, + backend="dynamo", + ) + input_sample = { + "idx": 0, + "message_log": [ + { + "role": "user", + "content": "prompt", + "token_ids": torch.tensor([10, 11, 12]), + } + ], + "extra_env_info": None, + "task_name": "test", + } + + def env_output(*, terminated: bool) -> EnvironmentReturn: + return EnvironmentReturn( + observations=[{"role": "user", "content": "environment"}], + metadata=[None], + next_stop_strings=[None], + rewards=torch.tensor([0.0]), + terminateds=torch.tensor([terminated]), + answers=[None], + ) + + with ( + patch( + "nemo_rl.experience.rollout_manager.calculate_rewards", + side_effect=[ + env_output(terminated=False), + env_output(terminated=True), + env_output(terminated=True), + ], + ), + patch( + "nemo_rl.experience.rollouts.uuid4", + side_effect=["attempt-a", "sibling-b"], + ), + ): + asyncio.run(impl._run_single_rollout(input_sample, traj_idx=0)) + asyncio.run(impl._run_single_rollout(input_sample, traj_idx=1)) + + generation = impl._policy_generation + assert [call["session_ids"][0] for call in generation.calls] == [ + "attempt-a", + "attempt-a", + "sibling-b", + ] diff --git a/tests/unit/experience/test_rollouts.py b/tests/unit/experience/test_rollouts.py index 4a810da668b..864c7071926 100644 --- a/tests/unit/experience/test_rollouts.py +++ b/tests/unit/experience/test_rollouts.py @@ -18,6 +18,7 @@ import tempfile from copy import deepcopy from dataclasses import asdict +from uuid import UUID import pytest import ray @@ -52,6 +53,7 @@ ) from nemo_rl.experience.rollouts import ( _add_multimodal_generation_payload, + _create_dynamo_session_id, _reattach_original_multimodal_payloads, async_generate_response_for_sample_turn, generate_responses_async, @@ -773,6 +775,7 @@ async def fake_generate( max_seq_len, greedy=False, *, + session_id=None, sample_multimodal_data=None, deduplicate_multimodal_data=False, ): @@ -858,6 +861,49 @@ def fake_rewards(batch, task_to_env): class _DummyDynamoGeneration(_DummySGLangGeneration): def __init__(self): self.cfg = {"backend": "dynamo"} + self.generation_input = None + + async def generate_async(self, data, greedy=False): + self.generation_input = data + async for item in super().generate_async(data, greedy=greedy): + yield item + + +def test_create_dynamo_session_id_mints_a_real_uuid_string() -> None: + session_id = _create_dynamo_session_id(_DummyDynamoGeneration()) + + assert isinstance(session_id, str) + assert UUID(session_id).version == 4 + assert session_id != _create_dynamo_session_id(_DummyDynamoGeneration()) + assert _create_dynamo_session_id(_CapturingAsyncVllmGeneration()) is None + assert _create_dynamo_session_id(object()) is None + + +def test_direct_session_id_is_added_only_for_dynamo() -> None: + message_log = [ + { + "role": "user", + "content": "prompt", + "token_ids": torch.tensor([1]), + } + ] + dynamo_generation = _DummyDynamoGeneration() + vllm_generation = _CapturingAsyncVllmGeneration() + + for generation in (dynamo_generation, vllm_generation): + asyncio.run( + async_generate_response_for_sample_turn( + generation, + message_log, + None, + _DummyTokenizer(), + max_seq_len=32, + session_id="trajectory-session", + ) + ) + + assert dynamo_generation.generation_input["session_ids"] == ["trajectory-session"] + assert "session_ids" not in vllm_generation.generation_input def test_generate_responses_async_requires_sglang_opt_in(): diff --git a/tests/unit/models/generation/test_dynamo_generation.py b/tests/unit/models/generation/test_dynamo_generation.py index 7a53e212f89..aa718e9f8b2 100644 --- a/tests/unit/models/generation/test_dynamo_generation.py +++ b/tests/unit/models/generation/test_dynamo_generation.py @@ -128,14 +128,17 @@ def shutdown(self): monkeypatch.setattr(generation_module, "ManagedDynamoRuntime", FakeRuntime) -def _data() -> BatchedDataDict: - return BatchedDataDict( +def _data(*, session_id: str | None = None) -> BatchedDataDict: + data = BatchedDataDict( { "input_ids": torch.tensor([[1, 2, 3, 0]], dtype=torch.long), "input_lengths": torch.tensor([3], dtype=torch.long), "stop_strings": [["stop"]], } ) + if session_id is not None: + data["session_ids"] = [session_id] + return data def _completion_response(token_ids: list[int]) -> dict[str, Any]: @@ -181,8 +184,8 @@ def test_blocking_generate_is_rejected_and_async_generation_uses_http( _patch_runtime(monkeypatch) requests = [] - async def fake_post(url, payload, timeout_s): - requests.append((url, payload, timeout_s)) + async def fake_post(url, payload, timeout_s, **kwargs): + requests.append((url, payload, timeout_s, kwargs)) return _completion_response([8, 9]) monkeypatch.setattr(generation_module, "async_http_post_json", fake_post) @@ -191,7 +194,12 @@ async def fake_post(url, payload, timeout_s): generation.generate(_data()) async def collect(): - return [item async for item in generation.generate_async(_data())] + return [ + item + async for item in generation.generate_async( + _data(session_id="trajectory-session") + ) + ] outputs = asyncio.run(collect()) assert outputs[0][0] == 0 @@ -200,6 +208,32 @@ async def collect(): assert requests[0][1]["max_tokens"] == 2 assert requests[0][1]["stop"] == ["stop"] assert "return_tokens_as_token_ids" not in requests[0][1] + assert requests[0][3]["headers"] == {"X-Dynamo-Session-ID": "trajectory-session"} + + +@pytest.mark.parametrize( + ("session_ids", "expected_message"), + [ + ([], "one value for each input sample"), + (["a", "b"], "one value for each input sample"), + ([""], "must be non-empty strings"), + ([" "], "must be non-empty strings"), + ([None], "must be non-empty strings"), + ], +) +def test_async_generation_rejects_invalid_session_ids( + monkeypatch, session_ids, expected_message +) -> None: + _patch_runtime(monkeypatch) + generation = DynamoGeneration(cluster=object(), config=_config()) + data = _data() + data["session_ids"] = session_ids + + async def collect(): + return [item async for item in generation.generate_async(data)] + + with pytest.raises(ValueError, match=expected_message): + asyncio.run(collect()) def test_prompt_at_context_limit_is_rejected(monkeypatch) -> None: @@ -447,8 +481,8 @@ def test_completion_retry_eventually_succeeds(monkeypatch) -> None: ) calls = [] - async def fake_post(*args): - calls.append(args) + async def fake_post(*args, **kwargs): + calls.append((args, kwargs)) return next(responses) async def no_sleep(_): @@ -464,11 +498,16 @@ async def no_sleep(_): greedy=False, stop_strings=None, max_new_tokens=1, + session_id="trajectory-session", ) ) assert token_ids == [8] assert len(calls) == 2 + assert [call[1]["headers"] for call in calls] == [ + {"X-Dynamo-Session-ID": "trajectory-session"}, + {"X-Dynamo-Session-ID": "trajectory-session"}, + ] @pytest.mark.parametrize("status", [400, 503]) @@ -478,7 +517,7 @@ def test_completion_retry_stops_on_nonretryable_or_exhaustion( _patch_runtime(monkeypatch) calls = [] - async def fake_post(*args): + async def fake_post(*args, **kwargs): calls.append(args) return {"status": "error", "http_status": status} @@ -511,7 +550,7 @@ def test_direct_completions_are_not_limited_by_default_thread_pool( all_entered = asyncio.Event() release = asyncio.Event() - async def fake_post(*args): + async def fake_post(*args, **kwargs): nonlocal entered_count entered_count += 1 if entered_count == request_count: diff --git a/tests/unit/models/generation/test_dynamo_http_client.py b/tests/unit/models/generation/test_dynamo_http_client.py index 195db9bb55a..f34fca4a052 100644 --- a/tests/unit/models/generation/test_dynamo_http_client.py +++ b/tests/unit/models/generation/test_dynamo_http_client.py @@ -65,8 +65,8 @@ async def __aenter__(self): async def __aexit__(self, *args): return None - def post(self, url, *, json): - self.requests.append((url, json)) + def post(self, url, *, json, headers=None): + self.requests.append((url, json, headers)) if self._error is not None: raise self._error return self._response @@ -191,7 +191,34 @@ def test_async_http_post_json_uses_same_error_contract( assert response == expected assert generation_module._is_retryable_http_response(response) is retryable - assert session.requests == [("http://worker/route", {"value": 1})] + assert session.requests == [("http://worker/route", {"value": 1}, None)] + + +def test_async_http_post_json_forwards_headers(monkeypatch) -> None: + session = _AsyncSession(_AsyncResponse(b'{"status":"ok"}')) + monkeypatch.setattr( + http_client.aiohttp, + "ClientSession", + lambda **kwargs: session, + ) + + response = asyncio.run( + http_client.async_http_post_json( + "http://worker/route", + {"value": 1}, + timeout_s=3, + headers={"X-Dynamo-Session-ID": "trajectory-session"}, + ) + ) + + assert response == {"status": "ok"} + assert session.requests == [ + ( + "http://worker/route", + {"value": 1}, + {"X-Dynamo-Session-ID": "trajectory-session"}, + ) + ] @pytest.mark.parametrize( diff --git a/tests/unit/models/generation/test_dynamo_token_wrapper.py b/tests/unit/models/generation/test_dynamo_token_wrapper.py index 53eb020bc2b..4bbbd4d4b5f 100644 --- a/tests/unit/models/generation/test_dynamo_token_wrapper.py +++ b/tests/unit/models/generation/test_dynamo_token_wrapper.py @@ -638,9 +638,15 @@ def post(self, url, *, json, headers): async def forward_twice(): await server._forward_chat_completion({}, authorization=None) - await server._forward_chat_completion({}, authorization="Bearer token") + await server._forward_chat_completion( + {}, + authorization="Bearer token", + session_id="trajectory-session", + ) asyncio.run(forward_twice()) assert len(session.calls) == 2 + assert "X-Dynamo-Session-ID" not in session.calls[0][2] assert session.calls[1][2]["Authorization"] == "Bearer token" + assert session.calls[1][2]["X-Dynamo-Session-ID"] == "trajectory-session"