use super::*;
struct PendingRequest {
response_tx: oneshot::Sender<Result<serde_json::Value>>,
}
pub(super) enum SessionServiceSource {
Router {
router: McpRouter,
factory: ServiceFactory,
},
Boxed(std::sync::Mutex<McpBoxService>),
}
pub(super) struct Session {
pub(super) id: String,
pub(super) service_source: SessionServiceSource,
pub(super) notifications_tx: broadcast::Sender<String>,
created_at: Instant,
last_accessed: RwLock<Instant>,
pending_requests: Mutex<HashMap<RequestId, PendingRequest>>,
pub(super) request_id_allocator: Option<Arc<AtomicI64>>,
pub(super) protocol_version: RwLock<String>,
pub(super) client_info: RwLock<Option<Implementation>>,
pub(super) client_capabilities: RwLock<Option<ClientCapabilities>>,
event_counter: AtomicU64,
event_store: Arc<dyn crate::event_store::EventStore>,
pub(super) initialized_notification_received: std::sync::atomic::AtomicBool,
}
impl Session {
pub(super) fn new(
router: McpRouter,
sampling_enabled: bool,
service_factory: ServiceFactory,
event_store: Arc<dyn crate::event_store::EventStore>,
) -> Self {
let (notifications_tx, _) = broadcast::channel(100);
let (notif_sender, mut notif_receiver) = notification_channel(256);
let router = router.with_notification_sender(notif_sender);
let broadcast_tx = notifications_tx.clone();
tokio::spawn(async move {
while let Some(notification) = notif_receiver.recv().await {
if let Some(json) = crate::transport::stdio::serialize_notification(¬ification) {
let _ = broadcast_tx.send(json);
}
}
});
let request_id_allocator = if sampling_enabled {
Some(Arc::new(AtomicI64::new(1)))
} else {
None
};
let now = Instant::now();
Self {
id: uuid::Uuid::new_v4().to_string(),
service_source: SessionServiceSource::Router {
router,
factory: service_factory,
},
notifications_tx,
created_at: now,
last_accessed: RwLock::new(now),
pending_requests: Mutex::new(HashMap::new()),
request_id_allocator,
protocol_version: RwLock::new(LATEST_PROTOCOL_VERSION.to_string()),
client_info: RwLock::new(None),
client_capabilities: RwLock::new(None),
event_counter: AtomicU64::new(0),
event_store,
initialized_notification_received: std::sync::atomic::AtomicBool::new(false),
}
}
fn from_service(
service: McpBoxService,
event_store: Arc<dyn crate::event_store::EventStore>,
) -> Self {
let (notifications_tx, _) = broadcast::channel(100);
let now = Instant::now();
Self {
id: uuid::Uuid::new_v4().to_string(),
service_source: SessionServiceSource::Boxed(std::sync::Mutex::new(service)),
notifications_tx,
created_at: now,
last_accessed: RwLock::new(now),
pending_requests: Mutex::new(HashMap::new()),
request_id_allocator: None,
protocol_version: RwLock::new(LATEST_PROTOCOL_VERSION.to_string()),
client_info: RwLock::new(None),
client_capabilities: RwLock::new(None),
event_counter: AtomicU64::new(0),
event_store,
initialized_notification_received: std::sync::atomic::AtomicBool::new(false),
}
}
fn restored(
record: &crate::session_store::SessionRecord,
router: McpRouter,
sampling_enabled: bool,
service_factory: ServiceFactory,
event_store: Arc<dyn crate::event_store::EventStore>,
) -> Self {
router.session().mark_preinitialized();
let (notifications_tx, _) = broadcast::channel(100);
let (notif_sender, mut notif_receiver) = notification_channel(256);
let router = router.with_notification_sender(notif_sender);
let broadcast_tx = notifications_tx.clone();
tokio::spawn(async move {
while let Some(notification) = notif_receiver.recv().await {
if let Some(json) = crate::transport::stdio::serialize_notification(¬ification) {
let _ = broadcast_tx.send(json);
}
}
});
let request_id_allocator = if sampling_enabled {
Some(Arc::new(AtomicI64::new(1)))
} else {
None
};
let now = Instant::now();
Self {
id: record.id.clone(),
service_source: SessionServiceSource::Router {
router,
factory: service_factory,
},
notifications_tx,
created_at: now,
last_accessed: RwLock::new(now),
pending_requests: Mutex::new(HashMap::new()),
request_id_allocator,
protocol_version: RwLock::new(record.protocol_version.clone()),
client_info: RwLock::new(record.client_info.clone()),
client_capabilities: RwLock::new(record.client_capabilities.clone()),
event_counter: AtomicU64::new(0),
event_store,
initialized_notification_received: std::sync::atomic::AtomicBool::new(true),
}
}
fn from_service_restored(
service: McpBoxService,
record: &crate::session_store::SessionRecord,
event_store: Arc<dyn crate::event_store::EventStore>,
) -> Self {
let (notifications_tx, _) = broadcast::channel(100);
let now = Instant::now();
Self {
id: record.id.clone(),
service_source: SessionServiceSource::Boxed(std::sync::Mutex::new(service)),
notifications_tx,
created_at: now,
last_accessed: RwLock::new(now),
pending_requests: Mutex::new(HashMap::new()),
request_id_allocator: None,
protocol_version: RwLock::new(record.protocol_version.clone()),
client_info: RwLock::new(record.client_info.clone()),
client_capabilities: RwLock::new(record.client_capabilities.clone()),
event_counter: AtomicU64::new(0),
event_store,
initialized_notification_received: std::sync::atomic::AtomicBool::new(true),
}
}
pub(super) fn make_service(&self) -> McpBoxService {
match &self.service_source {
SessionServiceSource::Router { router, factory } => (factory)(router.clone()),
SessionServiceSource::Boxed(mutex) => mutex.lock().unwrap().clone(),
}
}
pub(super) fn handle_notification(&self, notification: McpNotification) {
match &self.service_source {
SessionServiceSource::Router { router, .. } => {
router.handle_notification(notification);
}
SessionServiceSource::Boxed(_) => {
tracing::debug!(
notification = ?notification,
"Notification received on service-based session (not forwarded)"
);
}
}
}
pub(super) fn next_event_id(&self) -> u64 {
self.event_counter.fetch_add(1, Ordering::SeqCst)
}
pub(super) async fn buffer_event(&self, id: u64, data: String) {
let record = crate::event_store::EventRecord::new(id, data);
if let Err(e) = self.event_store.append(&self.id, record).await {
tracing::warn!(session_id = %self.id, event_id = id, error = %e, "Failed to append event to event store");
}
}
pub(super) async fn get_events_after(
&self,
after_id: u64,
) -> Vec<crate::event_store::EventRecord> {
match self.event_store.replay_after(&self.id, after_id).await {
Ok(events) => events,
Err(e) => {
tracing::warn!(session_id = %self.id, error = %e, "Failed to replay events from event store");
Vec::new()
}
}
}
async fn touch(&self) {
*self.last_accessed.write().await = Instant::now();
}
async fn is_expired(&self, ttl: Duration) -> bool {
self.last_accessed.read().await.elapsed() > ttl
}
pub(super) async fn add_pending_request(
&self,
id: RequestId,
response_tx: oneshot::Sender<Result<serde_json::Value>>,
) {
let mut pending = self.pending_requests.lock().await;
pending.insert(id, PendingRequest { response_tx });
}
pub(super) async fn complete_pending_request(
&self,
id: &RequestId,
result: Result<serde_json::Value>,
) -> bool {
let pending = {
let mut pending_requests = self.pending_requests.lock().await;
pending_requests.remove(id)
};
match pending {
Some(pending) => {
let _ = pending.response_tx.send(result);
true
}
None => false,
}
}
pub(super) async fn fail_pending_requests(&self, ids: &[RequestId], message: &str) {
let removed = {
let mut pending = self.pending_requests.lock().await;
ids.iter()
.filter_map(|id| pending.remove(id))
.collect::<Vec<_>>()
};
for pending in removed {
let _ = pending
.response_tx
.send(Err(Error::Transport(message.to_string())));
}
}
}
pub const DEFAULT_SESSION_TTL: Duration = Duration::from_secs(30 * 60);
const DEFAULT_CLEANUP_INTERVAL: Duration = Duration::from_secs(60);
#[derive(Debug, Clone)]
pub struct SessionConfig {
pub ttl: Duration,
pub max_sessions: Option<usize>,
pub cleanup_interval: Duration,
pub strict_initialization: bool,
}
impl Default for SessionConfig {
fn default() -> Self {
Self {
ttl: DEFAULT_SESSION_TTL,
max_sessions: None,
cleanup_interval: DEFAULT_CLEANUP_INTERVAL,
strict_initialization: true,
}
}
}
impl SessionConfig {
pub fn with_ttl(ttl: Duration) -> Self {
Self {
ttl,
..Default::default()
}
}
pub fn max_sessions(mut self, max: usize) -> Self {
self.max_sessions = Some(max);
self
}
pub fn cleanup_interval(mut self, interval: Duration) -> Self {
self.cleanup_interval = interval;
self
}
pub fn strict_initialization(mut self, enabled: bool) -> Self {
self.strict_initialization = enabled;
self
}
}
pub(super) struct SessionRegistry {
pub(super) sessions: RwLock<HashMap<String, Arc<Session>>>,
config: SessionConfig,
sampling_enabled: bool,
persistent: Arc<dyn crate::session_store::SessionStore>,
events: Arc<dyn crate::event_store::EventStore>,
service_source: ServiceSource,
auto_reinit: bool,
}
impl SessionRegistry {
pub(super) fn new(
config: SessionConfig,
sampling_enabled: bool,
persistent: Arc<dyn crate::session_store::SessionStore>,
events: Arc<dyn crate::event_store::EventStore>,
service_source: ServiceSource,
auto_reinit: bool,
) -> Self {
Self {
sessions: RwLock::new(HashMap::new()),
config,
sampling_enabled,
persistent,
events,
service_source,
auto_reinit,
}
}
async fn record_for(&self, session: &Session) -> crate::session_store::SessionRecord {
let protocol_version = session.protocol_version.read().await.clone();
let last_accessed = session.last_accessed.read().await;
let mut record = crate::session_store::SessionRecord::new(
session.id.clone(),
protocol_version,
self.config.ttl,
);
record.client_info = session.client_info.read().await.clone();
record.client_capabilities = session.client_capabilities.read().await.clone();
let now = std::time::SystemTime::now();
let created_ago = session.created_at.elapsed();
let last_accessed_ago = last_accessed.elapsed();
record.created_at = now.checked_sub(created_ago).unwrap_or(now);
record.last_accessed = now.checked_sub(last_accessed_ago).unwrap_or(now);
record.expires_at = record.last_accessed + self.config.ttl;
record
}
async fn persist_new(&self, session: &Session) {
let record = self.record_for(session).await;
if let Err(e) = self.persistent.create(&mut record.clone()).await {
tracing::warn!(session_id = %session.id, error = %e, "Failed to persist session record");
}
}
pub(super) async fn save_record(&self, session: &Session) {
let record = self.record_for(session).await;
if let Err(e) = self.persistent.save(&record).await {
tracing::warn!(session_id = %session.id, error = %e, "Failed to save session record");
}
}
pub(super) async fn create(
&self,
router: McpRouter,
service_factory: ServiceFactory,
) -> Option<Arc<Session>> {
router.session().mark_handshake_started();
let session = {
let mut sessions = self.sessions.write().await;
if let Some(max) = self.config.max_sessions
&& sessions.len() >= max
{
tracing::warn!(
max_sessions = max,
current = sessions.len(),
"Session limit reached, rejecting new session"
);
return None;
}
let session = Arc::new(Session::new(
router,
self.sampling_enabled,
service_factory,
self.events.clone(),
));
sessions.insert(session.id.clone(), session.clone());
tracing::debug!(session_id = %session.id, sampling = self.sampling_enabled, "Created new session");
session
};
self.persist_new(&session).await;
Some(session)
}
pub(super) async fn create_from_service(&self, service: McpBoxService) -> Option<Arc<Session>> {
let session = {
let mut sessions = self.sessions.write().await;
if let Some(max) = self.config.max_sessions
&& sessions.len() >= max
{
tracing::warn!(
max_sessions = max,
current = sessions.len(),
"Session limit reached, rejecting new session"
);
return None;
}
let session = Arc::new(Session::from_service(service, self.events.clone()));
sessions.insert(session.id.clone(), session.clone());
tracing::debug!(session_id = %session.id, "Created new session from service");
session
};
self.persist_new(&session).await;
Some(session)
}
pub(super) async fn create_initialized(
&self,
router: McpRouter,
service_factory: ServiceFactory,
) -> Option<Arc<Session>> {
router.session().mark_preinitialized();
let session = {
let mut sessions = self.sessions.write().await;
if let Some(max) = self.config.max_sessions
&& sessions.len() >= max
{
return None;
}
let session = Arc::new(Session::new(
router,
self.sampling_enabled,
service_factory,
self.events.clone(),
));
session
.initialized_notification_received
.store(true, Ordering::Release);
sessions.insert(session.id.clone(), session.clone());
tracing::debug!(session_id = %session.id, "Created pre-initialized session (optional_sessions)");
session
};
self.persist_new(&session).await;
Some(session)
}
pub(super) async fn create_initialized_from_service(
&self,
service: McpBoxService,
) -> Option<Arc<Session>> {
let session = {
let mut sessions = self.sessions.write().await;
if let Some(max) = self.config.max_sessions
&& sessions.len() >= max
{
return None;
}
let session = Arc::new(Session::from_service(service, self.events.clone()));
session
.initialized_notification_received
.store(true, Ordering::Release);
sessions.insert(session.id.clone(), session.clone());
tracing::debug!(session_id = %session.id, "Created pre-initialized session from service (optional_sessions)");
session
};
self.persist_new(&session).await;
Some(session)
}
pub(super) async fn get(&self, id: &str) -> Option<Arc<Session>> {
{
let sessions = self.sessions.read().await;
if let Some(s) = sessions.get(id).cloned() {
s.touch().await;
return Some(s);
}
}
match self.persistent.load(id).await {
Ok(Some(record)) => {
tracing::info!(session_id = %id, "Restoring session from persistent store");
if let Some(session) = self.restore_from_record(record).await {
return Some(session);
}
}
Ok(None) => {}
Err(e) => {
tracing::warn!(session_id = %id, error = %e, "Failed to load session record");
}
}
if self.auto_reinit {
tracing::info!(session_id = %id, "Auto-reinitializing unknown session");
return self.auto_reinitialize(id).await;
}
None
}
async fn restore_from_record(
&self,
record: crate::session_store::SessionRecord,
) -> Option<Arc<Session>> {
let session = {
let mut sessions = self.sessions.write().await;
if let Some(max) = self.config.max_sessions
&& sessions.len() >= max
{
tracing::warn!(
max_sessions = max,
"Session limit reached, cannot restore session"
);
return None;
}
if let Some(existing) = sessions.get(&record.id).cloned() {
existing.touch().await;
return Some(existing);
}
let session: Arc<Session> = match &self.service_source {
ServiceSource::Router { router, factory } => Arc::new(Session::restored(
&record,
router.with_fresh_session(),
self.sampling_enabled,
factory.clone(),
self.events.clone(),
)),
ServiceSource::Service(svc) => {
let service = svc.lock().unwrap().clone();
Arc::new(Session::from_service_restored(
service,
&record,
self.events.clone(),
))
}
};
sessions.insert(record.id.clone(), session.clone());
tracing::debug!(session_id = %session.id, "Restored session into local registry");
session
};
if let Ok(events) = self.events.replay_after(&record.id, 0).await
&& let Some(max_id) = events.iter().map(|e| e.id).max()
{
session
.event_counter
.store(max_id + 1, std::sync::atomic::Ordering::SeqCst);
}
let mut refreshed = record;
refreshed.touch(self.config.ttl);
if let Err(e) = self.persistent.save(&refreshed).await {
tracing::warn!(session_id = %refreshed.id, error = %e, "Failed to refresh restored session record");
}
Some(session)
}
async fn auto_reinitialize(&self, id: &str) -> Option<Arc<Session>> {
let mut record = crate::session_store::SessionRecord::new(
id.to_string(),
LATEST_PROTOCOL_VERSION.to_string(),
self.config.ttl,
);
record.client_info = Some(crate::protocol::Implementation {
name: "auto-recovered".into(),
version: "unknown".into(),
title: None,
description: None,
icons: None,
website_url: None,
meta: None,
});
record.client_capabilities = Some(crate::protocol::ClientCapabilities::default());
if let Err(e) = self.persistent.create(&mut record).await {
tracing::warn!(session_id = %id, error = %e, "Failed to persist auto-reinitialized session");
}
self.restore_from_record(record).await
}
pub(super) async fn remove(&self, id: &str) -> bool {
let removed = {
let mut sessions = self.sessions.write().await;
sessions.remove(id).is_some()
};
if removed {
tracing::debug!(session_id = %id, "Removed session");
if let Err(e) = self.persistent.delete(id).await {
tracing::warn!(session_id = %id, error = %e, "Failed to delete session record");
}
if let Err(e) = self.events.purge_session(id).await {
tracing::warn!(session_id = %id, error = %e, "Failed to purge session events");
}
}
removed
}
pub(super) async fn broadcast_to_all(&self, json: &str) {
let sessions = self.sessions.read().await;
for session in sessions.values() {
let _ = session.notifications_tx.send(json.to_string());
}
}
pub(super) async fn cleanup_expired(&self) -> usize {
let expired = {
let mut sessions = self.sessions.write().await;
let ttl = self.config.ttl;
let mut expired = Vec::new();
for (id, session) in sessions.iter() {
if session.is_expired(ttl).await {
expired.push(id.clone());
}
}
for id in &expired {
sessions.remove(id);
tracing::debug!(session_id = %id, "Expired session removed");
}
if !expired.is_empty() {
tracing::info!(
expired_count = expired.len(),
remaining = sessions.len(),
"Session cleanup completed"
);
}
expired
};
for id in &expired {
if let Err(e) = self.persistent.delete(id).await {
tracing::warn!(session_id = %id, error = %e, "Failed to delete expired session record");
}
if let Err(e) = self.events.purge_session(id).await {
tracing::warn!(session_id = %id, error = %e, "Failed to purge expired session events");
}
}
expired.len()
}
}
#[derive(Debug, Clone)]
pub struct SessionInfo {
pub id: String,
pub created_at: Duration,
pub last_activity: Duration,
}
#[derive(Clone)]
pub struct SessionHandle {
pub(super) store: Arc<SessionRegistry>,
#[cfg(feature = "stateless")]
pub(super) modern_subscriptions: Arc<ModernSubscriptionRegistry>,
}
impl SessionHandle {
pub async fn session_count(&self) -> usize {
self.store.sessions.read().await.len()
}
pub async fn list_sessions(&self) -> Vec<SessionInfo> {
let sessions = self.store.sessions.read().await;
let mut infos = Vec::with_capacity(sessions.len());
for session in sessions.values() {
let last_accessed = session.last_accessed.read().await;
infos.push(SessionInfo {
id: session.id.clone(),
created_at: session.created_at.elapsed(),
last_activity: last_accessed.elapsed(),
});
}
infos
}
pub async fn terminate_session(&self, id: &str) -> bool {
self.store.remove(id).await
}
#[cfg(feature = "stateless")]
pub fn subscription_count(&self) -> usize {
self.modern_subscriptions.len()
}
#[cfg(feature = "stateless")]
pub fn close_subscriptions(&self) -> usize {
self.modern_subscriptions.close_all()
}
}