use std::collections::HashMap;
use std::time::{Duration, Instant};
use tokio::sync::{Mutex, broadcast};
use tracing::{info, warn};
use uuid::Uuid;
use crate::protocol::McpProtocolVersion;
pub struct Session {
pub sender: broadcast::Sender<String>,
pub created: Instant,
pub version: McpProtocolVersion,
}
impl Session {
pub fn touch(&mut self) {
self.created = Instant::now();
}
}
pub struct SessionHandle {
pub session_id: String,
pub receiver: broadcast::Receiver<String>,
}
pub type SessionMap = Mutex<HashMap<String, Session>>;
lazy_static::lazy_static! {
static ref SESSIONS: SessionMap = Mutex::new(HashMap::new());
}
pub async fn new_session(mcp_version: McpProtocolVersion) -> SessionHandle {
let session_id = Uuid::now_v7().as_simple().to_string();
let (sender, receiver) = broadcast::channel(128);
let session = Session {
sender: sender.clone(),
created: Instant::now(),
version: mcp_version,
};
SESSIONS.lock().await.insert(session_id.clone(), session);
SessionHandle {
session_id,
receiver,
}
}
pub async fn get_sender(session_id: &str) -> Option<broadcast::Sender<String>> {
let mut sessions = SESSIONS.lock().await;
if let Some(session) = sessions.get_mut(session_id) {
session.touch();
return Some(session.sender.clone());
}
None
}
pub async fn get_receiver(session_id: &str) -> Option<broadcast::Receiver<String>> {
let mut sessions = SESSIONS.lock().await;
if let Some(session) = sessions.get_mut(session_id) {
session.touch();
return Some(session.sender.subscribe());
}
None
}
pub async fn session_exists(session_id: &str) -> bool {
let mut sessions = SESSIONS.lock().await;
if let Some(session) = sessions.get_mut(session_id) {
session.touch();
true
} else {
false
}
}
pub async fn remove_session(session_id: &str) -> bool {
SESSIONS.lock().await.remove(session_id).is_some()
}
pub async fn expire_old(max_age: Duration) {
let cutoff = Instant::now() - max_age;
let mut sessions = SESSIONS.lock().await;
sessions.retain(|sid, session| {
let alive = session.created >= cutoff;
if !alive {
info!("Session {} expired", sid);
}
alive
});
}
pub async fn send_to_session(session_id: &str, message: String) -> bool {
if let Some(sender) = get_sender(session_id).await {
sender.send(message).is_ok()
} else {
false
}
}
pub async fn broadcast_to_all(message: String) {
let sessions = SESSIONS.lock().await;
for (sid, session) in sessions.iter() {
warn!("Sending message: {} to session {}", message, sid);
let _ = session.sender.send(message.clone());
}
}
pub async fn disconnect_all() {
let mut sessions = SESSIONS.lock().await;
sessions.clear();
info!("All sessions have been disconnected");
}
pub async fn session_count() -> usize {
let sessions = SESSIONS.lock().await;
sessions.len()
}
pub fn spawn_session_cleanup() {
tokio::spawn(async {
let mut interval = tokio::time::interval(Duration::from_secs(60));
loop {
interval.tick().await;
expire_old(Duration::from_secs(30 * 60)).await; }
});
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_session_lifecycle() {
let handle = new_session(McpProtocolVersion::V2025_06_18).await;
let session_id = handle.session_id.clone();
assert!(session_exists(&session_id).await);
assert!(get_sender(&session_id).await.is_some());
assert!(get_receiver(&session_id).await.is_some());
assert!(remove_session(&session_id).await);
assert!(!session_exists(&session_id).await);
}
#[tokio::test]
async fn test_session_messaging() {
let handle = new_session(McpProtocolVersion::V2025_06_18).await;
let session_id = handle.session_id.clone();
let message = r#"{"method":"test","params":{}}"#.to_string();
assert!(send_to_session(&session_id, message.clone()).await);
let mut receiver = handle.receiver;
let received = tokio::time::timeout(Duration::from_millis(100), receiver.recv()).await;
assert!(received.is_ok());
assert_eq!(received.unwrap().unwrap(), message);
}
}