use async_trait::async_trait;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::path::PathBuf;
use crate::types::AgentStep;
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
pub enum ResumeError {
#[error("resume store I/O error: {0}")]
Io(String),
#[error("resume store serialization error: {0}")]
Serialize(String),
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct PendingApproval {
pub tool_name: String,
pub arguments: serde_json::Value,
pub tool_id: String,
pub inputs: HashMap<String, String>,
pub steps: Vec<AgentStep>,
pub iteration: usize,
pub tool_calls_consumed: usize,
pub tokens_consumed: Option<usize>,
pub trace_id: Option<String>,
}
#[async_trait]
pub trait ResumeStore: Send + Sync {
async fn save_pending(&self, pending: &PendingApproval) -> Result<(), ResumeError>;
async fn load_pending(&self) -> Result<Option<PendingApproval>, ResumeError>;
async fn clear_pending(&self) -> Result<(), ResumeError>;
}
pub struct FileResumeStore {
dir: PathBuf,
}
impl FileResumeStore {
pub fn new(dir: impl Into<PathBuf>) -> Result<Self, ResumeError> {
let dir = dir.into();
std::fs::create_dir_all(&dir)
.map_err(|e| ResumeError::Io(format!("create dir {}: {}", dir.display(), e)))?;
Ok(Self { dir })
}
fn pending_path(&self) -> PathBuf {
self.dir.join("pending.json")
}
}
#[async_trait]
impl ResumeStore for FileResumeStore {
async fn save_pending(&self, pending: &PendingApproval) -> Result<(), ResumeError> {
let bytes = serde_json::to_vec_pretty(pending)
.map_err(|e| ResumeError::Serialize(e.to_string()))?;
let tmp = self.dir.join("pending.json.tmp");
tokio::fs::write(&tmp, &bytes)
.await
.map_err(|e| ResumeError::Io(format!("write {}: {}", tmp.display(), e)))?;
tokio::fs::rename(&tmp, self.pending_path())
.await
.map_err(|e| ResumeError::Io(format!("rename {}: {}", tmp.display(), e)))?;
Ok(())
}
async fn load_pending(&self) -> Result<Option<PendingApproval>, ResumeError> {
let path = self.pending_path();
let bytes = match tokio::fs::read(&path).await {
Ok(bytes) => bytes,
Err(e) if e.kind() == std::io::ErrorKind::NotFound => return Ok(None),
Err(e) => {
return Err(ResumeError::Io(format!("read {}: {}", path.display(), e)));
}
};
let pending = serde_json::from_slice(&bytes)
.map_err(|e| ResumeError::Serialize(format!("parse {}: {}", path.display(), e)))?;
Ok(Some(pending))
}
async fn clear_pending(&self) -> Result<(), ResumeError> {
let path = self.pending_path();
match tokio::fs::remove_file(&path).await {
Ok(()) => Ok(()),
Err(e) if e.kind() == std::io::ErrorKind::NotFound => Ok(()),
Err(e) => Err(ResumeError::Io(format!("remove {}: {}", path.display(), e))),
}
}
}
#[derive(Default)]
pub struct MemoryResumeStore {
pending: tokio::sync::Mutex<Option<PendingApproval>>,
}
impl MemoryResumeStore {
pub fn new() -> Self {
Self::default()
}
}
#[async_trait]
impl ResumeStore for MemoryResumeStore {
async fn save_pending(&self, pending: &PendingApproval) -> Result<(), ResumeError> {
*self.pending.lock().await = Some(pending.clone());
Ok(())
}
async fn load_pending(&self) -> Result<Option<PendingApproval>, ResumeError> {
Ok(self.pending.lock().await.clone())
}
async fn clear_pending(&self) -> Result<(), ResumeError> {
*self.pending.lock().await = None;
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::types::{AgentAction, ToolInput};
fn sample_pending() -> PendingApproval {
let mut inputs = HashMap::new();
inputs.insert("input".to_string(), "compute".to_string());
PendingApproval {
tool_name: "calculator".to_string(),
arguments: serde_json::json!({"expression": "2 + 3"}),
tool_id: "call_1".to_string(),
inputs,
steps: vec![AgentStep::new(
AgentAction {
tool: "other".to_string(),
tool_input: ToolInput::String {
value: "x".to_string(),
},
log: String::new(),
},
"obs".to_string(),
)],
iteration: 3,
tool_calls_consumed: 4,
tokens_consumed: Some(128),
trace_id: Some("trace-1".to_string()),
}
}
#[tokio::test]
async fn test_file_store_roundtrip() {
let dir = tempfile::tempdir().unwrap();
let store = FileResumeStore::new(dir.path()).unwrap();
assert!(store.load_pending().await.unwrap().is_none());
let pending = sample_pending();
store.save_pending(&pending).await.unwrap();
let loaded = store.load_pending().await.unwrap().unwrap();
assert_eq!(loaded.tool_name, "calculator");
assert_eq!(loaded.arguments, serde_json::json!({"expression": "2 + 3"}));
assert_eq!(loaded.iteration, 3);
assert_eq!(loaded.tool_calls_consumed, 4);
assert_eq!(loaded.tokens_consumed, Some(128));
assert_eq!(loaded.inputs.get("input").unwrap(), "compute");
assert_eq!(loaded.steps.len(), 1);
assert_eq!(loaded.trace_id.as_deref(), Some("trace-1"));
store.clear_pending().await.unwrap();
assert!(store.load_pending().await.unwrap().is_none());
}
#[tokio::test]
async fn test_file_store_clear_idempotent() {
let dir = tempfile::tempdir().unwrap();
let store = FileResumeStore::new(dir.path()).unwrap();
store.clear_pending().await.unwrap();
store.clear_pending().await.unwrap();
}
#[tokio::test]
async fn test_file_store_atomic_no_tmp_left() {
let dir = tempfile::tempdir().unwrap();
let store = FileResumeStore::new(dir.path()).unwrap();
store.save_pending(&sample_pending()).await.unwrap();
let entries: Vec<String> = std::fs::read_dir(dir.path())
.unwrap()
.map(|e| e.unwrap().file_name().to_string_lossy().to_string())
.collect();
assert!(
!entries.iter().any(|n| n == "pending.json.tmp"),
"tmp file should be renamed away, got {entries:?}"
);
assert!(entries.contains(&"pending.json".to_string()));
}
#[tokio::test]
async fn test_memory_store_roundtrip() {
let store = MemoryResumeStore::new();
assert!(store.load_pending().await.unwrap().is_none());
store.save_pending(&sample_pending()).await.unwrap();
assert!(store.load_pending().await.unwrap().is_some());
store.clear_pending().await.unwrap();
assert!(store.load_pending().await.unwrap().is_none());
}
#[test]
fn test_pending_approval_serde_roundtrip() {
let pending = sample_pending();
let bytes = serde_json::to_vec(&pending).unwrap();
let back: PendingApproval = serde_json::from_slice(&bytes).unwrap();
assert_eq!(back.tool_name, pending.tool_name);
assert_eq!(back.arguments, pending.arguments);
assert_eq!(back.steps.len(), pending.steps.len());
assert_eq!(back.trace_id, pending.trace_id);
}
#[test]
fn test_resume_error_display() {
let e = ResumeError::Io("disk full".to_string());
assert!(e.to_string().contains("disk full"));
let e = ResumeError::Serialize("bad json".to_string());
assert!(e.to_string().contains("bad json"));
}
#[tokio::test]
async fn test_file_store_creates_dir() {
let dir = tempfile::tempdir().unwrap();
let nested = dir.path().join("a").join("b");
let store = FileResumeStore::new(&nested).unwrap();
assert!(nested.is_dir());
assert!(store.load_pending().await.is_ok());
}
}