diff --git a/recipe/__init__.py b/recipe/__init__.py
deleted file mode 100644
index e69de29bb2..0000000000
diff --git a/recipe/claude_code/calculator_tool.py b/recipe/claude_code/calculator_tool.py
new file mode 100644
index 0000000000..ecb28eeb62
--- /dev/null
+++ b/recipe/claude_code/calculator_tool.py
@@ -0,0 +1,140 @@
+from __future__ import annotations
+
+import json
+import re
+import sys
+from pathlib import Path
+
+from xtuner.v1.data_proto import RolloutState
+from xtuner.v1.rl.judger.native import Judger
+
+
+CALCULATOR_TOOL_NAME = "mcp__calculator__calculator"
+CALCULATOR_PROMPT = """You are NOT allowed to do arithmetic yourself.
+
+You MUST use the calculator tool to compute the result.
+The calculator tool is named mcp__calculator__calculator.
+Your first assistant response MUST be exactly one structured tool call and nothing else:
+
+
+
+23 + 19
+
+
+
+
+Question:
+What is 23 + 19?
+
+Return only the final answer."""
+
+CALCULATOR_SYSTEM_PROMPT = """You are testing an agent loop tool-calling path.
+The only successful behavior is:
+1. Call the mcp__calculator__calculator tool with {"expression": "23 + 19"}.
+2. Read the tool result.
+3. Return only the final answer as plain text: 42.
+
+Do not solve arithmetic directly. Do not describe a tool call in prose.
+Do not generate a title. Do not use boxed answer formatting.
+For Qwen-style tool calls, use this exact XML form:
+
+
+
+23 + 19
+
+
+"""
+
+
+class CalculatorJudger(Judger):
+ async def judge(self, rollout_state: RolloutState) -> RolloutState:
+ stdout = rollout_state.extra_fields.get("claudecode_cli_stdout") or ""
+ answer = ""
+ if stdout:
+ try:
+ answer = normalize_answer(json.loads(stdout).get("result"))
+ except Exception:
+ answer = normalize_answer(stdout)
+ rollout_state.reward = {
+ "score": 1.0 if answer == "42" else 0.0,
+ "answer": answer,
+ }
+ return rollout_state
+
+
+def normalize_answer(value: object) -> str:
+ text = "" if value is None else str(value)
+ text = text.strip().strip("`").strip()
+ boxed = re.search(r"\\boxed\{([^{}]+)\}", text)
+ if boxed:
+ return boxed.group(1).strip()
+ final_answer = re.search(r"final answer\s*:\s*([^\n]+)", text, flags=re.IGNORECASE)
+ if final_answer:
+ text = final_answer.group(1).strip()
+ return text.strip().strip("`").strip()
+
+
+def write_calculator_mcp_server(work_dir: Path) -> tuple[Path, Path]:
+ mcp_server_path = work_dir / "calculator_mcp_server.py"
+ mcp_config_path = work_dir / "calculator_mcp_config.json"
+ mcp_server_path.write_text(
+ """
+from __future__ import annotations
+
+import ast
+import operator
+
+from fastmcp import FastMCP
+
+
+mcp = FastMCP("calculator")
+
+OPS = {
+ ast.Add: operator.add,
+ ast.Sub: operator.sub,
+ ast.Mult: operator.mul,
+ ast.Div: operator.truediv,
+ ast.USub: operator.neg,
+}
+
+
+def _eval(node):
+ if isinstance(node, ast.Expression):
+ return _eval(node.body)
+ if isinstance(node, ast.Constant) and isinstance(node.value, (int, float)):
+ return node.value
+ if isinstance(node, ast.BinOp) and type(node.op) in OPS:
+ return OPS[type(node.op)](_eval(node.left), _eval(node.right))
+ if isinstance(node, ast.UnaryOp) and type(node.op) in OPS:
+ return OPS[type(node.op)](_eval(node.operand))
+ raise ValueError("Only simple arithmetic expressions are supported.")
+
+
+@mcp.tool(name="calculator", description="Evaluate a simple arithmetic expression")
+def calculator(expression: str) -> str:
+ value = _eval(ast.parse(expression, mode="eval"))
+ if isinstance(value, float) and value.is_integer():
+ value = int(value)
+ return str(value)
+
+
+if __name__ == "__main__":
+ mcp.run()
+""".lstrip(),
+ encoding="utf-8",
+ )
+ mcp_config_path.write_text(
+ json.dumps(
+ {
+ "mcpServers": {
+ "calculator": {
+ "command": sys.executable,
+ "args": [str(mcp_server_path)],
+ }
+ }
+ },
+ indent=2,
+ ),
+ encoding="utf-8",
+ )
+ return mcp_server_path, mcp_config_path
diff --git a/recipe/claude_code/claudecode_agent_loop.py b/recipe/claude_code/claudecode_agent_loop.py
new file mode 100644
index 0000000000..0d0ce24379
--- /dev/null
+++ b/recipe/claude_code/claudecode_agent_loop.py
@@ -0,0 +1,350 @@
+from __future__ import annotations
+
+import asyncio
+import copy
+import os
+from pathlib import Path
+from typing import Any
+from uuid import uuid4
+
+import httpx
+from pydantic import ConfigDict, Field
+
+from xtuner.v1.data_proto import RolloutState, SampleParams, Status
+from xtuner.v1.rl.agent_loop.agent_loop import AgentLoop, AgentLoopConfig
+from xtuner.v1.rl.judger.native import Judger
+from xtuner.v1.rl.rollout import RolloutController
+from xtuner.v1.rl.utils import chat_trace_records_to_rollout_states
+
+
+DEFAULT_READONLY_INSTRUCTION = (
+ "You are running inside an automated rollout collection job. "
+ "Work in read-only mode: inspect files and report findings, but do not edit, create, delete, move, "
+ "format, commit, push, install dependencies, or run commands that write to the repository or external services."
+)
+
+
+class ClaudeCodeAgentLoopConfig(AgentLoopConfig):
+ model_config = ConfigDict(extra="forbid", arbitrary_types_allowed=True)
+
+ claude_command: list[str] = Field(default_factory=lambda: ["$HOME/.local/bin/claude"])
+ cwd: str | None = None
+ timeout_s: float = 600.0
+ api_timeout_ms: int = 600000
+ max_turns: int = 5
+ output_format: str = "json"
+ permission_mode: str = "plan"
+ tools: str | None = "Read,Grep,Glob,LS,Bash"
+ allowed_tools: str | None = None
+ disallowed_tools: str | None = "Edit,Write,MultiEdit,NotebookEdit"
+ mcp_config: list[str] = Field(default_factory=list)
+ strict_mcp_config: bool = False
+ system_prompt: str | None = None
+ append_system_prompt: str | None = None
+ readonly_instruction: str = DEFAULT_READONLY_INSTRUCTION
+ bare: bool = True
+ extra_env: dict[str, str] = Field(default_factory=dict)
+
+ def build_local(
+ self,
+ rollout_controller,
+ judger: Judger | None = None,
+ logger=None,
+ ) -> ClaudeCodeAgentLoop:
+ return ClaudeCodeAgentLoop(
+ claude_command=self.claude_command,
+ cwd=self.cwd,
+ timeout_s=self.timeout_s,
+ api_timeout_ms=self.api_timeout_ms,
+ max_turns=self.max_turns,
+ output_format=self.output_format,
+ permission_mode=self.permission_mode,
+ tools=self.tools,
+ allowed_tools=self.allowed_tools,
+ disallowed_tools=self.disallowed_tools,
+ mcp_config=self.mcp_config,
+ strict_mcp_config=self.strict_mcp_config,
+ system_prompt=self.system_prompt,
+ append_system_prompt=self.append_system_prompt,
+ readonly_instruction=self.readonly_instruction,
+ bare=self.bare,
+ extra_env=self.extra_env,
+ rollout_ctl=rollout_controller,
+ sample_params=self.sample_params,
+ hf_checkpoint=self.hf_checkpoint,
+ judger=judger,
+ logger=logger,
+ )
+
+
+class ClaudeCodeAgentLoop(AgentLoop):
+ def __init__(
+ self,
+ claude_command: list[str],
+ cwd: str | None,
+ timeout_s: float,
+ api_timeout_ms: int,
+ max_turns: int,
+ output_format: str,
+ permission_mode: str,
+ tools: str | None,
+ allowed_tools: str | None,
+ disallowed_tools: str | None,
+ mcp_config: list[str],
+ strict_mcp_config: bool,
+ system_prompt: str | None,
+ append_system_prompt: str | None,
+ readonly_instruction: str,
+ bare: bool,
+ extra_env: dict[str, str],
+ rollout_ctl: RolloutController,
+ sample_params: SampleParams,
+ hf_checkpoint: str,
+ judger: Judger | None = None,
+ logger=None,
+ ) -> None:
+ super().__init__(
+ rollout_ctl=rollout_ctl,
+ sample_params=sample_params,
+ hf_checkpoint=hf_checkpoint,
+ judger=judger,
+ logger=logger,
+ )
+ self.claude_command = claude_command
+ self.cwd = cwd
+ self.timeout_s = timeout_s
+ self.api_timeout_ms = api_timeout_ms
+ self.max_turns = max_turns
+ self.output_format = output_format
+ self.permission_mode = permission_mode
+ self.tools = tools
+ self.allowed_tools = allowed_tools
+ self.disallowed_tools = disallowed_tools
+ self.mcp_config = mcp_config
+ self.strict_mcp_config = strict_mcp_config
+ self.system_prompt = system_prompt
+ self.append_system_prompt = append_system_prompt
+ self.readonly_instruction = readonly_instruction
+ self.bare = bare
+ self.extra_env = extra_env
+
+ async def generate_sample( # type: ignore[override]
+ self, rollout_state: RolloutState, **kwargs
+ ) -> list[RolloutState]:
+ try:
+ metadata = await self.rollout_ctl.get_rollout_metadata.remote() # type: ignore[attr-defined]
+ gateway_url = metadata.get("api_server_url")
+ rollout_config = metadata.get("rollout_config")
+ model_name = getattr(rollout_config, "model_name", None) or "rollout-controller"
+ if not gateway_url:
+ return [
+ self._failed_state(
+ rollout_state,
+ "Gateway is not started. Configure GatewayConfig(auto_start=True) "
+ "before using ClaudeCodeAgentLoop.",
+ )
+ ]
+
+ api_key = f"claudecode_{uuid4().hex}"
+ command = self._build_command(rollout_state, model_name=model_name)
+ returncode, stdout, stderr = await self._run_claude(command, gateway_url, model_name, api_key)
+ records = await self._pop_trace_store_records(gateway_url, api_key)
+ rollout_extra_fields = {
+ "claudecode_api_key": api_key,
+ "claudecode_cli_returncode": returncode,
+ "claudecode_cli_stdout": self._truncate(stdout),
+ "claudecode_cli_stderr": self._truncate(stderr),
+ }
+
+ if not records:
+ reason = "Claude Code finished without trace store records for this api_key."
+ if returncode != 0:
+ reason += f" returncode={returncode}, stderr={self._truncate(stderr)}"
+ return [self._failed_state(rollout_state, reason, extra_fields=rollout_extra_fields)]
+
+ reward = None
+ if self.judger is not None:
+ judge_state = rollout_state.model_copy(deep=True)
+ judge_state.extra_fields = {
+ **copy.deepcopy(rollout_state.extra_fields),
+ **copy.deepcopy(rollout_extra_fields),
+ }
+ judged_state = await self.judger.judge(judge_state)
+ if judged_state.reward is None:
+ return [
+ self._failed_state(
+ rollout_state,
+ "Judger completed without setting reward.",
+ extra_fields=rollout_extra_fields,
+ )
+ ]
+ reward = copy.deepcopy(judged_state.reward)
+
+ states = chat_trace_records_to_rollout_states(
+ rollout_state=rollout_state,
+ records=records,
+ tokenizer=self.tokenizer,
+ extra_fields=rollout_extra_fields,
+ )
+ if not states:
+ return [
+ self._failed_state(
+ rollout_state,
+ "Gateway trace records did not contain trainable turns.",
+ extra_fields=rollout_extra_fields,
+ )
+ ]
+
+ completed_states = [state for state in states if state.status == Status.COMPLETED]
+ if reward is not None:
+ for state in completed_states:
+ state.reward = copy.deepcopy(reward)
+ return states
+ except Exception as exc:
+ return [self._failed_state(rollout_state, f"ClaudeCodeAgentLoop failed: {exc}")]
+
+ def _build_command(self, rollout_state: RolloutState, *, model_name: str) -> list[str]:
+ command = [os.path.expandvars(os.path.expanduser(part)) for part in self.claude_command]
+ prompt = self._build_prompt(rollout_state)
+ if self.bare:
+ command.append("--bare")
+ if self.system_prompt:
+ command.extend(["--system-prompt", self.system_prompt])
+ if self.append_system_prompt:
+ command.extend(["--append-system-prompt", self.append_system_prompt])
+ for config in self.mcp_config:
+ command.extend(["--mcp-config", os.path.expandvars(os.path.expanduser(config))])
+ if self.strict_mcp_config:
+ command.append("--strict-mcp-config")
+ command.extend(
+ [
+ "-p",
+ prompt,
+ "--output-format",
+ self.output_format,
+ "--permission-mode",
+ self.permission_mode,
+ "--model",
+ model_name,
+ "--max-turns",
+ str(self.max_turns),
+ "--no-session-persistence",
+ ]
+ )
+ if self.tools is not None:
+ command.extend(["--tools", self.tools])
+ if self.allowed_tools:
+ command.extend(["--allowedTools", self.allowed_tools])
+ if self.disallowed_tools:
+ command.extend(["--disallowedTools", self.disallowed_tools])
+ return command
+
+ def _build_prompt(self, rollout_state: RolloutState) -> str:
+ content = ""
+ for message in reversed(rollout_state.message):
+ if message.get("role") == "user":
+ content = self._message_content_to_text(message.get("content"))
+ break
+ if not content and rollout_state.message:
+ content = self._message_content_to_text(rollout_state.message[-1].get("content"))
+ if not self.readonly_instruction:
+ return content
+ return f"{self.readonly_instruction}\n\nTask:\n{content}"
+
+ def _message_content_to_text(self, content: Any) -> str:
+ if content is None:
+ return ""
+ if isinstance(content, str):
+ return content
+ if isinstance(content, list):
+ parts = []
+ for item in content:
+ if isinstance(item, dict):
+ if "text" in item:
+ parts.append(str(item["text"]))
+ elif item.get("type") == "text":
+ parts.append(str(item.get("text", "")))
+ else:
+ parts.append(str(item))
+ else:
+ parts.append(str(item))
+ return "\n".join(part for part in parts if part)
+ return str(content)
+
+ async def _run_claude(
+ self,
+ command: list[str],
+ gateway_url: str,
+ model_name: str,
+ api_key: str,
+ ) -> tuple[int, str, str]:
+ env = os.environ.copy()
+ env.update(
+ {
+ "ANTHROPIC_BASE_URL": gateway_url,
+ "ANTHROPIC_AUTH_TOKEN": api_key,
+ "ANTHROPIC_API_KEY": api_key,
+ "ANTHROPIC_MODEL": model_name,
+ "API_TIMEOUT_MS": str(self.api_timeout_ms),
+ "PATH": f"{Path.home() / '.local' / 'bin'}:{env.get('PATH', '')}",
+ }
+ )
+ env.update({key: os.path.expandvars(os.path.expanduser(value)) for key, value in self.extra_env.items()})
+
+ process = await asyncio.create_subprocess_exec(
+ *command,
+ cwd=str(Path(self.cwd or os.getcwd()).resolve()),
+ env=env,
+ stdout=asyncio.subprocess.PIPE,
+ stderr=asyncio.subprocess.PIPE,
+ )
+ try:
+ stdout_bytes, stderr_bytes = await asyncio.wait_for(process.communicate(), timeout=self.timeout_s)
+ except asyncio.TimeoutError:
+ process.kill()
+ stdout_bytes, stderr_bytes = await process.communicate()
+ returncode = process.returncode if process.returncode is not None else -9
+ stderr = stderr_bytes.decode("utf-8", errors="replace")
+ stderr = f"Claude Code timed out after {self.timeout_s}s.\n{stderr}"
+ return returncode, stdout_bytes.decode("utf-8", errors="replace"), stderr
+
+ return (
+ process.returncode if process.returncode is not None else 0,
+ stdout_bytes.decode("utf-8", errors="replace"),
+ stderr_bytes.decode("utf-8", errors="replace"),
+ )
+
+ async def _pop_trace_store_records(self, gateway_url: str, api_key: str) -> list[dict[str, Any]]:
+ async with httpx.AsyncClient(timeout=30.0) as client:
+ response = await client.post(
+ f"{gateway_url.rstrip('/')}/trace_store/pop",
+ headers={"Authorization": f"Bearer {api_key}"},
+ )
+ response.raise_for_status()
+ payload = response.json()
+ records = payload.get("records", [])
+ if not isinstance(records, list):
+ return []
+ return records
+
+ def _failed_state(
+ self,
+ rollout_state: RolloutState,
+ error_msg: str,
+ *,
+ extra_fields: dict[str, Any] | None = None,
+ ) -> RolloutState:
+ failed = rollout_state.model_copy(deep=True)
+ failed.status = Status.FAILED
+ failed.error_msg = error_msg
+ if extra_fields:
+ failed.extra_fields = {
+ **copy.deepcopy(rollout_state.extra_fields),
+ **copy.deepcopy(extra_fields),
+ }
+ return failed
+
+ def _truncate(self, text: str, max_chars: int = 4096) -> str:
+ if len(text) <= max_chars:
+ return text
+ return text[:max_chars] + "..."
diff --git a/recipe/claude_code/run_claudecode_tool_e2e.sh b/recipe/claude_code/run_claudecode_tool_e2e.sh
new file mode 100755
index 0000000000..8d2dd87774
--- /dev/null
+++ b/recipe/claude_code/run_claudecode_tool_e2e.sh
@@ -0,0 +1,52 @@
+#!/usr/bin/env bash
+set -euo pipefail
+
+SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
+REPO_ROOT="$(cd "${SCRIPT_DIR}/../.." && pwd)"
+
+XTUNER_USE_LMDEPLOY="${XTUNER_USE_LMDEPLOY:-1}"
+XTUNER_CLAUDECODE_TOOL_MAX_TURNS="${XTUNER_CLAUDECODE_TOOL_MAX_TURNS:-4}"
+XTUNER_CLAUDECODE_TOOL_MAX_TOKENS="${XTUNER_CLAUDECODE_TOOL_MAX_TOKENS:-1024}"
+XTUNER_CLAUDECODE_TOOL_CONTEXT_LENGTH="${XTUNER_CLAUDECODE_TOOL_CONTEXT_LENGTH:-32768}"
+XTUNER_CLAUDECODE_TOOL_TIMEOUT_S="${XTUNER_CLAUDECODE_TOOL_TIMEOUT_S:-600}"
+XTUNER_CLAUDECODE_TOOL_OUTPUT_DIR="${XTUNER_CLAUDECODE_TOOL_OUTPUT_DIR:-/tmp/xtuner_claudecode_tool_e2e}"
+XTUNER_CLAUDECODE_TOOL_CALL_PARSER="${XTUNER_CLAUDECODE_TOOL_CALL_PARSER:-qwen3p5}"
+XTUNER_CLAUDECODE_REASONING_PARSER="${XTUNER_CLAUDECODE_REASONING_PARSER:-qwen3}"
+
+if [[ -z "${ROLLOUT_MODEL_PATH:-}" ]]; then
+ echo "ROLLOUT_MODEL_PATH must be set." >&2
+ exit 1
+fi
+
+if [[ ! -e "${ROLLOUT_MODEL_PATH}" ]]; then
+ echo "ROLLOUT_MODEL_PATH does not exist: ${ROLLOUT_MODEL_PATH}" >&2
+ exit 1
+fi
+
+mkdir -p "${XTUNER_CLAUDECODE_TOOL_OUTPUT_DIR}"
+
+export PYTHONPATH="${REPO_ROOT}:${PYTHONPATH:-}"
+
+export ROLLOUT_MODEL_PATH
+export XTUNER_USE_LMDEPLOY
+export XTUNER_CLAUDECODE_TOOL_MAX_TURNS
+export XTUNER_CLAUDECODE_TOOL_MAX_TOKENS
+export XTUNER_CLAUDECODE_TOOL_CONTEXT_LENGTH
+export XTUNER_CLAUDECODE_TOOL_TIMEOUT_S
+export XTUNER_CLAUDECODE_TOOL_OUTPUT_DIR
+export XTUNER_CLAUDECODE_TOOL_CALL_PARSER
+export XTUNER_CLAUDECODE_REASONING_PARSER
+
+cd "${REPO_ROOT}"
+
+echo "Running Claude Code calculator tool E2E"
+echo " repo root: ${REPO_ROOT}"
+echo " model: ${ROLLOUT_MODEL_PATH}"
+echo " output: ${XTUNER_CLAUDECODE_TOOL_OUTPUT_DIR}"
+echo " max_turns: ${XTUNER_CLAUDECODE_TOOL_MAX_TURNS}"
+echo " max_tokens: ${XTUNER_CLAUDECODE_TOOL_MAX_TOKENS}"
+echo " context_length: ${XTUNER_CLAUDECODE_TOOL_CONTEXT_LENGTH}"
+echo " tool_call_parser: ${XTUNER_CLAUDECODE_TOOL_CALL_PARSER}"
+echo " reasoning_parser: ${XTUNER_CLAUDECODE_REASONING_PARSER}"
+
+python "${SCRIPT_DIR}/test_claude_code_with_calculator.py"
diff --git a/recipe/claude_code/test_claude_code_with_calculator.py b/recipe/claude_code/test_claude_code_with_calculator.py
new file mode 100644
index 0000000000..d702355f55
--- /dev/null
+++ b/recipe/claude_code/test_claude_code_with_calculator.py
@@ -0,0 +1,247 @@
+from __future__ import annotations
+
+import asyncio
+import json
+import os
+import tempfile
+from pathlib import Path
+from typing import Any
+from uuid import uuid4
+
+from calculator_tool import (
+ CalculatorJudger,
+ CALCULATOR_PROMPT,
+ CALCULATOR_SYSTEM_PROMPT,
+ CALCULATOR_TOOL_NAME,
+ normalize_answer,
+ write_calculator_mcp_server,
+)
+from claudecode_agent_loop import ClaudeCodeAgentLoopConfig
+from xtuner.v1.data_proto import RolloutState, SampleParams, Status
+from xtuner.v1.rl.gateway import wait_for_gateway_ready
+from xtuner.v1.rl.utils import find_free_ports
+
+
+RESOURCE_MAP = {
+ "npu": "NPU",
+ "cuda": "GPU",
+}
+
+
+async def test_claude_code_with_calculator(model_path: str) -> list[RolloutState]:
+ import ray
+
+ os.environ.setdefault("XTUNER_USE_FA3", "1")
+ os.environ.setdefault("LMD_SKIP_WARMUP", "1")
+ os.environ.pop("RAY_ADDRESS", None)
+ ray.init(address="local", ignore_reinit_error=True)
+
+ temp_dir = tempfile.TemporaryDirectory()
+ work_dir = Path(temp_dir.name)
+ worker_log_dir = work_dir / "work_dirs"
+ output_dir = Path(os.environ.get("XTUNER_CLAUDECODE_TOOL_OUTPUT_DIR", work_dir / "outputs"))
+ output_dir.mkdir(parents=True, exist_ok=True)
+ controller = None
+ placement_group = None
+
+ try:
+ _, mcp_config_path = write_calculator_mcp_server(work_dir)
+ gateway_url, controller, placement_group = _start_rollout_controller_and_gateway(
+ model_path=model_path,
+ worker_log_dir=worker_log_dir,
+ )
+ cfg = ClaudeCodeAgentLoopConfig(
+ hf_checkpoint=model_path,
+ sample_params=SampleParams(
+ max_tokens=int(os.environ.get("XTUNER_CLAUDECODE_TOOL_MAX_TOKENS", "1024")),
+ temperature=0.0,
+ ),
+ claude_command=[os.environ.get("XTUNER_CLAUDE_BIN", str(Path.home() / ".local" / "bin" / "claude"))],
+ cwd=str(work_dir),
+ timeout_s=float(os.environ.get("XTUNER_CLAUDECODE_TOOL_TIMEOUT_S", "600")),
+ api_timeout_ms=int(os.environ.get("XTUNER_CLAUDECODE_TOOL_API_TIMEOUT_MS", "600000")),
+ max_turns=int(os.environ.get("XTUNER_CLAUDECODE_TOOL_MAX_TURNS", "4")),
+ output_format="json",
+ permission_mode=os.environ.get("XTUNER_CLAUDECODE_PERMISSION_MODE", "bypassPermissions"),
+ tools=None,
+ allowed_tools=CALCULATOR_TOOL_NAME,
+ disallowed_tools="Bash,Edit,Read,Grep,Glob,LS,WebFetch,WebSearch",
+ mcp_config=[str(mcp_config_path)],
+ strict_mcp_config=True,
+ system_prompt=CALCULATOR_SYSTEM_PROMPT,
+ readonly_instruction="",
+ )
+ agent_loop = cfg.build(rollout_controller=controller, judger=CalculatorJudger())
+ rollout_state = RolloutState(
+ message=[{"role": "user", "content": CALCULATOR_PROMPT}],
+ task_name="calculator_tool_call",
+ extra_fields={"gateway_url": gateway_url},
+ )
+
+ states = await agent_loop.generate_sample(rollout_state)
+ _dump_rollout_states(output_dir, "calculator", states)
+ failed = [state for state in states if state.status == Status.FAILED]
+ if failed:
+ raise AssertionError(f"ClaudeCodeAgentLoop returned failed states: {[state.error_msg for state in failed]}")
+
+ completed = [state for state in states if state.status == Status.COMPLETED]
+ if len(completed) < 2:
+ raise AssertionError("Expected one tool-call turn and one final-answer turn.")
+
+ api_keys = {state.extra_fields["claudecode_api_key"] for state in completed}
+ if len(api_keys) != 1:
+ raise AssertionError(f"Expected one Claude Code api key, got {api_keys}.")
+ for index, state in enumerate(completed):
+ if state.prompt_ids is None:
+ raise AssertionError(f"State {index} is missing prompt_ids.")
+ if state.tokens != state.prompt_ids:
+ raise AssertionError(f"State {index} tokens must equal prompt_ids.")
+ if not state.response_ids:
+ raise AssertionError(f"State {index} is missing response_ids.")
+ if state.logprobs is None:
+ raise AssertionError(f"State {index} is missing logprobs.")
+ if len(state.logprobs) != len(state.response_ids):
+ raise AssertionError(f"State {index} logprobs length does not match response_ids length.")
+ if state.response_mask != [1] * len(state.response_ids):
+ raise AssertionError(f"State {index} response_mask does not match response_ids length.")
+ if state.response is None:
+ raise AssertionError(f"State {index} is missing response text.")
+ if "gateway_trace_records" not in state.extra_fields:
+ raise AssertionError(f"State {index} is missing gateway_trace_records.")
+ if state.extra_fields["gateway_trace_count"] != len(states):
+ raise AssertionError(f"State {index} gateway_trace_count does not match trace count.")
+ if state.extra_fields["claudecode_cli_returncode"] != 0:
+ raise AssertionError(f"State {index} Claude Code returncode is not 0.")
+
+ tool_blocks = []
+ for state in completed:
+ snapshot = state.extra_fields.get("gateway_response_snapshot") or {}
+ for block in snapshot.get("content") or []:
+ if block.get("type") == "tool_use":
+ tool_blocks.append(block)
+ if not tool_blocks:
+ raise AssertionError("Expected at least one calculator tool_use block.")
+ calculator_blocks = [block for block in tool_blocks if "calculator" in str(block.get("name"))]
+ if not calculator_blocks:
+ raise AssertionError(f"Expected calculator tool call, got: {tool_blocks}")
+ if calculator_blocks[0].get("input", {}).get("expression") != "23 + 19":
+ raise AssertionError(f"Unexpected calculator input: {calculator_blocks[0].get('input')}")
+
+ final_answer = normalize_answer(completed[-1].response)
+ if final_answer != "42":
+ raise AssertionError(f"Expected final answer 42, got {final_answer!r}.")
+ if completed[-1].reward != {"score": 1.0, "answer": "42"}:
+ raise AssertionError(f"Expected reward score 1.0 for answer 42, got {completed[-1].reward!r}.")
+ return states
+ finally:
+ _cleanup_ray(controller=controller, placement_group=placement_group)
+ temp_dir.cleanup()
+
+
+def _start_rollout_controller_and_gateway(
+ *,
+ model_path: str,
+ worker_log_dir: Path,
+) -> tuple[str, Any, Any]:
+ import ray
+ import torch
+
+ from xtuner.v1.rl.gateway.config import GatewayConfig
+ from xtuner.v1.rl.rollout.worker import RolloutConfig
+ from xtuner.v1.rl.utils import AcceleratorResourcesConfig, AutoAcceleratorWorkers
+
+ accelerator = RESOURCE_MAP[torch.accelerator.current_accelerator().type]
+ tensor_parallel_size = int(os.environ.get("XTUNER_CLAUDECODE_TOOL_TP", "1"))
+ num_workers = int(os.environ.get("XTUNER_CLAUDECODE_TOOL_NUM_WORKERS", str(tensor_parallel_size)))
+ resource_config = AcceleratorResourcesConfig(
+ accelerator=accelerator,
+ num_workers=num_workers,
+ num_cpus_per_worker=int(os.environ.get("XTUNER_CLAUDECODE_TOOL_CPUS_PER_WORKER", "8")),
+ cpu_memory_per_worker=int(os.environ.get("XTUNER_CLAUDECODE_TOOL_CPU_MEMORY", str(16 * 1024**3))),
+ )
+ placement_group = AutoAcceleratorWorkers.build_placement_group(
+ resource_config,
+ name=f"claudecode_tool_pg_{uuid4().hex[:8]}",
+ )
+ rollout_config = RolloutConfig(
+ env=f"claudecode_tool_{uuid4().hex[:8]}",
+ model_path=model_path,
+ model_name=os.path.basename(model_path).lower(),
+ tokenizer_path=model_path,
+ context_length=int(os.environ.get("XTUNER_CLAUDECODE_TOOL_CONTEXT_LENGTH", "32768")),
+ worker_log_dir=worker_log_dir / "rollout",
+ tensor_parallel_size=tensor_parallel_size,
+ expert_parallel_size=1,
+ dist_port_base=int(
+ os.environ.get("XTUNER_CLAUDECODE_TOOL_DIST_PORT_BASE", str(find_free_ports(nums=8, contiguous=True)[0]))
+ ),
+ tool_call_parser=os.environ.get("XTUNER_CLAUDECODE_TOOL_CALL_PARSER", "qwen3p5"),
+ reasoning_parser=os.environ.get("XTUNER_CLAUDECODE_REASONING_PARSER", "qwen3"),
+ api_host="127.0.0.1",
+ api_port=find_free_ports()[0],
+ )
+ controller = rollout_config.build(placement_group)
+ gateway_host = ray.util.get_node_ip_address()
+ gateway_config = GatewayConfig(
+ host=gateway_host,
+ port=find_free_ports(host=gateway_host)[0],
+ capture_folder=str(worker_log_dir / "gateway_captures"),
+ )
+ gateway_url = ray.get(controller.start_gateway.remote(gateway_config), timeout=1800)
+ wait_for_gateway_ready(gateway_url)
+ return gateway_url, controller, placement_group
+
+
+def _dump_rollout_states(output_dir: Path, case_name: str, states: list[RolloutState]) -> None:
+ output_path = output_dir / f"rollout-states-{case_name}.json"
+ payload = [_redact_rollout_state_for_dump(state) for state in states]
+ output_path.write_text(json.dumps(payload, ensure_ascii=False, indent=2), encoding="utf-8")
+ print(f"Claude Code rollout states written to {output_path}")
+
+
+def _redact_rollout_state_for_dump(state: RolloutState) -> dict:
+ payload = state.model_dump(mode="json")
+ for key in ("prompt_ids", "response_ids"):
+ payload.pop(key, None)
+ extra_fields = payload.get("extra_fields")
+ if isinstance(extra_fields, dict):
+ records = extra_fields.get("gateway_trace_records")
+ if isinstance(records, list):
+ for record in records:
+ if isinstance(record, dict):
+ record.pop("prompt_ids", None)
+ record.pop("response_ids", None)
+ record.pop("tokens", None)
+ return payload
+
+
+def _cleanup_ray(*, controller: Any, placement_group: Any) -> None:
+ import ray
+
+ if controller is not None:
+ try:
+ ray.get(controller.shutdown.remote(), timeout=300)
+ except Exception:
+ pass
+ try:
+ ray.kill(controller, no_restart=True)
+ except Exception:
+ pass
+ if placement_group is not None:
+ ray.util.remove_placement_group(placement_group)
+ if ray.is_initialized():
+ ray.shutdown()
+
+
+async def _main_async() -> None:
+ states = await test_claude_code_with_calculator(model_path=os.environ["ROLLOUT_MODEL_PATH"])
+ completed_count = sum(state.status == Status.COMPLETED for state in states)
+ print(f"Claude Code calculator tool E2E passed with {completed_count} completed rollout states.")
+
+
+def main() -> None:
+ asyncio.run(_main_async())
+
+
+if __name__ == "__main__":
+ main()
diff --git a/xtuner/v1/rl/gateway/__init__.py b/xtuner/v1/rl/gateway/__init__.py
index 26779e3a9e..e5a8dfafdb 100644
--- a/xtuner/v1/rl/gateway/__init__.py
+++ b/xtuner/v1/rl/gateway/__init__.py
@@ -1,6 +1,12 @@
from .backend.local_backend import LocalRolloutBackend
from .config import GatewayConfig
-from .server import build_gateway_app, build_local_gateway_app, serve_gateway, serve_gateway_in_thread
+from .server import (
+ build_gateway_app,
+ build_local_gateway_app,
+ serve_gateway,
+ serve_gateway_in_thread,
+ wait_for_gateway_ready,
+)
__all__ = [
@@ -10,4 +16,5 @@
"build_local_gateway_app",
"serve_gateway",
"serve_gateway_in_thread",
+ "wait_for_gateway_ready",
]
diff --git a/xtuner/v1/rl/gateway/adapters/anthropic.py b/xtuner/v1/rl/gateway/adapters/anthropic.py
index aa478543e1..81705986e3 100644
--- a/xtuner/v1/rl/gateway/adapters/anthropic.py
+++ b/xtuner/v1/rl/gateway/adapters/anthropic.py
@@ -271,7 +271,10 @@ def canonical_response_to_protocol_response(
canonical_response: CanonicalGenerateResponse,
request: AnthropicMessagesRequest,
) -> AnthropicMessagesResponse:
- content = self._canonical_response_to_anthropic_blocks(canonical_response)
+ content = self._canonical_response_to_anthropic_blocks(
+ canonical_response,
+ tools=self._anthropic_tools_to_canonical(request.tools),
+ )
stop_reason = canonical_response.finish_reason or "stop"
if any(block.get("type") == "tool_use" for block in content):
stop_reason = "tool_use"
@@ -554,6 +557,7 @@ def _anthropic_tool_choice_to_canonical(
def _canonical_response_to_anthropic_blocks(
self,
response: CanonicalGenerateResponse,
+ tools: list[CanonicalToolDefinition] | None = None,
) -> list[dict[str, Any]]:
blocks: list[dict[str, Any]] = []
for block in response.output.content:
@@ -561,12 +565,13 @@ def _canonical_response_to_anthropic_blocks(
if block.text:
blocks.append({"type": "text", "text": block.text})
elif isinstance(block, CanonicalToolCallBlock):
+ tool_call = self._sanitize_tool_call_for_request(block.tool_call, tools=tools or [])
blocks.append(
{
"type": "tool_use",
- "id": block.tool_call.id,
- "name": block.tool_call.name,
- "input": block.tool_call.arguments if block.tool_call.arguments is not None else {},
+ "id": tool_call.id,
+ "name": tool_call.name,
+ "input": tool_call.arguments if tool_call.arguments is not None else {},
}
)
elif isinstance(block, CanonicalToolResultBlock):
@@ -587,6 +592,58 @@ def _canonical_response_to_anthropic_blocks(
blocks.append({"type": "thinking", "thinking": reasoning_text})
return blocks or [{"type": "text", "text": ""}]
+ def _sanitize_tool_call_for_request(
+ self,
+ tool_call: CanonicalToolCall,
+ *,
+ tools: list[CanonicalToolDefinition],
+ ) -> CanonicalToolCall:
+ tool_definition = next((tool for tool in tools if tool.name == tool_call.name), None)
+ if tool_definition is None:
+ return tool_call
+
+ properties = tool_definition.parameters_json_schema.get("properties")
+ if not isinstance(properties, dict):
+ return tool_call
+
+ arguments = tool_call.arguments
+ normalized_arguments = False
+ if not isinstance(arguments, dict):
+ normalized_arguments = True
+ if tool_call.raw_arguments_text is not None:
+ try:
+ decoded = json.loads(tool_call.raw_arguments_text)
+ except Exception:
+ decoded = {"raw": tool_call.raw_arguments_text}
+ arguments = decoded if isinstance(decoded, dict) else {"value": decoded}
+ elif arguments is None:
+ arguments = {}
+ elif isinstance(arguments, str):
+ try:
+ decoded = json.loads(arguments)
+ except Exception:
+ decoded = {"raw": arguments}
+ arguments = decoded if isinstance(decoded, dict) else {"value": decoded}
+ else:
+ arguments = {"value": arguments}
+
+ allowed_keys = set(properties)
+ cleaned_arguments = {key: value for key, value in arguments.items() if key in allowed_keys}
+ if cleaned_arguments == arguments and not normalized_arguments:
+ return tool_call
+
+ dropped_keys = sorted(set(arguments) - set(cleaned_arguments))
+ metadata = dict(tool_call.metadata)
+ if dropped_keys:
+ metadata["dropped_arguments"] = dropped_keys
+ return CanonicalToolCall(
+ id=tool_call.id,
+ name=tool_call.name,
+ arguments=cleaned_arguments,
+ raw_arguments_text=None,
+ metadata=metadata,
+ )
+
def _reasoning_to_text(self, reasoning: CanonicalReasoning) -> str:
return "\n".join(step.text for step in reasoning.steps if step.text).strip()
diff --git a/xtuner/v1/rl/gateway/backend/local_backend.py b/xtuner/v1/rl/gateway/backend/local_backend.py
index a2bab2c946..28555027e8 100644
--- a/xtuner/v1/rl/gateway/backend/local_backend.py
+++ b/xtuner/v1/rl/gateway/backend/local_backend.py
@@ -40,9 +40,10 @@ def __init__(
self,
controller: ActorHandle,
tokenizer: PreTrainedTokenizer | PreTrainedTokenizerFast | str | None = None,
+ rollout_config: RolloutConfig | None = None,
):
self._controller = controller
- self._config = self._resolve_rollout_config(controller)
+ self._config = rollout_config or self._resolve_rollout_config(controller)
if isinstance(tokenizer, str):
tokenizer = AutoTokenizer.from_pretrained(tokenizer, trust_remote_code=True)
resolved_tokenizer = tokenizer
@@ -297,7 +298,7 @@ def _build_sample_params(
) -> SampleParams:
kwargs = {
"return_token_ids": True,
- "return_logprob": False,
+ "return_logprob": True,
"stream": canonical_request.stream,
"stops": canonical_request.stop,
**{
diff --git a/xtuner/v1/rl/gateway/server/__init__.py b/xtuner/v1/rl/gateway/server/__init__.py
index f4f09a5377..722863cd83 100644
--- a/xtuner/v1/rl/gateway/server/__init__.py
+++ b/xtuner/v1/rl/gateway/server/__init__.py
@@ -1,4 +1,16 @@
-from .app import build_gateway_app, build_local_gateway_app, serve_gateway, serve_gateway_in_thread
+from .app import (
+ build_gateway_app,
+ build_local_gateway_app,
+ serve_gateway,
+ serve_gateway_in_thread,
+ wait_for_gateway_ready,
+)
-__all__ = ["build_gateway_app", "build_local_gateway_app", "serve_gateway", "serve_gateway_in_thread"]
+__all__ = [
+ "build_gateway_app",
+ "build_local_gateway_app",
+ "serve_gateway",
+ "serve_gateway_in_thread",
+ "wait_for_gateway_ready",
+]
diff --git a/xtuner/v1/rl/gateway/server/app.py b/xtuner/v1/rl/gateway/server/app.py
index 56cb3ea1ad..a7ca63389a 100644
--- a/xtuner/v1/rl/gateway/server/app.py
+++ b/xtuner/v1/rl/gateway/server/app.py
@@ -2,8 +2,10 @@
import socket
import threading
+import time
from typing import Union
+import httpx
import ray
import uvicorn
from fastapi import FastAPI, Request
@@ -149,11 +151,13 @@ def build_gateway_app(
def build_local_gateway_app(
controller: ActorHandle,
config: GatewayConfig | None = None,
+ rollout_config: RolloutConfig | None = None,
) -> FastAPI:
"""Build a gateway app backed by a Ray-actor RolloutController."""
cfg = config or GatewayConfig(port=8080)
- rollout_metadata = ray.get(controller.get_rollout_metadata.remote())
- rollout_config: RolloutConfig = rollout_metadata["rollout_config"]
+ if rollout_config is None:
+ rollout_metadata = ray.get(controller.get_rollout_metadata.remote())
+ rollout_config = rollout_metadata["rollout_config"]
tokenizer = AutoTokenizer.from_pretrained(rollout_config.tokenizer_path, trust_remote_code=True)
model_name = rollout_config.model_name
@@ -163,7 +167,7 @@ def build_local_gateway_app(
if context_length is None:
raise ValueError("controller.config.context_length must be set when building a local gateway app")
- backend = LocalRolloutBackend(controller, tokenizer=tokenizer)
+ backend = LocalRolloutBackend(controller, tokenizer=tokenizer, rollout_config=rollout_config)
return build_gateway_app(
backend,
tokenizer=tokenizer,
@@ -243,6 +247,22 @@ def serve_gateway_in_thread(app: FastAPI, config: GatewayConfig) -> threading.Th
return thread
+def wait_for_gateway_ready(base_url: str, *, timeout_seconds: float = 180.0) -> None:
+ """Block until a gateway server responds successfully on ``/livez``."""
+ deadline = time.time() + timeout_seconds
+ last_error = None
+ while time.time() < deadline:
+ try:
+ response = httpx.get(f"{base_url}/livez", timeout=5.0)
+ if response.status_code == 200:
+ return
+ last_error = response.text
+ except Exception as exc:
+ last_error = repr(exc)
+ time.sleep(1.0)
+ raise AssertionError(f"Gateway did not become ready at {base_url}: {last_error}")
+
+
def _ensure_gateway_port_available(config: GatewayConfig) -> None:
try:
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as sock:
diff --git a/xtuner/v1/rl/rollout/controller.py b/xtuner/v1/rl/rollout/controller.py
index df00323449..c421626cae 100644
--- a/xtuner/v1/rl/rollout/controller.py
+++ b/xtuner/v1/rl/rollout/controller.py
@@ -127,10 +127,12 @@ def start_gateway(self, config: "GatewayConfig") -> str:
from xtuner.v1.rl.gateway import build_local_gateway_app, serve_gateway_in_thread
config.capture_folder = str(Path(self.config.worker_log_dir) / config._CAPTURE_PATH_FOLDER)
- app = build_local_gateway_app(ray.get_runtime_context().current_actor, config=config)
+ app = build_local_gateway_app(
+ ray.get_runtime_context().current_actor, config=config, rollout_config=self.config
+ )
serve_gateway_in_thread(app, config)
- node_ip = ray.util.get_node_ip_address()
- url = f"http://{node_ip}:{config.port}"
+ host = ray.util.get_node_ip_address() if config.host in ("", "0.0.0.0") else config.host
+ url = f"http://{host}:{config.port}"
self._gateway_url = url
self.logger.info(f"Gateway server started at {url}, capture_folder: {config.capture_folder}")
return url
diff --git a/xtuner/v1/rl/utils/__init__.py b/xtuner/v1/rl/utils/__init__.py
index d87ef41407..326b642241 100644
--- a/xtuner/v1/rl/utils/__init__.py
+++ b/xtuner/v1/rl/utils/__init__.py
@@ -11,6 +11,8 @@
ScalarOperator,
SetNode,
SetOperator,
+ chat_trace_records_to_rollout_states,
+ find_free_ports,
gather_logprobs,
get_eos_token,
load_function,
@@ -66,4 +68,6 @@
"LogicOperator",
"Operators",
"get_eos_token",
+ "chat_trace_records_to_rollout_states",
+ "find_free_ports",
]
diff --git a/xtuner/v1/rl/utils/misc.py b/xtuner/v1/rl/utils/misc.py
index 7868fc913f..89d2317875 100644
--- a/xtuner/v1/rl/utils/misc.py
+++ b/xtuner/v1/rl/utils/misc.py
@@ -1,13 +1,17 @@
import importlib
import json
+import random
import socket
import typing
from abc import ABC
+from copy import deepcopy
+from dataclasses import asdict, is_dataclass
from pathlib import Path
from typing import Any, List, Literal, Union
import torch.nn.functional as F
+from xtuner.v1.data_proto import RolloutState, Status
from xtuner.v1.utils.logger import get_logger
@@ -122,13 +126,93 @@ def load_function(path):
return getattr(module, attr)
-def _is_port_available(check_socket: socket.socket, port: int) -> bool:
- try:
- check_socket.bind(("", port))
- check_socket.listen(1)
- return True
- except OSError:
- return False
+def find_free_ports(
+ *,
+ nums: int = 1,
+ host: str = "127.0.0.1",
+ start_port: int | None = None,
+ end_port: int | None = None,
+ contiguous: bool = False,
+) -> list[int]:
+ """Return available TCP ports on the given host.
+
+ The candidate sockets are kept open until all requested ports are found so
+ one call cannot return duplicate ports. Set ``contiguous=True`` to require
+ the returned ports to be a continuous range.
+ """
+ if nums < 1:
+ raise ValueError("nums must be greater than 0.")
+ if start_port is not None:
+ if end_port is None:
+ raise ValueError("end_port must be set when start_port is set.")
+ if end_port - start_port < nums:
+ raise ValueError("The port range must contain at least nums ports.")
+
+ def try_bind_ports(candidate_ports: list[int]) -> list[int] | None:
+ ports: list[int] = []
+ sockets: list[socket.socket] = []
+ try:
+ for candidate_port in candidate_ports:
+ sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
+ sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
+ try:
+ sock.bind((host, candidate_port))
+ sock.listen(1)
+ except OSError:
+ sock.close()
+ return None
+
+ sockets.append(sock)
+ ports.append(int(sock.getsockname()[1]))
+ return ports
+ finally:
+ for sock in sockets:
+ sock.close()
+
+ if contiguous:
+ if start_port is None:
+ for _ in range(100):
+ candidate = random.randint(20000, 60000 - nums)
+ bound_ports = try_bind_ports(list(range(candidate, candidate + nums)))
+ if bound_ports is not None:
+ return bound_ports
+ else:
+ assert end_port is not None
+ for candidate in range(start_port, end_port - nums + 1):
+ bound_ports = try_bind_ports(list(range(candidate, candidate + nums)))
+ if bound_ports is not None:
+ return bound_ports
+ else:
+ available_ports: list[int] = []
+ sockets: list[socket.socket] = []
+ try:
+ if start_port is None:
+ candidates: range | list[int] = [0] * nums
+ else:
+ assert end_port is not None
+ candidates = range(start_port, end_port)
+
+ for candidate_port in candidates:
+ sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
+ sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
+ try:
+ sock.bind((host, candidate_port))
+ sock.listen(1)
+ except OSError:
+ sock.close()
+ continue
+
+ sockets.append(sock)
+ available_ports.append(int(sock.getsockname()[1]))
+ if len(available_ports) >= nums:
+ return available_ports
+ finally:
+ for sock in sockets:
+ sock.close()
+
+ if start_port is None:
+ raise RuntimeError(f"Could not find {nums} available ports.")
+ raise RuntimeError(f"Could not find {nums} available ports from {start_port} to {end_port}.")
def get_eos_token(model_path: str) -> int | List[int]:
@@ -146,3 +230,99 @@ def get_eos_token(model_path: str) -> int | List[int]:
f"eos_token_id is not found in {generation_config_path}. You must provide eos_token manually."
)
return eos_token_id
+
+
+def chat_trace_records_to_rollout_states(
+ rollout_state: RolloutState,
+ records: list[Any],
+ *,
+ tokenizer: Any | None = None,
+ extra_fields: dict[str, Any] | None = None,
+) -> list[RolloutState]:
+ """Convert Gateway chat trace records into trainable rollout states.
+
+ The records may be ``ChatTraceRecord`` dataclass instances or serialized
+ dictionaries returned by ``/trace_store``.
+ """
+ normalized_records = []
+ for record in records:
+ if isinstance(record, dict):
+ normalized_records.append(record)
+ elif not isinstance(record, type) and is_dataclass(record):
+ normalized_records.append(asdict(record))
+ elif hasattr(record, "__dict__"):
+ normalized_records.append(dict(record.__dict__))
+ else:
+ raise TypeError(f"Unsupported chat trace record type: {type(record)}")
+
+ trace_count = len(normalized_records)
+ trace_summary = [
+ {
+ "request_id": record.get("request_id"),
+ "finish_reason": record.get("finish_reason"),
+ "status": record.get("status"),
+ "prompt_ids": record.get("prompt_ids", []),
+ "response_ids": record.get("response_ids", []),
+ }
+ for record in normalized_records
+ ]
+
+ states: list[RolloutState] = []
+ for index, record in enumerate(normalized_records):
+ prompt_ids = record.get("prompt_ids")
+ response_ids = record.get("response_ids")
+ if not prompt_ids or not response_ids:
+ raise RuntimeError(f"Gateway trace record {index} is missing prompt_ids or response_ids.")
+
+ logprobs = record.get("logprobs")
+ if not isinstance(logprobs, list) or len(logprobs) != len(response_ids):
+ logprobs = None
+
+ status_value = record.get("status")
+ if isinstance(status_value, Status):
+ status = status_value
+ elif isinstance(status_value, str):
+ try:
+ status = Status(status_value)
+ except ValueError:
+ status = Status.FAILED
+ else:
+ status = Status.FAILED
+
+ request_id = record.get("request_id")
+ try:
+ uid = int(request_id) if request_id is not None else None
+ except (TypeError, ValueError):
+ uid = None
+
+ response = record.get("output_text")
+ if response is None and tokenizer is not None:
+ try:
+ response = tokenizer.decode(response_ids)
+ except Exception:
+ response = None
+
+ normalized = rollout_state.model_copy(deep=True)
+ normalized.uid = uid
+ normalized.prompt_ids = list(prompt_ids)
+ normalized.tokens = list(prompt_ids)
+ normalized.response_ids = list(response_ids)
+ normalized.response_mask = [1] * len(response_ids)
+ normalized.logprobs = logprobs
+ normalized.response = response
+ normalized.finish_reason = record.get("finish_reason")
+ normalized.status = status
+ normalized.error_msg = None if status == Status.COMPLETED else f"Gateway trace status={status.value}"
+ normalized.reward = None
+ normalized.extra_fields = {
+ **deepcopy(rollout_state.extra_fields),
+ "gateway_trace_index": index,
+ "gateway_trace_count": trace_count,
+ "gateway_trace_records": deepcopy(trace_summary),
+ "gateway_request_id": record.get("request_id"),
+ "gateway_request_snapshot": record.get("request_snapshot"),
+ "gateway_response_snapshot": record.get("response_snapshot"),
+ **deepcopy(extra_fields or {}),
+ }
+ states.append(normalized)
+ return states
diff --git a/xtuner/v1/rl/utils/ray_utils.py b/xtuner/v1/rl/utils/ray_utils.py
index 14d94323d7..40dc69e415 100644
--- a/xtuner/v1/rl/utils/ray_utils.py
+++ b/xtuner/v1/rl/utils/ray_utils.py
@@ -1,6 +1,5 @@
import atexit
import signal
-import socket
import subprocess
from typing import TYPE_CHECKING, cast
@@ -8,7 +7,7 @@
from xtuner.v1.utils.logger import get_logger
-from .misc import _is_port_available
+from .misc import find_free_ports
if TYPE_CHECKING:
@@ -37,43 +36,7 @@ def find_master_addr_and_port(nums=1, start_port=None, end_port=None):
or a list of ports if `nums` is greater than 1.
"""
addr = ray.util.get_node_ip_address()
- ports: list[int] = []
- sockets: list[socket.socket] = []
-
- if start_port is None:
- for _ in range(nums):
- s = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
- # if the port is binded and listened by this socket and then we close it,
- # socket.SO_REUSEADDR would make the port be reusable even it's in TIME_WAIT state.
- s.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
- sockets.append(s)
- if _is_port_available(check_socket=s, port=0):
- ports.append(s.getsockname()[1])
- else:
- assert isinstance(start_port, int), "If start_port isn't None, it must be an integer."
- assert isinstance(end_port, int), "If start_port isn't None, end_port must be an integer."
- assert end_port - start_port >= nums, (
- "If start_port isn't None, the range between start_port and end_port must be at least nums."
- )
-
- for candidate_port in range(start_port, end_port):
- s = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
- # if the port is binded and listened by this socket and then we close it,
- # socket.SO_REUSEADDR would make the port be reusable even it's in TIME_WAIT state.
- s.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
- sockets.append(s)
- if _is_port_available(check_socket=s, port=candidate_port):
- ports.append(candidate_port)
- # enough ports found
- if len(ports) >= nums:
- break
-
- if len(ports) < nums:
- raise RuntimeError(f"Could not find {nums} available ports starting from port {start_port} to {end_port}.")
-
- # close all sockets, no matter available or not
- for s in sockets:
- s.close()
+ ports = find_free_ports(nums=nums, host="", start_port=start_port, end_port=end_port)
if len(ports) == 1:
return addr, ports[0]