Skip to main content

vv_agent/memory/session/
mod.rs

1mod config;
2mod entry;
3mod parse;
4mod prompt;
5mod state;
6mod storage;
7
8use std::panic::{catch_unwind, AssertUnwindSafe};
9use std::path::PathBuf;
10
11use crate::memory::token_utils::estimate_tokens;
12use crate::types::Message;
13
14pub use config::{SessionMemoryConfig, SessionMemoryExtractionCallback};
15pub use entry::SessionMemoryEntry;
16use entry::{entry_key, SESSION_MEMORY_CATEGORIES};
17use prompt::{build_extraction_prompt, should_skip_message};
18pub use state::SessionMemoryState;
19
20#[derive(Debug, Clone)]
21pub struct SessionMemory {
22    pub config: SessionMemoryConfig,
23    pub state: SessionMemoryState,
24    pub(super) workspace: Option<PathBuf>,
25    pub(super) storage_scope: Option<String>,
26}
27
28impl SessionMemory {
29    pub fn new(config: SessionMemoryConfig) -> Self {
30        Self::with_workspace(config, None, None)
31    }
32
33    pub fn with_workspace(
34        config: SessionMemoryConfig,
35        workspace: Option<PathBuf>,
36        storage_scope: Option<String>,
37    ) -> Self {
38        Self {
39            config,
40            state: SessionMemoryState::default(),
41            workspace,
42            storage_scope: storage_scope.map(|scope| scope.trim().to_string()),
43        }
44    }
45
46    pub fn should_extract(&self, current_tokens: u64, message_count: usize) -> bool {
47        if self.config.extraction_callback.is_none() || current_tokens == 0 || message_count == 0 {
48            return false;
49        }
50
51        if !self.state.initialized {
52            return current_tokens >= self.config.min_tokens_before_extraction
53                && message_count >= self.config.min_text_messages;
54        }
55
56        let growth_threshold = ((self.config.min_tokens_before_extraction as f64)
57            * self.config.growth_ratio)
58            .floor()
59            .max(1.0) as u64;
60        let growth = if current_tokens >= self.state.tokens_at_last_extraction {
61            current_tokens - self.state.tokens_at_last_extraction
62        } else {
63            current_tokens
64        };
65        growth >= growth_threshold
66    }
67
68    pub fn extract(
69        &mut self,
70        messages: &[Message],
71        current_cycle: i32,
72        current_tokens: u64,
73    ) -> usize {
74        let Some(callback) = self.config.extraction_callback.as_ref().cloned() else {
75            return 0;
76        };
77        if messages.is_empty() {
78            return 0;
79        }
80
81        let start_index = if self.state.last_extracted_message_index >= 0
82            && (self.state.last_extracted_message_index as usize) < messages.len()
83        {
84            self.state.last_extracted_message_index as usize + 1
85        } else {
86            0
87        };
88        let new_messages = messages
89            .iter()
90            .enumerate()
91            .filter_map(|(index, message)| {
92                (index >= start_index && !should_skip_message(message)).then_some(message)
93            })
94            .collect::<Vec<_>>();
95
96        if new_messages.is_empty() {
97            self.record_extraction(messages.len() as i32 - 1, current_tokens);
98            return 0;
99        }
100
101        let prompt = build_extraction_prompt(&new_messages);
102        let raw_result = catch_unwind(AssertUnwindSafe(|| {
103            callback(
104                &prompt,
105                self.config.extraction_backend.as_deref(),
106                self.config.extraction_model.as_deref(),
107            )
108        }));
109        let Some(raw_result) = raw_result.ok().flatten() else {
110            return 0;
111        };
112
113        let entries = self.parse_extraction_result(&raw_result, current_cycle);
114        let merged_count = self.merge_entries(entries);
115        self.prune_to_budget();
116        self.record_extraction(messages.len() as i32 - 1, current_tokens);
117        self.save();
118        merged_count
119    }
120
121    pub fn render_as_system_context(&self) -> String {
122        if self.state.entries.is_empty() {
123            return String::new();
124        }
125
126        let mut parts = vec!["<Session Memory>".to_string()];
127        for category in SESSION_MEMORY_CATEGORIES {
128            let entries = self
129                .state
130                .entries
131                .iter()
132                .filter(|entry| entry.category == *category)
133                .collect::<Vec<_>>();
134            if entries.is_empty() {
135                continue;
136            }
137            parts.push(format!("## {category}"));
138            for entry in entries {
139                parts.push(format!("- {}", entry.content));
140            }
141        }
142        parts.push("</Session Memory>".to_string());
143        parts.join("\n")
144    }
145
146    pub fn on_compaction(&mut self, current_tokens: Option<u64>) {
147        self.state.last_extracted_message_index = -1;
148        if let Some(current_tokens) = current_tokens {
149            self.state.tokens_at_last_extraction = current_tokens;
150            self.state.initialized = true;
151        }
152        self.save();
153    }
154
155    pub fn load(&mut self) {
156        let Some(path) = self.storage_path() else {
157            return;
158        };
159        let Ok(content) = std::fs::read_to_string(path) else {
160            return;
161        };
162        let Ok(state) = serde_json::from_str::<SessionMemoryState>(&content) else {
163            return;
164        };
165        self.state = state;
166    }
167
168    pub fn save(&self) {
169        let Some(path) = self.storage_path() else {
170            return;
171        };
172        let Some(parent) = path.parent() else {
173            return;
174        };
175        if std::fs::create_dir_all(parent).is_err() {
176            return;
177        }
178        let Ok(content) = serde_json::to_string_pretty(&self.state) else {
179            return;
180        };
181        let _ = std::fs::write(path, content);
182    }
183
184    pub fn merge_entries(&mut self, entries: Vec<SessionMemoryEntry>) -> usize {
185        let mut merged = 0;
186        for entry in entries {
187            let key = entry_key(&entry);
188            if let Some(existing) = self
189                .state
190                .entries
191                .iter_mut()
192                .find(|candidate| entry_key(candidate) == key)
193            {
194                existing.importance = existing.importance.max(entry.importance);
195                existing.source_cycle = existing.source_cycle.max(entry.source_cycle);
196                continue;
197            }
198            self.state.entries.push(entry);
199            merged += 1;
200        }
201        merged
202    }
203
204    pub fn prune_to_budget(&mut self) {
205        if self.config.max_tokens == 0 || self.state.entries.is_empty() {
206            return;
207        }
208        let mut current_tokens =
209            estimate_tokens(&self.render_as_system_context(), &self.config.token_model);
210        while current_tokens > self.config.max_tokens && !self.state.entries.is_empty() {
211            let Some((drop_index, _)) = self
212                .state
213                .entries
214                .iter()
215                .enumerate()
216                .min_by_key(|(index, entry)| (entry.importance, entry.source_cycle, *index as i32))
217            else {
218                break;
219            };
220            self.state.entries.remove(drop_index);
221            current_tokens =
222                estimate_tokens(&self.render_as_system_context(), &self.config.token_model);
223        }
224    }
225
226    fn record_extraction(&mut self, last_message_index: i32, current_tokens: u64) {
227        self.state.last_extracted_message_index = last_message_index;
228        self.state.tokens_at_last_extraction = current_tokens;
229        self.state.initialized = true;
230    }
231}