From f53a9844888ca6b0de653b70b68618ebce34a0df Mon Sep 17 00:00:00 2001 From: braisedpork1964 <497494458@qq.com> Date: Fri, 5 Jun 2026 09:43:33 +0000 Subject: [PATCH 1/9] refactor SessionServer --- xtuner/v1/rl/rollout/chat_template.py | 21 +- xtuner/v1/rl/rollout/session_server.py | 724 ++++++++++++++++++------- 2 files changed, 550 insertions(+), 195 deletions(-) diff --git a/xtuner/v1/rl/rollout/chat_template.py b/xtuner/v1/rl/rollout/chat_template.py index afe5c1915c..0fa21e9df8 100644 --- a/xtuner/v1/rl/rollout/chat_template.py +++ b/xtuner/v1/rl/rollout/chat_template.py @@ -2,7 +2,6 @@ import json from typing import Any - _RAW_ARGUMENTS_KEY = "__xtuner_raw_arguments__" @@ -15,10 +14,28 @@ def canonicalize_messages_for_chat_template(messages: list[dict]) -> list[dict]: template-renderable shape before ``apply_chat_template``. It must not mutate model responses returned to the agent/client, nor the generated token ids, labels, or raw rollout artifacts used for training. + + Two transforms: + + - ``tool_calls[].function.arguments`` strings are parsed to dicts so the + template's ``tojson`` filter sees a real object. + - Any ``role='system'`` entry past the first message is demoted to + ``role='user'``. Claude Agent SDK's ``context_management`` injects a JIT + system reminder at the tail before each assistant turn (and those persist + mid-conversation as the rollout progresses); Anthropic's API allows it, + but Qwen-style chat templates require ``system`` only at index 0. + Demoting (rather than dropping) keeps the dynamic CWD/todo/time hints + visible to the policy at the cost of a small semantic shift. """ messages = copy.deepcopy(messages) - for message in messages: + for idx, message in enumerate(messages): + # TODO(wzy): drop this demote once the Qwen Jinja chat template + # accepts ``role='system'`` at any position. Until then we degrade SDK + # context_management injections to ``role='user'`` to keep the template + # renderable. + if idx > 0 and isinstance(message, dict) and message.get("role") == "system": + message["role"] = "user" tool_calls = message.get("tool_calls") if not isinstance(tool_calls, list): continue diff --git a/xtuner/v1/rl/rollout/session_server.py b/xtuner/v1/rl/rollout/session_server.py index 92f65265ad..03c15ce3db 100644 --- a/xtuner/v1/rl/rollout/session_server.py +++ b/xtuner/v1/rl/rollout/session_server.py @@ -1,18 +1,32 @@ +import copy import json from functools import reduce +from http import HTTPStatus from operator import add -from typing import Any, Optional +from typing import Any, List, Optional import numpy as np import ray from aiohttp import ClientConnectionResetError, ClientSession, ClientTimeout, web - from transformers import AutoTokenizer + from xtuner.v1.utils import get_logger from .chat_template import canonicalize_messages_for_chat_template from .trace_store import TokenizedSegment, get_store +FMT_OPENAI = "openai" +FMT_ANTHROPIC = "anthropic" + +# Fields the SessionServer consumes locally and never forwards upstream. +_SESSION_SERVER_ONLY_KEYS = {"session_id"} + + +def _detect_format(req_path: str) -> str: + if req_path.endswith("/messages") or "/v1/messages" in req_path: + return FMT_ANTHROPIC + return FMT_OPENAI + def _is_error_payload(payload: dict) -> bool: return payload.get("error") is not None or payload.get("type") == "error" or payload.get("object") == "error" @@ -27,80 +41,172 @@ def _lmdeploy_error_payload(message: str, status: int = 500, error_type: str = " } -def _stream_has_traceable_choices(raw: bytes) -> bool: - text = raw.decode("utf-8", errors="replace") - has_choices = False - saw_done = False - saw_terminal_finish = False - for line in text.split("\n"): - line = line.strip() - if line == "data: [DONE]": - saw_done = True - continue - if line.startswith("data: "): - try: - event = json.loads(line[6:]) - except json.JSONDecodeError: - return False - if _is_error_payload(event): - return False - if event.get("choices"): - if any(choice.get("finish_reason") == "error" for choice in event.get("choices", [])): - return False - if any(choice.get("finish_reason") for choice in event.get("choices", [])): - saw_terminal_finish = True - has_choices = True - return has_choices and saw_done and saw_terminal_finish - - -def _extract_output_logprobs(choice: dict, output_token_ids: list[int]) -> list[float]: +def _bool_request_value(value: Any, default: bool = False) -> bool: + """Coerce an opaque request-body value (bool / int / str / None) to bool. + + Mirrors how lmdeploy validates flag-style fields — accepts ``"true"``, + ``"1"`` etc. as truthy and ``"false"``, ``"0"`` as falsy. + """ + if value is None: + return default + if isinstance(value, str): + return value.strip().lower() not in {"", "0", "false", "no", "off"} + return bool(value) + + +def _request_uses_trace_store(req_body: dict) -> bool: + """Whether this request should drive the trace_store / token-in-token-out path. + + Evaluate-mode callers opt out by setting ``return_token_ids=False`` on the + request body; tracing then becomes a passthrough proxy with parameter-name + hygiene (``logprobs`` → ``return_logprob``) and the upstream worker handles + its own tokenization. + """ + if req_body.get("session_id") is None or "messages" not in req_body: + return False + return _bool_request_value(req_body.get("return_token_ids"), True) + + +def _extract_output_logprobs(output_token_logprobs: Optional[list], output_token_ids: list[int]) -> list[float]: if not output_token_ids: return [] - output_token_logprobs = choice.get("output_token_logprobs") if output_token_logprobs is None: raise RuntimeError( - "SessionServer response choice has no output_token_logprobs; " + "SessionServer response has no output_token_logprobs; " "the return_logprob protocol is required for training traces." ) logprob_token_ids = [item[1] for item in output_token_logprobs] if logprob_token_ids != output_token_ids: raise RuntimeError( - "SessionServer response choice has mismatched output_token_logprobs: " + "SessionServer response has mismatched output_token_logprobs: " f"output_ids_len={len(output_token_ids)}, logprob_ids_len={len(logprob_token_ids)}" ) return [item[0] for item in output_token_logprobs] -_SESSION_SERVER_ONLY_KEYS = {"session_id"} +def _maybe_json_loads(value): + """Parse a JSON string, falling back to the original value on error.""" + if not isinstance(value, str): + return value + try: + return json.loads(value) + except json.JSONDecodeError: + return value -def _bool_request_value(value: Any, default: bool = False) -> bool: - if value is None: - return default - if isinstance(value, str): - return value.strip().lower() not in {"", "0", "false", "no", "off"} - return bool(value) +def _normalize_tool_call_arguments(messages) -> None: + """In-place: parse ``tool_calls[].function.arguments`` strings into dicts. + lmdeploy stringifies tool-call arguments; the standard OpenAI/chat_template + path expects dicts. Normalize so both paths share one shape. + """ + for msg in messages: + if not isinstance(msg, dict): + continue + for tc in msg.get("tool_calls") or []: + if not isinstance(tc, dict): + continue + fn = tc.get("function") + if isinstance(fn, dict) and "arguments" in fn: + fn["arguments"] = _maybe_json_loads(fn["arguments"]) -def _request_uses_trace_store(req_body: dict) -> bool: - if req_body.get("session_id") is None or "messages" not in req_body: - return False - return _bool_request_value(req_body.get("return_token_ids"), True) + +def _strip_stop_word(content, stop_word: str): + """Remove ``stop_word`` from a message ``content`` field (any OpenAI shape). + + OpenAI ``content`` accepts both a plain string and a list of typed parts + (``[{"type": "text", "text": "..."}, ...]``); the latter shows up after we + pass anthropic messages through ``to_openai_messages`` — multiple anthropic + text blocks inside one assistant turn become multiple ``text`` parts. The + stop_word can land in any of those ``text`` fields, so we walk both shapes. + + Returns ``(content, modified)``. List parts are mutated in place; for str + a new string is returned (the original is left untouched). + """ + if not stop_word or content is None: + return content, False + if isinstance(content, str): + if stop_word in content: + return content.replace(stop_word, ""), True + return content, False + if isinstance(content, list): + modified = False + for part in content: + if isinstance(part, dict) and isinstance(part.get("text"), str) and stop_word in part["text"]: + part["text"] = part["text"].replace(stop_word, "") + modified = True + return content, modified + return content, False + + +def _anthropic_response_to_assistant_message(response: dict) -> dict: + """Wrap an Anthropic ``MessagesResponse`` body as an anthropic-shaped assistant message. + + The returned dict can be appended to ``MessagesRequest.messages`` so the + full conversation (input + this assistant turn) survives a round-trip + through ``MessagesRequest.model_validate`` + ``to_openai_messages``. + """ + blocks = response.get("content") or [] + normalized: List[dict] = [] + for block in blocks: + btype = block.get("type") + if btype == "text": + normalized.append({"type": "text", "text": block.get("text", "")}) + elif btype == "thinking": + nb: dict = {"type": "thinking", "thinking": block.get("thinking", "")} + if "signature" in block: + nb["signature"] = block["signature"] + normalized.append(nb) + elif btype == "tool_use": + tool_id = block.get("id") + name = block.get("name") + if not tool_id or not name: + raise ValueError(f"tool_use block missing id/name: {block}") + normalized.append({"type": "tool_use", "id": tool_id, "name": name, "input": block.get("input", {})}) + return {"role": "assistant", "content": normalized} + + +def _anthropic_request_to_openai(req_body: dict) -> tuple[list[dict], Optional[list[dict]]]: + """Anthropic request body → OpenAI (messages, tools). + + Uses lmdeploy's ``MessagesRequest`` + ``to_openai_messages`` / + ``to_openai_tools`` to keep all anthropic content-block / tool-use / + tool-result / system handling consistent with what the worker would do. + """ + from lmdeploy.serve.anthropic.adapter import to_openai_messages, to_openai_tools + from lmdeploy.serve.anthropic.protocol import MessagesRequest + + req = MessagesRequest.model_validate(req_body) + messages = to_openai_messages(req) + converted_tools = to_openai_tools(req.tools) + tools = [t.model_dump() for t in converted_tools] if converted_tools else None + _normalize_tool_call_arguments(messages) + return messages, tools class SessionServer: - """SessionServer intercepts and records requests sent to a remote LLM API - worker. + """SessionServer intercepts and records requests sent to a remote LLM API worker. - It acts as a reverse-proxy (or interceptor) in front of an already running - worker (like lmdeploy, sglang, or vllm). It binds to a specific (host, port) - and relays any received traffic to the actual worker URL. + It acts as a reverse-proxy in front of an already running worker (lmdeploy, + sglang, vllm). Each supported endpoint is transparently forwarded along + its native path with its native response shape — only the on_request and + on_response hooks normalize messages to OpenAI form so the tokenizer + + ``trace_store`` always see one shape. - You can optionally provide before_request and after_response hooks to - perform extra logging, trace state mutations, or message cleanup before/after - routing the request to the worker backend. + Supported endpoints: + + - ``POST /v1/chat/completions`` — passthrough + - ``POST /v1/messages`` (Anthropic) — passthrough; hooks render the + request via the lmdeploy adapter to drive prefix-cache lookup, swap + ``messages``/``system``/``tools``/``tool_choice`` out for raw + ``input_ids`` on the wire (lmdeploy requires this when ``input_ids`` is + set), then re-assemble the OpenAI-shaped trace from the anthropic + response. + + The ``/v1/responses`` endpoint is intentionally rejected (501) because the + upstream worker does not implement it. Args: worker_base_url (str): The base URL of the real worker (e.g. "http://127.0.0.1:8000") @@ -134,9 +240,29 @@ def __init__( self._site: Optional[web.TCPSite] = None self._lmdeploy_actor: Optional[ray.actor.ActorHandle] = None - async def on_request(self, req_body: dict, *, trace_enabled: bool = True) -> dict: - """Hook for processing/modifying the request before forwarding.""" - + async def on_request(self, req_body: dict, fmt: str, *, trace_enabled: bool = True) -> dict: + """Normalize the request to drive the prefix cache + inject extension fields. + + When ``trace_enabled`` is False (evaluate mode), forward the request + as-is with parameter-name hygiene only — no token-in-token-out + rewrite — so the upstream worker handles tokenization on its own. + + Args: + req_body (dict): The original client request body. + fmt (str): The detected request format (one of ``FMT_OPENAI`` / + ``FMT_ANTHROPIC``). + trace_enabled (bool): Whether the request should populate the + trace_store. Disabled when the caller sets + ``return_token_ids=False`` (evaluate path). + + Returns: + dict: The forward-ready request body. The wire shape still + matches ``fmt``; in trace mode it carries ``input_ids`` plus the + ``return_*`` extension flags. + """ + # Evaluate path: forward unchanged except for the legacy ``logprobs`` + # → ``return_logprob`` rename and stripping the SessionServer-only + # ``session_id`` field. if not trace_enabled: worker_req = {k: v for k, v in req_body.items() if k not in _SESSION_SERVER_ONLY_KEYS} if "logprobs" in worker_req: @@ -149,15 +275,33 @@ async def on_request(self, req_body: dict, *, trace_enabled: bool = True) -> dic return worker_req session_id = req_body["session_id"] - # 1. chat_template render 出完整 prompt string,不 tokenize 全量 + + # 1. Render the request as OpenAI chat-template input to drive the + # string-prefix cache. For non-OpenAI formats we convert in-memory; + # we do NOT mutate req_body["messages"] in those cases — the wire + # body keeps its native shape. + if fmt == FMT_ANTHROPIC: + openai_messages, openai_tools = _anthropic_request_to_openai(req_body) + else: + openai_messages = req_body["messages"] + openai_tools = req_body.get("tools", None) + + # Defensive scrub: an earlier round may have rendered ``<|im_end|>`` into + # an assistant text block (e.g. via ``include_stop_str_in_output``); if + # the client echoes that block back to us, the chat template would emit + # the stop_word twice and break tokenization. + for msg in openai_messages: + if isinstance(msg, dict) and "content" in msg: + msg["content"], _ = _strip_stop_word(msg["content"], self.stop_word) + prompt_text = self.tokenizer.apply_chat_template( - canonicalize_messages_for_chat_template(req_body["messages"]), - tools=req_body.get("tools", None), + canonicalize_messages_for_chat_template(openai_messages), + tools=openai_tools, add_generation_prompt=True, tokenize=False, ) - # 2. Store 做 string prefix match。 + # 2. Prefix-cache lookup; insert the delta tail back into the store. prefix, nodes = await self.store.search.remote(session_id, prompt_text, filter_none=True) if prefix: get_logger().debug(f"Hit prefix cache for session {session_id}") @@ -167,50 +311,91 @@ async def on_request(self, req_body: dict, *, trace_enabled: bool = True) -> dic await self.store.insert.remote(session_id, prompt_text, TokenizedSegment(text=delta, token_ids=delta_ids)) input_ids = reduce(add, [node.value.token_ids for node in nodes] + [delta_ids]) - # 3. 组装 OpenAI chat completions 请求。 + # 3. Inject extension fields on the forwarded request. The body still + # follows ``fmt``'s native schema. ``logprobs``/``top_logprobs`` are + # dropped because we always force ``return_logprob=True`` to drive + # the trace_store; their original values would just shadow ours. worker_req = { - **{ - k: v - for k, v in req_body.items() - if k not in _SESSION_SERVER_ONLY_KEYS | {"messages", "logprobs", "top_logprobs"} - }, - "messages": [], - "input_ids": input_ids, - "return_token_ids": True, - "return_routed_experts": True, - "return_logprob": True, - "include_stop_str_in_output": True, + k: v for k, v in req_body.items() if k not in _SESSION_SERVER_ONLY_KEYS | {"logprobs", "top_logprobs"} } - return worker_req - - async def on_response(self, worker_resp: dict, *, trace_enabled: bool = True) -> dict: - """Hook for processing the parsed response received from the worker.""" + worker_req["input_ids"] = input_ids + worker_req["return_token_ids"] = True + worker_req["return_routed_experts"] = True + worker_req["return_logprob"] = True + worker_req["include_stop_str_in_output"] = True + + if fmt == FMT_ANTHROPIC: + # lmdeploy's /v1/messages requires messages to be empty AND + # system / tools / tool_choice to be unset whenever input_ids is + # supplied (raw input_ids bypass message rendering). + worker_req["messages"] = [] + worker_req.pop("system", None) + worker_req.pop("tools", None) + worker_req.pop("tool_choice", None) + else: + worker_req["messages"] = [] - if not trace_enabled: - return {k: v for k, v in worker_resp.items() if k not in {"messages", "tools"}} + return worker_req - session_id = worker_resp["session_id"] - messages = worker_resp["messages"] - tools = worker_resp["tools"] - choice = worker_resp["choices"][0] + async def on_response(self, worker_resp: dict, fmt: str, orig_req_body: dict) -> dict: + """Trace the assistant turn into ``trace_store`` (always in OpenAI shape). + + Args: + worker_resp (dict): The parsed upstream response (still in + ``fmt`` shape). + fmt (str): The request format that produced this response. + orig_req_body (dict): The original client request body, before + ``on_request`` rewrote it. Needed to reconstruct the full + conversation for OpenAI-shape tracing because the wire body + forwarded to the worker no longer carries ``messages`` / + ``system`` / ``tools``. + + Returns: + dict: The (possibly trimmed) response — currently unused by the + caller but kept for symmetry. + """ + session_id = orig_req_body["session_id"] + + if fmt == FMT_ANTHROPIC: + output_token_ids = worker_resp.get("output_ids") + output_token_logprobs = worker_resp.get("output_token_logprobs") + raw_routed_expert = worker_resp.get("routed_experts") + assistant_anthropic_msg = _anthropic_response_to_assistant_message(worker_resp) + full_req = copy.deepcopy(orig_req_body) + full_req.setdefault("messages", []).append(assistant_anthropic_msg) + openai_messages, openai_tools = _anthropic_request_to_openai(full_req) + assistant_msg = openai_messages[-1] + messages = openai_messages[:-1] + tools = openai_tools + else: + choice = worker_resp["choices"][0] + output_token_ids = choice.get("output_ids") + output_token_logprobs = choice.get("output_token_logprobs") + raw_routed_expert = choice.get("routed_experts") + assistant_msg = choice["message"] + messages = orig_req_body["messages"] + tools = orig_req_body.get("tools", None) - output_token_ids = choice.get("output_ids") # len = N_out if output_token_ids is None: raise RuntimeError( - "SessionServer response choice has no output_ids; " - "cannot export a training trace for this assistant turn." + "SessionServer response has no output_ids; cannot export a training trace for this assistant turn." ) - output_logprobs = _extract_output_logprobs(choice, output_token_ids) - raw_routed_expert = choice.get("routed_experts") # 本次 call 的 raw routed_expert,可为 None + output_logprobs = _extract_output_logprobs(output_token_logprobs, output_token_ids) - # 2. Store 把 input_delta / assistant_output 两个节点补齐字段。 + if isinstance(assistant_msg.get("content"), (str, list)): + assistant_msg["content"], _ = _strip_stop_word(assistant_msg["content"], self.stop_word) + for msg in messages: + if isinstance(msg, dict) and "content" in msg: + msg["content"], _ = _strip_stop_word(msg["content"], self.stop_word) + + # Render OpenAI prompts for the prefix-cache key boundary. old_prompt = self.tokenizer.apply_chat_template( canonicalize_messages_for_chat_template(messages), tools=tools, add_generation_prompt=True, tokenize=False ) - messages = [*messages, choice["message"]] + full_messages = [*messages, assistant_msg] new_prompt = ( self.tokenizer.apply_chat_template( - canonicalize_messages_for_chat_template(messages), + canonicalize_messages_for_chat_template(full_messages), tools=tools, add_generation_prompt=False, tokenize=False, @@ -235,14 +420,11 @@ async def on_response(self, worker_resp: dict, *, trace_enabled: bool = True) -> prefix_len = sum(len(n.value.token_ids) for n in nodes[:-1]) assert prefix_len + delta_len + len(output_token_ids) == len(raw_routed_expert) - # split raw_routed_expert - # raw_routed_expert target shape mapping: [prefix_len + delta_len + response_len, ...] delta_expert = raw_routed_expert[prefix_len : prefix_len + delta_len] response_expert = raw_routed_expert[prefix_len + delta_len :] if delta_len > 0: delta_node_val.expert_key = ray.put(delta_expert) - # update delta node in store await self.store.insert.remote(session_id, old_prompt, delta_node_val) raw_routed_expert = ray.put(response_expert) @@ -262,9 +444,7 @@ async def on_response(self, worker_resp: dict, *, trace_enabled: bool = True) -> ), ) - # 3. 返回标准 OpenAI response,session_id 由 SessionClient 层再剥 - resp = {k: v for k, v in worker_resp.items() if k != "messages"} - return resp + return worker_resp @property def url(self) -> str: @@ -295,31 +475,46 @@ async def stop(self): get_logger().info("SessionServer stopped.") async def _handle_request(self, request: web.Request) -> web.Response: - """Proxy handler for the worker API.""" + """Proxy handler: detect format, run hooks, forward, stream back.""" + + req_path = request.match_info["path"] + + # Reject /v1/responses outright — upstream worker doesn't implement it. + if req_path.endswith("/responses") or "/v1/responses" in req_path: + return web.json_response( + _lmdeploy_error_payload( + "/v1/responses is not supported by SessionServer (upstream worker has no responses endpoint).", + status=HTTPStatus.NOT_IMPLEMENTED, + error_type="not_implemented", + ), + status=HTTPStatus.NOT_IMPLEMENTED, + ) + + fmt = _detect_format(req_path) # Read the request body request_body = await request.read() - request_data = session_id = messages = None + request_data = None + orig_req_body: Optional[dict] = None trace_enabled = False - orig_return_logprob = orig_return_token_ids = orig_return_routed_experts = False + orig_return_logprob = orig_return_token_ids = False + orig_return_routed_experts = True if request_body: try: request_data = json.loads(request_body) + orig_req_body = copy.deepcopy(request_data) trace_enabled = _request_uses_trace_store(request_data) + # Accept either ``return_logprob`` (canonical) or the legacy + # ``logprobs`` alias when deciding whether the client wanted + # logprobs forwarded back. orig_return_logprob = _bool_request_value( request_data.get("return_logprob", request_data.get("logprobs")), False ) orig_return_token_ids = _bool_request_value(request_data.get("return_token_ids"), False) orig_return_routed_experts = _bool_request_value(request_data.get("return_routed_experts"), True) - session_id = request_data.get("session_id") - messages = request_data.get("messages") - tools = request_data.get("tools", None) - - # Apply purely abstract on_request processing - request_data = await self.on_request(request_data, trace_enabled=trace_enabled) - # Re-serialize the modified payload back to bytes + request_data = await self.on_request(request_data, fmt, trace_enabled=trace_enabled) request_body = json.dumps(request_data).encode("utf-8") except json.JSONDecodeError: pass @@ -328,66 +523,39 @@ async def _handle_request(self, request: web.Request) -> web.Response: get_logger().error(message) return web.json_response(_lmdeploy_error_payload(message), status=500) - # Build forwarding headers, dropping original Host + # Build forwarding headers, dropping original Host / Content-Length. forward_headers = dict(request.headers) forward_headers.pop("Host", None) forward_headers.pop("host", None) forward_headers.pop("Content-Length", None) forward_headers.pop("content-length", None) + if fmt == FMT_ANTHROPIC: + forward_headers["anthropic-version"] = "2023-06-01" - # Re-build Path - req_path = request.match_info["path"] + # Path is forwarded verbatim — each endpoint keeps its native shape. target_url = f"{self.worker_base_url}/{req_path.lstrip('/')}" if request.query_string: target_url += f"?{request.query_string}" is_stream = request_data.get("stream", False) if request_data else False - def _clean_data(data: dict) -> bool: - modified = False - for key, drop in [ - ("output_token_logprobs", not orig_return_logprob), - ("output_ids", not orig_return_token_ids), - ("routed_experts", not orig_return_routed_experts), - ]: - if drop and key in data: - data.pop(key) - modified = True - if drop: - for c in data.get("choices", []): - if key in c: - c.pop(key) - modified = True - - for c in data.get("choices", []): - if "logprobs" in c: - c.pop("logprobs") - modified = True - - for c in data.get("choices", []): - if c.get("message") and isinstance(c["message"].get("content"), str): - if self.stop_word in c["message"]["content"]: - c["message"]["content"] = c["message"]["content"].replace(self.stop_word, "") - modified = True - if c.get("delta") and isinstance(c["delta"].get("content"), str): - if self.stop_word in c["delta"]["content"]: - c["delta"]["content"] = c["delta"]["content"].replace(self.stop_word, "") - modified = True - - return modified + # Build the per-format stream cleaner (strips lmdeploy-injected + # extension fields that the client didn't ask for, and removes the + # stop word from any user-visible text). + clean_data = self._build_data_cleaner( + fmt, + orig_return_logprob=orig_return_logprob, + orig_return_token_ids=orig_return_token_ids, + orig_return_routed_experts=orig_return_routed_experts, + ) - # Forward the request to the upstream worker - # read_bufsize controls StreamReader's line buffer limit; SSE events with large - # tool_calls/reasoning_content payloads can exceed the 64KB default and trigger - # "Chunk too big" from readuntil(b"\n"). timeout = ClientTimeout(total=self.request_timeout, sock_connect=30) async with ClientSession(read_bufsize=self.read_bufsize, timeout=timeout) as client: async with client.request( method=request.method, url=target_url, headers=forward_headers, data=request_body ) as resp: - # Setup proper stream vs sync response objects if is_stream: - response_chunks = [] + response_chunks: list[bytes] = [] response = web.StreamResponse( status=resp.status, headers={ @@ -404,16 +572,17 @@ def _clean_data(data: dict) -> bool: # in full but stop attempting to write to the closed socket. client_alive = True async for line in resp.content: - # Keep unmodified line for trace store parsing + # Only retain chunks when we'll actually need to parse + # them for tracing; evaluate-mode requests skip this + # so memory does not grow with stream length. if trace_enabled: response_chunks.append(line) - # Dynamically prune added fields before writing to client if request_data is not None and line.startswith(b"data: ") and line.strip() != b"data: [DONE]": try: text = line.decode("utf-8") data = json.loads(text[6:]) - if _clean_data(data): + if clean_data(data): line = ("data: " + json.dumps(data) + "\n").encode("utf-8") except Exception: pass @@ -432,9 +601,9 @@ def _clean_data(data: dict) -> bool: if request_data is not None: try: - clean_data = json.loads(raw_response) - if _clean_data(clean_data): - final_raw_response = json.dumps(clean_data).encode("utf-8") + parsed = json.loads(raw_response) + if clean_data(parsed): + final_raw_response = json.dumps(parsed).encode("utf-8") except Exception: pass @@ -445,21 +614,25 @@ def _clean_data(data: dict) -> bool: for k, v in resp.headers.items() if k.lower() not in ("transfer-encoding", "content-length", "content-encoding") }, - body=final_raw_response, # Modified raw response without our injected trace params + body=final_raw_response, ) # Apply abstract on_response processing - response_data = None + response_data: Optional[dict] = None skip_done = bool(is_stream and not trace_enabled) - session_error_msg = None - if request_data and trace_enabled: + session_error_msg: Optional[str] = None + if request_data and trace_enabled and orig_req_body is not None: if is_stream: - skip_done = not _stream_has_traceable_choices(raw_response) - if not skip_done: - try: - response_data = self._parse_stream_response(raw_response) - except Exception as exc: - session_error_msg = f"SessionServer stream trace failed: {type(exc).__name__}: {exc}" + try: + response_data = self._parse_stream_response(raw_response, fmt) + if response_data is None: + # Upstream emitted no traceable content — suppress the + # synthetic [DONE] line we usually append; the real + # stream content has already been forwarded as-is. + skip_done = True + except Exception as exc: + session_error_msg = f"SessionServer stream trace failed: {type(exc).__name__}: {exc}" + skip_done = True else: try: response_data = json.loads(raw_response) @@ -470,14 +643,7 @@ def _clean_data(data: dict) -> bool: if response_data is not None: try: - for c in response_data.get("choices", []): - if c.get("message") and isinstance(c["message"].get("content"), str): - c["message"]["content"] = c["message"]["content"].replace(self.stop_word, "") - - response_data["session_id"] = session_id - response_data["messages"] = messages - response_data["tools"] = tools - await self.on_response(response_data, trace_enabled=trace_enabled) + await self.on_response(response_data, fmt, orig_req_body) except Exception as exc: session_error_msg = f"SessionServer response hook failed: {type(exc).__name__}: {exc}" @@ -496,7 +662,6 @@ def _clean_data(data: dict) -> bool: await response.write(b"data: [DONE]\n\n") await response.write_eof() except (ConnectionError, ClientConnectionResetError): - # Client already gone; trace was still recorded above. pass elif session_error_msg: return web.json_response(_lmdeploy_error_payload(session_error_msg), status=500) @@ -512,11 +677,87 @@ async def _decode_routed_experts(self, routed_experts: Any) -> np.ndarray: return np.asarray(routed_experts_data) return np.asarray(routed_experts) + def _build_data_cleaner( + self, fmt: str, *, orig_return_logprob: bool, orig_return_token_ids: bool, orig_return_routed_experts: bool + ): + """Return a per-format ``data -> bool`` cleaner stripping injected fields.""" + drops = { + "output_token_logprobs": not orig_return_logprob, + "output_ids": not orig_return_token_ids, + "routed_experts": not orig_return_routed_experts, + } + + if fmt == FMT_ANTHROPIC: + + def clean(data: dict) -> bool: + modified = False + # Non-stream MessagesResponse and stream MessageDeltaEvent put + # extension fields at the top level. ContentBlockDeltaEvent + # carries output_ids/output_token_logprobs on the event itself. + for key, drop in drops.items(): + if drop and key in data: + data.pop(key) + modified = True + # Stop-word scrubbing for text deltas / text content blocks. + delta = data.get("delta") + if isinstance(delta, dict) and isinstance(delta.get("text"), str) and self.stop_word in delta["text"]: + delta["text"] = delta["text"].replace(self.stop_word, "") + modified = True + for block in data.get("content") or []: + if ( + isinstance(block, dict) + and isinstance(block.get("text"), str) + and self.stop_word in block["text"] + ): + block["text"] = block["text"].replace(self.stop_word, "") + modified = True + return modified + + return clean + + # FMT_OPENAI + def clean_openai(data: dict) -> bool: + modified = False + for key, drop in drops.items(): + if drop and key in data: + data.pop(key) + modified = True + if drop: + for c in data.get("choices", []): + if key in c: + c.pop(key) + modified = True + + for c in data.get("choices", []): + if "logprobs" in c: + c.pop("logprobs") + modified = True + + for c in data.get("choices", []): + msg = c.get("message") + if isinstance(msg, dict) and "content" in msg: + msg["content"], changed = _strip_stop_word(msg["content"], self.stop_word) + if changed: + modified = True + delta = c.get("delta") + if isinstance(delta, dict) and "content" in delta: + delta["content"], changed = _strip_stop_word(delta["content"], self.stop_word) + if changed: + modified = True + + return modified + + return clean_openai + @staticmethod - def _parse_stream_response(raw: bytes) -> Optional[dict]: - """Parse SSE stream to reconstruct the complete final message state.""" + def _parse_stream_response(raw: bytes, fmt: str) -> Optional[dict]: + """Parse an SSE stream and reconstruct the final response object. + + Returns ``None`` when the stream carried no traceable content (so the + caller skips the trace hook entirely instead of recording garbage). + """ text = raw.decode("utf-8", errors="replace") - events = [] + events: list[dict] = [] saw_done = False for line in text.split("\n"): line = line.strip() @@ -531,10 +772,16 @@ def _parse_stream_response(raw: bytes) -> Optional[dict]: if not events: return None + + if fmt == FMT_ANTHROPIC: + return SessionServer._parse_anthropic_stream(events) + return SessionServer._parse_openai_stream(events, saw_done=saw_done) + + @staticmethod + def _parse_openai_stream(events: list[dict], *, saw_done: bool) -> Optional[dict]: if not any(event.get("choices") for event in events): raise RuntimeError(f"Upstream SSE stream ended without choices: {json.dumps(events, ensure_ascii=False)}") - # Reconstruct standard stream output (Assuming OpenAI format here) message: dict[str, Any] = {"choices": [{"message": {"role": "assistant", "content": ""}}]} content_parts: list[str] = [] tool_calls_map: dict[int, dict[str, Any]] = {} @@ -546,46 +793,33 @@ def _parse_stream_response(raw: bytes) -> Optional[dict]: if event.get("model"): message["model"] = event["model"] - choices = event.get("choices", []) - for choice in choices: + for choice in event.get("choices", []): if choice.get("finish_reason") == "error": raise RuntimeError( f"Upstream SSE choice finished with error: {json.dumps(event, ensure_ascii=False)}" ) delta = choice.get("delta", {}) - # Check text content if delta.get("content"): content_parts.append(delta["content"]) - # Check output ids if choice.get("output_ids") is not None: - assistant_choice = message["choices"][0] - if "output_ids" not in assistant_choice: - assistant_choice["output_ids"] = [] - assistant_choice["output_ids"].extend(choice["output_ids"]) + message["choices"][0].setdefault("output_ids", []).extend(choice["output_ids"]) - # Check routed experts. LMDeploy only emits this in the final - # chunk, often as a Ray shared-store key string. if choice.get("routed_experts") is not None: - assistant_choice = message["choices"][0] - assistant_choice["routed_experts"] = choice["routed_experts"] + message["choices"][0]["routed_experts"] = choice["routed_experts"] - # Check raw output logprobs from LMDeploy return_logprob protocol. if choice.get("output_token_logprobs") is not None: - assistant_choice = message["choices"][0] - if "output_token_logprobs" not in assistant_choice: - assistant_choice["output_token_logprobs"] = [] - assistant_choice["output_token_logprobs"].extend(choice["output_token_logprobs"]) + message["choices"][0].setdefault("output_token_logprobs", []).extend( + choice["output_token_logprobs"] + ) - # Check reasoning content if delta.get("reasoning_content"): assistant_msg = message["choices"][0]["message"] assistant_msg["reasoning_content"] = ( assistant_msg.get("reasoning_content", "") + delta["reasoning_content"] ) - # Check tool calls for tc_delta in delta.get("tool_calls") or []: idx = tc_delta.get("index", 0) if idx not in tool_calls_map: @@ -624,6 +858,110 @@ def _parse_stream_response(raw: bytes) -> Optional[dict]: return message + @staticmethod + def _parse_anthropic_stream(events: list[dict]) -> Optional[dict]: + """Aggregate Anthropic SSE events into a ``MessagesResponse``-shaped dict. + + Extension fields: + - ``output_ids`` / ``output_token_logprobs`` live on each + ``content_block_delta`` event and are concatenated in order. + - ``routed_experts`` lives on the terminal ``message_delta`` event. + + Returns ``None`` when no ``message_stop`` was seen — that means the + upstream stream was aborted before producing a complete turn. + """ + message: dict[str, Any] = {"role": "assistant", "type": "message", "content": []} + content_blocks: list[dict] = [] + current_block: Optional[dict] = None + output_ids: list[int] = [] + output_token_logprobs: list = [] + routed_experts = None + usage: dict[str, Any] = {} + saw_message_stop = False + + for event in events: + etype = event.get("type", "") + + if etype == "message_start": + msg = event.get("message", {}) + if msg.get("id"): + message["id"] = msg["id"] + if msg.get("model"): + message["model"] = msg["model"] + if msg.get("usage"): + usage = dict(msg["usage"]) + if msg.get("role"): + message["role"] = msg["role"] + + elif etype == "content_block_start": + current_block = dict(event.get("content_block") or {}) + + elif etype == "content_block_delta": + delta = event.get("delta") or {} + dtype = delta.get("type") + if current_block is None: + current_block = {} + if dtype == "text_delta": + current_block.setdefault("type", "text") + current_block["text"] = current_block.get("text", "") + delta.get("text", "") + elif dtype == "thinking_delta": + current_block.setdefault("type", "thinking") + current_block["thinking"] = current_block.get("thinking", "") + delta.get("thinking", "") + elif dtype == "signature_delta": + current_block["signature"] = current_block.get("signature", "") + delta.get("signature", "") + elif dtype == "input_json_delta": + current_block.setdefault("type", "tool_use") + current_block["_partial_json"] = current_block.get("_partial_json", "") + delta.get( + "partial_json", "" + ) + + if event.get("output_ids"): + output_ids.extend(event["output_ids"]) + if event.get("output_token_logprobs"): + output_token_logprobs.extend(event["output_token_logprobs"]) + + elif etype == "content_block_stop": + if current_block is not None: + if "_partial_json" in current_block: + raw = current_block.pop("_partial_json") + try: + current_block["input"] = json.loads(raw) if raw else {} + except json.JSONDecodeError: + current_block["input"] = {} + content_blocks.append(current_block) + current_block = None + + elif etype == "message_delta": + d = event.get("delta") or {} + if d.get("stop_reason"): + message["stop_reason"] = d["stop_reason"] + if d.get("stop_sequence") is not None: + message["stop_sequence"] = d["stop_sequence"] + u = event.get("usage") or {} + for k, v in u.items(): + if isinstance(v, (int, float)): + usage[k] = usage.get(k, 0) + v + else: + usage[k] = v + if event.get("routed_experts") is not None: + routed_experts = event["routed_experts"] + + elif etype == "message_stop": + saw_message_stop = True + + if not saw_message_stop: + return None + if not output_ids: + raise RuntimeError("Upstream anthropic SSE stream ended without output_ids.") + + message["content"] = content_blocks + message["output_ids"] = output_ids + message["output_token_logprobs"] = output_token_logprobs or None + message["routed_experts"] = routed_experts + if usage: + message["usage"] = usage + return message + class SessionServerActor: """Ray actor wrapper that owns one SessionServer instance.""" From b8926282f149390d4deb71951a0c018e57aae948 Mon Sep 17 00:00:00 2001 From: braisedpork1964 <497494458@qq.com> Date: Mon, 8 Jun 2026 02:48:01 +0000 Subject: [PATCH 2/9] drop anthropic server-side tools --- xtuner/v1/rl/rollout/session_server.py | 32 ++++++++++++++++++++++++++ 1 file changed, 32 insertions(+) diff --git a/xtuner/v1/rl/rollout/session_server.py b/xtuner/v1/rl/rollout/session_server.py index 03c15ce3db..a79f1b5225 100644 --- a/xtuner/v1/rl/rollout/session_server.py +++ b/xtuner/v1/rl/rollout/session_server.py @@ -186,6 +186,28 @@ def _anthropic_request_to_openai(req_body: dict) -> tuple[list[dict], Optional[l return messages, tools +def _filter_anthropic_user_tools(tools) -> Optional[list]: + """Drop Anthropic server-side tools (``web_search_*`` / ``computer_*`` / ``bash_*`` / ``text_editor_*``). + + ``ToolParam`` in lmdeploy's anthropic protocol requires ``input_schema``; + server-side built-ins are managed by Anthropic itself and only carry + ``type`` + ``name``, so they fail ``MessagesRequest.model_validate`` and + the worker can't execute them anyway. Keep only entries that have an + ``input_schema`` (= user-defined tools the worker can route to). + + Returns the filtered list (possibly empty), or ``None`` if input was None. + """ + if tools is None: + return None + if not isinstance(tools, list): + return tools + kept = [t for t in tools if isinstance(t, dict) and "input_schema" in t] + if len(kept) != len(tools): + dropped = [t.get("type") or t.get("name") for t in tools if t not in kept] + get_logger().debug(f"Dropped {len(tools) - len(kept)} anthropic server-side tool(s): {dropped}") + return kept + + class SessionServer: """SessionServer intercepts and records requests sent to a remote LLM API worker. @@ -502,6 +524,16 @@ async def _handle_request(self, request: web.Request) -> web.Response: if request_body: try: request_data = json.loads(request_body) + # Anthropic server-side built-ins (web_search_*, computer_*, ...) + # don't carry ``input_schema`` and the lmdeploy worker can't + # route them anyway — drop them before anything downstream + # (orig_req_body / on_request / wire forward) sees them. + if fmt == FMT_ANTHROPIC and "tools" in request_data: + filtered = _filter_anthropic_user_tools(request_data["tools"]) + if filtered: + request_data["tools"] = filtered + else: + request_data.pop("tools", None) orig_req_body = copy.deepcopy(request_data) trace_enabled = _request_uses_trace_store(request_data) From a8d9618214f92d4adb3bfd88084909844cb20d2c Mon Sep 17 00:00:00 2001 From: braisedpork1964 <497494458@qq.com> Date: Tue, 9 Jun 2026 02:48:04 +0000 Subject: [PATCH 3/9] format --- xtuner/v1/rl/rollout/chat_template.py | 1 + xtuner/v1/rl/rollout/session_server.py | 30 +++++++++++++++++--------- 2 files changed, 21 insertions(+), 10 deletions(-) diff --git a/xtuner/v1/rl/rollout/chat_template.py b/xtuner/v1/rl/rollout/chat_template.py index 0fa21e9df8..ec7b9d2465 100644 --- a/xtuner/v1/rl/rollout/chat_template.py +++ b/xtuner/v1/rl/rollout/chat_template.py @@ -2,6 +2,7 @@ import json from typing import Any + _RAW_ARGUMENTS_KEY = "__xtuner_raw_arguments__" diff --git a/xtuner/v1/rl/rollout/session_server.py b/xtuner/v1/rl/rollout/session_server.py index a79f1b5225..0abe5a91df 100644 --- a/xtuner/v1/rl/rollout/session_server.py +++ b/xtuner/v1/rl/rollout/session_server.py @@ -8,13 +8,14 @@ import numpy as np import ray from aiohttp import ClientConnectionResetError, ClientSession, ClientTimeout, web -from transformers import AutoTokenizer +from transformers import AutoTokenizer from xtuner.v1.utils import get_logger from .chat_template import canonicalize_messages_for_chat_template from .trace_store import TokenizedSegment, get_store + FMT_OPENAI = "openai" FMT_ANTHROPIC = "anthropic" @@ -55,7 +56,8 @@ def _bool_request_value(value: Any, default: bool = False) -> bool: def _request_uses_trace_store(req_body: dict) -> bool: - """Whether this request should drive the trace_store / token-in-token-out path. + """Whether this request should drive the trace_store / token-in-token-out + path. Evaluate-mode callers opt out by setting ``return_token_ids=False`` on the request body; tracing then becomes a passthrough proxy with parameter-name @@ -114,7 +116,8 @@ def _normalize_tool_call_arguments(messages) -> None: def _strip_stop_word(content, stop_word: str): - """Remove ``stop_word`` from a message ``content`` field (any OpenAI shape). + """Remove ``stop_word`` from a message ``content`` field (any OpenAI + shape). OpenAI ``content`` accepts both a plain string and a list of typed parts (``[{"type": "text", "text": "..."}, ...]``); the latter shows up after we @@ -142,7 +145,8 @@ def _strip_stop_word(content, stop_word: str): def _anthropic_response_to_assistant_message(response: dict) -> dict: - """Wrap an Anthropic ``MessagesResponse`` body as an anthropic-shaped assistant message. + """Wrap an Anthropic ``MessagesResponse`` body as an anthropic-shaped + assistant message. The returned dict can be appended to ``MessagesRequest.messages`` so the full conversation (input + this assistant turn) survives a round-trip @@ -187,7 +191,8 @@ def _anthropic_request_to_openai(req_body: dict) -> tuple[list[dict], Optional[l def _filter_anthropic_user_tools(tools) -> Optional[list]: - """Drop Anthropic server-side tools (``web_search_*`` / ``computer_*`` / ``bash_*`` / ``text_editor_*``). + """Drop Anthropic server-side tools (``web_search_*`` / ``computer_*`` / + ``bash_*`` / ``text_editor_*``). ``ToolParam`` in lmdeploy's anthropic protocol requires ``input_schema``; server-side built-ins are managed by Anthropic itself and only carry @@ -209,7 +214,8 @@ def _filter_anthropic_user_tools(tools) -> Optional[list]: class SessionServer: - """SessionServer intercepts and records requests sent to a remote LLM API worker. + """SessionServer intercepts and records requests sent to a remote LLM API + worker. It acts as a reverse-proxy in front of an already running worker (lmdeploy, sglang, vllm). Each supported endpoint is transparently forwarded along @@ -263,7 +269,8 @@ def __init__( self._lmdeploy_actor: Optional[ray.actor.ActorHandle] = None async def on_request(self, req_body: dict, fmt: str, *, trace_enabled: bool = True) -> dict: - """Normalize the request to drive the prefix cache + inject extension fields. + """Normalize the request to drive the prefix cache + inject extension + fields. When ``trace_enabled`` is False (evaluate mode), forward the request as-is with parameter-name hygiene only — no token-in-token-out @@ -360,7 +367,8 @@ async def on_request(self, req_body: dict, fmt: str, *, trace_enabled: bool = Tr return worker_req async def on_response(self, worker_resp: dict, fmt: str, orig_req_body: dict) -> dict: - """Trace the assistant turn into ``trace_store`` (always in OpenAI shape). + """Trace the assistant turn into ``trace_store`` (always in OpenAI + shape). Args: worker_resp (dict): The parsed upstream response (still in @@ -712,7 +720,8 @@ async def _decode_routed_experts(self, routed_experts: Any) -> np.ndarray: def _build_data_cleaner( self, fmt: str, *, orig_return_logprob: bool, orig_return_token_ids: bool, orig_return_routed_experts: bool ): - """Return a per-format ``data -> bool`` cleaner stripping injected fields.""" + """Return a per-format ``data -> bool`` cleaner stripping injected + fields.""" drops = { "output_token_logprobs": not orig_return_logprob, "output_ids": not orig_return_token_ids, @@ -892,7 +901,8 @@ def _parse_openai_stream(events: list[dict], *, saw_done: bool) -> Optional[dict @staticmethod def _parse_anthropic_stream(events: list[dict]) -> Optional[dict]: - """Aggregate Anthropic SSE events into a ``MessagesResponse``-shaped dict. + """Aggregate Anthropic SSE events into a ``MessagesResponse``-shaped + dict. Extension fields: - ``output_ids`` / ``output_token_logprobs`` live on each From aba7530b7a3725fdf0120f148d425789104f7827 Mon Sep 17 00:00:00 2001 From: braisedpork1964 <497494458@qq.com> Date: Wed, 10 Jun 2026 06:54:30 +0000 Subject: [PATCH 4/9] disable role demotion --- xtuner/v1/rl/rollout/chat_template.py | 21 +++------------------ 1 file changed, 3 insertions(+), 18 deletions(-) diff --git a/xtuner/v1/rl/rollout/chat_template.py b/xtuner/v1/rl/rollout/chat_template.py index ec7b9d2465..06cd2b4d30 100644 --- a/xtuner/v1/rl/rollout/chat_template.py +++ b/xtuner/v1/rl/rollout/chat_template.py @@ -16,27 +16,12 @@ def canonicalize_messages_for_chat_template(messages: list[dict]) -> list[dict]: mutate model responses returned to the agent/client, nor the generated token ids, labels, or raw rollout artifacts used for training. - Two transforms: - - - ``tool_calls[].function.arguments`` strings are parsed to dicts so the - template's ``tojson`` filter sees a real object. - - Any ``role='system'`` entry past the first message is demoted to - ``role='user'``. Claude Agent SDK's ``context_management`` injects a JIT - system reminder at the tail before each assistant turn (and those persist - mid-conversation as the rollout progresses); Anthropic's API allows it, - but Qwen-style chat templates require ``system`` only at index 0. - Demoting (rather than dropping) keeps the dynamic CWD/todo/time hints - visible to the policy at the cost of a small semantic shift. + ``tool_calls[].function.arguments`` strings are parsed to dicts so the + template's ``tojson`` filter sees a real object. """ messages = copy.deepcopy(messages) - for idx, message in enumerate(messages): - # TODO(wzy): drop this demote once the Qwen Jinja chat template - # accepts ``role='system'`` at any position. Until then we degrade SDK - # context_management injections to ``role='user'`` to keep the template - # renderable. - if idx > 0 and isinstance(message, dict) and message.get("role") == "system": - message["role"] = "user" + for message in messages: tool_calls = message.get("tool_calls") if not isinstance(tool_calls, list): continue From a0723f5dd4e12ec8812e4559fcbe789c2ce04420 Mon Sep 17 00:00:00 2001 From: braisedpork1964 <497494458@qq.com> Date: Thu, 11 Jun 2026 07:10:21 +0000 Subject: [PATCH 5/9] drop the orphaned placeholder from trie on exception --- xtuner/v1/rl/rollout/session_server.py | 128 +++++++++++++++---------- xtuner/v1/rl/rollout/trace_store.py | 11 ++- xtuner/v1/rl/rollout/worker.py | 1 + 3 files changed, 84 insertions(+), 56 deletions(-) diff --git a/xtuner/v1/rl/rollout/session_server.py b/xtuner/v1/rl/rollout/session_server.py index 0abe5a91df..5cdcfbc6db 100644 --- a/xtuner/v1/rl/rollout/session_server.py +++ b/xtuner/v1/rl/rollout/session_server.py @@ -253,6 +253,7 @@ def __init__( port: int = 8080, request_timeout: float = 1200.0, read_bufsize: int = 2**26, + enable_return_routed_experts: bool = False, ): self.worker_base_url = worker_base_url.rstrip("/") self.tokenizer = AutoTokenizer.from_pretrained(tokenizer_path, trust_remote_code=True) @@ -260,6 +261,7 @@ def __init__( self.port = port self.request_timeout = request_timeout self.read_bufsize = read_bufsize + self.enable_return_routed_experts = enable_return_routed_experts self.store = get_store() self.stop_word = self.tokenizer.eos_token or "" @@ -422,57 +424,71 @@ async def on_response(self, worker_resp: dict, fmt: str, orig_req_body: dict) -> old_prompt = self.tokenizer.apply_chat_template( canonicalize_messages_for_chat_template(messages), tools=tools, add_generation_prompt=True, tokenize=False ) - full_messages = [*messages, assistant_msg] - new_prompt = ( - self.tokenizer.apply_chat_template( - canonicalize_messages_for_chat_template(full_messages), - tools=tools, - add_generation_prompt=False, - tokenize=False, + try: + full_messages = [*messages, assistant_msg] + new_prompt = ( + self.tokenizer.apply_chat_template( + canonicalize_messages_for_chat_template(full_messages), + tools=tools, + add_generation_prompt=False, + tokenize=False, + ) + ).rstrip() + assert new_prompt.startswith(old_prompt) and new_prompt.endswith(self.stop_word) + + if raw_routed_expert is not None: + raw_routed_expert = await self._decode_routed_experts(raw_routed_expert) + if len(raw_routed_expert) > 0: + num_layers = raw_routed_expert.shape[1] + topk_experts = raw_routed_expert.shape[2] + dummy_expert = np.full((1, num_layers, topk_experts), 0, dtype=raw_routed_expert.dtype) + raw_routed_expert = np.concatenate([dummy_expert, raw_routed_expert], axis=0) + + _, nodes = await self.store.search.remote(session_id, old_prompt, filter_none=True) + + # last node in nodes corresponds to the delta inserted in on_request (if any) + if nodes: + delta_node_val: TokenizedSegment = nodes[-1].value + delta_len = len(delta_node_val.token_ids) + prefix_len = sum(len(n.value.token_ids) for n in nodes[:-1]) + assert prefix_len + delta_len + len(output_token_ids) == len(raw_routed_expert) + + delta_expert = raw_routed_expert[prefix_len : prefix_len + delta_len] + response_expert = raw_routed_expert[prefix_len + delta_len :] + + if delta_len > 0: + delta_node_val.expert_key = ray.put(delta_expert) + await self.store.insert.remote(session_id, old_prompt, delta_node_val) + + raw_routed_expert = ray.put(response_expert) + else: + raw_routed_expert = ray.put(raw_routed_expert) + elif self.enable_return_routed_experts and _bool_request_value( + orig_req_body.get("return_routed_experts"), True + ): + raise RuntimeError( + "SessionServer response is missing routed_experts, but enable_return_routed_experts is True." + " The upstream worker may encounter an LLM call error." + ) + + await self.store.insert.remote( + session_id, + key=new_prompt, + value=TokenizedSegment( + text=new_prompt[len(old_prompt) :], + token_ids=output_token_ids, + logprobs=output_logprobs, + labels=output_token_ids, + expert_key=raw_routed_expert, + length=len(output_token_ids), + ), ) - ).rstrip() - assert new_prompt.startswith(old_prompt) and new_prompt.endswith(self.stop_word) - - if raw_routed_expert is not None: - raw_routed_expert = await self._decode_routed_experts(raw_routed_expert) - if len(raw_routed_expert) > 0: - num_layers = raw_routed_expert.shape[1] - topk_experts = raw_routed_expert.shape[2] - dummy_expert = np.full((1, num_layers, topk_experts), 0, dtype=raw_routed_expert.dtype) - raw_routed_expert = np.concatenate([dummy_expert, raw_routed_expert], axis=0) - - _, nodes = await self.store.search.remote(session_id, old_prompt, filter_none=True) - - # last node in nodes corresponds to the delta inserted in on_request (if any) - if nodes: - delta_node_val: TokenizedSegment = nodes[-1].value - delta_len = len(delta_node_val.token_ids) - prefix_len = sum(len(n.value.token_ids) for n in nodes[:-1]) - assert prefix_len + delta_len + len(output_token_ids) == len(raw_routed_expert) - - delta_expert = raw_routed_expert[prefix_len : prefix_len + delta_len] - response_expert = raw_routed_expert[prefix_len + delta_len :] - - if delta_len > 0: - delta_node_val.expert_key = ray.put(delta_expert) - await self.store.insert.remote(session_id, old_prompt, delta_node_val) - - raw_routed_expert = ray.put(response_expert) - else: - raw_routed_expert = ray.put(raw_routed_expert) - - await self.store.insert.remote( - session_id, - key=new_prompt, - value=TokenizedSegment( - text=new_prompt[len(old_prompt) :], - token_ids=output_token_ids, - logprobs=output_logprobs, - labels=output_token_ids, - expert_key=raw_routed_expert, - length=len(output_token_ids), - ), - ) + except Exception: + if self.enable_return_routed_experts: + matched_key, nodes = await self.store.search.remote(session_id, old_prompt, filter_none=True) + if matched_key == old_prompt and nodes and getattr(nodes[-1].value, "expert_key", None) is None: + await self.store.release.remote(session_id, old_prompt) + raise return worker_resp @@ -1008,12 +1024,21 @@ def _parse_anthropic_stream(events: list[dict]) -> Optional[dict]: class SessionServerActor: """Ray actor wrapper that owns one SessionServer instance.""" - def __init__(self, worker_base_url: str, tokenizer_path: str, host: str, port: int, request_timeout: float): + def __init__( + self, + worker_base_url: str, + tokenizer_path: str, + host: str, + port: int, + request_timeout: float, + enable_return_routed_experts: bool = False, + ): self.worker_base_url = worker_base_url self.tokenizer_path = tokenizer_path self.host = host self.port = port self.request_timeout = request_timeout + self.enable_return_routed_experts = enable_return_routed_experts self.server: SessionServer | None = None @property @@ -1030,6 +1055,7 @@ async def start(self) -> str: host=self.host, port=self.port, request_timeout=self.request_timeout, + enable_return_routed_experts=self.enable_return_routed_experts, ) await self.server.start() return self.server.url diff --git a/xtuner/v1/rl/rollout/trace_store.py b/xtuner/v1/rl/rollout/trace_store.py index 719590bd80..1ed9937ade 100644 --- a/xtuner/v1/rl/rollout/trace_store.py +++ b/xtuner/v1/rl/rollout/trace_store.py @@ -295,16 +295,17 @@ def search(self, session_id: str, text: str, filter_none: bool = False): trie = self.get_or_create(session_id) return trie.search(text, filter_none) - def release(self, session_id: str): - """Release the Trie and free associated resources for a specific - session. + def release(self, session_id: str, key: str | None = None): + """Release Trie nodes and free associated resources for a session. Args: session_id (str): The session identifier. + key (str | None): If None, release the whole session Trie and drop the session. Otherwise release only + the node at ``key`` (used by the SessionServer to drop an orphaned placeholder), keeping the session. """ assert session_id in self.sessions, f"Session ID '{session_id}' not found for release." - trie = self.sessions.pop(session_id) - trie.release() + trie = self.sessions.pop(session_id) if key is None else self.sessions[session_id] + trie.release(key) def release_all(self): """Release all sessions and free associated resources.""" diff --git a/xtuner/v1/rl/rollout/worker.py b/xtuner/v1/rl/rollout/worker.py index 4d0ab16664..dde51ed32e 100644 --- a/xtuner/v1/rl/rollout/worker.py +++ b/xtuner/v1/rl/rollout/worker.py @@ -698,6 +698,7 @@ def _start_session_server(self) -> None: host=self.host, port=self.session_server_port, request_timeout=self.config.session_server_timeout, + enable_return_routed_experts=self.enable_return_routed_experts, ) ) self.session_server_url = ray.get( From 30a17cd0776691f08aecfd11b03b7b23f40a0678 Mon Sep 17 00:00:00 2001 From: braisedpork1964 <497494458@qq.com> Date: Thu, 25 Jun 2026 12:30:53 +0000 Subject: [PATCH 6/9] remove the orphaned placeholder on null output_token_ids --- xtuner/v1/rl/rollout/session_server.py | 12 ++++++------ 1 file changed, 6 insertions(+), 6 deletions(-) diff --git a/xtuner/v1/rl/rollout/session_server.py b/xtuner/v1/rl/rollout/session_server.py index 5cdcfbc6db..c086eb329d 100644 --- a/xtuner/v1/rl/rollout/session_server.py +++ b/xtuner/v1/rl/rollout/session_server.py @@ -408,12 +408,6 @@ async def on_response(self, worker_resp: dict, fmt: str, orig_req_body: dict) -> messages = orig_req_body["messages"] tools = orig_req_body.get("tools", None) - if output_token_ids is None: - raise RuntimeError( - "SessionServer response has no output_ids; cannot export a training trace for this assistant turn." - ) - output_logprobs = _extract_output_logprobs(output_token_logprobs, output_token_ids) - if isinstance(assistant_msg.get("content"), (str, list)): assistant_msg["content"], _ = _strip_stop_word(assistant_msg["content"], self.stop_word) for msg in messages: @@ -425,6 +419,12 @@ async def on_response(self, worker_resp: dict, fmt: str, orig_req_body: dict) -> canonicalize_messages_for_chat_template(messages), tools=tools, add_generation_prompt=True, tokenize=False ) try: + if output_token_ids is None: + raise RuntimeError( + "SessionServer response has no output_ids; cannot export a training trace for this assistant turn." + ) + output_logprobs = _extract_output_logprobs(output_token_logprobs, output_token_ids) + full_messages = [*messages, assistant_msg] new_prompt = ( self.tokenizer.apply_chat_template( From 52446c03ea5378506f4f7d905ac58e38c6dde3fd Mon Sep 17 00:00:00 2001 From: braisedpork1964 <497494458@qq.com> Date: Thu, 9 Jul 2026 06:02:03 +0000 Subject: [PATCH 7/9] ensure trace-store nodes never carry a null expert_key via single-phase write --- xtuner/v1/rl/rollout/session_server.py | 147 +++++++++++++------------ xtuner/v1/rl/rollout/trace_store.py | 38 ++++++- 2 files changed, 110 insertions(+), 75 deletions(-) diff --git a/xtuner/v1/rl/rollout/session_server.py b/xtuner/v1/rl/rollout/session_server.py index c086eb329d..69d658de06 100644 --- a/xtuner/v1/rl/rollout/session_server.py +++ b/xtuner/v1/rl/rollout/session_server.py @@ -332,14 +332,17 @@ async def on_request(self, req_body: dict, fmt: str, *, trace_enabled: bool = Tr tokenize=False, ) - # 2. Prefix-cache lookup; insert the delta tail back into the store. + # 2. Prefix-cache lookup. The delta tail is NOT persisted here: writing a + # node before generation would leave it with ``expert_key=None`` (routed + # experts don't exist yet), which is exactly the orphan the trace store + # must never keep. ``delta_ids`` is only needed locally to build + # ``input_ids``; ``on_response`` re-derives and persists the delta node + # once its routed experts are known (single-phase write). prefix, nodes = await self.store.search.remote(session_id, prompt_text, filter_none=True) if prefix: get_logger().debug(f"Hit prefix cache for session {session_id}") - delta, delta_ids = prompt_text[len(prefix) :], [] - if delta: - delta_ids = self.tokenizer.encode(delta, add_special_tokens=False) - await self.store.insert.remote(session_id, prompt_text, TokenizedSegment(text=delta, token_ids=delta_ids)) + delta = prompt_text[len(prefix) :] + delta_ids = self.tokenizer.encode(delta, add_special_tokens=False) if delta else [] input_ids = reduce(add, [node.value.token_ids for node in nodes] + [delta_ids]) # 3. Inject extension fields on the forwarded request. The body still @@ -418,77 +421,77 @@ async def on_response(self, worker_resp: dict, fmt: str, orig_req_body: dict) -> old_prompt = self.tokenizer.apply_chat_template( canonicalize_messages_for_chat_template(messages), tools=tools, add_generation_prompt=True, tokenize=False ) - try: - if output_token_ids is None: - raise RuntimeError( - "SessionServer response has no output_ids; cannot export a training trace for this assistant turn." - ) - output_logprobs = _extract_output_logprobs(output_token_logprobs, output_token_ids) - - full_messages = [*messages, assistant_msg] - new_prompt = ( - self.tokenizer.apply_chat_template( - canonicalize_messages_for_chat_template(full_messages), - tools=tools, - add_generation_prompt=False, - tokenize=False, - ) - ).rstrip() - assert new_prompt.startswith(old_prompt) and new_prompt.endswith(self.stop_word) - - if raw_routed_expert is not None: - raw_routed_expert = await self._decode_routed_experts(raw_routed_expert) - if len(raw_routed_expert) > 0: - num_layers = raw_routed_expert.shape[1] - topk_experts = raw_routed_expert.shape[2] - dummy_expert = np.full((1, num_layers, topk_experts), 0, dtype=raw_routed_expert.dtype) - raw_routed_expert = np.concatenate([dummy_expert, raw_routed_expert], axis=0) - - _, nodes = await self.store.search.remote(session_id, old_prompt, filter_none=True) - - # last node in nodes corresponds to the delta inserted in on_request (if any) - if nodes: - delta_node_val: TokenizedSegment = nodes[-1].value - delta_len = len(delta_node_val.token_ids) - prefix_len = sum(len(n.value.token_ids) for n in nodes[:-1]) - assert prefix_len + delta_len + len(output_token_ids) == len(raw_routed_expert) - - delta_expert = raw_routed_expert[prefix_len : prefix_len + delta_len] - response_expert = raw_routed_expert[prefix_len + delta_len :] - - if delta_len > 0: - delta_node_val.expert_key = ray.put(delta_expert) - await self.store.insert.remote(session_id, old_prompt, delta_node_val) - - raw_routed_expert = ray.put(response_expert) - else: - raw_routed_expert = ray.put(raw_routed_expert) - elif self.enable_return_routed_experts and _bool_request_value( - orig_req_body.get("return_routed_experts"), True - ): - raise RuntimeError( - "SessionServer response is missing routed_experts, but enable_return_routed_experts is True." - " The upstream worker may encounter an LLM call error." - ) + if output_token_ids is None: + raise RuntimeError( + "SessionServer response has no output_ids; cannot export a training trace for this assistant turn." + ) + output_logprobs = _extract_output_logprobs(output_token_logprobs, output_token_ids) + + full_messages = [*messages, assistant_msg] + new_prompt = ( + self.tokenizer.apply_chat_template( + canonicalize_messages_for_chat_template(full_messages), + tools=tools, + add_generation_prompt=False, + tokenize=False, + ) + ).rstrip() + assert new_prompt.startswith(old_prompt) and new_prompt.endswith(self.stop_word) + + # Re-derive the input delta the worker processed. on_request no longer + # persists it (that would leave a None-expert placeholder in the store), + # so recompute it from the current prefix and split the routed experts + # across (prefix, delta, response) using the same boundaries on_request + # used to build ``input_ids``. + matched_prefix, prefix_nodes = await self.store.search.remote(session_id, old_prompt, filter_none=True) + delta_text = old_prompt[len(matched_prefix) :] + delta_ids = self.tokenizer.encode(delta_text, add_special_tokens=False) if delta_text else [] + prefix_len = sum(len(n.value.token_ids) for n in prefix_nodes) + delta_len = len(delta_ids) + + delta_expert: Any = None + response_expert: Any = None + if raw_routed_expert is not None: + raw_routed_expert = await self._decode_routed_experts(raw_routed_expert) + if len(raw_routed_expert) > 0: + num_layers = raw_routed_expert.shape[1] + topk_experts = raw_routed_expert.shape[2] + dummy_expert = np.full((1, num_layers, topk_experts), 0, dtype=raw_routed_expert.dtype) + raw_routed_expert = np.concatenate([dummy_expert, raw_routed_expert], axis=0) + assert prefix_len + delta_len + len(output_token_ids) == len(raw_routed_expert) + if delta_len > 0: + delta_expert = ray.put(raw_routed_expert[prefix_len : prefix_len + delta_len]) + response_expert = ray.put(raw_routed_expert[prefix_len + delta_len :]) + elif self.enable_return_routed_experts and _bool_request_value( + orig_req_body.get("return_routed_experts"), True + ): + raise RuntimeError( + "SessionServer response is missing routed_experts, but enable_return_routed_experts is True." + " The upstream worker may encounter an LLM call error." + ) + # Single-phase write: nodes are persisted only now, already carrying their + # final ``expert_key`` (a valid ref for MoE, None for dense). Every assert / + # raise above runs before any write, so a dropped turn leaves nothing + # half-written in the store (and no orphan placeholder to clean up). + if delta_ids: await self.store.insert.remote( session_id, - key=new_prompt, - value=TokenizedSegment( - text=new_prompt[len(old_prompt) :], - token_ids=output_token_ids, - logprobs=output_logprobs, - labels=output_token_ids, - expert_key=raw_routed_expert, - length=len(output_token_ids), - ), + key=old_prompt, + value=TokenizedSegment(text=delta_text, token_ids=delta_ids, expert_key=delta_expert), ) - except Exception: - if self.enable_return_routed_experts: - matched_key, nodes = await self.store.search.remote(session_id, old_prompt, filter_none=True) - if matched_key == old_prompt and nodes and getattr(nodes[-1].value, "expert_key", None) is None: - await self.store.release.remote(session_id, old_prompt) - raise + await self.store.insert.remote( + session_id, + key=new_prompt, + value=TokenizedSegment( + text=new_prompt[len(old_prompt) :], + token_ids=output_token_ids, + logprobs=output_logprobs, + labels=output_token_ids, + expert_key=response_expert, + length=len(output_token_ids), + ), + ) return worker_resp diff --git a/xtuner/v1/rl/rollout/trace_store.py b/xtuner/v1/rl/rollout/trace_store.py index 1ed9937ade..8fed1119a8 100644 --- a/xtuner/v1/rl/rollout/trace_store.py +++ b/xtuner/v1/rl/rollout/trace_store.py @@ -66,6 +66,34 @@ def _mismatch_window(left: str, right: str, *, radius: int = 160) -> dict[str, A } +def _resolve_routed_experts(expert_keys: List[Any], session_id: str) -> List[Any] | None: + """Collapse per-node ``expert_key`` values into the exported + ``routed_experts``. + + The SessionServer's single-phase write guarantees that, within one session, + either every trace node carries routed experts (MoE) or none of them do + (dense). This helper enforces that invariant at export time. + + Args: + expert_keys (List[Any]): One ``expert_key`` per trace node, in path order. + session_id (str): The session identifier, used only for error messages. + + Returns: + List[Any] | None: The per-node expert refs for a MoE trace, or ``None`` + when the trace carries no routed experts at all (dense model), so the + trainer skips routed-expert consumption entirely. + """ + present = [expert_key is not None for expert_key in expert_keys] + if not any(present): + return None + if all(present): + return expert_keys + raise ValueError( + f"Inconsistent routed_experts in session {session_id!r}: {sum(present)}/{len(present)} trace nodes carry " + "routed experts; the trace mixes MoE turns with missing-expert turns and cannot be exported." + ) + + class TokenizedSegment(BaseModel): model_config = ConfigDict(extra="forbid") @@ -324,7 +352,9 @@ def export_training_trace(self, session_id: str, prompt_text: str) -> dict: Returns: dict: The trace dictionary containing `input_ids`, `labels`, `logprobs`, - and `routed_experts`. + and `routed_experts`. `routed_experts` is a per-node list of expert-id + refs for a MoE trace, or `None` when the trace carries no routed + experts at all (dense model). Raises: ValueError: If the prompt_text does not completely match the trace keys in the session. @@ -354,7 +384,8 @@ def export_training_trace(self, session_id: str, prompt_text: str) -> dict: f"prompt_len={len(prompt_text)} matched_len={len(key)} key_count={len(session_keys)}. " "See the logged '[TraceStore] prompt mismatch' report for the full diff." ) - trace: dict[str, list[Any]] = {"input_ids": [], "labels": [], "logprobs": [], "routed_experts": []} + trace: dict[str, Any] = {"input_ids": [], "labels": [], "logprobs": []} + expert_keys: list[Any] = [] for node in nodes: node_val = node.value if not isinstance(node_val, TokenizedSegment): @@ -364,7 +395,8 @@ def export_training_trace(self, session_id: str, prompt_text: str) -> dict: trace["input_ids"].extend(node_val.token_ids) trace["labels"].extend(node_val.labels) trace["logprobs"].extend(node_val.logprobs) - trace["routed_experts"].append(node_val.expert_key) + expert_keys.append(node_val.expert_key) + trace["routed_experts"] = _resolve_routed_experts(expert_keys, session_id) return trace def get_objects(self, keys: list[str]) -> list[ray.ObjectRef]: From a6fc67d8a8e1b6a7e3b34702d7e4f9fcc8a0ce1f Mon Sep 17 00:00:00 2001 From: braisedpork1964 <497494458@qq.com> Date: Thu, 9 Jul 2026 06:07:11 +0000 Subject: [PATCH 8/9] emit format-native error payload so dropped turns aren't recorded downstream --- xtuner/v1/rl/rollout/session_server.py | 41 ++++++++++++++++---------- 1 file changed, 25 insertions(+), 16 deletions(-) diff --git a/xtuner/v1/rl/rollout/session_server.py b/xtuner/v1/rl/rollout/session_server.py index 69d658de06..16177ecf03 100644 --- a/xtuner/v1/rl/rollout/session_server.py +++ b/xtuner/v1/rl/rollout/session_server.py @@ -33,13 +33,19 @@ def _is_error_payload(payload: dict) -> bool: return payload.get("error") is not None or payload.get("type") == "error" or payload.get("object") == "error" -def _lmdeploy_error_payload(message: str, status: int = 500, error_type: str = "internal_server_error") -> dict: - return { - "message": message, - "type": error_type, - "code": status, - "object": "error", - } +def _error_payload(fmt: str, message: str, status: int = 500, error_type: str = "internal_server_error") -> dict: + """Error body in the caller's native shape. + + The downstream SessionClient recorder drops a turn only when it recognizes the + error shape: Anthropic keys on top-level ``type == "error"``, OpenAI/lmdeploy on + ``object == "error"``. Emitting the matching shape keeps a turn the SessionServer + dropped (e.g. an ``on_response`` that raised on missing routed experts) out of the + client's records, so the trace store stays a superset of them and export never + sees a trajectory whose nodes were never written. + """ + if fmt == FMT_ANTHROPIC: + return {"type": "error", "error": {"type": error_type, "message": message}} + return {"message": message, "type": error_type, "code": status, "object": "error"} def _bool_request_value(value: Any, default: bool = False) -> bool: @@ -527,11 +533,13 @@ async def _handle_request(self, request: web.Request) -> web.Response: """Proxy handler: detect format, run hooks, forward, stream back.""" req_path = request.match_info["path"] + fmt = _detect_format(req_path) # Reject /v1/responses outright — upstream worker doesn't implement it. if req_path.endswith("/responses") or "/v1/responses" in req_path: return web.json_response( - _lmdeploy_error_payload( + _error_payload( + fmt, "/v1/responses is not supported by SessionServer (upstream worker has no responses endpoint).", status=HTTPStatus.NOT_IMPLEMENTED, error_type="not_implemented", @@ -539,8 +547,6 @@ async def _handle_request(self, request: web.Request) -> web.Response: status=HTTPStatus.NOT_IMPLEMENTED, ) - fmt = _detect_format(req_path) - # Read the request body request_body = await request.read() request_data = None @@ -580,7 +586,7 @@ async def _handle_request(self, request: web.Request) -> web.Response: except Exception as exc: message = f"SessionServer request hook failed: {type(exc).__name__}: {exc}" get_logger().error(message) - return web.json_response(_lmdeploy_error_payload(message), status=500) + return web.json_response(_error_payload(fmt, message), status=500) # Build forwarding headers, dropping original Host / Content-Length. forward_headers = dict(request.headers) @@ -712,10 +718,13 @@ async def _handle_request(self, request: web.Request) -> web.Response: if is_stream: try: if session_error_msg: - error_payload = _lmdeploy_error_payload(session_error_msg) - await response.write( - ("data: " + json.dumps(error_payload, ensure_ascii=False) + "\n\n").encode("utf-8") - ) + error_payload = _error_payload(fmt, session_error_msg) + error_line = "data: " + json.dumps(error_payload, ensure_ascii=False) + "\n\n" + # Anthropic SSE carries a named ``event: error`` line; the recorder keys on the ``data:`` JSON + # either way, but a real Anthropic SDK consumer needs the event name. + if fmt == FMT_ANTHROPIC: + error_line = "event: error\n" + error_line + await response.write(error_line.encode("utf-8")) skip_done = True if not skip_done: await response.write(b"data: [DONE]\n\n") @@ -723,7 +732,7 @@ async def _handle_request(self, request: web.Request) -> web.Response: except (ConnectionError, ClientConnectionResetError): pass elif session_error_msg: - return web.json_response(_lmdeploy_error_payload(session_error_msg), status=500) + return web.json_response(_error_payload(fmt, session_error_msg), status=500) return response From bad6d5fd8162e513867bcdbd2bd6b3ffdc705cae Mon Sep 17 00:00:00 2001 From: braisedpork1964 <497494458@qq.com> Date: Mon, 13 Jul 2026 07:19:41 +0000 Subject: [PATCH 9/9] fix typos --- xtuner/v1/rl/rollout/session_server.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/xtuner/v1/rl/rollout/session_server.py b/xtuner/v1/rl/rollout/session_server.py index 16177ecf03..9914d734af 100644 --- a/xtuner/v1/rl/rollout/session_server.py +++ b/xtuner/v1/rl/rollout/session_server.py @@ -51,8 +51,8 @@ def _error_payload(fmt: str, message: str, status: int = 500, error_type: str = def _bool_request_value(value: Any, default: bool = False) -> bool: """Coerce an opaque request-body value (bool / int / str / None) to bool. - Mirrors how lmdeploy validates flag-style fields — accepts ``"true"``, - ``"1"`` etc. as truthy and ``"false"``, ``"0"`` as falsy. + Mirrors how lmdeploy validates flag-style fields — accepts ``"true"``, ``"1"`` etc. as true, + ``"false"``, ``"0"``, ``""`` etc. as false, and falls back to the default value for unrecognized strings. """ if value is None: return default