跳转至

search_r1_like_utils.py — verl/tools/utils/search_r1_like_utils.py

文件路径

verl/tools/utils/search_r1_like_utils.py

文件概述

这个文件实现了类似 DeepSeek-R1 风格的搜索工具底层逻辑,包括调用远程搜索 API 和处理搜索结果。它是 SearchTool 的核心依赖,提供了健壮的 HTTP 请求、重试机制和结果格式化功能。

关键代码讲解

1. 常量定义

DEFAULT_TIMEOUT = 30   # 默认搜索请求超时(秒)
MAX_RETRIES = 10       # 最大重试次数
INITIAL_RETRY_DELAY = 1  # 初始重试延迟(秒)
API_TIMEOUT = 10       # API 超时(秒)

2. 调用搜索 API(核心函数)

def call_search_api(
    retrieval_service_url: str,
    query_list: list[str],
    topk: int = 3,
    return_scores: bool = True,
    timeout: int = DEFAULT_TIMEOUT,
) -> tuple[Optional[dict[str, Any]], Optional[str]]:
    request_id = str(uuid.uuid4())
    payload = {"queries": query_list, "topk": topk, "return_scores": return_scores}
    headers = {"Content-Type": "application/json", "Accept": "application/json"}

    for attempt in range(MAX_RETRIES):
        try:
            response = requests.post(retrieval_service_url, headers=headers, json=payload, timeout=timeout)

            # 服务器错误时重试
            if response.status_code in [500, 502, 503, 504]:
                delay = INITIAL_RETRY_DELAY * (attempt + 1)
                time.sleep(delay)
                continue

            response.raise_for_status()
            return response.json(), None

        except requests.exceptions.ConnectionError as e:
            delay = INITIAL_RETRY_DELAY * (attempt + 1)
            time.sleep(delay)
            continue
        except requests.exceptions.Timeout as e:
            delay = INITIAL_RETRY_DELAY * (attempt + 1)
            time.sleep(delay)
            continue
        except requests.exceptions.RequestException as e:
            break  # 其他请求错误不重试
        except json.JSONDecodeError as e:
            break  # JSON 解析错误不重试

    return None, last_error

关键设计: - 线性递增重试:每次重试的延迟为 INITIAL_RETRY_DELAY * (attempt + 1),即 1s、2s、3s... - 可重试错误:服务器错误(5xx)、连接错误、超时错误会重试。 - 不可重试错误:客户端错误(4xx)、JSON 解析错误立即终止。 - 请求 ID:每次请求生成唯一 ID,方便日志追踪。

3. 格式化搜索结果

def _passages2string(retrieval_result):
    format_reference = ""
    for idx, doc_item in enumerate(retrieval_result):
        content = doc_item["document"]["contents"]
        title = content.split("\n")[0]
        text = "\n".join(content.split("\n")[1:])
        format_reference += f"Doc {idx + 1} (Title: {title})\n{text}\n\n"
    return format_reference.strip()

将检索结果转换为可读格式:

Doc 1 (Title: 量子计算简介)
量子计算是一种利用量子力学原理...

Doc 2 (Title: 经典计算与量子计算对比)
与经典计算机不同...

4. 批量搜索

def perform_single_search_batch(
    retrieval_service_url, query_list, topk=3, concurrent_semaphore=None, timeout=DEFAULT_TIMEOUT
) -> tuple[str, dict[str, Any]]:
    # 可选的并发控制
    if concurrent_semaphore:
        with concurrent_semaphore:
            api_response, error_msg = call_search_api(...)
    else:
        api_response, error_msg = call_search_api(...)

    # 处理结果
    if error_msg:
        metadata["status"] = "api_error"
        result_text = json.dumps({"result": f"Search error: {error_msg}"})
    elif api_response:
        raw_results = api_response.get("result", [])
        if raw_results:
            pretty_results = []
            for retrieval in raw_results:
                formatted = _passages2string(retrieval)
                pretty_results.append(formatted)
            final_result = "\n---\n".join(pretty_results)
            result_text = json.dumps({"result": final_result})
            metadata["status"] = "success"

    return result_text, metadata

这是 SearchTool.execute_search 调用的函数: 1. 可选地使用信号量控制并发(但在 Ray 框架下,并发由 Ray 管理,所以通常传 None)。 2. 调用搜索 API。 3. 格式化结果,用 \n---\n 分隔不同查询的结果。 4. 返回 JSON 字符串和 metadata。

核心类/函数列表

函数 作用
call_search_api 调用远程搜索 API(带重试)
_passages2string 将检索结果格式化为可读文本
perform_single_search_batch 执行批量搜索并处理结果

与其他模块的关系

  • 被 SearchTool(search_tool.py)的 execute_search 方法调用。
  • 不依赖 verl 的其他模块(是一个相对独立的工具函数集合)。
  • 仅使用标准库(requests、json、threading)。

小结

search_r1_like_utils.py 提供了健壮的搜索 API 调用能力,包括重试机制、错误处理和结果格式化。它是搜索工具的底层实现,确保了在网络不稳定时也能可靠地获取搜索结果。