From 796f8a1d4b87144ec30457bc0badfe59e2097021 Mon Sep 17 00:00:00 2001 From: lucarlig Date: Wed, 3 Jun 2026 14:05:09 +0100 Subject: [PATCH 1/3] Add session cleanup primitives Signed-off-by: lucarlig --- .../src/gateway/mcp_gateway.rs | 28 +++++++++++++++++-- .../session_store/local_session_store.rs | 5 ++++ .../src/gateway/session_store/mod.rs | 5 ++++ .../session_store/redis_session_store.rs | 14 ++++++++++ 4 files changed, 50 insertions(+), 2 deletions(-) diff --git a/crates/contextforge-gateway-rs-lib/src/gateway/mcp_gateway.rs b/crates/contextforge-gateway-rs-lib/src/gateway/mcp_gateway.rs index 6be8859..4de533d 100644 --- a/crates/contextforge-gateway-rs-lib/src/gateway/mcp_gateway.rs +++ b/crates/contextforge-gateway-rs-lib/src/gateway/mcp_gateway.rs @@ -35,6 +35,30 @@ use crate::{ }, }; +pub type BackendTransports = Arc>>; + +pub fn new_backend_transports() -> BackendTransports { + Arc::new(Mutex::new(HashMap::new())) +} + +#[derive(Clone)] +pub struct BackendTransportCleanup { + transports: BackendTransports, +} + +impl BackendTransportCleanup { + pub fn new(transports: BackendTransports) -> Self { + Self { transports } + } + + pub async fn remove_backends(&self, session_id: &str, backend_names: Vec) { + let mut transports = self.transports.lock().await; + for backend_name in backend_names { + transports.remove(&BackendTransportKey::from((backend_name.as_str(), session_id))); + } + } +} + #[derive(Clone, TypedBuilder)] #[builder(field_defaults(setter(prefix = "with_")))] pub struct McpService @@ -43,8 +67,8 @@ where { #[builder(default = Arc::new(Mutex::new(HashSet::new())))] subscriptions: Arc>>, - #[builder(default = Arc::new(Mutex::new(HashMap::new())))] - transports: Arc>>, + #[builder(default = new_backend_transports())] + transports: BackendTransports, #[builder(default = Arc::new(Mutex::new(LoggingLevel::Debug)))] log_level: Arc>, http_client: reqwest::Client, diff --git a/crates/contextforge-gateway-rs-lib/src/gateway/session_store/local_session_store.rs b/crates/contextforge-gateway-rs-lib/src/gateway/session_store/local_session_store.rs index 56cb69a..d9c4d13 100644 --- a/crates/contextforge-gateway-rs-lib/src/gateway/session_store/local_session_store.rs +++ b/crates/contextforge-gateway-rs-lib/src/gateway/session_store/local_session_store.rs @@ -40,4 +40,9 @@ impl UserSessionStore for LocalUserSessionStore { self.cache.lock().await.insert(session_key.clone(), mapping.clone()); Ok(()) } + + async fn remove_session<'a>(&self, session_key: &'a UserSession) -> Result<(), SessionStoreError> { + self.cache.lock().await.remove(session_key); + Ok(()) + } } diff --git a/crates/contextforge-gateway-rs-lib/src/gateway/session_store/mod.rs b/crates/contextforge-gateway-rs-lib/src/gateway/session_store/mod.rs index 4b6b94f..448ac84 100644 --- a/crates/contextforge-gateway-rs-lib/src/gateway/session_store/mod.rs +++ b/crates/contextforge-gateway-rs-lib/src/gateway/session_store/mod.rs @@ -40,6 +40,10 @@ impl SessionMapping { pub fn get<'a>(&'a self, host: &'a str) -> Option<&'a SessionMap> { self.session_mapping.iter().find(|m| m.backend_name == host) } + + pub fn backend_names(&self) -> Vec { + self.session_mapping.iter().map(|m| m.backend_name.clone()).collect() + } } #[derive(Debug, Clone, Deserialize, Serialize, thiserror::Error)] @@ -77,4 +81,5 @@ pub trait UserSessionStore: Send + Sync { key: &'a UserSession, session_mapping: &'a SessionMapping, ) -> Result<(), SessionStoreError>; + async fn remove_session<'a>(&self, key: &'a UserSession) -> Result<(), SessionStoreError>; } diff --git a/crates/contextforge-gateway-rs-lib/src/gateway/session_store/redis_session_store.rs b/crates/contextforge-gateway-rs-lib/src/gateway/session_store/redis_session_store.rs index 71de4ab..61ef96b 100644 --- a/crates/contextforge-gateway-rs-lib/src/gateway/session_store/redis_session_store.rs +++ b/crates/contextforge-gateway-rs-lib/src/gateway/session_store/redis_session_store.rs @@ -94,4 +94,18 @@ impl UserSessionStore for RedisUserSessionStore { return Err(SessionStoreError::CantWriteData); } } + + async fn remove_session<'a>(&self, session_key: &'a UserSession) -> Result<(), SessionStoreError> { + let Ok(key) = rmp_serde::encode::to_vec::(session_key) else { + return Err(SessionStoreError::DataEncoding); + }; + + let mut connection = self.connection.clone(); + if redis::cmd("DEL").arg(key).query_async::<()>(&mut connection).await.is_ok() { + self.cache.lock().await.remove(session_key); + Ok(()) + } else { + Err(SessionStoreError::CantWriteData) + } + } } From 3e5387f2cfbe8fe3ebe0577c85fec232ab75d139 Mon Sep 17 00:00:00 2001 From: lucarlig Date: Wed, 3 Jun 2026 14:05:14 +0100 Subject: [PATCH 2/3] Clean up MCP sessions on DELETE Signed-off-by: lucarlig --- Cargo.lock | 1 - Cargo.toml | 1 - crates/contextforge-gateway-rs-lib/Cargo.toml | 1 - .../src/gateway/mcp_call_validator.rs | 18 +++- .../src/gateway/mcp_gateway.rs | 24 +++-- .../src/gateway/mod.rs | 3 +- .../src/gateway/session_store/mod.rs | 4 - .../src/layers/session_id.rs | 101 +++++++++++------- crates/contextforge-gateway-rs-lib/src/lib.rs | 14 ++- 9 files changed, 105 insertions(+), 62 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index b93d05f..074d372 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -590,7 +590,6 @@ dependencies = [ "tokio-rustls", "tower", "tower-http", - "tower-layer", "tracing", "typed-builder", "uuid", diff --git a/Cargo.toml b/Cargo.toml index 7e24082..87ca950 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -41,7 +41,6 @@ tracing-subscriber = { version = "0.3", features = ["env-filter"] } tokio = { version = "1.48.0", features = ["full"] } axum = "0.8" tower-http = { version = "0.6.8", features = ["full"] } -tower-layer = "0.3.3" tower = "0.5.3" http = "1.4.0" futures = { version = "0.3", features = ["std", "alloc"] } diff --git a/crates/contextforge-gateway-rs-lib/Cargo.toml b/crates/contextforge-gateway-rs-lib/Cargo.toml index 8f0861b..e044d7a 100644 --- a/crates/contextforge-gateway-rs-lib/Cargo.toml +++ b/crates/contextforge-gateway-rs-lib/Cargo.toml @@ -18,7 +18,6 @@ tracing.workspace = true tokio.workspace = true axum.workspace = true tower-http.workspace = true -tower-layer.workspace = true tower.workspace = true http.workspace = true futures.workspace = true diff --git a/crates/contextforge-gateway-rs-lib/src/gateway/mcp_call_validator.rs b/crates/contextforge-gateway-rs-lib/src/gateway/mcp_call_validator.rs index bf49a7f..87061ea 100644 --- a/crates/contextforge-gateway-rs-lib/src/gateway/mcp_call_validator.rs +++ b/crates/contextforge-gateway-rs-lib/src/gateway/mcp_call_validator.rs @@ -7,7 +7,10 @@ use rmcp::{ }; use tracing::info; -use crate::layers::{session_id::SessionId, virtual_host_id::VirtualHostId}; +use crate::{ + common::ContextForgeClaims, + layers::{session_id::SessionId, virtual_host_id::VirtualHostId}, +}; pub struct AuthorizedCallValidator<'a> { call_name: &'a str, @@ -73,12 +76,13 @@ impl<'a> InitializeCallValidator<'a> { pub fn new(ctx: &'a RequestContext) -> Self { Self { ctx } } - pub fn validate(self) -> Result<(&'a VirtualHost, &'a DownstreamSessionId), ErrorData> { + pub fn validate(self) -> Result<(&'a VirtualHost, &'a DownstreamSessionId, &'a ContextForgeClaims), ErrorData> { let maybe_parts = self.ctx.extensions.get::(); let maybe_downstream_session = self.ctx.extensions.get::(); let maybe_user_config = maybe_parts.and_then(|parts| parts.extensions.get::()); let maybe_virtual_host_id = maybe_parts.and_then(|parts| parts.extensions.get::()); + let maybe_claims = maybe_parts.and_then(|parts| parts.extensions.get::()); info!( "intialize user_config = {maybe_user_config:#?} downstream_session_id = {maybe_downstream_session:#?} virtual_host_id = {maybe_virtual_host_id:#?}" ); @@ -115,6 +119,14 @@ impl<'a> InitializeCallValidator<'a> { }); }; - Ok((virtual_host, downstream_session_id)) + let Some(claims) = maybe_claims else { + return Err(ErrorData { + code: ErrorCode::INTERNAL_ERROR, + message: "Routing problem... claims not found".into(), + data: None, + }); + }; + + Ok((virtual_host, downstream_session_id, claims)) } } diff --git a/crates/contextforge-gateway-rs-lib/src/gateway/mcp_gateway.rs b/crates/contextforge-gateway-rs-lib/src/gateway/mcp_gateway.rs index 4de533d..ca23cd5 100644 --- a/crates/contextforge-gateway-rs-lib/src/gateway/mcp_gateway.rs +++ b/crates/contextforge-gateway-rs-lib/src/gateway/mcp_gateway.rs @@ -51,11 +51,9 @@ impl BackendTransportCleanup { Self { transports } } - pub async fn remove_backends(&self, session_id: &str, backend_names: Vec) { + pub async fn remove_session(&self, session_id: &str) { let mut transports = self.transports.lock().await; - for backend_name in backend_names { - transports.remove(&BackendTransportKey::from((backend_name.as_str(), session_id))); - } + transports.retain(|key, _| key.session_id != session_id); } } @@ -132,10 +130,10 @@ where cx: RequestContext, ) -> Result { let call_validator = InitializeCallValidator::new(&cx); - let (virtual_host, downstream_session_id) = call_validator.validate()?; + let (virtual_host, downstream_session_id, claims) = call_validator.validate()?; let session_mapping = if let Ok(maybe_session_mapping) = self .user_session_store - .get_session(&UserSession::new(String::new(), Arc::clone(&downstream_session_id.session_id))) + .get_session(&UserSession::new(claims.sub.clone(), Arc::clone(&downstream_session_id.session_id))) .await { maybe_session_mapping.unwrap_or_default() @@ -208,13 +206,21 @@ where }) .unzip(); - let _ = self + if self .user_session_store .set_session( - &UserSession::new(String::new(), Arc::clone(&downstream_session_id.session_id)), + &UserSession::new(claims.sub.clone(), Arc::clone(&downstream_session_id.session_id)), &session_mapping, ) - .await; + .await + .is_err() + { + return Err(ErrorData { + code: ErrorCode::INTERNAL_ERROR, + message: "Internal problem... session store can't be written".into(), + data: None, + }); + } let mut transports = self.transports.lock().await; for (name, svc) in backend_services { diff --git a/crates/contextforge-gateway-rs-lib/src/gateway/mod.rs b/crates/contextforge-gateway-rs-lib/src/gateway/mod.rs index e94cdc9..74a09da 100644 --- a/crates/contextforge-gateway-rs-lib/src/gateway/mod.rs +++ b/crates/contextforge-gateway-rs-lib/src/gateway/mod.rs @@ -3,4 +3,5 @@ pub(crate) mod mcp_gateway; mod session_manager; mod session_store; -pub use mcp_gateway::{LocalUserSessionStore, McpService}; +pub use mcp_gateway::{BackendTransportCleanup, LocalUserSessionStore, McpService, new_backend_transports}; +pub use session_store::{UserSession, UserSessionStore}; diff --git a/crates/contextforge-gateway-rs-lib/src/gateway/session_store/mod.rs b/crates/contextforge-gateway-rs-lib/src/gateway/session_store/mod.rs index 448ac84..6b88be4 100644 --- a/crates/contextforge-gateway-rs-lib/src/gateway/session_store/mod.rs +++ b/crates/contextforge-gateway-rs-lib/src/gateway/session_store/mod.rs @@ -40,10 +40,6 @@ impl SessionMapping { pub fn get<'a>(&'a self, host: &'a str) -> Option<&'a SessionMap> { self.session_mapping.iter().find(|m| m.backend_name == host) } - - pub fn backend_names(&self) -> Vec { - self.session_mapping.iter().map(|m| m.backend_name.clone()).collect() - } } #[derive(Debug, Clone, Deserialize, Serialize, thiserror::Error)] diff --git a/crates/contextforge-gateway-rs-lib/src/layers/session_id.rs b/crates/contextforge-gateway-rs-lib/src/layers/session_id.rs index b5b6fb5..31c36e1 100644 --- a/crates/contextforge-gateway-rs-lib/src/layers/session_id.rs +++ b/crates/contextforge-gateway-rs-lib/src/layers/session_id.rs @@ -1,59 +1,82 @@ -use std::task::{Context, Poll}; +use std::sync::Arc; -use axum::http::Request; -use tower::Service; -use tower_layer::Layer; +use axum::{body::Body, extract::State, http::Request, middleware::Next, response::Response}; +use http::{Method, StatusCode, header}; use tracing::info; -use crate::const_values::MCP_SESSION_ID; - -#[derive(Debug, Clone)] -pub struct SessionIdLayer; - -impl Layer for SessionIdLayer { - type Service = SessionIdService; - - fn layer(&self, service: S) -> Self::Service { - SessionIdService { service } - } -} - -#[derive(Debug, Clone)] -pub struct SessionIdService { - service: S, -} +use crate::{ + common::ContextForgeClaims, + const_values::MCP_SESSION_ID, + gateway::{BackendTransportCleanup, UserSession, UserSessionStore}, + layers::virtual_host_id::VirtualHostId, +}; #[derive(Debug, Clone)] pub struct SessionId { value: String, } + impl SessionId { pub fn value(&self) -> &String { &self.value } } -impl Service> for SessionIdService -where - S: Service>, -{ - type Response = S::Response; - type Error = S::Error; - type Future = S::Future; +#[derive(Clone)] +pub struct SessionIdState { + pub user_session_store: Arc, + pub backend_transport_cleanup: BackendTransportCleanup, +} + +pub async fn session_id_layer(State(state): State, mut request: Request, next: Next) -> Response { + let session_id = + request.headers().get(MCP_SESSION_ID).and_then(|session_id| session_id.to_str().ok()).map(str::to_owned); - fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll> { - self.service.poll_ready(cx) + if request.method() != Method::DELETE { + if let Some(session_id) = session_id { + info!("MCP Session ID {session_id}"); + request.extensions_mut().insert(SessionId { value: session_id }); + } + return next.run(request).await; } - fn call(&mut self, mut request: Request) -> Self::Future { - let maybe_session = request.headers().get(MCP_SESSION_ID).cloned(); - if let Some(session_id_header_value) = maybe_session { - info!("MCP Session ID {:?}", session_id_header_value.to_str()); - if let Ok(session_id) = session_id_header_value.to_str() { - request.extensions_mut().insert(SessionId { value: session_id.to_owned() }); - } - } + let Some(session_id) = session_id else { + return response(StatusCode::BAD_REQUEST); + }; + let Some(claims) = request.extensions().get::() else { + return response(StatusCode::BAD_REQUEST); + }; + let Some(_virtual_host_id) = request.extensions().get::() else { + return response(StatusCode::BAD_REQUEST); + }; - self.service.call(request) + let user_session = UserSession::new(claims.sub.clone(), Arc::from(session_id.as_str())); + match state.user_session_store.get_session(&user_session).await { + Ok(Some(_)) => {}, + Ok(None) => return response(StatusCode::NOT_FOUND), + Err(_) => return response(StatusCode::INTERNAL_SERVER_ERROR), } + + request.extensions_mut().insert(SessionId { value: session_id.clone() }); + let rmcp_response = next.run(request).await; + if !rmcp_response.status().is_success() { + return rmcp_response; + } + + let remove_result = state.user_session_store.remove_session(&user_session).await; + state.backend_transport_cleanup.remove_session(&session_id).await; + + if remove_result.is_err() { + return response(StatusCode::INTERNAL_SERVER_ERROR); + } + + rmcp_response +} + +fn response(status: StatusCode) -> Response { + Response::builder() + .status(status) + .header(header::CONTENT_TYPE, "text/plain") + .body(Body::empty()) + .expect("response should build") } diff --git a/crates/contextforge-gateway-rs-lib/src/lib.rs b/crates/contextforge-gateway-rs-lib/src/lib.rs index 338bbf9..1a8598b 100644 --- a/crates/contextforge-gateway-rs-lib/src/lib.rs +++ b/crates/contextforge-gateway-rs-lib/src/lib.rs @@ -19,7 +19,7 @@ mod tools; mod user_config_store; pub use common::{RedisClient, RedisConfig, UpstreamConnectionMode}; -use gateway::McpService; +use gateway::{BackendTransportCleanup, McpService, new_backend_transports}; use layers::session_id::SessionId; use tower_http::cors::{Any, CorsLayer}; use transports::{DownstreamTls, Tcp}; @@ -36,7 +36,9 @@ use crate::{ common::{ContextForgeGatewayAppState, JwtTokenDecoders}, gateway::LocalUserSessionStore, layers::{ - claims_id::claims_layer, session_id::SessionIdLayer, user_config_store::user_config_store_layer, + claims_id::claims_layer, + session_id::{SessionIdState, session_id_layer}, + user_config_store::user_config_store_layer, virtual_host_id::virtual_host_id_layer, }, }; @@ -68,6 +70,11 @@ impl Gateway { let user_config_store = user_config_store as Arc; let user_session_store = LocalUserSessionStore::new(); + let backend_transports = new_backend_transports(); + let session_id_state = SessionIdState { + user_session_store: Arc::new(user_session_store.clone()), + backend_transport_cleanup: BackendTransportCleanup::new(Arc::clone(&backend_transports)), + }; let mcp_plugin_runtime = self.plugin_runtime; let streamable_config = StreamableHttpServerConfig::default().disable_allowed_hosts(); @@ -81,6 +88,7 @@ impl Gateway { Ok(McpService::builder() .with_user_session_store(user_session_store.clone()) .with_http_client(reqwest_backend_client.clone()) + .with_transports(Arc::clone(&backend_transports)) .with_plugin_runtime(mcp_plugin_runtime.clone()) .build()) }, @@ -119,8 +127,8 @@ impl Gateway { let app = axum::Router::new() .nest_service("/servers/{virtual_host_name}/mcp", mcp_service) .layer(middleware::from_fn_with_state(mcp_add_state.clone(), user_config_store_layer)) + .layer(middleware::from_fn_with_state(session_id_state, session_id_layer)) .layer(middleware::from_fn_with_state(mcp_add_state.clone(), claims_layer)) - .layer(SessionIdLayer) .layer(middleware::from_fn(virtual_host_id_layer)) .layer(cors_layer); From 7db1d3c77e2c1ed3818dc0edec71384f98b4bea1 Mon Sep 17 00:00:00 2001 From: lucarlig Date: Wed, 3 Jun 2026 16:47:27 +0100 Subject: [PATCH 3/3] Simplify session delete middleware flow Signed-off-by: lucarlig --- .../src/gateway/mcp_call_validator.rs | 13 +++- .../src/gateway/mcp_gateway.rs | 67 ++++++++++--------- .../src/gateway/mod.rs | 2 +- .../src/gateway/session_manager.rs | 25 ++++--- .../src/layers/session_id.rs | 63 ++++++++--------- crates/contextforge-gateway-rs-lib/src/lib.rs | 8 +-- 6 files changed, 89 insertions(+), 89 deletions(-) diff --git a/crates/contextforge-gateway-rs-lib/src/gateway/mcp_call_validator.rs b/crates/contextforge-gateway-rs-lib/src/gateway/mcp_call_validator.rs index 87061ea..ae9f0aa 100644 --- a/crates/contextforge-gateway-rs-lib/src/gateway/mcp_call_validator.rs +++ b/crates/contextforge-gateway-rs-lib/src/gateway/mcp_call_validator.rs @@ -21,10 +21,11 @@ impl<'a> AuthorizedCallValidator<'a> { pub fn new(call_name: &'a str, ctx: &'a RequestContext) -> Self { Self { call_name, ctx } } - pub fn validate(self) -> Result<(&'a VirtualHost, &'a SessionId), ErrorData> { + pub fn validate(self) -> Result<(&'a VirtualHost, &'a SessionId, &'a ContextForgeClaims), ErrorData> { let maybe_parts = self.ctx.extensions.get::(); let maybe_session_id = maybe_parts.and_then(|parts| parts.extensions.get::()); let maybe_user_config = maybe_parts.and_then(|parts| parts.extensions.get::()); + let maybe_claims = maybe_parts.and_then(|parts| parts.extensions.get::()); let maybe_virtual_host_id = maybe_parts.and_then(|parts| parts.extensions.get::()); info!( @@ -64,7 +65,15 @@ impl<'a> AuthorizedCallValidator<'a> { }); }; - Ok((virtual_host, session_id)) + let Some(claims) = maybe_claims else { + return Err(ErrorData { + code: ErrorCode::INTERNAL_ERROR, + message: "Routing problem... claims not found".into(), + data: None, + }); + }; + + Ok((virtual_host, session_id, claims)) } } diff --git a/crates/contextforge-gateway-rs-lib/src/gateway/mcp_gateway.rs b/crates/contextforge-gateway-rs-lib/src/gateway/mcp_gateway.rs index ca23cd5..a7fc094 100644 --- a/crates/contextforge-gateway-rs-lib/src/gateway/mcp_gateway.rs +++ b/crates/contextforge-gateway-rs-lib/src/gateway/mcp_gateway.rs @@ -35,25 +35,17 @@ use crate::{ }, }; -pub type BackendTransports = Arc>>; +#[derive(Clone, Default)] +pub struct BackendTransports(Arc>>); -pub fn new_backend_transports() -> BackendTransports { - Arc::new(Mutex::new(HashMap::new())) -} - -#[derive(Clone)] -pub struct BackendTransportCleanup { - transports: BackendTransports, -} - -impl BackendTransportCleanup { - pub fn new(transports: BackendTransports) -> Self { - Self { transports } +impl BackendTransports { + pub async fn remove_session(&self, principal: &str, session_id: &str) { + let mut transports = self.0.lock().await; + transports.retain(|key, _| key.principal != principal || key.session_id != session_id); } - pub async fn remove_session(&self, session_id: &str) { - let mut transports = self.transports.lock().await; - transports.retain(|key, _| key.session_id != session_id); + pub fn inner(&self) -> &Arc>> { + &self.0 } } @@ -65,7 +57,7 @@ where { #[builder(default = Arc::new(Mutex::new(HashSet::new())))] subscriptions: Arc>>, - #[builder(default = new_backend_transports())] + #[builder(default = BackendTransports::default())] transports: BackendTransports, #[builder(default = Arc::new(Mutex::new(LoggingLevel::Debug)))] log_level: Arc>, @@ -77,6 +69,7 @@ where #[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)] pub struct BackendTransportKey { + principal: String, backend_name: String, session_id: String, } @@ -102,15 +95,23 @@ pub struct BackendTransportService { pub(crate) service: Option, } -impl From<(&str, &str)> for BackendTransportKey { - fn from((backend_name, session_name): (&str, &str)) -> Self { - Self { backend_name: backend_name.to_owned(), session_id: session_name.to_owned() } +impl From<(&str, &str, &str)> for BackendTransportKey { + fn from((backend_name, session_name, principal): (&str, &str, &str)) -> Self { + Self { + principal: principal.to_owned(), + backend_name: backend_name.to_owned(), + session_id: session_name.to_owned(), + } } } -impl From<(&String, &SessionId)> for BackendTransportKey { - fn from((backend_name, session_name): (&String, &SessionId)) -> Self { - Self { backend_name: backend_name.to_owned(), session_id: session_name.value().to_owned() } +impl From<(&String, &SessionId, &str)> for BackendTransportKey { + fn from((backend_name, session_name, principal): (&String, &SessionId, &str)) -> Self { + Self { + principal: principal.to_owned(), + backend_name: backend_name.to_owned(), + session_id: session_name.value().to_owned(), + } } } @@ -222,10 +223,10 @@ where }); } - let mut transports = self.transports.lock().await; + let mut transports = self.transports.inner().lock().await; for (name, svc) in backend_services { transports - .entry(BackendTransportKey::from((name.as_str(), downstream_session_id.value()))) + .entry(BackendTransportKey::from((name.as_str(), downstream_session_id.value(), claims.sub.as_str()))) .insert_entry(svc); } drop(transports); @@ -245,9 +246,9 @@ where cx: RequestContext, ) -> Result { let mcp_call_validator = AuthorizedCallValidator::new("list_tools", &cx); - let (virtual_host, session_id) = mcp_call_validator.validate()?; + let (virtual_host, session_id, claims) = mcp_call_validator.validate()?; - let session_manager = SessionManager::new(virtual_host, session_id, &self.transports); + let session_manager = SessionManager::new(virtual_host, session_id, claims.sub.as_str(), &self.transports); let backend_transports: Vec<_> = session_manager.borrow_transports().await; let list_tools_tasks = backend_transports @@ -294,8 +295,8 @@ where cx: RequestContext, ) -> Result { let mcp_call_validator = AuthorizedCallValidator::new("call_tool", &cx); - let (virtual_host, session_id) = mcp_call_validator.validate()?; - let session_manager = SessionManager::new(virtual_host, session_id, &self.transports); + let (virtual_host, session_id, claims) = mcp_call_validator.validate()?; + let session_manager = SessionManager::new(virtual_host, session_id, claims.sub.as_str(), &self.transports); let backend_names = session_manager.get_backend_names(); @@ -385,9 +386,9 @@ where cx: RequestContext, ) -> Result { let mcp_call_validator = AuthorizedCallValidator::new("list_resources", &cx); - let (virtual_host, session_id) = mcp_call_validator.validate()?; + let (virtual_host, session_id, claims) = mcp_call_validator.validate()?; - let session_manager = SessionManager::new(virtual_host, session_id, &self.transports); + let session_manager = SessionManager::new(virtual_host, session_id, claims.sub.as_str(), &self.transports); let backend_transports: Vec<_> = session_manager.borrow_transports().await; let list_resources_tasks = backend_transports @@ -434,8 +435,8 @@ where cx: RequestContext, ) -> Result { let mcp_call_validator = AuthorizedCallValidator::new("read_resource", &cx); - let (virtual_host, session_id) = mcp_call_validator.validate()?; - let session_manager = SessionManager::new(virtual_host, session_id, &self.transports); + let (virtual_host, session_id, claims) = mcp_call_validator.validate()?; + let session_manager = SessionManager::new(virtual_host, session_id, claims.sub.as_str(), &self.transports); let backend_names = session_manager.get_backend_names(); diff --git a/crates/contextforge-gateway-rs-lib/src/gateway/mod.rs b/crates/contextforge-gateway-rs-lib/src/gateway/mod.rs index 74a09da..5339216 100644 --- a/crates/contextforge-gateway-rs-lib/src/gateway/mod.rs +++ b/crates/contextforge-gateway-rs-lib/src/gateway/mod.rs @@ -3,5 +3,5 @@ pub(crate) mod mcp_gateway; mod session_manager; mod session_store; -pub use mcp_gateway::{BackendTransportCleanup, LocalUserSessionStore, McpService, new_backend_transports}; +pub use mcp_gateway::{BackendTransports, LocalUserSessionStore, McpService}; pub use session_store::{UserSession, UserSessionStore}; diff --git a/crates/contextforge-gateway-rs-lib/src/gateway/session_manager.rs b/crates/contextforge-gateway-rs-lib/src/gateway/session_manager.rs index 15f1b61..fdd69fe 100644 --- a/crates/contextforge-gateway-rs-lib/src/gateway/session_manager.rs +++ b/crates/contextforge-gateway-rs-lib/src/gateway/session_manager.rs @@ -1,25 +1,24 @@ -use std::{collections::HashMap, sync::Arc}; - use contextforge_gateway_rs_apis::user_store::VirtualHost; -use tokio::sync::Mutex; use tracing::{debug, info}; -use super::mcp_gateway::{BackendTransportKey, BackendTransportService, ServiceHolder}; +use super::mcp_gateway::{BackendTransportKey, BackendTransports, ServiceHolder}; use crate::layers::session_id::SessionId; pub struct SessionManager<'a> { virtual_host: &'a VirtualHost, session_id: &'a SessionId, - transports: &'a Arc>>, + principal: &'a str, + transports: &'a BackendTransports, } impl<'a> SessionManager<'a> { pub fn new( virtual_host: &'a VirtualHost, session_id: &'a SessionId, - transports: &'a Arc>>, + principal: &'a str, + transports: &'a BackendTransports, ) -> Self { - Self { virtual_host, session_id, transports } + Self { virtual_host, session_id, principal, transports } } pub fn get_backend_names(&self) -> Vec<&str> { @@ -28,12 +27,12 @@ impl<'a> SessionManager<'a> { pub async fn borrow_transports(&self) -> Vec { let names: Vec<_> = self.virtual_host.backends.keys().cloned().collect(); - let mut transports = self.transports.lock().await; + let mut transports = self.transports.inner().lock().await; names .into_iter() .filter_map(|name| { transports - .get_mut(&BackendTransportKey::from((&name, self.session_id))) + .get_mut(&BackendTransportKey::from((&name, self.session_id, self.principal))) .map(|b| ServiceHolder::new(name, b.service.clone())) }) .collect() @@ -42,10 +41,10 @@ impl<'a> SessionManager<'a> { // pub async fn return_transports(&self, backend_transports: impl Iterator) { // let backend_transports = backend_transports.collect::>(); // info!("Returning transports {:?} {backend_transports:?}", self.session_id); - // let mut transports = self.transports.lock().await; + // let mut transports = self.transports.inner().lock().await; // for svc_holder in backend_transports { // transports - // .entry(BackendTransportKey::from((&svc_holder.name, self.session_id))) + // .entry(BackendTransportKey::from((&svc_holder.name, self.session_id, self.principal))) // .and_modify(|e| e.service = svc_holder.running_service); // } // } @@ -53,9 +52,9 @@ impl<'a> SessionManager<'a> { pub async fn cleanup_backends(&self, reason: &'static str) { let names: Vec<_> = self.virtual_host.backends.keys().cloned().collect(); info!("Cleaning up backends {:?}", self.session_id); - let mut transports = self.transports.lock().await; + let mut transports = self.transports.inner().lock().await; for name in names { - let key = BackendTransportKey::from((&name, self.session_id)); + let key = BackendTransportKey::from((&name, self.session_id, self.principal)); debug!("session_manager: removing transport for {key:?} {reason}"); transports.remove(&key); } diff --git a/crates/contextforge-gateway-rs-lib/src/layers/session_id.rs b/crates/contextforge-gateway-rs-lib/src/layers/session_id.rs index 31c36e1..7ba42b0 100644 --- a/crates/contextforge-gateway-rs-lib/src/layers/session_id.rs +++ b/crates/contextforge-gateway-rs-lib/src/layers/session_id.rs @@ -7,8 +7,7 @@ use tracing::info; use crate::{ common::ContextForgeClaims, const_values::MCP_SESSION_ID, - gateway::{BackendTransportCleanup, UserSession, UserSessionStore}, - layers::virtual_host_id::VirtualHostId, + gateway::{BackendTransports, UserSession, UserSessionStore}, }; #[derive(Debug, Clone)] @@ -25,52 +24,44 @@ impl SessionId { #[derive(Clone)] pub struct SessionIdState { pub user_session_store: Arc, - pub backend_transport_cleanup: BackendTransportCleanup, + pub backend_transports: BackendTransports, } pub async fn session_id_layer(State(state): State, mut request: Request, next: Next) -> Response { let session_id = request.headers().get(MCP_SESSION_ID).and_then(|session_id| session_id.to_str().ok()).map(str::to_owned); - if request.method() != Method::DELETE { - if let Some(session_id) = session_id { - info!("MCP Session ID {session_id}"); - request.extensions_mut().insert(SessionId { value: session_id }); - } - return next.run(request).await; + if let Some(session_id) = &session_id { + info!("MCP Session ID {session_id}"); + request.extensions_mut().insert(SessionId { value: session_id.clone() }); } - let Some(session_id) = session_id else { - return response(StatusCode::BAD_REQUEST); - }; - let Some(claims) = request.extensions().get::() else { - return response(StatusCode::BAD_REQUEST); - }; - let Some(_virtual_host_id) = request.extensions().get::() else { - return response(StatusCode::BAD_REQUEST); - }; + match request.method() { + &Method::DELETE => { + let subject = request.extensions().get::().map(|claims| claims.sub.clone()); + let rmcp_response = next.run(request).await; + if !rmcp_response.status().is_success() { + return rmcp_response; + } - let user_session = UserSession::new(claims.sub.clone(), Arc::from(session_id.as_str())); - match state.user_session_store.get_session(&user_session).await { - Ok(Some(_)) => {}, - Ok(None) => return response(StatusCode::NOT_FOUND), - Err(_) => return response(StatusCode::INTERNAL_SERVER_ERROR), - } - - request.extensions_mut().insert(SessionId { value: session_id.clone() }); - let rmcp_response = next.run(request).await; - if !rmcp_response.status().is_success() { - return rmcp_response; - } + let Some(session_id) = session_id else { + return rmcp_response; + }; + let Some(subject) = subject else { + return response(StatusCode::BAD_REQUEST); + }; + let user_session = UserSession::new(subject.clone(), Arc::from(session_id.as_str())); + let remove_result = state.user_session_store.remove_session(&user_session).await; + state.backend_transports.remove_session(&subject, &session_id).await; - let remove_result = state.user_session_store.remove_session(&user_session).await; - state.backend_transport_cleanup.remove_session(&session_id).await; + if remove_result.is_err() { + return response(StatusCode::INTERNAL_SERVER_ERROR); + } - if remove_result.is_err() { - return response(StatusCode::INTERNAL_SERVER_ERROR); + rmcp_response + }, + _ => next.run(request).await, } - - rmcp_response } fn response(status: StatusCode) -> Response { diff --git a/crates/contextforge-gateway-rs-lib/src/lib.rs b/crates/contextforge-gateway-rs-lib/src/lib.rs index 1a8598b..698d4dc 100644 --- a/crates/contextforge-gateway-rs-lib/src/lib.rs +++ b/crates/contextforge-gateway-rs-lib/src/lib.rs @@ -19,7 +19,7 @@ mod tools; mod user_config_store; pub use common::{RedisClient, RedisConfig, UpstreamConnectionMode}; -use gateway::{BackendTransportCleanup, McpService, new_backend_transports}; +use gateway::{BackendTransports, McpService}; use layers::session_id::SessionId; use tower_http::cors::{Any, CorsLayer}; use transports::{DownstreamTls, Tcp}; @@ -70,10 +70,10 @@ impl Gateway { let user_config_store = user_config_store as Arc; let user_session_store = LocalUserSessionStore::new(); - let backend_transports = new_backend_transports(); + let backend_transports = BackendTransports::default(); let session_id_state = SessionIdState { user_session_store: Arc::new(user_session_store.clone()), - backend_transport_cleanup: BackendTransportCleanup::new(Arc::clone(&backend_transports)), + backend_transports: backend_transports.clone(), }; let mcp_plugin_runtime = self.plugin_runtime; @@ -88,7 +88,7 @@ impl Gateway { Ok(McpService::builder() .with_user_session_store(user_session_store.clone()) .with_http_client(reqwest_backend_client.clone()) - .with_transports(Arc::clone(&backend_transports)) + .with_transports(backend_transports.clone()) .with_plugin_runtime(mcp_plugin_runtime.clone()) .build()) },