use crate::config::{AgentConfig, SharedAgentConfig};
use crate::types::Tool;
use crate::cm_benchmark::adapter::{BenchmarkAdapter, create_adapter};
use crate::cm_benchmark::metrics::{BatchSummary, TaskMetrics};
use crate::cm_benchmark::types::{
BatchRunConfig, BenchmarkResult, BenchmarkTask, TaskStatus, parse_task_jsonl_line,
};
use log::{error, info, warn};
use std::collections::HashSet;
use std::io::{BufRead, Write};
use std::path::Path;
use std::sync::Arc;
use std::time::Instant;
pub async fn run_batch(
cfg: &SharedAgentConfig,
client: &reqwest::Client,
api_key: &str,
tools: &[Tool],
batch_cfg: &BatchRunConfig,
) -> Result<(), Box<dyn std::error::Error>> {
let adapter = create_adapter(batch_cfg.benchmark);
info!(
target: "crabmate::benchmark",
"批量测评开始 benchmark={} input={} output={}",
batch_cfg.benchmark.as_str(),
batch_cfg.input_path,
batch_cfg.output_path,
);
let tasks = load_tasks(&batch_cfg.input_path)?;
if tasks.is_empty() {
eprintln!("输入文件为空或无有效任务");
return Ok(());
}
eprintln!(
"[benchmark] 已加载 {} 条任务 ({})",
tasks.len(),
batch_cfg.benchmark.as_str()
);
let existing_ids = if batch_cfg.resume_from_existing {
load_existing_ids(&batch_cfg.output_path)
} else {
HashSet::new()
};
let samples = batch_cfg.samples.max(1);
let mut results: Vec<BenchmarkResult> = Vec::new();
let base_work_dir = {
let g = cfg.read().await;
std::path::Path::new(&g.command_exec.run_command_working_dir)
.canonicalize()
.unwrap_or_else(|_| std::path::PathBuf::from(&g.command_exec.run_command_working_dir))
};
let mut out_file = open_output_file(&batch_cfg.output_path, batch_cfg.resume_from_existing)?;
for (idx, task) in tasks.iter().enumerate() {
let base_snap = {
let g = cfg.read().await;
Arc::new(g.clone())
};
for sample in 0..samples {
if existing_ids.contains(&(task.instance_id.clone(), sample)) {
eprintln!(
"[benchmark] [{}/{}] 跳过已有结果: {} sample={}",
idx + 1,
tasks.len(),
task.instance_id,
sample
);
continue;
}
eprintln!(
"[benchmark] [{}/{}] 开始: {} sample={}",
idx + 1,
tasks.len(),
task.instance_id,
sample
);
let mut result = run_single_task(
&base_snap,
client,
api_key,
tools,
adapter.as_ref(),
task,
&base_work_dir,
batch_cfg,
)
.await;
result.sample_index = sample;
eprintln!(
"[benchmark] [{}/{}] 完成: {} sample={} status={:?} time={:.1}s",
idx + 1,
tasks.len(),
task.instance_id,
sample,
result.status,
result.metrics.wall_time_secs,
);
write_result_line(&mut out_file, &result)?;
results.push(result);
}
}
let summary = BatchSummary::from_results(&results);
let summary_path = summary_path_from_output(&batch_cfg.output_path);
write_summary(&summary_path, &summary)?;
eprintln!("\n[benchmark] 批量测评完成");
eprintln!(
" 总任务: {} 成功: {} 超时: {} 错误: {} 达到轮次上限: {}",
summary.total_tasks,
summary.success_count,
summary.timeout_count,
summary.error_count,
summary.max_rounds_count,
);
eprintln!(
" 平均耗时: {:.1}s 总工具调用: {}",
summary.avg_wall_time_secs, summary.total_tool_calls,
);
eprintln!(" 结果: {}", batch_cfg.output_path);
eprintln!(" 汇总: {}", summary_path);
Ok(())
}
#[allow(clippy::too_many_arguments)]
async fn run_single_task(
cfg: &Arc<AgentConfig>,
client: &reqwest::Client,
api_key: &str,
tools: &[Tool],
adapter: &dyn BenchmarkAdapter,
task: &BenchmarkTask,
base_work_dir: &Path,
batch_cfg: &BatchRunConfig,
) -> BenchmarkResult {
if let Err(e) = adapter.validate_task(task) {
return adapter.extract_result(
task,
None,
base_work_dir,
TaskStatus::Error,
TaskMetrics::default(),
&cfg.llm.model,
Some(format!("输入校验失败: {e}")),
);
}
let work_dir = match adapter.setup_workspace(task, base_work_dir) {
Ok(d) => d,
Err(e) => {
error!(
target: "crabmate::benchmark",
"工作区初始化失败 {}: {e}",
task.instance_id,
);
return adapter.extract_result(
task,
None,
base_work_dir,
TaskStatus::Error,
TaskMetrics::default(),
&cfg.llm.model,
Some(format!("工作区初始化失败: {e}")),
);
}
};
let task_cfg = build_task_config(cfg, adapter, batch_cfg);
let user_prompt = adapter.build_user_prompt(task);
let mut messages =
crate::types::messages_chat_seed(&task_cfg.roles_prompts.system_prompt, &user_prompt);
let start = Instant::now();
let cancel = Arc::new(std::sync::atomic::AtomicBool::new(false));
let timeout_secs = batch_cfg.task_timeout_secs;
let run_fut = crate::run_agent_turn(crate::RunAgentTurnParams::benchmark_batch(
client,
api_key,
&task_cfg,
tools,
&mut messages,
&work_dir,
cancel.clone(),
));
let (status, agent_error) = if timeout_secs > 0 {
match tokio::time::timeout(std::time::Duration::from_secs(timeout_secs), run_fut).await {
Ok(Ok(())) => (TaskStatus::Success, None),
Ok(Err(e)) => {
let msg = e.to_string();
warn!(
target: "crabmate::benchmark",
"任务执行出错 {}: {msg}",
task.instance_id,
);
(TaskStatus::Error, Some(msg))
}
Err(_elapsed) => {
cancel.store(true, std::sync::atomic::Ordering::Relaxed);
warn!(
target: "crabmate::benchmark",
"任务超时 {}: {timeout_secs}s",
task.instance_id,
);
(TaskStatus::Timeout, Some(format!("超时 ({timeout_secs}s)")))
}
}
} else {
match run_fut.await {
Ok(()) => (TaskStatus::Success, None),
Err(e) => {
let msg = e.to_string();
(TaskStatus::Error, Some(msg))
}
}
};
let wall_time = start.elapsed().as_secs_f64();
let raw_reply = messages
.iter()
.rev()
.find(|m| m.role == "assistant")
.and_then(|m| crate::types::message_content_as_str(&m.content));
let tool_calls_count = messages
.iter()
.filter(|m| m.role == "assistant" && m.tool_calls.is_some())
.map(|m| m.tool_calls.as_ref().map_or(0, |tc| tc.len()))
.sum();
let agent_rounds = messages.iter().filter(|m| m.role == "assistant").count();
let metrics = TaskMetrics {
wall_time_secs: wall_time,
tool_calls_count,
agent_rounds,
};
adapter.extract_result(
task,
raw_reply,
&work_dir,
status,
metrics,
&cfg.llm.model,
agent_error,
)
}
fn build_task_config(
base: &Arc<AgentConfig>,
adapter: &dyn BenchmarkAdapter,
batch_cfg: &BatchRunConfig,
) -> Arc<AgentConfig> {
let mut cfg = (**base).clone();
if let Some(ref override_prompt) = batch_cfg.system_prompt_override {
cfg.roles_prompts.system_prompt = override_prompt.clone();
}
if let Some(suffix) = adapter.system_prompt_suffix() {
if !cfg.roles_prompts.system_prompt.is_empty() {
cfg.roles_prompts.system_prompt.push('\n');
}
cfg.roles_prompts.system_prompt.push_str(&suffix);
}
if batch_cfg.max_tool_rounds > 0 {
let estimated_max = 2 + batch_cfg.max_tool_rounds * 3;
if estimated_max < cfg.session_ui.max_message_history
|| cfg.session_ui.max_message_history == 0
{
cfg.session_ui.max_message_history = estimated_max;
}
}
Arc::new(cfg)
}
fn load_tasks(path: &str) -> Result<Vec<BenchmarkTask>, Box<dyn std::error::Error>> {
let file = std::fs::File::open(path).map_err(|e| format!("无法打开输入文件 {path}: {e}"))?;
let reader = std::io::BufReader::new(file);
let mut tasks = Vec::new();
for (line_num, line) in reader.lines().enumerate() {
let line = line.map_err(|e| format!("读取第 {} 行失败: {e}", line_num + 1))?;
match parse_task_jsonl_line(&line) {
Ok(None) => continue,
Ok(Some(task)) => tasks.push(task),
Err(e) => {
warn!(
target: "crabmate::benchmark",
"跳过第 {} 行(JSON 解析失败): {e}",
line_num + 1,
);
}
}
}
Ok(tasks)
}
fn load_existing_ids(path: &str) -> HashSet<(String, usize)> {
let mut ids = HashSet::new();
let Ok(file) = std::fs::File::open(path) else {
return ids;
};
let reader = std::io::BufReader::new(file);
for line in reader.lines().map_while(Result::ok) {
let trimmed = line.trim();
if let Ok(r) = serde_json::from_str::<BenchmarkResult>(trimmed) {
ids.insert((r.instance_id, r.sample_index));
}
}
ids
}
fn open_output_file(path: &str, append: bool) -> Result<std::fs::File, Box<dyn std::error::Error>> {
let file = std::fs::OpenOptions::new()
.create(true)
.append(append)
.truncate(!append)
.write(true)
.open(path)
.map_err(|e| format!("无法打开输出文件 {path}: {e}"))?;
Ok(file)
}
fn write_result_line(
file: &mut std::fs::File,
result: &BenchmarkResult,
) -> Result<(), Box<dyn std::error::Error>> {
let json = serde_json::to_string(result)?;
writeln!(file, "{json}")?;
file.flush()?;
Ok(())
}
fn summary_path_from_output(output_path: &str) -> String {
let p = Path::new(output_path);
let stem = p.file_stem().unwrap_or_default().to_string_lossy();
let parent = p.parent().unwrap_or(Path::new("."));
parent
.join(format!("{stem}_summary.json"))
.to_string_lossy()
.to_string()
}
fn write_summary(path: &str, summary: &BatchSummary) -> Result<(), Box<dyn std::error::Error>> {
let json = serde_json::to_string_pretty(summary)?;
std::fs::write(path, json)?;
eprintln!("[benchmark] 汇总已写入 {path}");
Ok(())
}