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()
将检索结果转换为可读格式:
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 调用能力,包括重试机制、错误处理和结果格式化。它是搜索工具的底层实现,确保了在网络不稳定时也能可靠地获取搜索结果。