use axum::{extract::State, response::Json};
use super::state::MultiUserMemoryManager;
use super::types::{
BackupResponse, CleanupCorruptedRequest, CleanupCorruptedResponse, ConsolidateRequest,
ConsolidateResponse, CreateBackupRequest, ListBackupsRequest, ListBackupsResponse, MemoryEvent,
MigrateLegacyRequest, MigrateLegacyResponse, PurgeBackupsRequest, PurgeBackupsResponse,
RebuildIndexRequest, RebuildIndexResponse, RepairIndexRequest, RepairIndexResponse,
RestoreBackupRequest, RestoreBackupResponse, VerifyBackupRequest, VerifyBackupResponse,
VerifyIndexRequest,
};
use crate::errors::{AppError, ValidationErrorExt};
use crate::memory;
use crate::metrics;
use crate::validation;
pub type AppState = std::sync::Arc<MultiUserMemoryManager>;
#[tracing::instrument(skip(state), fields(user_id = %req.user_id))]
pub async fn consolidate_memories(
State(state): State<AppState>,
Json(req): Json<ConsolidateRequest>,
) -> Result<Json<ConsolidateResponse>, AppError> {
validation::validate_user_id(&req.user_id).map_validation_err("user_id")?;
let _ = state
.get_user_memory(&req.user_id)
.map_err(AppError::Internal)?;
let user_id = req.user_id.clone();
let min_support = req.min_support;
let min_age_days = req.min_age_days;
let state_clone = state.clone();
tokio::task::spawn(async move {
let op_start = std::time::Instant::now();
let memory = match state_clone.get_user_memory(&user_id) {
Ok(m) => m,
Err(e) => {
tracing::error!(user_id = %user_id, "Consolidation: failed to get memory: {e}");
return;
}
};
let result = {
let memory = memory.clone();
let uid = user_id.clone();
match tokio::task::spawn_blocking(move || {
let memory_guard = memory.read();
memory_guard.distill_facts(&uid, min_support, min_age_days)
})
.await
{
Ok(Ok(r)) => r,
Ok(Err(e)) => {
tracing::error!(user_id = %user_id, "Consolidation fact extraction failed: {e}");
return;
}
Err(e) => {
tracing::error!(user_id = %user_id, "Consolidation fact extraction panicked: {e}");
return;
}
}
};
let decay_factor = state_clone.server_config().activation_decay_factor;
let maintenance_result = {
let memory = memory.clone();
let uid = user_id.clone();
match tokio::task::spawn_blocking(move || {
let memory_guard = memory.read();
memory_guard.run_maintenance(decay_factor, &uid, true)
})
.await
{
Ok(Ok(r)) => r,
Ok(Err(e)) => {
tracing::error!(user_id = %user_id, "Consolidation maintenance failed: {e}");
return;
}
Err(e) => {
tracing::error!(user_id = %user_id, "Consolidation maintenance panicked: {e}");
return;
}
}
};
let mut edges_strengthened: usize = 0;
let mut entity_edges_strengthened: usize = 0;
if !maintenance_result.edge_boosts.is_empty() {
if let Ok(graph) = state_clone.get_user_graph(&user_id) {
let graph_guard = graph.read();
match graph_guard.strengthen_memory_edges(&maintenance_result.edge_boosts) {
Ok((count, promotion_boosts)) => {
edges_strengthened += count;
if !promotion_boosts.is_empty() {
let memory_guard = memory.read();
let _ = memory_guard.apply_edge_promotion_boosts(&promotion_boosts);
}
}
Err(e) => {
tracing::debug!("On-demand edge boost failed: {e}");
}
}
}
}
if !maintenance_result.replay_memory_ids.is_empty() {
if let Ok(graph) = state_clone.get_user_graph(&user_id) {
let graph_guard = graph.read();
for mem_id_str in &maintenance_result.replay_memory_ids {
if let Ok(uuid) = uuid::Uuid::parse_str(mem_id_str) {
match graph_guard.strengthen_episode_entity_edges(&uuid) {
Ok(count) => entity_edges_strengthened += count,
Err(e) => {
tracing::debug!(
"Entity edge strengthening failed for {mem_id_str}: {e}"
);
}
}
}
}
}
}
if let Ok(graph) = state_clone.get_user_graph(&user_id) {
let graph_guard = graph.read();
let _ = graph_guard.flush_pending_maintenance();
}
let duration = op_start.elapsed().as_secs_f64();
metrics::CONSOLIDATE_DURATION.observe(duration);
metrics::CONSOLIDATE_TOTAL
.with_label_values(&["success"])
.inc();
tracing::info!(
user_id = %user_id,
memories_processed = result.memories_processed,
facts_extracted = result.facts_extracted,
facts_reinforced = result.facts_reinforced,
memories_replayed = maintenance_result.replay_memory_ids.len(),
edges_strengthened,
entity_edges_strengthened,
memories_decayed = maintenance_result.decayed_count,
duration_secs = format!("{:.1}", duration),
"Consolidation complete (background)"
);
});
Ok(Json(ConsolidateResponse {
memories_analyzed: 0,
facts_extracted: 0,
facts_reinforced: 0,
fact_ids: vec![],
memories_replayed: 0,
edges_strengthened: 0,
entity_edges_strengthened: 0,
memories_decayed: 0,
warnings: vec![
"Consolidation started in background. Check /api/consolidation/report for results."
.to_string(),
],
}))
}
pub async fn verify_index_integrity(
State(state): State<AppState>,
Json(req): Json<VerifyIndexRequest>,
) -> Result<Json<memory::IndexIntegrityReport>, AppError> {
validation::validate_user_id(&req.user_id).map_validation_err("user_id")?;
let memory_sys = state
.get_user_memory(&req.user_id)
.map_err(AppError::Internal)?;
let memory_guard = memory_sys.read();
let report = memory_guard
.verify_index_integrity()
.map_err(AppError::Internal)?;
Ok(Json(report))
}
pub async fn repair_vector_index(
State(state): State<AppState>,
Json(req): Json<RepairIndexRequest>,
) -> Result<Json<RepairIndexResponse>, AppError> {
validation::validate_user_id(&req.user_id).map_validation_err("user_id")?;
let memory_sys = state
.get_user_memory(&req.user_id)
.map_err(AppError::Internal)?;
let memory_guard = memory_sys.read();
let (total_storage, total_indexed, repaired, failed) = memory_guard
.repair_vector_index()
.map_err(AppError::Internal)?;
Ok(Json(RepairIndexResponse {
success: failed == 0,
total_storage,
total_indexed,
repaired,
failed,
is_healthy: total_storage == total_indexed,
}))
}
pub async fn cleanup_corrupted(
State(state): State<AppState>,
Json(req): Json<CleanupCorruptedRequest>,
) -> Result<Json<CleanupCorruptedResponse>, AppError> {
validation::validate_user_id(&req.user_id).map_validation_err("user_id")?;
let memory_sys = state
.get_user_memory(&req.user_id)
.map_err(AppError::Internal)?;
let memory_guard = memory_sys.read();
let deleted_count = memory_guard
.cleanup_corrupted()
.map_err(AppError::Internal)?;
if deleted_count > 0 {
state.emit_event(MemoryEvent {
event_type: "DELETE".to_string(),
timestamp: chrono::Utc::now(),
user_id: req.user_id.clone(),
memory_id: None,
content_preview: Some(format!("cleanup: {} corrupted entries", deleted_count)),
memory_type: None,
importance: None,
count: Some(deleted_count),
entities: None,
results: None,
});
}
Ok(Json(CleanupCorruptedResponse {
success: true,
deleted_count,
}))
}
pub async fn migrate_legacy(
State(state): State<AppState>,
Json(req): Json<MigrateLegacyRequest>,
) -> Result<Json<MigrateLegacyResponse>, AppError> {
validation::validate_user_id(&req.user_id).map_validation_err("user_id")?;
let memory_sys = state
.get_user_memory(&req.user_id)
.map_err(AppError::Internal)?;
let memory_guard = memory_sys.read();
let (migrated, already_current, failed) =
memory_guard.migrate_legacy().map_err(AppError::Internal)?;
if migrated > 0 {
state.emit_event(MemoryEvent {
event_type: "MIGRATE".to_string(),
timestamp: chrono::Utc::now(),
user_id: req.user_id.clone(),
memory_id: None,
content_preview: Some(format!(
"migrated {} memories, {} already current, {} failed",
migrated, already_current, failed
)),
memory_type: None,
importance: None,
count: Some(migrated),
entities: None,
results: None,
});
}
Ok(Json(MigrateLegacyResponse {
success: true,
migrated_count: migrated,
already_current_count: already_current,
failed_count: failed,
}))
}
pub async fn rebuild_index(
State(state): State<AppState>,
Json(req): Json<RebuildIndexRequest>,
) -> Result<Json<RebuildIndexResponse>, AppError> {
validation::validate_user_id(&req.user_id).map_validation_err("user_id")?;
let memory_sys = state
.get_user_memory(&req.user_id)
.map_err(AppError::Internal)?;
let memory_guard = memory_sys.read();
let (storage_count, indexed_count) =
memory_guard.rebuild_index().map_err(AppError::Internal)?;
Ok(Json(RebuildIndexResponse {
success: true,
storage_count,
indexed_count,
is_healthy: storage_count == indexed_count,
}))
}
pub async fn create_backup(
State(state): State<AppState>,
Json(req): Json<CreateBackupRequest>,
) -> Result<Json<BackupResponse>, AppError> {
validation::validate_user_id(&req.user_id).map_validation_err("user_id")?;
let memory_sys = state
.get_user_memory(&req.user_id)
.map_err(AppError::Internal)?;
let memory_guard = memory_sys.read();
let db = memory_guard.get_db();
let secondary_refs = state.collect_secondary_store_refs();
let store_refs: Vec<crate::backup::SecondaryStoreRef<'_>> = secondary_refs
.iter()
.map(|(name, db)| crate::backup::SecondaryStoreRef { name, db })
.collect();
let graph_lock = state.get_user_graph(&req.user_id).ok();
let graph_guard = graph_lock.as_ref().map(|g| g.read());
let graph_db_ref = graph_guard.as_ref().map(|g| g.get_db());
let result = if store_refs.is_empty() && graph_db_ref.is_none() {
state.backup_engine().create_backup(&db, &req.user_id)
} else {
state
.backup_engine()
.create_comprehensive_backup_with_graph(&db, &req.user_id, &store_refs, graph_db_ref)
};
match result {
Ok(metadata) => {
let secondary_count = metadata.secondary_stores.len();
state.log_event(
&req.user_id,
"BACKUP_CREATED",
&metadata.backup_id.to_string(),
&format!(
"Backup created: {} bytes + {} secondary stores ({} bytes)",
metadata.size_bytes, secondary_count, metadata.secondary_size_bytes
),
);
Ok(Json(BackupResponse {
success: true,
backup: Some(metadata),
message: "Backup created successfully".to_string(),
}))
}
Err(e) => Ok(Json(BackupResponse {
success: false,
backup: None,
message: format!("Backup failed: {}", e),
})),
}
}
pub async fn list_backups(
State(state): State<AppState>,
Json(req): Json<ListBackupsRequest>,
) -> Result<Json<ListBackupsResponse>, AppError> {
validation::validate_user_id(&req.user_id).map_validation_err("user_id")?;
match state.backup_engine().list_backups(&req.user_id) {
Ok(backups) => {
let count = backups.len();
Ok(Json(ListBackupsResponse {
success: true,
backups,
count,
}))
}
Err(e) => Err(AppError::Internal(e)),
}
}
pub async fn verify_backup(
State(state): State<AppState>,
Json(req): Json<VerifyBackupRequest>,
) -> Result<Json<VerifyBackupResponse>, AppError> {
validation::validate_user_id(&req.user_id).map_validation_err("user_id")?;
match state
.backup_engine()
.verify_backup(&req.user_id, req.backup_id)
{
Ok(is_valid) => Ok(Json(VerifyBackupResponse {
success: true,
is_valid,
message: if is_valid {
"Backup integrity verified".to_string()
} else {
"Backup checksum mismatch - may be corrupted".to_string()
},
})),
Err(e) => Ok(Json(VerifyBackupResponse {
success: false,
is_valid: false,
message: format!("Verification failed: {}", e),
})),
}
}
pub async fn purge_backups(
State(state): State<AppState>,
Json(req): Json<PurgeBackupsRequest>,
) -> Result<Json<PurgeBackupsResponse>, AppError> {
validation::validate_user_id(&req.user_id).map_validation_err("user_id")?;
match state
.backup_engine()
.purge_old_backups(&req.user_id, req.keep_count)
{
Ok(purged_count) => {
if purged_count > 0 {
state.log_event(
&req.user_id,
"BACKUP_PURGE",
&format!("keep_{}", req.keep_count),
&format!("Purged {} old backups", purged_count),
);
}
Ok(Json(PurgeBackupsResponse {
success: true,
purged_count,
}))
}
Err(e) => Err(AppError::Internal(e)),
}
}
pub async fn restore_backup(
State(state): State<AppState>,
Json(req): Json<RestoreBackupRequest>,
) -> Result<Json<RestoreBackupResponse>, AppError> {
validation::validate_user_id(&req.user_id).map_validation_err("user_id")?;
let user_id = req.user_id.clone();
let memory_db_path = state.base_path().join(&user_id).join("storage");
let graph_path = state.base_path().join(&user_id).join("graph").join("graph");
state.evict_user(&user_id);
let secondary_restore_paths: Vec<(&str, &std::path::Path)> = vec![];
let restored_stores = state
.backup_engine()
.restore_comprehensive_backup(
&user_id,
req.backup_id,
&memory_db_path,
&secondary_restore_paths,
)
.map_err(AppError::Internal)?;
let resolved_backup_id = req.backup_id.unwrap_or_else(|| {
state
.backup_engine()
.list_backups(&user_id)
.ok()
.and_then(|b| b.last().map(|m| m.backup_id))
.unwrap_or(0)
});
let graph_checkpoint = state
.backup_engine()
.backup_path()
.join(&user_id)
.join(format!("secondary_{resolved_backup_id}"))
.join("graph");
let mut all_restored = restored_stores;
if graph_checkpoint.exists() {
if graph_path.exists() {
let _ = std::fs::remove_dir_all(&graph_path);
}
if let Err(e) = crate::backup::copy_dir_recursive_pub(&graph_checkpoint, &graph_path) {
tracing::warn!(error = %e, "Failed to restore graph DB from backup");
} else {
all_restored.push("graph".to_string());
tracing::info!("Graph DB restored from backup");
}
}
state.log_event(
&user_id,
"BACKUP_RESTORED",
&format!("backup_{}", req.backup_id.unwrap_or(0)),
&format!("Restored {} stores: {:?}", all_restored.len(), all_restored),
);
Ok(Json(RestoreBackupResponse {
success: true,
message: format!(
"Restore complete for user '{}'. Restored: {:?}. Server restart recommended to re-initialize all caches.",
user_id, all_restored
),
restored_stores: all_restored,
}))
}
use serde::Deserialize;
#[derive(Debug, Deserialize)]
pub struct ConsolidationReportRequest {
pub user_id: String,
#[serde(default)]
pub since: Option<String>,
#[serde(default)]
pub until: Option<String>,
}
#[derive(Debug, Deserialize)]
pub struct ConsolidationEventsRequest {
pub user_id: String,
#[serde(default)]
pub since: Option<String>,
}
#[tracing::instrument(skip(state), fields(user_id = %req.user_id))]
pub async fn get_consolidation_report(
State(state): State<AppState>,
Json(req): Json<ConsolidationReportRequest>,
) -> Result<Json<memory::ConsolidationReport>, AppError> {
validation::validate_user_id(&req.user_id).map_validation_err("user_id")?;
let memory = state
.get_user_memory(&req.user_id)
.map_err(AppError::Internal)?;
let now = chrono::Utc::now();
let since = if let Some(since_str) = &req.since {
chrono::DateTime::parse_from_rfc3339(since_str)
.map_err(|e| AppError::InvalidInput {
field: "since".to_string(),
reason: format!("Invalid timestamp: {}", e),
})?
.with_timezone(&chrono::Utc)
} else {
now - chrono::Duration::hours(1)
};
let until = if let Some(until_str) = &req.until {
Some(
chrono::DateTime::parse_from_rfc3339(until_str)
.map_err(|e| AppError::InvalidInput {
field: "until".to_string(),
reason: format!("Invalid timestamp: {}", e),
})?
.with_timezone(&chrono::Utc),
)
} else {
None
};
let report = {
let memory = memory.clone();
tokio::task::spawn_blocking(move || {
let memory_guard = memory.read();
memory_guard.get_consolidation_report(since, until)
})
.await
.map_err(|e| AppError::Internal(anyhow::anyhow!("Blocking task panicked: {e}")))?
};
Ok(Json(report))
}
#[tracing::instrument(skip(state), fields(user_id = %req.user_id))]
pub async fn get_consolidation_events(
State(state): State<AppState>,
axum::extract::Query(req): axum::extract::Query<ConsolidationEventsRequest>,
) -> Result<Json<Vec<memory::ConsolidationEvent>>, AppError> {
validation::validate_user_id(&req.user_id).map_validation_err("user_id")?;
let memory = state
.get_user_memory(&req.user_id)
.map_err(AppError::Internal)?;
let now = chrono::Utc::now();
let since = if let Some(since_str) = &req.since {
chrono::DateTime::parse_from_rfc3339(since_str)
.map_err(|e| AppError::InvalidInput {
field: "since".to_string(),
reason: format!("Invalid timestamp: {}", e),
})?
.with_timezone(&chrono::Utc)
} else {
now - chrono::Duration::hours(1)
};
let events = {
let memory = memory.clone();
tokio::task::spawn_blocking(move || {
let memory_guard = memory.read();
memory_guard.get_consolidation_events_since(since)
})
.await
.map_err(|e| AppError::Internal(anyhow::anyhow!("Blocking task panicked: {e}")))?
};
Ok(Json(events))
}