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
12 changes: 12 additions & 0 deletions miles/rollout/session/sessions.py
Original file line number Diff line number Diff line change
Expand Up @@ -716,6 +716,18 @@ async def chat_completions(request: Request, session_id: str):
completion_tokens_emit = -1
_stats["turns_completed"] += 1

# Strip routed_experts/token-id blobs (in meta_info + prompt_token_ids)
# from the agent's copy only; the trainer reads them from the recorded
# SessionRecord. Returning them every turn made harbor RSS grow unbounded.
result["response_body"] = orjson.dumps(
{
**response,
"choices": [
{key: value for key, value in choice.items() if key not in ("meta_info", "prompt_token_ids")}
for choice in response.get("choices", [])
],
}
)
return backend.build_proxy_response(result)
finally:
_inflight_chat["count"] -= 1
Expand Down
22 changes: 12 additions & 10 deletions tests/fast/router/test_session_pretokenized_e2e.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,8 +9,8 @@
prefix of full_str) and uses pretokenized_token_ids for prompt construction.
- We check mock server request_log to verify pretokenized fields are injected
on turn 2+.
- We verify pretokenized_token_ids is a valid prefix of the returned
prompt_token_ids.
- We verify pretokenized_token_ids is a valid prefix of the prompt_token_ids
retained in the session record.

Tests are parametrized across multiple models and chat template combinations.
Models whose native templates satisfy the prefix invariant (Qwen3.5,
Expand Down Expand Up @@ -239,13 +239,14 @@ def _run_trajectory_e2e(env):
return session_id, responses


def _verify_pretokenized_injection(env, responses):
def _verify_pretokenized_injection(env, session_id):
"""Verify input_ids in request_log and prompt consistency."""
backend = env.backend
records = requests.get(f"{env.url}/sessions/{session_id}", timeout=5.0).json()["records"]

for turn_idx in range(len(env.trajectory.turns)):
req = backend.request_log[turn_idx]
resp_prompt_ids = responses[turn_idx]["choices"][0]["prompt_token_ids"]
resp_prompt_ids = records[turn_idx]["response"]["choices"][0]["prompt_token_ids"]

if turn_idx == 0:
assert "input_ids" not in req, "Turn 0 should not have input_ids"
Expand All @@ -262,21 +263,22 @@ def _verify_pretokenized_injection(env, responses):
)


def _verify_pretokenized_NOT_prefix(env, responses):
def _verify_pretokenized_NOT_prefix(env, session_id):
"""Verify that input_ids does NOT match prompt on turn 2+.

This is the inverse of _verify_pretokenized_injection: it proves that a
broken native template causes token mismatch in the e2e flow.
"""
backend = env.backend
records = requests.get(f"{env.url}/sessions/{session_id}", timeout=5.0).json()["records"]
found_mismatch = False

for turn_idx in range(1, len(env.trajectory.turns)):
req = backend.request_log[turn_idx]
if "input_ids" not in req:
continue
input_ids = req["input_ids"]
resp_prompt_ids = responses[turn_idx]["choices"][0]["prompt_token_ids"]
resp_prompt_ids = records[turn_idx]["response"]["choices"][0]["prompt_token_ids"]

if resp_prompt_ids != input_ids:
found_mismatch = True
Expand Down Expand Up @@ -308,8 +310,8 @@ def test_all_turns_succeed(self):
assert len(responses) == 2

def test_pretokenized_injection(self):
session_id, responses = _run_trajectory_e2e(self._env)
_verify_pretokenized_injection(self._env, responses)
session_id, _responses = _run_trajectory_e2e(self._env)
_verify_pretokenized_injection(self._env, session_id)

def test_session_records(self):
session_id, responses = _run_trajectory_e2e(self._env)
Expand All @@ -336,8 +338,8 @@ def test_all_turns_succeed(self):

def test_pretokenized_injection(self):
self._env.backend.reset_stats()
session_id, responses = _run_trajectory_e2e(self._env)
_verify_pretokenized_injection(self._env, responses)
session_id, _responses = _run_trajectory_e2e(self._env)
_verify_pretokenized_injection(self._env, session_id)

def test_pretokenized_grows_across_turns(self):
self._env.backend.reset_stats()
Expand Down
13 changes: 13 additions & 0 deletions tests/fast/router/test_sessions.py
Original file line number Diff line number Diff line change
Expand Up @@ -141,6 +141,9 @@ def test_proxy_chat_appends_record(self, router_env):
body = resp.json()
assert "choices" in body
assert body["choices"]
agent_choice = body["choices"][0]
assert "meta_info" not in agent_choice
assert "prompt_token_ids" not in agent_choice

get_resp = requests.get(f"{router_env.url}/sessions/{session_id}", timeout=5.0)
records = get_resp.json()["records"]
Expand All @@ -150,6 +153,16 @@ def test_proxy_chat_appends_record(self, router_env):
record = records[0]
assert record["path"] == "/v1/chat/completions"
assert record["status_code"] == 200
recorded_choice = record["response"]["choices"][0]
assert recorded_choice["prompt_token_ids"]
assert recorded_choice["meta_info"]["output_token_logprobs"]

merged_resp = requests.get(f"{router_env.url}/sessions/{session_id}/merged", timeout=5.0)
assert merged_resp.status_code == 200
sample = merged_resp.json()["sample"]
assert sample is not None
assert sample["response_length"] > 0
assert len(sample["rollout_log_probs"]) == sample["response_length"]


class TestTokenizationOffload:
Expand Down
Loading