use axum::{extract::State, response::Json};
use serde::{Deserialize, Serialize};
use super::state::MultiUserMemoryManager;
use crate::errors::{AppError, ValidationErrorExt};
use crate::memory::{self, SemanticFact};
use crate::validation;
use std::sync::Arc;
type AppState = Arc<MultiUserMemoryManager>;
fn facts_default_limit() -> usize {
50
}
#[derive(Debug, Deserialize)]
pub struct FactsListRequest {
pub user_id: String,
#[serde(default = "facts_default_limit")]
pub limit: usize,
}
#[derive(Debug, Deserialize)]
pub struct FactsSearchRequest {
pub user_id: String,
pub query: String,
#[serde(default = "facts_default_limit")]
pub limit: usize,
}
#[derive(Debug, Deserialize)]
pub struct FactsByEntityRequest {
pub user_id: String,
pub entity: String,
#[serde(default = "facts_default_limit")]
pub limit: usize,
}
#[derive(Debug, Serialize)]
pub struct FactsResponse {
pub facts: Vec<SemanticFact>,
pub total: usize,
}
#[tracing::instrument(skip(state), fields(user_id = %req.user_id))]
pub async fn list_facts(
State(state): State<AppState>,
Json(req): Json<FactsListRequest>,
) -> Result<Json<FactsResponse>, 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 user_id = req.user_id.clone();
let limit = req.limit;
let facts = tokio::task::spawn_blocking(move || {
let memory_guard = memory.read();
memory_guard.get_facts(&user_id, limit)
})
.await
.map_err(|e| AppError::Internal(anyhow::anyhow!("Blocking task panicked: {e}")))?
.map_err(AppError::Internal)?;
let total = facts.len();
Ok(Json(FactsResponse { facts, total }))
}
#[tracing::instrument(skip(state), fields(user_id = %req.user_id, query = %req.query))]
pub async fn search_facts(
State(state): State<AppState>,
Json(req): Json<FactsSearchRequest>,
) -> Result<Json<FactsResponse>, 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 user_id = req.user_id.clone();
let query = req.query.clone();
let limit = req.limit;
let facts = tokio::task::spawn_blocking(move || {
let memory_guard = memory.read();
memory_guard.search_facts(&user_id, &query, limit)
})
.await
.map_err(|e| AppError::Internal(anyhow::anyhow!("Blocking task panicked: {e}")))?
.map_err(AppError::Internal)?;
let total = facts.len();
Ok(Json(FactsResponse { facts, total }))
}
#[tracing::instrument(skip(state), fields(user_id = %req.user_id, entity = %req.entity))]
pub async fn facts_by_entity(
State(state): State<AppState>,
Json(req): Json<FactsByEntityRequest>,
) -> Result<Json<FactsResponse>, 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 user_id = req.user_id.clone();
let entity = req.entity.clone();
let limit = req.limit;
let facts = tokio::task::spawn_blocking(move || {
let memory_guard = memory.read();
memory_guard.get_facts_by_entity(&user_id, &entity, limit)
})
.await
.map_err(|e| AppError::Internal(anyhow::anyhow!("Blocking task panicked: {e}")))?
.map_err(AppError::Internal)?;
let total = facts.len();
Ok(Json(FactsResponse { facts, total }))
}
#[tracing::instrument(skip(state), fields(user_id = %req.user_id))]
pub async fn get_facts_stats(
State(state): State<AppState>,
Json(req): Json<FactsListRequest>,
) -> Result<Json<memory::FactStats>, 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 user_id = req.user_id.clone();
let stats = tokio::task::spawn_blocking(move || {
let memory_guard = memory.read();
memory_guard.get_fact_stats(&user_id)
})
.await
.map_err(|e| AppError::Internal(anyhow::anyhow!("Blocking task panicked: {e}")))?
.map_err(AppError::Internal)?;
Ok(Json(stats))
}