use std::{collections::HashMap, sync::Arc}; use async_trait::async_trait; use crank_registry::PostgresPoolConfig; use sqlx::{ PgPool, Row, postgres::{PgConnectOptions, PgPoolOptions}, query, }; use thiserror::Error; use time::OffsetDateTime; use tokio::sync::{RwLock, mpsc}; use tracing::{info, warn}; use uuid::Uuid; const ACTIVE_SESSION_COUNT_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(2); #[derive(Clone, Debug, PartialEq, Eq)] pub struct SessionState { pub id: String, pub protocol_version: String, pub initialized: bool, pub supports_elicitation: bool, pub workspace_slug: String, pub agent_slug: String, pub created_at: OffsetDateTime, pub updated_at: OffsetDateTime, pub expires_at: Option, } #[derive(Debug, Error)] #[error("transport session store is unavailable: {details}")] pub struct SessionStoreError { pub details: String, } #[async_trait] pub trait TransportSessionStore: Send + Sync { async fn create( &self, protocol_version: &str, workspace_slug: &str, agent_slug: &str, supports_elicitation: bool, now: OffsetDateTime, expires_at: Option, ) -> Result; async fn get(&self, session_id: &str) -> Result, SessionStoreError>; async fn mark_initialized( &self, session_id: &str, now: OffsetDateTime, ) -> Result; async fn delete(&self, session_id: &str) -> Result; async fn cleanup_expired(&self, now: OffsetDateTime) -> Result; async fn active_count(&self, now: OffsetDateTime) -> Result; } pub type SharedSessionStore = Arc; #[derive(Clone)] pub struct ActiveSessionMetrics { refresh_tx: mpsc::Sender<()>, } impl ActiveSessionMetrics { pub fn start(sessions: SharedSessionStore) -> Self { let (refresh_tx, mut refresh_rx) = mpsc::channel(1); tokio::spawn(async move { while refresh_rx.recv().await.is_some() { match tokio::time::timeout( ACTIVE_SESSION_COUNT_TIMEOUT, sessions.active_count(OffsetDateTime::now_utc()), ) .await { Ok(Ok(count)) => { crank_metrics::set_mcp_active_sessions(count); crank_metrics::set_mcp_session_metrics_fresh(true); } Ok(Err(_)) | Err(_) => { crank_metrics::set_mcp_session_metrics_fresh(false); warn!( name: "mcp.active_session_metrics.refresh_failed", error_category = "session_store", "active session metrics refresh failed" ); } } } }); let metrics = Self { refresh_tx }; metrics.refresh(); metrics } pub fn refresh(&self) { let _ = self.refresh_tx.try_send(()); } } pub fn spawn_session_cleanup( sessions: SharedSessionStore, metrics: ActiveSessionMetrics, cleanup_interval: std::time::Duration, ) { tokio::spawn(async move { let mut interval = tokio::time::interval(cleanup_interval); loop { interval.tick().await; match sessions.cleanup_expired(OffsetDateTime::now_utc()).await { Ok(removed) if removed > 0 => { info!(name: "mcp.session_cleanup.completed", removed); } Ok(_) => {} Err(_) => { warn!(name: "mcp.session_cleanup.failed", error_category = "session_store"); } } metrics.refresh(); } }); } #[derive(Clone, Debug)] pub struct PostgresTransportSessionStore { pool: PgPool, } impl PostgresTransportSessionStore { pub async fn from_pool(pool: PgPool) -> Result { crank_registry::MigrationAuthority::require_current(&pool) .await .map_err(|error| SessionStoreError { details: error.to_string(), })?; Ok(Self { pool }) } pub async fn connect_with_options_and_pool_config( connect_options: PgConnectOptions, pool_config: PostgresPoolConfig, ) -> Result { let pool = PgPoolOptions::new() .max_connections(pool_config.max_connections) .min_connections(pool_config.min_connections) .acquire_timeout(std::time::Duration::from_millis( pool_config.acquire_timeout_ms, )) .idle_timeout(Some(std::time::Duration::from_millis( pool_config.idle_timeout_ms, ))) .max_lifetime(Some(std::time::Duration::from_millis( pool_config.max_lifetime_ms, ))) .connect_with(connect_options) .await .map_err(|error| SessionStoreError { details: error.to_string(), })?; Self::from_pool(pool).await } } #[derive(Clone, Default)] pub struct InMemorySessionStore { inner: Arc>>, } #[async_trait] impl TransportSessionStore for InMemorySessionStore { async fn create( &self, protocol_version: &str, workspace_slug: &str, agent_slug: &str, supports_elicitation: bool, now: OffsetDateTime, expires_at: Option, ) -> Result { let session_id = Uuid::now_v7().to_string(); let mut guard = self.inner.write().await; guard.insert( session_id.clone(), SessionState { id: session_id.clone(), protocol_version: protocol_version.to_owned(), initialized: false, supports_elicitation, workspace_slug: workspace_slug.to_owned(), agent_slug: agent_slug.to_owned(), created_at: now, updated_at: now, expires_at, }, ); Ok(session_id) } async fn get(&self, session_id: &str) -> Result, SessionStoreError> { { let guard = self.inner.read().await; if let Some(session) = guard.get(session_id) { if !is_expired(session, OffsetDateTime::now_utc()) { return Ok(Some(session.clone())); } } else { return Ok(None); } } let mut guard = self.inner.write().await; guard.remove(session_id); Ok(None) } async fn mark_initialized( &self, session_id: &str, now: OffsetDateTime, ) -> Result { let mut guard = self.inner.write().await; if let Some(session) = guard.get_mut(session_id) { session.initialized = true; session.updated_at = now; return Ok(true); } Ok(false) } async fn delete(&self, session_id: &str) -> Result { let mut guard = self.inner.write().await; Ok(guard.remove(session_id).is_some()) } async fn cleanup_expired(&self, now: OffsetDateTime) -> Result { let mut guard = self.inner.write().await; let before = guard.len(); guard.retain(|_, session| !is_expired(session, now)); Ok(u64::try_from(before.saturating_sub(guard.len())).unwrap_or(u64::MAX)) } async fn active_count(&self, now: OffsetDateTime) -> Result { let guard = self.inner.read().await; Ok(u64::try_from( guard .values() .filter(|session| !is_expired(session, now)) .count(), ) .unwrap_or(u64::MAX)) } } #[async_trait] impl TransportSessionStore for PostgresTransportSessionStore { async fn create( &self, protocol_version: &str, workspace_slug: &str, agent_slug: &str, supports_elicitation: bool, now: OffsetDateTime, expires_at: Option, ) -> Result { let session_id = Uuid::now_v7().to_string(); query( "insert into mcp_transport_sessions ( id, protocol_version, initialized, supports_elicitation, workspace_slug, agent_slug, created_at, updated_at, expires_at ) values ( $1, $2, false, $3, $4, $5, $6::timestamptz, $6::timestamptz, $7::timestamptz )", ) .bind(&session_id) .bind(protocol_version) .bind(supports_elicitation) .bind(workspace_slug) .bind(agent_slug) .bind(now) .bind(expires_at) .execute(&self.pool) .await .map_err(|error| SessionStoreError { details: error.to_string(), })?; Ok(session_id) } async fn get(&self, session_id: &str) -> Result, SessionStoreError> { let row = sqlx::query( "select id, protocol_version, initialized, supports_elicitation, workspace_slug, agent_slug, created_at, updated_at, expires_at from mcp_transport_sessions where id = $1", ) .bind(session_id) .fetch_optional(&self.pool) .await .map_err(|error| SessionStoreError { details: error.to_string(), })?; let Some(session) = row.map(|row| SessionState { id: row.get("id"), protocol_version: row.get("protocol_version"), initialized: row.get("initialized"), supports_elicitation: row.get("supports_elicitation"), workspace_slug: row.get("workspace_slug"), agent_slug: row.get("agent_slug"), created_at: row.get("created_at"), updated_at: row.get("updated_at"), expires_at: row.get("expires_at"), }) else { return Ok(None); }; if is_expired(&session, OffsetDateTime::now_utc()) { query("delete from mcp_transport_sessions where id = $1") .bind(session_id) .execute(&self.pool) .await .map_err(|error| SessionStoreError { details: error.to_string(), })?; return Ok(None); } Ok(Some(session)) } async fn mark_initialized( &self, session_id: &str, now: OffsetDateTime, ) -> Result { let result = query( "update mcp_transport_sessions set initialized = true, updated_at = $2::timestamptz where id = $1", ) .bind(session_id) .bind(now) .execute(&self.pool) .await .map_err(|error| SessionStoreError { details: error.to_string(), })?; Ok(result.rows_affected() > 0) } async fn delete(&self, session_id: &str) -> Result { let result = query("delete from mcp_transport_sessions where id = $1") .bind(session_id) .execute(&self.pool) .await .map_err(|error| SessionStoreError { details: error.to_string(), })?; Ok(result.rows_affected() > 0) } async fn cleanup_expired(&self, now: OffsetDateTime) -> Result { let result = query( "delete from mcp_transport_sessions where expires_at is not null and expires_at <= $1::timestamptz", ) .bind(now) .execute(&self.pool) .await .map_err(|error| SessionStoreError { details: error.to_string(), })?; Ok(result.rows_affected()) } async fn active_count(&self, now: OffsetDateTime) -> Result { let row = query( "select count(*)::bigint as active_count from mcp_transport_sessions where expires_at is null or expires_at > $1::timestamptz", ) .bind(now) .fetch_one(&self.pool) .await .map_err(|error| SessionStoreError { details: error.to_string(), })?; let count = row.get::("active_count"); Ok(u64::try_from(count).unwrap_or_default()) } } fn is_expired(session: &SessionState, now: OffsetDateTime) -> bool { session .expires_at .is_some_and(|expires_at| expires_at <= now) }