use std::collections::VecDeque;
use std::sync::Arc;
use std::time::SystemTime;
use serde::Deserialize;
use serde::Serialize;
use uuid::Uuid;
use crate::error::CodexErr;
use crate::error::Result as CodexResult;
use crate::models::ContentItem;
use crate::models::ResponseItem;
const MAX_UNDO_STATES: usize = 50;
const MAX_SNAPSHOT_SIZE: usize = 10 * 1024 * 1024;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ConversationSnapshot {
pub id: Uuid,
pub timestamp: SystemTime,
pub items: Vec<ResponseItem>,
pub metadata: SnapshotMetadata,
pub branch_info: Option<BranchInfo>,
pub size_bytes: usize,
pub compressed: bool,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SnapshotMetadata {
pub turn_number: usize,
pub _total_tokens: usize,
pub _model: String,
pub mode: String,
pub user: Option<String>,
pub tags: Vec<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct BranchInfo {
pub name: String,
pub parent_id: Uuid,
pub description: Option<String>,
pub is_active: bool,
}
pub struct UndoRedoManager {
undo_stack: VecDeque<Arc<ConversationSnapshot>>,
redo_stack: VecDeque<Arc<ConversationSnapshot>>,
current_state: Option<Arc<ConversationSnapshot>>,
branches: std::collections::HashMap<Uuid, Vec<Arc<ConversationSnapshot>>>,
total_memory_usage: usize,
max_memory_usage: usize,
}
impl UndoRedoManager {
pub fn new() -> Self {
Self {
undo_stack: VecDeque::with_capacity(MAX_UNDO_STATES),
redo_stack: VecDeque::new(),
current_state: None,
branches: std::collections::HashMap::new(),
total_memory_usage: 0,
max_memory_usage: 100 * 1024 * 1024, }
}
pub fn with_memory_limit(max_memory_mb: usize) -> Self {
Self {
undo_stack: VecDeque::with_capacity(MAX_UNDO_STATES),
redo_stack: VecDeque::new(),
current_state: None,
branches: std::collections::HashMap::new(),
total_memory_usage: 0,
max_memory_usage: max_memory_mb * 1024 * 1024,
}
}
pub fn save_state(
&mut self,
items: Vec<ResponseItem>,
metadata: SnapshotMetadata,
) -> CodexResult<Uuid> {
self.redo_stack.clear();
if let Some(current) = self.current_state.take() {
self.push_to_undo_stack(current);
}
let snapshot = self.create_snapshot(items, metadata)?;
let snapshot_id = snapshot.id;
let snapshot_arc = Arc::new(snapshot);
self.total_memory_usage += snapshot_arc.size_bytes;
self.enforce_memory_limit();
self.current_state = Some(snapshot_arc);
Ok(snapshot_id)
}
pub fn undo(&mut self) -> CodexResult<Option<ConversationSnapshot>> {
if self.undo_stack.is_empty() {
return Ok(None);
}
if let Some(current) = self.current_state.take() {
self.push_to_redo_stack(current);
}
if let Some(previous) = self.undo_stack.pop_back() {
let snapshot = (*previous).clone();
self.current_state = Some(previous);
Ok(Some(snapshot))
} else {
Ok(None)
}
}
pub fn redo(&mut self) -> CodexResult<Option<ConversationSnapshot>> {
if self.redo_stack.is_empty() {
return Ok(None);
}
if let Some(current) = self.current_state.take() {
self.push_to_undo_stack(current);
}
if let Some(next) = self.redo_stack.pop_back() {
let snapshot = (*next).clone();
self.current_state = Some(next);
Ok(Some(snapshot))
} else {
Ok(None)
}
}
pub fn create_branch(
&mut self,
branch_name: String,
description: Option<String>,
items: Vec<ResponseItem>,
metadata: SnapshotMetadata,
) -> CodexResult<Uuid> {
let parent_id = self
.current_state
.as_ref()
.map(|s| s.id)
.ok_or(CodexErr::NoBranchPointAvailable)?;
let branch_info = BranchInfo {
name: branch_name,
parent_id,
description,
is_active: false,
};
let mut snapshot = self.create_snapshot(items, metadata)?;
snapshot.branch_info = Some(branch_info);
let snapshot_id = snapshot.id;
let snapshot_arc = Arc::new(snapshot);
self.branches
.entry(parent_id)
.or_default()
.push(snapshot_arc.clone());
self.total_memory_usage += snapshot_arc.size_bytes;
self.enforce_memory_limit();
Ok(snapshot_id)
}
pub fn switch_to_branch(
&mut self,
branch_id: Uuid,
) -> CodexResult<Option<ConversationSnapshot>> {
let branch_snapshot = self
.branches
.values()
.flatten()
.find(|s| s.id == branch_id)
.cloned();
if let Some(snapshot) = branch_snapshot {
if let Some(current) = self.current_state.take() {
self.push_to_undo_stack(current);
}
self.redo_stack.clear();
let result = (*snapshot).clone();
self.current_state = Some(snapshot);
Ok(Some(result))
} else {
Ok(None)
}
}
pub fn get_branches(&self) -> Vec<(Uuid, BranchInfo)> {
self.branches
.values()
.flatten()
.filter_map(|s| s.branch_info.as_ref().map(|b| (s.id, b.clone())))
.collect()
}
pub fn current_state(&self) -> Option<&ConversationSnapshot> {
self.current_state.as_deref()
}
pub fn undo_history(&self) -> Vec<&ConversationSnapshot> {
self.undo_stack.iter().map(|s| s.as_ref()).collect()
}
pub fn redo_history(&self) -> Vec<&ConversationSnapshot> {
self.redo_stack.iter().map(|s| s.as_ref()).collect()
}
pub fn clear(&mut self) {
self.undo_stack.clear();
self.redo_stack.clear();
self.current_state = None;
self.branches.clear();
self.total_memory_usage = 0;
}
pub fn create_checkpoint(&mut self, name: String) -> CodexResult<Uuid> {
if let Some(current) = &self.current_state {
let mut checkpoint = (**current).clone();
checkpoint.id = Uuid::new_v4();
checkpoint.timestamp = SystemTime::now();
checkpoint
.metadata
.tags
.push(format!("checkpoint:{}", name));
let checkpoint_id = checkpoint.id;
let checkpoint_arc = Arc::new(checkpoint);
self.branches
.entry(current.id)
.or_default()
.push(checkpoint_arc);
Ok(checkpoint_id)
} else {
Err(CodexErr::NoCurrentStateForCheckpoint)
}
}
pub fn restore_checkpoint(
&mut self,
checkpoint_id: Uuid,
) -> CodexResult<Option<ConversationSnapshot>> {
self.switch_to_branch(checkpoint_id)
}
pub fn memory_info(&self) -> MemoryInfo {
MemoryInfo {
total_usage_bytes: self.total_memory_usage,
max_usage_bytes: self.max_memory_usage,
undo_stack_size: self.undo_stack.len(),
redo_stack_size: self.redo_stack.len(),
branch_count: self.branches.values().map(|v| v.len()).sum(),
usage_percentage: (self.total_memory_usage as f64 / self.max_memory_usage as f64)
* 100.0,
}
}
fn create_snapshot(
&self,
items: Vec<ResponseItem>,
metadata: SnapshotMetadata,
) -> CodexResult<ConversationSnapshot> {
let size_bytes = Self::estimate_size(&items);
let compressed = size_bytes > MAX_SNAPSHOT_SIZE;
let snapshot = ConversationSnapshot {
id: Uuid::new_v4(),
timestamp: SystemTime::now(),
items,
metadata,
branch_info: None,
size_bytes,
compressed,
};
Ok(snapshot)
}
fn push_to_undo_stack(&mut self, snapshot: Arc<ConversationSnapshot>) {
while self.undo_stack.len() >= MAX_UNDO_STATES {
if let Some(removed) = self.undo_stack.pop_front() {
self.total_memory_usage =
self.total_memory_usage.saturating_sub(removed.size_bytes);
}
}
self.undo_stack.push_back(snapshot);
}
fn push_to_redo_stack(&mut self, snapshot: Arc<ConversationSnapshot>) {
self.redo_stack.push_back(snapshot);
}
fn enforce_memory_limit(&mut self) {
while self.total_memory_usage > self.max_memory_usage && !self.undo_stack.is_empty() {
if let Some(removed) = self.undo_stack.pop_front() {
self.total_memory_usage =
self.total_memory_usage.saturating_sub(removed.size_bytes);
}
}
}
fn estimate_size(items: &[ResponseItem]) -> usize {
items
.iter()
.map(|item| match item {
ResponseItem::Message { content, .. } => {
content
.iter()
.map(|c| match c {
ContentItem::InputText { text } | ContentItem::OutputText { text } => {
text.len()
}
ContentItem::InputImage { .. } => 1024, })
.sum::<usize>()
}
ResponseItem::Reasoning {
summary, content, ..
} => {
let summary_size: usize = summary.iter().map(|_| 100).sum(); let content_size = content.as_ref().map(|c| c.len() * 50).unwrap_or(0); summary_size + content_size
}
ResponseItem::FunctionCall { arguments, .. } => arguments.len(),
ResponseItem::FunctionCallOutput { output, .. } => {
output.content.len() + 100 }
_ => 256, })
.sum()
}
}
impl Default for UndoRedoManager {
fn default() -> Self {
Self::new()
}
}
#[derive(Debug, Clone)]
pub struct MemoryInfo {
pub total_usage_bytes: usize,
pub max_usage_bytes: usize,
pub undo_stack_size: usize,
pub redo_stack_size: usize,
pub branch_count: usize,
pub usage_percentage: f64,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ConversationDiff {
pub added: Vec<ResponseItem>,
pub removed: Vec<usize>,
pub modified: Vec<(usize, ResponseItem)>,
}
impl ConversationDiff {
pub fn create(old: &[ResponseItem], new: &[ResponseItem]) -> Self {
let mut added = Vec::new();
let mut removed = Vec::new();
let mut modified = Vec::new();
let min_len = old.len().min(new.len());
for i in 0..min_len {
if !Self::items_equal(&old[i], &new[i]) {
modified.push((i, new[i].clone()));
}
}
if new.len() > old.len() {
added.extend(new[old.len()..].iter().cloned());
}
if old.len() > new.len() {
for i in new.len()..old.len() {
removed.push(i);
}
}
Self {
added,
removed,
modified,
}
}
pub fn apply(&self, items: &mut Vec<ResponseItem>) {
for (index, new_item) in &self.modified {
if *index < items.len() {
items[*index] = new_item.clone();
}
}
for &index in self.removed.iter().rev() {
if index < items.len() {
items.remove(index);
}
}
items.extend(self.added.iter().cloned());
}
fn items_equal(a: &ResponseItem, b: &ResponseItem) -> bool {
match (a, b) {
(
ResponseItem::Message {
role: r1,
content: c1,
..
},
ResponseItem::Message {
role: r2,
content: c2,
..
},
) => r1 == r2 && Self::content_equal(c1, c2),
_ => false, }
}
fn content_equal(a: &[ContentItem], b: &[ContentItem]) -> bool {
if a.len() != b.len() {
return false;
}
a.iter()
.zip(b.iter())
.all(|(a_item, b_item)| match (a_item, b_item) {
(
ContentItem::InputText { text: t1 } | ContentItem::OutputText { text: t1 },
ContentItem::InputText { text: t2 } | ContentItem::OutputText { text: t2 },
) => t1 == t2,
_ => false,
})
}
}
#[cfg(test)]
mod tests {
use super::*;
fn create_test_message(role: &str, content: &str) -> ResponseItem {
ResponseItem::Message {
id: None,
role: role.to_string(),
content: vec![ContentItem::OutputText {
text: content.to_string(),
}],
}
}
fn create_test_metadata(turn: usize) -> SnapshotMetadata {
SnapshotMetadata {
turn_number: turn,
_total_tokens: turn * 100,
_model: "test-model".to_string(),
mode: "Build".to_string(),
user: None,
tags: Vec::new(),
}
}
#[test]
fn test_save_and_undo() {
let mut manager = UndoRedoManager::new();
let items1 = vec![create_test_message("user", "Hello")];
let _id1 = manager
.save_state(items1.clone(), create_test_metadata(1))
.unwrap();
assert!(manager.current_state().is_some());
let items2 = vec![
create_test_message("user", "Hello"),
create_test_message("assistant", "Hi there"),
];
let _id2 = manager
.save_state(items2.clone(), create_test_metadata(2))
.unwrap();
let undone = manager.undo().unwrap();
assert!(undone.is_some());
assert_eq!(undone.unwrap().items.len(), 1);
}
#[test]
fn test_undo_redo() {
let mut manager = UndoRedoManager::new();
let items1 = vec![create_test_message("user", "1")];
manager.save_state(items1, create_test_metadata(1)).unwrap();
let items2 = vec![create_test_message("user", "2")];
manager.save_state(items2, create_test_metadata(2)).unwrap();
let items3 = vec![create_test_message("user", "3")];
manager.save_state(items3, create_test_metadata(3)).unwrap();
manager.undo().unwrap();
let state = manager.undo().unwrap().unwrap();
assert_eq!(state.metadata.turn_number, 1);
let state = manager.redo().unwrap().unwrap();
assert_eq!(state.metadata.turn_number, 2);
}
#[test]
fn test_branching() {
let mut manager = UndoRedoManager::new();
let items1 = vec![create_test_message("user", "main")];
manager.save_state(items1, create_test_metadata(1)).unwrap();
let branch_items = vec![create_test_message("user", "branch")];
let branch_id = manager
.create_branch(
"Alternative".to_string(),
Some("Testing branch".to_string()),
branch_items,
create_test_metadata(2),
)
.unwrap();
let branches = manager.get_branches();
assert_eq!(branches.len(), 1);
assert_eq!(branches[0].1.name, "Alternative");
let switched = manager.switch_to_branch(branch_id).unwrap();
assert!(switched.is_some());
}
#[test]
fn test_memory_limit() {
let mut manager = UndoRedoManager::with_memory_limit(1);
for i in 0..100 {
let items = vec![create_test_message("user", &"x".repeat(20000))]; manager.save_state(items, create_test_metadata(i)).unwrap();
}
let info = manager.memory_info();
assert!(info.total_usage_bytes <= info.max_usage_bytes);
assert!(manager.undo_stack.len() < 100);
}
#[test]
fn test_checkpoint() {
let mut manager = UndoRedoManager::new();
let items = vec![create_test_message("user", "checkpoint test")];
manager.save_state(items, create_test_metadata(1)).unwrap();
let checkpoint_id = manager
.create_checkpoint("test_checkpoint".to_string())
.unwrap();
let items2 = vec![create_test_message("user", "after checkpoint")];
manager.save_state(items2, create_test_metadata(2)).unwrap();
let restored = manager.restore_checkpoint(checkpoint_id).unwrap();
assert!(restored.is_some());
assert!(
restored
.unwrap()
.metadata
.tags
.contains(&"checkpoint:test_checkpoint".to_string())
);
}
#[test]
fn test_conversation_diff() {
let old = vec![
create_test_message("user", "Hello"),
create_test_message("assistant", "Hi"),
];
let new = vec![
create_test_message("user", "Hello"),
create_test_message("assistant", "Hi there!"),
create_test_message("user", "How are you?"),
];
let diff = ConversationDiff::create(&old, &new);
assert_eq!(diff.modified.len(), 1);
assert_eq!(diff.added.len(), 1);
assert_eq!(diff.removed.len(), 0);
let mut result = old.clone();
diff.apply(&mut result);
assert_eq!(result.len(), new.len());
}
}