use std::path::{Path, PathBuf};
use std::sync::atomic::{AtomicUsize, Ordering};
use async_trait::async_trait;
use serde::{Deserialize, Serialize};
use crate::cm_types::{ChatRequest, Message};
use crate::cm_llm::backend::ChatCompletionsBackend;
use crate::cm_llm::chat_params::StreamChatParams;
use crate::cm_llm::fingerprint::RequestFingerprint;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum E2eMode {
Real,
Record,
Replay,
}
impl E2eMode {
pub fn as_str(self) -> &'static str {
match self {
Self::Real => "real",
Self::Record => "record",
Self::Replay => "replay",
}
}
}
pub fn detect_mode_from_env() -> E2eMode {
let real = std::env::var("REAL_LLM_E2E").is_ok();
let record = std::env::var("CM_E2E_RECORD").is_ok();
match (real, record) {
(true, true) => E2eMode::Record,
(true, false) => E2eMode::Real,
_ => E2eMode::Replay,
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RecordingManifest {
pub test_name: String,
pub recorded_at: String,
pub model: String,
pub rounds: usize,
#[serde(skip_serializing_if = "Option::is_none")]
pub crabmate_version: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RecordedRequest {
pub round: usize,
pub fingerprint: String,
pub model: String,
pub messages: serde_json::Value,
pub tools_count: usize,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RecordedResponse {
pub round: usize,
pub fingerprint: String,
pub finish_reason: String,
pub elapsed_ms: u64,
pub message: Message,
}
pub struct RecordingBackend {
inner: Box<dyn ChatCompletionsBackend>,
recordings_dir: PathBuf,
test_name: String,
call_seq: AtomicUsize,
manifest_written: std::sync::atomic::AtomicBool,
model_for_manifest: std::sync::Mutex<Option<String>>,
}
impl RecordingBackend {
pub fn new(
inner: Box<dyn ChatCompletionsBackend>,
recordings_dir: impl Into<PathBuf>,
test_name: impl Into<String>,
) -> Self {
Self {
inner,
recordings_dir: recordings_dir.into(),
test_name: test_name.into(),
call_seq: AtomicUsize::new(0),
manifest_written: std::sync::atomic::AtomicBool::new(false),
model_for_manifest: std::sync::Mutex::new(None),
}
}
fn test_dir(&self) -> PathBuf {
self.recordings_dir.join(&self.test_name)
}
fn write_json(&self, filename: &str, value: &impl Serialize) {
let dir = self.test_dir();
let _ = std::fs::create_dir_all(&dir);
let path = dir.join(filename);
if let Ok(json) = serde_json::to_string_pretty(value) {
let _ = std::fs::write(path, json);
}
}
fn maybe_write_manifest(&self, model: &str) {
if self.manifest_written.swap(true, Ordering::SeqCst) {
return;
}
*self.model_for_manifest.lock().unwrap() = Some(model.to_string());
let manifest = RecordingManifest {
test_name: self.test_name.clone(),
recorded_at: chrono_now_iso(),
model: model.to_string(),
rounds: 0, crabmate_version: None,
};
self.write_json("manifest.json", &manifest);
}
pub fn finalize_manifest(&self) {
let rounds = self.call_seq.load(Ordering::SeqCst);
let model = self.model_for_manifest.lock().unwrap().clone();
let manifest = RecordingManifest {
test_name: self.test_name.clone(),
recorded_at: chrono_now_iso(),
model: model.unwrap_or_default(),
rounds,
crabmate_version: None,
};
self.write_json("manifest.json", &manifest);
}
}
#[async_trait]
impl ChatCompletionsBackend for RecordingBackend {
async fn stream_chat(
&self,
params: &StreamChatParams<'_>,
req: &mut ChatRequest,
) -> Result<(Message, String), Box<dyn std::error::Error + Send + Sync>> {
let round = self.call_seq.fetch_add(1, Ordering::SeqCst);
let fp = RequestFingerprint::from_request(req, round);
self.maybe_write_manifest(&req.model);
let messages_value = serde_json::to_value(&req.messages).unwrap_or(serde_json::Value::Null);
let recorded_req = RecordedRequest {
round,
fingerprint: fp.hash.clone(),
model: req.model.clone(),
messages: messages_value,
tools_count: req.tools.as_ref().map_or(0, |t| t.len()),
};
self.write_json(&format!("round_{round}_req.json"), &recorded_req);
let start = std::time::Instant::now();
let result = self.inner.stream_chat(params, req).await;
let elapsed_ms = start.elapsed().as_millis() as u64;
match &result {
Ok((msg, finish_reason)) => {
let recorded_resp = RecordedResponse {
round,
fingerprint: fp.hash.clone(),
finish_reason: finish_reason.clone(),
elapsed_ms,
message: msg.clone(),
};
self.write_json(&format!("round_{round}_resp.json"), &recorded_resp);
}
Err(e) => {
let err_snapshot = serde_json::json!({
"round": round,
"fingerprint": fp.hash,
"error": e.to_string(),
});
self.write_json(&format!("round_{round}_error.json"), &err_snapshot);
}
}
result
}
}
#[derive(Debug)]
pub struct ReplayBackend {
responses: Vec<ReplayEntry>,
call_seq: AtomicUsize,
}
#[derive(Debug)]
struct ReplayEntry {
#[allow(dead_code)]
fingerprint: String,
message: Message,
finish_reason: String,
}
fn parse_round_resp_filename(filename: &str) -> Option<usize> {
let rest = filename.strip_prefix("round_")?;
let n_str = rest.strip_suffix("_resp.json")?;
n_str.parse().ok()
}
fn load_replay_entry(path: &Path) -> Result<(usize, ReplayEntry), String> {
let filename = path
.file_name()
.map(|s| s.to_string_lossy().into_owned())
.unwrap_or_default();
let Some(n) = parse_round_resp_filename(&filename) else {
return Err(format!("不是 round_N_resp.json: {}", path.display()));
};
let content =
std::fs::read_to_string(path).map_err(|e| format!("读取 {} 失败: {e}", path.display()))?;
let resp: RecordedResponse =
serde_json::from_str(&content).map_err(|e| format!("解析 {} 失败: {e}", path.display()))?;
Ok((
n,
ReplayEntry {
fingerprint: resp.fingerprint,
message: resp.message,
finish_reason: resp.finish_reason,
},
))
}
fn scan_round_resp_entries(test_dir: &Path) -> Result<Vec<(usize, ReplayEntry)>, String> {
let mut entries = Vec::new();
for entry in std::fs::read_dir(test_dir)
.map_err(|e| format!("读取录制目录失败 {}: {e}", test_dir.display()))?
{
let entry = entry.map_err(|e| format!("读取目录项失败: {e}"))?;
let path = entry.path();
let filename = entry.file_name().to_string_lossy().into_owned();
if parse_round_resp_filename(&filename).is_none() {
continue;
}
entries.push(load_replay_entry(&path)?);
}
Ok(entries)
}
impl ReplayBackend {
pub fn load(recordings_dir: &Path, test_name: &str) -> Result<Self, String> {
let test_dir = recordings_dir.join(test_name);
if !test_dir.is_dir() {
return Err(format!(
"录制目录不存在: {}(test_name={test_name})。\n\
提示:请先用 `REAL_LLM_E2E=1 CM_E2E_RECORD=1 cargo test --test <name>` 录制一次。",
test_dir.display()
));
}
let mut entries = scan_round_resp_entries(&test_dir)?;
if entries.is_empty() {
return Err(format!(
"录制目录为空(无 round_N_resp.json): {}(test_name={test_name})",
test_dir.display()
));
}
entries.sort_by_key(|(n, _)| *n);
let responses: Vec<ReplayEntry> = entries.into_iter().map(|(_, e)| e).collect();
Ok(Self {
responses,
call_seq: AtomicUsize::new(0),
})
}
pub fn recorded_rounds(&self) -> usize {
self.responses.len()
}
}
#[async_trait]
impl ChatCompletionsBackend for ReplayBackend {
async fn stream_chat(
&self,
_params: &StreamChatParams<'_>,
_req: &mut ChatRequest,
) -> Result<(Message, String), Box<dyn std::error::Error + Send + Sync>> {
let round = self.call_seq.fetch_add(1, Ordering::SeqCst);
let entry = self.responses.get(round).ok_or_else(
|| -> Box<dyn std::error::Error + Send + Sync> {
format!(
"ReplayBackend: round {round} 超出录制范围(共 {} 轮录制)。\n\
提示:agent 实际调用了更多轮 LLM,请重新录制。",
self.responses.len()
)
.into()
},
)?;
Ok((entry.message.clone(), entry.finish_reason.clone()))
}
}
pub fn build_e2e_backend(
mode: E2eMode,
real_backend: Box<dyn ChatCompletionsBackend>,
recordings_dir: &Path,
test_name: &str,
) -> Result<Box<dyn ChatCompletionsBackend>, String> {
match mode {
E2eMode::Real => Ok(real_backend),
E2eMode::Record => Ok(Box::new(RecordingBackend::new(
real_backend,
recordings_dir,
test_name,
))),
E2eMode::Replay => Ok(Box::new(ReplayBackend::load(recordings_dir, test_name)?)),
}
}
fn chrono_now_iso() -> String {
use std::time::{SystemTime, UNIX_EPOCH};
let secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|d| d.as_secs())
.unwrap_or(0);
format!("unix:{secs}")
}
#[cfg(test)]
mod tests {
use super::*;
use async_trait::async_trait;
use crate::cm_types::{ChatRequest, Message};
struct StubBackend {
responses: Vec<(Message, String)>,
call_seq: AtomicUsize,
}
impl StubBackend {
fn new(responses: Vec<(Message, String)>) -> Self {
Self {
responses,
call_seq: AtomicUsize::new(0),
}
}
}
#[async_trait]
impl ChatCompletionsBackend for StubBackend {
async fn stream_chat(
&self,
_params: &StreamChatParams<'_>,
_req: &mut ChatRequest,
) -> Result<(Message, String), Box<dyn std::error::Error + Send + Sync>> {
let idx = self.call_seq.fetch_add(1, Ordering::SeqCst);
self.responses
.get(idx)
.cloned()
.ok_or_else(|| "StubBackend: 响应序列耗尽".to_string().into())
}
}
#[test]
fn detect_mode_defaults_to_replay() {
if std::env::var("REAL_LLM_E2E").is_err() {
assert_eq!(detect_mode_from_env(), E2eMode::Replay);
}
}
#[test]
fn replay_errors_when_recordings_missing() {
let tmp = tempfile::tempdir().unwrap();
let result = ReplayBackend::load(tmp.path(), "nonexistent");
assert!(result.is_err());
assert!(result.unwrap_err().contains("录制目录不存在"));
}
#[test]
fn replay_errors_when_empty_dir() {
let tmp = tempfile::tempdir().unwrap();
let test_dir = tmp.path().join("empty_test");
std::fs::create_dir_all(&test_dir).unwrap();
let result = ReplayBackend::load(tmp.path(), "empty_test");
assert!(result.is_err());
assert!(result.unwrap_err().contains("录制目录为空"));
}
#[test]
fn replay_loads_from_manual_files_in_order() {
let tmp = tempfile::tempdir().unwrap();
let test_dir = tmp.path().join("manual_test");
std::fs::create_dir_all(&test_dir).unwrap();
for (round, content) in [(2usize, "third"), (0, "first"), (1, "second")] {
let resp = RecordedResponse {
round,
fingerprint: format!("fp_{round}"),
finish_reason: "stop".to_string(),
elapsed_ms: 100,
message: Message::assistant_only(content.to_string()),
};
std::fs::write(
test_dir.join(format!("round_{round}_resp.json")),
serde_json::to_string_pretty(&resp).unwrap(),
)
.unwrap();
}
std::fs::write(
test_dir.join("round_0_req.json"),
r#"{"round":0,"fingerprint":"x","model":"m","messages":[],"tools_count":0}"#,
)
.unwrap();
std::fs::write(
test_dir.join("manifest.json"),
r#"{"test_name":"manual_test"}"#,
)
.unwrap();
let replay = ReplayBackend::load(tmp.path(), "manual_test").unwrap();
assert_eq!(replay.recorded_rounds(), 3);
}
#[test]
fn build_e2e_backend_real_returns_inner() {
let stub: Box<dyn ChatCompletionsBackend> = Box::new(StubBackend::new(vec![]));
let tmp = tempfile::tempdir().unwrap();
let _ = build_e2e_backend(E2eMode::Real, stub, tmp.path(), "any").unwrap();
}
#[test]
fn build_e2e_backend_record_constructs() {
let stub: Box<dyn ChatCompletionsBackend> = Box::new(StubBackend::new(vec![]));
let tmp = tempfile::tempdir().unwrap();
let _ = build_e2e_backend(E2eMode::Record, stub, tmp.path(), "wrap_test").unwrap();
assert!(!tmp.path().join("wrap_test").exists());
}
#[test]
fn build_e2e_backend_replay_loads_files() {
let tmp = tempfile::tempdir().unwrap();
let test_dir = tmp.path().join("replay_build");
std::fs::create_dir_all(&test_dir).unwrap();
let resp = RecordedResponse {
round: 0,
fingerprint: "fake".to_string(),
finish_reason: "stop".to_string(),
elapsed_ms: 10,
message: Message::assistant_only("manual".to_string()),
};
std::fs::write(
test_dir.join("round_0_resp.json"),
serde_json::to_string_pretty(&resp).unwrap(),
)
.unwrap();
let stub: Box<dyn ChatCompletionsBackend> = Box::new(StubBackend::new(vec![]));
let _ = build_e2e_backend(E2eMode::Replay, stub, tmp.path(), "replay_build").unwrap();
}
}