From 7ac32200f0302a772ff995089e73cc519cf2d2df Mon Sep 17 00:00:00 2001 From: Blue <3067670134@qq.com> Date: Sat, 3 Oct 2026 21:17:04 +0800 Subject: [PATCH 1/2] fix: prevent incomplete Search-R1 data cache files --- examples/search_r1/data_process.sh | 21 +++- tests/examples/test_search_r1_data.py | 151 ++++++++++++++++++++++++++ 2 files changed, 168 insertions(+), 4 deletions(-) create mode 100644 tests/examples/test_search_r1_data.py diff --git a/examples/search_r1/data_process.sh b/examples/search_r1/data_process.sh index 06bf8dc2e..bf83ccad0 100755 --- a/examples/search_r1/data_process.sh +++ b/examples/search_r1/data_process.sh @@ -19,17 +19,30 @@ TEST_URL="https://huggingface.co/datasets/PeterJinGo/nq_hotpotqa_train/resolve/m download() { local url="$1" local output="$2" + local temporary if [[ -s "$output" ]]; then echo "skip existing $output" return fi + temporary="$(mktemp "${output}.tmp.XXXXXX")" if command -v curl >/dev/null 2>&1; then - curl -L --fail --retry 5 -o "$output" "$url" + if ! curl -L --fail --retry 5 -o "$temporary" "$url"; then + rm -f "$temporary" + return 1 + fi elif command -v wget >/dev/null 2>&1; then - wget -O "$output" "$url" + if ! wget -O "$temporary" "$url"; then + rm -f "$temporary" + return 1 + fi else echo "curl or wget is required to download $url" >&2 - exit 1 + rm -f "$temporary" + return 1 + fi + if ! mv -f "$temporary" "$output"; then + rm -f "$temporary" + return 1 fi } @@ -66,7 +79,7 @@ download "$TRAIN_URL" "$DATA_DIR/train.parquet" download "$TEST_URL" "$DATA_DIR/test.parquet" if [[ ! -s "$DATA_DIR/e5_Flat.index" ]]; then - cat "$DATA_DIR"/part_* > "$DATA_DIR/e5_Flat.index" + cat "$DATA_DIR/part_aa" "$DATA_DIR/part_ab" > "$DATA_DIR/e5_Flat.index" fi if [[ ! -s "$DATA_DIR/wiki-18.jsonl" ]]; then diff --git a/tests/examples/test_search_r1_data.py b/tests/examples/test_search_r1_data.py new file mode 100644 index 000000000..eb6739a88 --- /dev/null +++ b/tests/examples/test_search_r1_data.py @@ -0,0 +1,151 @@ +# Copyright (c) Microsoft. All rights reserved. + +"""Regression tests for the Search-R1 data preparation cache.""" + +from __future__ import annotations + +import gzip +import os +import shutil +import subprocess +from pathlib import Path + +import pytest + +SCRIPT = Path(__file__).parents[2] / "examples" / "search_r1" / "data_process.sh" +BASH = shutil.which("bash") +LINUX_ONLY = pytest.mark.skipif(os.name == "nt" or BASH is None, reason="requires Linux bash") + + +def _write_executable(path: Path, content: str) -> None: + path.write_text(content, encoding="utf-8") + path.chmod(0o755) + + +def _tool_path(tmp_path: Path, downloader: str) -> Path: + tools = tmp_path / "bin" + tools.mkdir() + for name in ("awk", "cat", "dirname", "grep", "gzip", "mkdir", "mktemp", "mv", "rm"): + target = shutil.which(name) + assert target is not None + (tools / name).symlink_to(target) + + _write_executable( + tools / "conda", + """#!/bin/sh +if [ "$1 $2 $3" = "shell.bash hook" ]; then + printf ':\\n' +elif [ "$1 $2" = "env list" ]; then + printf 'retriever /mock/retriever\\n' +fi +exit 0 +""", + ) + _write_executable( + tools / downloader, + """#!/bin/sh +set -eu +url='' +output='' +while [ "$#" -gt 0 ]; do + case "$1" in + -o|-O) output="$2"; shift 2 ;; + http*) url="$1"; shift ;; + *) shift ;; + esac +done +printf '%s\\n' "$url" >> "$MOCK_DOWNLOAD_LOG" +if [ "$MOCK_DOWNLOAD_MODE" = fail ]; then + printf partial > "$output" + exit 22 +fi +cp "$MOCK_FIXTURE_DIR/${url##*/}" "$output" +""", + ) + cp = shutil.which("cp") + assert cp is not None + (tools / "cp").symlink_to(cp) + return tools + + +def _fixtures(tmp_path: Path) -> Path: + fixtures = tmp_path / "fixtures" + fixtures.mkdir() + with gzip.open(fixtures / "wiki-18.jsonl.gz", "wb") as output: + output.write(b'{"id": 1}\n') + (fixtures / "part_aa").write_bytes(b"AA") + (fixtures / "part_ab").write_bytes(b"AB") + (fixtures / "train.parquet").write_bytes(b"TRAIN") + (fixtures / "test.parquet").write_bytes(b"TEST") + return fixtures + + +def _run(script_env: dict[str, str]) -> subprocess.CompletedProcess[str]: + assert BASH is not None + # Git may materialize shell files with CRLF in the Windows checkout used to + # drive WSL tests. Execute an LF copy of the exact source text. + script = Path(script_env["SEARCH_R1_DATA_DIR"]).parent / "data_process.sh" + script.write_text(SCRIPT.read_text(encoding="utf-8"), encoding="utf-8", newline="\n") + return subprocess.run( + [BASH, str(script)], + env=script_env, + capture_output=True, + text=True, + check=False, + ) + + +def _environment(tmp_path: Path, downloader: str) -> tuple[dict[str, str], Path, Path]: + data = tmp_path / "data" + log = tmp_path / "downloads.log" + tools = _tool_path(tmp_path, downloader) + env = { + **os.environ, + "PATH": str(tools), + "SEARCH_R1_DATA_DIR": str(data), + "SEARCH_R1_SKIP_RETRIEVER_INSTALL": "1", + "MOCK_DOWNLOAD_MODE": "success", + "MOCK_DOWNLOAD_LOG": str(log), + "MOCK_FIXTURE_DIR": str(_fixtures(tmp_path)), + } + return env, data, log + + +@LINUX_ONLY +@pytest.mark.parametrize("downloader", ["curl", "wget"]) +def test_failed_download_does_not_poison_cache_and_can_retry( + tmp_path: Path, + downloader: str, +) -> None: + env, data, log = _environment(tmp_path, downloader) + env["MOCK_DOWNLOAD_MODE"] = "fail" + + failed = _run(env) + + assert failed.returncode != 0 + assert not (data / "wiki-18.jsonl.gz").exists() + assert list(data.glob("*.tmp.*")) == [] + + env["MOCK_DOWNLOAD_MODE"] = "success" + retried = _run(env) + assert retried.returncode == 0, retried.stderr + assert gzip.decompress((data / "wiki-18.jsonl.gz").read_bytes()) == b'{"id": 1}\n' + assert (data / "wiki-18.jsonl").read_bytes() == b'{"id": 1}\n' + + calls_after_success = log.read_text(encoding="utf-8").splitlines() + env["MOCK_DOWNLOAD_MODE"] = "fail" + cached = _run(env) + assert cached.returncode == 0, cached.stderr + assert log.read_text(encoding="utf-8").splitlines() == calls_after_success + + +@LINUX_ONLY +def test_index_ignores_unrelated_part_files(tmp_path: Path) -> None: + env, data, _ = _environment(tmp_path, "curl") + data.mkdir() + (data / "part_backup").write_bytes(b"STALE") + + result = _run(env) + + assert result.returncode == 0, result.stderr + assert (data / "e5_Flat.index").read_bytes() == b"AAAB" From f3fd4ffe7c0e18ab0d0d1dcf7ffbdf0543576373 Mon Sep 17 00:00:00 2001 From: Blue <3067670134@qq.com> Date: Sat, 3 Oct 2026 22:41:34 +0800 Subject: [PATCH 2/2] fix(examples): publish Search-R1 index only after complete assembly --- examples/search_r1/data_process.sh | 10 +++++++- tests/examples/test_search_r1_data.py | 34 ++++++++++++++++++++++++++- 2 files changed, 42 insertions(+), 2 deletions(-) diff --git a/examples/search_r1/data_process.sh b/examples/search_r1/data_process.sh index bf83ccad0..b8ab372b6 100755 --- a/examples/search_r1/data_process.sh +++ b/examples/search_r1/data_process.sh @@ -79,7 +79,15 @@ download "$TRAIN_URL" "$DATA_DIR/train.parquet" download "$TEST_URL" "$DATA_DIR/test.parquet" if [[ ! -s "$DATA_DIR/e5_Flat.index" ]]; then - cat "$DATA_DIR/part_aa" "$DATA_DIR/part_ab" > "$DATA_DIR/e5_Flat.index" + INDEX_TEMPORARY="$(mktemp "$DATA_DIR/e5_Flat.index.tmp.XXXXXX")" + if ! cat "$DATA_DIR/part_aa" "$DATA_DIR/part_ab" > "$INDEX_TEMPORARY"; then + rm -f "$INDEX_TEMPORARY" + exit 1 + fi + if ! mv -f "$INDEX_TEMPORARY" "$DATA_DIR/e5_Flat.index"; then + rm -f "$INDEX_TEMPORARY" + exit 1 + fi fi if [[ ! -s "$DATA_DIR/wiki-18.jsonl" ]]; then diff --git a/tests/examples/test_search_r1_data.py b/tests/examples/test_search_r1_data.py index eb6739a88..7e1150249 100644 --- a/tests/examples/test_search_r1_data.py +++ b/tests/examples/test_search_r1_data.py @@ -25,11 +25,25 @@ def _write_executable(path: Path, content: str) -> None: def _tool_path(tmp_path: Path, downloader: str) -> Path: tools = tmp_path / "bin" tools.mkdir() - for name in ("awk", "cat", "dirname", "grep", "gzip", "mkdir", "mktemp", "mv", "rm"): + for name in ("awk", "dirname", "grep", "gzip", "mkdir", "mktemp", "mv", "rm"): target = shutil.which(name) assert target is not None (tools / name).symlink_to(target) + cat = shutil.which("cat") + assert cat is not None + _write_executable( + tools / "cat", + f"""#!/bin/sh +if [ "${{MOCK_CAT_MODE:-success}}" = "fail-index" ] && [ "$#" -eq 2 ] && \ + [ "${{1##*/}}" = "part_aa" ] && [ "${{2##*/}}" = "part_ab" ]; then + {cat} "$1" + exit 23 +fi +exec {cat} "$@" +""", + ) + _write_executable( tools / "conda", """#!/bin/sh @@ -107,6 +121,7 @@ def _environment(tmp_path: Path, downloader: str) -> tuple[dict[str, str], Path, "MOCK_DOWNLOAD_MODE": "success", "MOCK_DOWNLOAD_LOG": str(log), "MOCK_FIXTURE_DIR": str(_fixtures(tmp_path)), + "MOCK_CAT_MODE": "success", } return env, data, log @@ -149,3 +164,20 @@ def test_index_ignores_unrelated_part_files(tmp_path: Path) -> None: assert result.returncode == 0, result.stderr assert (data / "e5_Flat.index").read_bytes() == b"AAAB" + + +@LINUX_ONLY +def test_failed_index_concatenation_does_not_poison_cache_and_can_retry(tmp_path: Path) -> None: + env, data, _ = _environment(tmp_path, "curl") + env["MOCK_CAT_MODE"] = "fail-index" + + failed = _run(env) + + assert failed.returncode != 0 + assert not (data / "e5_Flat.index").exists() + assert list(data.glob("e5_Flat.index.tmp.*")) == [] + + env["MOCK_CAT_MODE"] = "success" + retried = _run(env) + assert retried.returncode == 0, retried.stderr + assert (data / "e5_Flat.index").read_bytes() == b"AAAB"