use std::collections::HashSet;
use std::path::Path;
use std::sync::Arc;
use anyhow::{Context, Result, bail};
use tokio::sync::Semaphore;
use super::{FactsInput, GroundedFactRecord, PersistentMemoryWriteReport, RuntimeAgentConfig, VTCodeConfig};
const FACTS_PER_SESSION: usize = 32;
#[derive(Debug, Clone, serde::Serialize)]
pub struct BatchMemoryReport {
pub sessions_scanned: usize,
pub sessions_with_facts: usize,
pub candidate_facts: usize,
pub active_sessions_skipped: usize,
pub write_report: Option<PersistentMemoryWriteReport>,
}
pub async fn run_batch_memory_extraction(
runtime_config: &RuntimeAgentConfig,
vt_cfg: Option<&VTCodeConfig>,
workspace_root: &Path,
max_sessions: usize,
concurrency: usize,
) -> Result<BatchMemoryReport> {
if max_sessions == 0 {
bail!("batch memory extraction requires max_sessions > 0");
}
let config = super::effective_generated_memory_config(vt_cfg);
if !config.enabled {
bail!(
"persistent memory generation is disabled; enable `[agent.persistent_memory] enabled` (and memories) first"
);
}
let summaries = vtcode_memory::recent_sessions(workspace_root, max_sessions);
let semaphore = Arc::new(Semaphore::new(concurrency.max(1)));
let mut handles = Vec::with_capacity(summaries.len());
for summary in summaries {
let permit = semaphore
.clone()
.acquire_owned()
.await
.context("batch memory extraction semaphore closed")?;
let workspace = workspace_root.to_path_buf();
let session_id = summary.session_id.clone();
let active = summary.status == "active";
handles.push(tokio::task::spawn_blocking(move || {
let _permit = permit;
match vtcode_memory::session_memory_facts(&workspace, &session_id, FACTS_PER_SESSION) {
Ok(facts) => (session_id, active, facts),
Err(error) => {
tracing::warn!(session_id = %session_id, error = %error, "batch memory: failed to read session facts; skipping");
(session_id, active, Vec::new())
}
}
}));
}
let mut seen: HashSet<String> = HashSet::new();
let mut candidates: Vec<GroundedFactRecord> = Vec::new();
let mut sessions_scanned = 0usize;
let mut sessions_with_facts = 0usize;
let mut active_sessions_skipped = 0usize;
for handle in handles {
let Ok((session_id, active, facts)) = handle.await else {
continue;
};
sessions_scanned += 1;
if active {
active_sessions_skipped += 1;
continue;
}
let mut contributed = false;
for fact in facts {
let normalized = super::normalize_whitespace(&fact.fact).to_ascii_lowercase();
if normalized.is_empty() || !seen.insert(normalized) {
continue;
}
candidates.push(GroundedFactRecord {
fact: fact.fact,
source: format!("session:{session_id}"),
});
contributed = true;
}
if contributed {
sessions_with_facts += 1;
}
}
let candidate_count = candidates.len();
let write_report = if candidate_count == 0 {
None
} else {
super::persist_memory_internal(
&config,
workspace_root,
Some(runtime_config),
vt_cfg,
FactsInput::Candidates(&candidates),
true,
false,
)
.await
.context("batch memory consolidation")?
};
Ok(BatchMemoryReport {
sessions_scanned,
sessions_with_facts,
candidate_facts: candidate_count,
active_sessions_skipped,
write_report,
})
}
pub fn batch_parameters(config: &crate::config::PersistentMemoryConfig) -> (usize, usize) {
let sessions = config.memories.batch_sessions.max(1);
let concurrency = config.memories.batch_concurrency.clamp(1, 16);
(sessions, concurrency)
}
pub fn batch_parameters_from_config(config: Option<&crate::config::PersistentMemoryConfig>) -> (usize, usize) {
match config {
Some(config) => batch_parameters(config),
None => batch_parameters(&crate::config::PersistentMemoryConfig::default()),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn batch_parameters_clamp_degenerate_config() {
let mut config = crate::config::PersistentMemoryConfig::default();
config.memories.batch_sessions = 0;
config.memories.batch_concurrency = 999;
let (sessions, concurrency) = batch_parameters(&config);
assert_eq!(sessions, 1);
assert_eq!(concurrency, 16);
}
#[test]
fn batch_parameters_use_defaults() {
let config = crate::config::PersistentMemoryConfig::default();
let (sessions, concurrency) = batch_parameters(&config);
assert_eq!(sessions, 50);
assert_eq!(concurrency, 8);
}
}