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::memory::{
13    RuntimeMemoryCallback, RuntimeMemoryCallbackError, SessionMemoryDiagnosticCallback,
14    SessionMemoryOutputDiagnostic,
15};
16use crate::types::Message;
17
18pub use config::{SessionMemoryConfig, SessionMemoryExtractionCallback};
19pub use entry::SessionMemoryEntry;
20use entry::{entry_key, SESSION_MEMORY_CATEGORIES};
21use prompt::{build_extraction_prompt, should_skip_message};
22pub use state::SessionMemoryState;
23
24#[derive(Debug, Clone)]
25pub struct SessionMemory {
26    pub config: SessionMemoryConfig,
27    pub state: SessionMemoryState,
28    pub(super) workspace: Option<PathBuf>,
29    pub(super) storage_scope: Option<String>,
30}
31
32impl SessionMemory {
33    pub fn new(config: SessionMemoryConfig) -> Self {
34        Self::with_workspace(config, None, None)
35    }
36
37    pub fn with_workspace(
38        config: SessionMemoryConfig,
39        workspace: Option<PathBuf>,
40        storage_scope: Option<String>,
41    ) -> Self {
42        Self {
43            config,
44            state: SessionMemoryState::default(),
45            workspace,
46            storage_scope: storage_scope.map(|scope| scope.trim().to_string()),
47        }
48    }
49
50    pub fn should_extract(&self, current_tokens: u64, message_count: usize) -> bool {
51        if self.config.extraction_callback.is_none() {
52            return false;
53        }
54        self.should_extract_with_runtime_callback(current_tokens, message_count)
55    }
56
57    pub(crate) fn should_extract_with_runtime_callback(
58        &self,
59        current_tokens: u64,
60        message_count: usize,
61    ) -> bool {
62        if current_tokens == 0 || message_count == 0 {
63            return false;
64        }
65
66        if !self.state.initialized {
67            return current_tokens >= self.config.min_tokens_before_extraction
68                && message_count >= self.config.min_text_messages;
69        }
70
71        let growth_threshold = ((self.config.min_tokens_before_extraction as f64)
72            * self.config.growth_ratio)
73            .floor()
74            .max(1.0) as u64;
75        let growth = if current_tokens >= self.state.tokens_at_last_extraction {
76            current_tokens - self.state.tokens_at_last_extraction
77        } else {
78            current_tokens
79        };
80        growth >= growth_threshold
81    }
82
83    pub fn extract(
84        &mut self,
85        messages: &[Message],
86        current_cycle: i32,
87        current_tokens: u64,
88    ) -> usize {
89        let Some(callback) = self.config.extraction_callback.as_ref().cloned() else {
90            return 0;
91        };
92        if messages.is_empty() {
93            return 0;
94        }
95
96        let start_index = if self.state.last_extracted_message_index >= 0
97            && (self.state.last_extracted_message_index as usize) < messages.len()
98        {
99            self.state.last_extracted_message_index as usize + 1
100        } else {
101            0
102        };
103        let new_messages = messages
104            .iter()
105            .enumerate()
106            .filter_map(|(index, message)| {
107                (index >= start_index && !should_skip_message(message)).then_some(message)
108            })
109            .collect::<Vec<_>>();
110
111        if new_messages.is_empty() {
112            self.record_extraction(messages.len() as i32 - 1, current_tokens);
113            return 0;
114        }
115
116        let prompt = build_extraction_prompt(&new_messages);
117        let raw_result = catch_unwind(AssertUnwindSafe(|| {
118            callback(
119                &prompt,
120                self.config.extraction_backend.as_deref(),
121                self.config.extraction_model.as_deref(),
122            )
123        }));
124        let Some(raw_result) = raw_result.ok().flatten() else {
125            return 0;
126        };
127
128        let entries = self.parse_extraction_result(&raw_result, current_cycle);
129        let merged_count = self.merge_entries(entries);
130        self.prune_to_budget();
131        self.record_extraction(messages.len() as i32 - 1, current_tokens);
132        self.save();
133        merged_count
134    }
135
136    pub(crate) fn extract_with_runtime_callback(
137        &mut self,
138        messages: &[Message],
139        current_cycle: i32,
140        current_tokens: u64,
141        callback: &RuntimeMemoryCallback,
142        diagnostic_callback: Option<&SessionMemoryDiagnosticCallback>,
143    ) -> Result<usize, RuntimeMemoryCallbackError> {
144        if messages.is_empty() {
145            return Ok(0);
146        }
147        let new_messages = self.new_extraction_messages(messages);
148        if new_messages.is_empty() {
149            self.record_extraction(messages.len() as i32 - 1, current_tokens);
150            return Ok(0);
151        }
152
153        let prompt = build_extraction_prompt(&new_messages);
154        let Some(raw_result) = callback(
155            &prompt,
156            self.config.extraction_backend.as_deref(),
157            self.config.extraction_model.as_deref(),
158            current_cycle.max(1) as u32,
159        )?
160        else {
161            return Ok(0);
162        };
163        let entries = match self.parse_extraction_result_checked(&raw_result, current_cycle) {
164            Ok(entries) => entries,
165            Err(reason) => {
166                if let Some(diagnostic_callback) = diagnostic_callback {
167                    diagnostic_callback(&SessionMemoryOutputDiagnostic {
168                        cycle_index: current_cycle.max(1) as u32,
169                        backend: self.config.extraction_backend.clone(),
170                        model: self.config.extraction_model.clone(),
171                        reason,
172                    });
173                }
174                return Ok(0);
175            }
176        };
177        let merged_count = self.merge_entries(entries);
178        self.prune_to_budget();
179        self.record_extraction(messages.len() as i32 - 1, current_tokens);
180        self.save();
181        Ok(merged_count)
182    }
183
184    fn new_extraction_messages<'a>(&self, messages: &'a [Message]) -> Vec<&'a Message> {
185        let start_index = if self.state.last_extracted_message_index >= 0
186            && (self.state.last_extracted_message_index as usize) < messages.len()
187        {
188            self.state.last_extracted_message_index as usize + 1
189        } else {
190            0
191        };
192        messages
193            .iter()
194            .enumerate()
195            .filter_map(|(index, message)| {
196                (index >= start_index && !should_skip_message(message)).then_some(message)
197            })
198            .collect()
199    }
200
201    pub fn render_as_system_context(&self) -> String {
202        if self.state.entries.is_empty() {
203            return String::new();
204        }
205
206        let mut parts = vec!["<Session Memory>".to_string()];
207        for category in SESSION_MEMORY_CATEGORIES {
208            let entries = self
209                .state
210                .entries
211                .iter()
212                .filter(|entry| entry.category == *category)
213                .collect::<Vec<_>>();
214            if entries.is_empty() {
215                continue;
216            }
217            parts.push(format!("## {category}"));
218            for entry in entries {
219                parts.push(format!("- {}", entry.content));
220            }
221        }
222        parts.push("</Session Memory>".to_string());
223        parts.join("\n")
224    }
225
226    pub fn on_compaction(&mut self, current_tokens: Option<u64>) {
227        self.state.last_extracted_message_index = -1;
228        if let Some(current_tokens) = current_tokens {
229            self.state.tokens_at_last_extraction = current_tokens;
230            self.state.initialized = true;
231        }
232        self.save();
233    }
234
235    pub fn load(&mut self) {
236        let Some(path) = self.storage_path() else {
237            return;
238        };
239        let Ok(content) = std::fs::read_to_string(path) else {
240            return;
241        };
242        let Ok(state) = serde_json::from_str::<SessionMemoryState>(&content) else {
243            return;
244        };
245        self.state = state;
246    }
247
248    pub fn save(&self) {
249        let Some(path) = self.storage_path() else {
250            return;
251        };
252        let Some(parent) = path.parent() else {
253            return;
254        };
255        if std::fs::create_dir_all(parent).is_err() {
256            return;
257        }
258        let Ok(content) = serde_json::to_string_pretty(&self.state) else {
259            return;
260        };
261        let _ = std::fs::write(path, content);
262    }
263
264    pub fn merge_entries(&mut self, entries: Vec<SessionMemoryEntry>) -> usize {
265        let mut merged = 0;
266        for entry in entries {
267            let key = entry_key(&entry);
268            if let Some(existing) = self
269                .state
270                .entries
271                .iter_mut()
272                .find(|candidate| entry_key(candidate) == key)
273            {
274                existing.importance = existing.importance.max(entry.importance);
275                existing.source_cycle = existing.source_cycle.max(entry.source_cycle);
276                continue;
277            }
278            self.state.entries.push(entry);
279            merged += 1;
280        }
281        merged
282    }
283
284    pub fn prune_to_budget(&mut self) {
285        if self.config.max_tokens == 0 || self.state.entries.is_empty() {
286            return;
287        }
288        let mut current_tokens =
289            estimate_tokens(&self.render_as_system_context(), &self.config.token_model);
290        while current_tokens > self.config.max_tokens && !self.state.entries.is_empty() {
291            let Some((drop_index, _)) = self
292                .state
293                .entries
294                .iter()
295                .enumerate()
296                .min_by_key(|(index, entry)| (entry.importance, entry.source_cycle, *index as i32))
297            else {
298                break;
299            };
300            self.state.entries.remove(drop_index);
301            current_tokens =
302                estimate_tokens(&self.render_as_system_context(), &self.config.token_model);
303        }
304    }
305
306    fn record_extraction(&mut self, last_message_index: i32, current_tokens: u64) {
307        self.state.last_extracted_message_index = last_message_index;
308        self.state.tokens_at_last_extraction = current_tokens;
309        self.state.initialized = true;
310    }
311}