use axum::{
extract::{Path, Query, State},
response::Json,
};
use serde::{Deserialize, Serialize};
use super::state::MultiUserMemoryManager;
use crate::errors::{AppError, ValidationErrorExt};
use crate::memory::{Session, SessionId, SessionStatus, SessionStoreStats, SessionSummary};
use crate::validation;
use std::sync::Arc;
type AppState = Arc<MultiUserMemoryManager>;
fn default_sessions_limit() -> usize {
10
}
fn default_end_reason() -> String {
"user_ended".to_string()
}
#[derive(Debug, Deserialize)]
pub struct ListSessionsRequest {
pub user_id: String,
#[serde(default = "default_sessions_limit")]
pub limit: usize,
}
#[derive(Debug, Serialize)]
pub struct ListSessionsResponse {
pub success: bool,
pub sessions: Vec<SessionSummary>,
pub count: usize,
}
#[derive(Debug, Deserialize)]
pub struct GetSessionRequest {
pub user_id: String,
}
#[derive(Debug, Serialize)]
pub struct GetSessionResponse {
pub success: bool,
pub session: Option<Session>,
}
#[derive(Debug, Deserialize)]
pub struct EndSessionRequest {
pub user_id: String,
#[serde(default = "default_end_reason")]
pub reason: String,
}
#[derive(Debug, Serialize)]
pub struct EndSessionResponse {
pub success: bool,
pub session: Option<Session>,
}
#[derive(Debug, Serialize)]
pub struct SessionStoreStatsResponse {
pub success: bool,
pub stats: SessionStoreStats,
}
pub async fn list_sessions(
State(state): State<AppState>,
Json(req): Json<ListSessionsRequest>,
) -> Result<Json<ListSessionsResponse>, AppError> {
validation::validate_user_id(&req.user_id).map_validation_err("user_id")?;
let sessions = state
.session_store
.get_user_sessions(&req.user_id, req.limit);
let count = sessions.len();
Ok(Json(ListSessionsResponse {
success: true,
sessions,
count,
}))
}
pub async fn get_session(
State(state): State<AppState>,
Path(session_id): Path<String>,
Query(req): Query<GetSessionRequest>,
) -> Result<Json<GetSessionResponse>, AppError> {
validation::validate_user_id(&req.user_id).map_validation_err("user_id")?;
let uuid = uuid::Uuid::parse_str(&session_id).map_err(|e| AppError::InvalidInput {
field: "session_id".to_string(),
reason: format!("Invalid UUID: {e}"),
})?;
let sid = SessionId(uuid);
let session = state.session_store.get_session(&sid);
Ok(Json(GetSessionResponse {
success: session.is_some(),
session,
}))
}
pub async fn end_session(
State(state): State<AppState>,
Json(req): Json<EndSessionRequest>,
) -> Result<Json<EndSessionResponse>, AppError> {
validation::validate_user_id(&req.user_id).map_validation_err("user_id")?;
let sessions = state.session_store.get_user_sessions(&req.user_id, 1);
let active_session = sessions
.into_iter()
.find(|s| matches!(s.status, SessionStatus::Active));
if let Some(summary) = active_session {
let session = state.session_store.end_session(&summary.id, &req.reason);
Ok(Json(EndSessionResponse {
success: session.is_some(),
session,
}))
} else {
Ok(Json(EndSessionResponse {
success: false,
session: None,
}))
}
}
pub async fn get_session_stats(
State(state): State<AppState>,
) -> Result<Json<SessionStoreStatsResponse>, AppError> {
let stats = state.session_store.stats();
Ok(Json(SessionStoreStatsResponse {
success: true,
stats,
}))
}