vv_agent/memory/session/
mod.rs1mod 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}