use std::collections::BTreeMap;
use std::path::Path;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use sha2::{Digest, Sha256};
use crate::provider::{CompletionRequest, CompletionResponse, ModelProvider, ProviderError, Usage};
use crate::tools::{ToolError, ToolHost, ToolInvocation};
pub const CASSETTE_VERSION: &str = "0.2";
pub const SUPPORTED_CASSETTE_VERSIONS: &[&str] = &["0.1", "0.2"];
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct Cassette {
pub cassette_version: String,
pub agent: String,
#[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
pub inputs: BTreeMap<String, Value>,
pub interactions: Vec<Interaction>,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub tool_calls: Vec<ToolExchange>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct Interaction {
pub index: usize,
pub node: String,
pub request_digest: String,
pub response_type: String,
pub value: Value,
pub usage: Usage,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub model: Option<String>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct ToolExchange {
pub index: usize,
pub node: String,
pub tool: String,
pub invocation_digest: String,
pub result_type: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub value: Option<Value>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub error: Option<String>,
}
impl Cassette {
pub fn new(agent: impl Into<String>) -> Cassette {
Cassette {
cassette_version: CASSETTE_VERSION.to_string(),
agent: agent.into(),
inputs: BTreeMap::new(),
interactions: Vec::new(),
tool_calls: Vec::new(),
}
}
pub fn to_canonical_json(&self) -> String {
let mut json =
serde_json::to_string_pretty(self).expect("a cassette is always serializable");
json.push('\n');
json
}
pub fn from_json(text: &str) -> Result<Cassette, String> {
let cassette: Cassette =
serde_json::from_str(text).map_err(|error| format!("invalid cassette: {error}"))?;
if !SUPPORTED_CASSETTE_VERSIONS.contains(&cassette.cassette_version.as_str()) {
return Err(format!(
"cassette version `{}` is not supported by this compiler (supported: {})",
cassette.cassette_version,
SUPPORTED_CASSETTE_VERSIONS.join(", ")
));
}
if cassette.cassette_version == "0.1" && !cassette.tool_calls.is_empty() {
return Err(
"the cassette states version `0.1` and carries tool calls, which 0.1 has no \
field for; re-record it"
.to_string(),
);
}
Ok(cassette)
}
pub fn load(path: impl AsRef<Path>) -> Result<Cassette, String> {
let path = path.as_ref();
let text = std::fs::read_to_string(path)
.map_err(|error| format!("cannot read {}: {error}", path.display()))?;
Cassette::from_json(&text)
}
pub fn save(&self, path: impl AsRef<Path>) -> Result<(), String> {
let path = path.as_ref();
if let Some(parent) = path.parent() {
if !parent.as_os_str().is_empty() {
std::fs::create_dir_all(parent)
.map_err(|error| format!("cannot create {}: {error}", parent.display()))?;
}
}
std::fs::write(path, self.to_canonical_json())
.map_err(|error| format!("cannot write {}: {error}", path.display()))
}
}
pub struct ReplayProvider {
cassette: Cassette,
position: usize,
strict: bool,
}
impl ReplayProvider {
pub fn new(cassette: Cassette) -> ReplayProvider {
ReplayProvider {
cassette,
position: 0,
strict: true,
}
}
pub fn lenient(mut self) -> ReplayProvider {
self.strict = false;
self
}
pub fn skipping(mut self, played: usize) -> ReplayProvider {
self.position = played.min(self.cassette.interactions.len());
self
}
pub fn remaining(&self) -> usize {
self.cassette
.interactions
.len()
.saturating_sub(self.position)
}
}
impl ModelProvider for ReplayProvider {
fn name(&self) -> &str {
"replay"
}
fn complete(
&mut self,
request: &CompletionRequest,
) -> Result<CompletionResponse, ProviderError> {
let Some(interaction) = self.cassette.interactions.get(self.position) else {
return Err(ProviderError::Cassette(format!(
"the cassette has {} interaction(s) but the run asked for another at node `{}`; \
re-record it",
self.cassette.interactions.len(),
request.node
)));
};
self.position += 1;
if interaction.response_type != request.response_type {
return Err(ProviderError::Cassette(format!(
"interaction {} recorded a `{}` response but node `{}` now asks for `{}`; \
re-record the cassette",
interaction.index, interaction.response_type, request.node, request.response_type
)));
}
if self.strict && interaction.request_digest != request.digest() {
return Err(ProviderError::Cassette(format!(
"interaction {} was recorded for a different request at node `{}`. \
The prompt or its context changed since recording — \
re-record the cassette and review the diff.",
interaction.index, request.node
)));
}
Ok(CompletionResponse {
value: interaction.value.clone(),
usage: interaction.usage,
model: interaction
.model
.clone()
.unwrap_or_else(|| "replay".to_string()),
})
}
}
pub struct RecordingProvider<P: ModelProvider> {
inner: P,
cassette: Cassette,
}
impl<P: ModelProvider> RecordingProvider<P> {
pub fn new(inner: P, agent: impl Into<String>) -> RecordingProvider<P> {
RecordingProvider {
inner,
cassette: Cassette::new(agent),
}
}
pub fn with_inputs(mut self, inputs: BTreeMap<String, Value>) -> RecordingProvider<P> {
self.cassette.inputs = inputs;
self
}
pub fn finish(self) -> Cassette {
self.cassette
}
fn record(&mut self, request: &CompletionRequest, response: &CompletionResponse) {
self.cassette.interactions.push(Interaction {
index: self.cassette.interactions.len(),
node: request.node.clone(),
request_digest: request.digest(),
response_type: request.response_type.clone(),
value: response.value.clone(),
usage: response.usage,
model: Some(response.model.clone()),
});
}
}
impl<P: ModelProvider> ModelProvider for RecordingProvider<P> {
fn name(&self) -> &str {
self.inner.name()
}
fn complete(
&mut self,
request: &CompletionRequest,
) -> Result<CompletionResponse, ProviderError> {
let response = self.inner.complete(request)?;
self.record(request, &response);
Ok(response)
}
fn streams(&self) -> bool {
self.inner.streams()
}
fn complete_streaming(
&mut self,
request: &CompletionRequest,
on_delta: crate::provider::DeltaSink<'_>,
) -> Result<CompletionResponse, ProviderError> {
let response = self.inner.complete_streaming(request, on_delta)?;
self.record(request, &response);
Ok(response)
}
}
pub fn invocation_digest(invocation: &ToolInvocation) -> String {
let mut hasher = Sha256::new();
hasher.update(invocation.agent.as_bytes());
hasher.update([0]);
hasher.update(invocation.reference.as_bytes());
hasher.update([0]);
hasher.update(invocation.result_type.as_bytes());
for (name, value) in &invocation.arguments {
hasher.update([0]);
hasher.update(name.as_bytes());
hasher.update([0]);
hasher.update(value.to_string().as_bytes());
}
format!("{:x}", hasher.finalize())
}
pub struct ReplayToolHost {
calls: Vec<ToolExchange>,
position: usize,
strict: bool,
}
impl ReplayToolHost {
pub fn new(calls: Vec<ToolExchange>) -> ReplayToolHost {
ReplayToolHost {
calls,
position: 0,
strict: true,
}
}
pub fn lenient(mut self) -> ReplayToolHost {
self.strict = false;
self
}
pub fn skipping(mut self, played: usize) -> ReplayToolHost {
self.position = played.min(self.calls.len());
self
}
pub fn remaining(&self) -> usize {
self.calls.len().saturating_sub(self.position)
}
}
impl ToolHost for ReplayToolHost {
fn name(&self) -> &str {
"replay"
}
fn provides(&self, _tool: &str) -> bool {
true
}
fn call(&mut self, invocation: &ToolInvocation) -> Result<Value, ToolError> {
let Some(recorded) = self.calls.get(self.position) else {
return Err(ToolError::Failed(format!(
"the cassette records {} tool call(s) and the run asked for another: `{}`; re-record it",
self.calls.len(),
invocation.name
)));
};
self.position += 1;
if recorded.tool != invocation.name {
return Err(ToolError::Failed(format!(
"tool call {} recorded `{}` and the run called `{}`; re-record the cassette",
recorded.index, recorded.tool, invocation.name
)));
}
if self.strict && recorded.invocation_digest != invocation_digest(invocation) {
return Err(ToolError::Failed(format!(
"tool call {} recorded different arguments for `{}`. The call changed since recording — re-record the cassette and review the diff.",
recorded.index, recorded.tool
)));
}
match (&recorded.value, &recorded.error) {
(Some(value), _) => Ok(value.clone()),
(None, Some(error)) => Err(ToolError::Failed(error.clone())),
(None, None) => Err(ToolError::InvalidResult(format!(
"tool call {} for `{}` recorded neither a value nor an error",
recorded.index, recorded.tool
))),
}
}
}
pub struct RecordingTools<H: ToolHost> {
inner: H,
calls: Vec<ToolExchange>,
}
impl<H: ToolHost> RecordingTools<H> {
pub fn new(inner: H) -> RecordingTools<H> {
RecordingTools {
inner,
calls: Vec::new(),
}
}
pub fn finish(self) -> Vec<ToolExchange> {
self.calls
}
}
impl<H: ToolHost> ToolHost for RecordingTools<H> {
fn name(&self) -> &str {
self.inner.name()
}
fn provides(&self, tool: &str) -> bool {
self.inner.provides(tool)
}
fn call(&mut self, invocation: &ToolInvocation) -> Result<Value, ToolError> {
let result = self.inner.call(invocation);
let (value, error) = match &result {
Ok(value) => (Some(value.clone()), None),
Err(error) => (None, Some(error.to_string())),
};
self.calls.push(ToolExchange {
index: self.calls.len(),
node: invocation.node.clone(),
tool: invocation.name.clone(),
invocation_digest: invocation_digest(invocation),
result_type: invocation.result_type.clone(),
value,
error,
});
result
}
}
pub struct ScriptedProvider {
answers: Vec<Value>,
position: usize,
usage: Usage,
}
impl ScriptedProvider {
pub fn new(answers: Vec<Value>) -> ScriptedProvider {
ScriptedProvider {
answers,
position: 0,
usage: Usage {
input_tokens: 10,
output_tokens: 5,
cache_read_tokens: 0,
},
}
}
pub fn with_usage(mut self, usage: Usage) -> ScriptedProvider {
self.usage = usage;
self
}
pub fn calls(&self) -> usize {
self.position
}
}
impl ModelProvider for ScriptedProvider {
fn name(&self) -> &str {
"scripted"
}
fn complete(
&mut self,
request: &CompletionRequest,
) -> Result<CompletionResponse, ProviderError> {
let Some(value) = self.answers.get(self.position).cloned() else {
return Err(ProviderError::Cassette(format!(
"the script has {} answer(s) but node `{}` asked for another",
self.answers.len(),
request.node
)));
};
self.position += 1;
Ok(CompletionResponse {
value,
usage: self.usage,
model: "scripted".to_string(),
})
}
}
pub fn load_directory(dir: impl AsRef<Path>) -> Result<BTreeMap<String, Cassette>, String> {
let dir = dir.as_ref();
let mut cassettes = BTreeMap::new();
let entries = std::fs::read_dir(dir)
.map_err(|error| format!("cannot read {}: {error}", dir.display()))?;
for entry in entries.flatten() {
let path = entry.path();
if path.extension().and_then(|ext| ext.to_str()) != Some("json") {
continue;
}
let name = path
.file_stem()
.and_then(|stem| stem.to_str())
.unwrap_or_default()
.to_string();
cassettes.insert(name, Cassette::load(&path)?);
}
Ok(cassettes)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::schema::ResponseShape;
use serde_json::json;
fn request(node: &str, prompt: &str) -> CompletionRequest {
CompletionRequest {
node: node.into(),
model: crate::provider::ModelSelection::Default,
system: None,
prompt: prompt.into(),
context: Vec::new(),
response_type: "markdown".into(),
shape: ResponseShape::Prose,
max_tokens: 1024,
}
}
fn recorded() -> Cassette {
let mut provider = RecordingProvider::new(
ScriptedProvider::new(vec![json!("first"), json!("second")]),
"test.Agent",
);
provider.complete(&request("n0", "one")).unwrap();
provider.complete(&request("n1", "two")).unwrap();
provider.finish()
}
#[test]
fn a_cassette_round_trips_through_canonical_json() {
let cassette = recorded();
let parsed = Cassette::from_json(&cassette.to_canonical_json()).unwrap();
assert_eq!(parsed, cassette);
}
#[test]
fn canonical_json_ends_with_a_newline() {
assert!(recorded().to_canonical_json().ends_with("}\n"));
}
#[test]
fn replay_serves_recorded_answers_in_order() {
let mut provider = ReplayProvider::new(recorded());
assert_eq!(
provider.complete(&request("n0", "one")).unwrap().value,
json!("first")
);
assert_eq!(
provider.complete(&request("n1", "two")).unwrap().value,
json!("second")
);
assert_eq!(provider.remaining(), 0);
}
#[test]
fn replay_rejects_a_changed_prompt() {
let mut provider = ReplayProvider::new(recorded());
let error = provider
.complete(&request("n0", "a different prompt"))
.unwrap_err();
let message = error.to_string();
assert!(
message.contains("recorded for a different request"),
"{message}"
);
assert!(message.contains("re-record"), "{message}");
}
#[test]
fn replay_rejects_a_changed_response_type() {
let mut provider = ReplayProvider::new(recorded());
let mut changed = request("n0", "one");
changed.response_type = "string".into();
let error = provider.complete(&changed).unwrap_err();
assert!(
error.to_string().contains("recorded a `markdown` response"),
"{error}"
);
}
#[test]
fn replay_reports_an_exhausted_cassette() {
let mut provider = ReplayProvider::new(recorded());
provider.complete(&request("n0", "one")).unwrap();
provider.complete(&request("n1", "two")).unwrap();
let error = provider.complete(&request("n2", "three")).unwrap_err();
assert!(error.to_string().contains("asked for another"), "{error}");
}
fn tool_call(name: &str, path: &str) -> ToolInvocation {
ToolInvocation {
node: "n0".into(),
agent: "test.Agent".into(),
reference: format!("mcp:{name}"),
name: name.into(),
transport: "mcp".into(),
arguments: [("path".to_string(), json!(path))].into(),
effects: vec!["filesystem_read".into()],
result_type: "text".into(),
}
}
struct FixedTool(Result<Value, &'static str>);
impl ToolHost for FixedTool {
fn name(&self) -> &str {
"fixed"
}
fn provides(&self, _tool: &str) -> bool {
true
}
fn call(&mut self, _invocation: &ToolInvocation) -> Result<Value, ToolError> {
self.0
.clone()
.map_err(|error| ToolError::Failed(error.into()))
}
}
#[test]
fn a_recorded_tool_call_replays_without_reaching_anything() {
let mut recorder = RecordingTools::new(FixedTool(Ok(json!("# Sample"))));
recorder
.call(&tool_call("fs.read_file", "README.md"))
.unwrap();
let recorded = recorder.finish();
assert_eq!(recorded.len(), 1);
assert_eq!(recorded[0].tool, "fs.read_file");
let mut replay = ReplayToolHost::new(recorded);
assert_eq!(
replay
.call(&tool_call("fs.read_file", "README.md"))
.unwrap(),
json!("# Sample")
);
assert_eq!(replay.remaining(), 0);
}
#[test]
fn a_recorded_failure_replays_as_a_failure() {
let mut recorder = RecordingTools::new(FixedTool(Err("no such file")));
assert!(recorder
.call(&tool_call("fs.read_file", "gone.md"))
.is_err());
let recorded = recorder.finish();
assert_eq!(recorded[0].value, None);
assert!(recorded[0]
.error
.as_deref()
.unwrap()
.contains("no such file"));
let mut replay = ReplayToolHost::new(recorded);
let error = replay
.call(&tool_call("fs.read_file", "gone.md"))
.unwrap_err();
assert!(error.to_string().contains("no such file"), "{error}");
}
#[test]
fn replay_refuses_a_call_whose_arguments_changed() {
let mut recorder = RecordingTools::new(FixedTool(Ok(json!("# Sample"))));
recorder
.call(&tool_call("fs.read_file", "README.md"))
.unwrap();
let mut replay = ReplayToolHost::new(recorder.finish());
let error = replay
.call(&tool_call("fs.read_file", "notes.md"))
.unwrap_err();
assert!(error.to_string().contains("re-record"), "{error}");
}
#[test]
fn replay_refuses_a_different_tool_by_name() {
let mut recorder = RecordingTools::new(FixedTool(Ok(json!("# Sample"))));
recorder
.call(&tool_call("fs.read_file", "README.md"))
.unwrap();
let mut replay = ReplayToolHost::new(recorder.finish());
let error = replay
.call(&tool_call("fs.list_dir", "README.md"))
.unwrap_err();
let message = error.to_string();
assert!(message.contains("fs.read_file"), "{message}");
assert!(message.contains("fs.list_dir"), "{message}");
}
#[test]
fn replay_reports_a_call_beyond_the_recording() {
let mut replay = ReplayToolHost::new(Vec::new());
let error = replay
.call(&tool_call("fs.read_file", "README.md"))
.unwrap_err();
assert!(error.to_string().contains("asked for another"), "{error}");
}
#[test]
fn the_invocation_digest_ignores_effects_and_notices_the_agent() {
let base = tool_call("fs.read_file", "README.md");
let mut other_effects = tool_call("fs.read_file", "README.md");
other_effects.effects = vec!["filesystem_write".into()];
assert_eq!(
invocation_digest(&base),
invocation_digest(&other_effects),
"effects say what a call may do, not what it answers"
);
let mut other_agent = tool_call("fs.read_file", "README.md");
other_agent.agent = "test.Other".into();
assert_ne!(
invocation_digest(&base),
invocation_digest(&other_agent),
"two agents hold different policies, so the same call from another is another call"
);
}
#[test]
fn a_zero_one_cassette_still_replays_and_a_lying_one_does_not() {
let mut cassette = recorded();
cassette.cassette_version = "0.1".into();
let parsed = Cassette::from_json(&cassette.to_canonical_json()).unwrap();
assert!(parsed.tool_calls.is_empty());
cassette.tool_calls.push(ToolExchange {
index: 0,
node: "n0".into(),
tool: "fs.read_file".into(),
invocation_digest: "x".into(),
result_type: "text".into(),
value: Some(json!("hi")),
error: None,
});
let error = Cassette::from_json(&cassette.to_canonical_json()).unwrap_err();
assert!(error.contains("re-record"), "{error}");
}
#[test]
fn a_future_cassette_version_is_rejected() {
let mut cassette = recorded();
cassette.cassette_version = "9.0".into();
let error = Cassette::from_json(&cassette.to_canonical_json()).unwrap_err();
assert!(error.contains("not supported"), "{error}");
}
#[test]
fn lenient_replay_tolerates_a_changed_prompt() {
let mut provider = ReplayProvider::new(recorded()).lenient();
assert!(provider.complete(&request("n0", "changed")).is_ok());
}
}