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]