From 0bf7a98d5fa73ff511c04f87d243507e3ac7e2d1 Mon Sep 17 00:00:00 2001 From: Dawid Nowak Date: Mon, 18 May 2026 12:53:52 +0100 Subject: [PATCH 1/2] Removing of option from mcp_gateways/BackendTransportService Signed-off-by: Dawid Nowak --- .../src/gateway/mcp_gateway.rs | 140 +++++++----------- .../src/gateway/session_manager.rs | 22 +-- 2 files changed, 65 insertions(+), 97 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 98fc8005..ec7b5137 100644 --- a/crates/contextforge-gateway-rs-lib/src/gateway/mcp_gateway.rs +++ b/crates/contextforge-gateway-rs-lib/src/gateway/mcp_gateway.rs @@ -19,7 +19,7 @@ use rmcp::{ service::{RequestContext, RunningService}, transport::{StreamableHttpClientTransport, streamable_http_client::StreamableHttpClientTransportConfig}, }; -use tokio::sync::Mutex; +use tokio::sync::{Mutex, RwLock}; use tracing::{debug, info, warn}; use typed_builder::TypedBuilder; @@ -61,13 +61,13 @@ pub type BackendService = RunningService; #[derive(Debug)] pub struct ServiceHolder { pub name: String, - pub running_service: Option>, + pub running_service: Option>>>, } impl ServiceHolder { pub fn new( name: String, - running_service: Option>, + running_service: Option>>>, ) -> ServiceHolder { Self { name, running_service } } @@ -77,7 +77,7 @@ impl ServiceHolder { pub struct BackendTransportService { #[expect(dead_code, reason = "stored backend capabilities are kept with transport state for future routing")] capabilities: Option, - pub(crate) service: Option, + pub(crate) service: Option>>, } impl From<(&str, &str)> for BackendTransportKey { @@ -92,13 +92,9 @@ impl From<(&String, &SessionId)> for BackendTransportKey { } } -impl From<(Option, Option)> for BackendTransportService { - fn from((capabilities, service): (Option, Option)) -> Self { - if let Some(service) = service { - Self { capabilities, service: Some(service) } - } else { - Self { capabilities, service: None } - } +impl From<(Option, Option>>)> for BackendTransportService { + fn from((capabilities, service): (Option, Option>>)) -> Self { + Self { capabilities, service } } } @@ -183,7 +179,7 @@ where .map(|pi| pi.capabilities.clone())); ( (name.clone(), server_capabilities.clone()), - (name.clone(), BackendTransportService::from((server_capabilities, running_service))), + (name.clone(), BackendTransportService::from((server_capabilities, running_service.map(|rs| Arc::new(RwLock::new(rs)))))), ) }) .unzip(); @@ -230,31 +226,25 @@ where let request = request.clone(); async move { if let Some(service) = service_holder.running_service { + let service = service.read().await; let response = service.list_tools(request).await; - (service_holder.name, Some(service), Some(response)) + (service_holder.name, Some(response)) } else { - (service_holder.name, None, None) + (service_holder.name, None) } } }) .collect::>(); - let list_tools_tasks_results: Vec<(String, Option<_>, Option<_>)> = - futures::future::join_all(list_tools_tasks).await; + let list_tools_tasks_results: Vec<(String, Option<_>)> = futures::future::join_all(list_tools_tasks).await; - let (backend_services, responses): (Vec<_>, Vec<_>) = list_tools_tasks_results + let responses: Vec<_> = list_tools_tasks_results .into_iter() - .map(|(name, service, response)| { + .map(|(name, response)| { info!("list_tools: backend {name} {response:?}"); - ((name.clone(), service), (name, response)) + (name, response) }) - .unzip(); - - let mut transports = self.transports.lock().await; - for (name, svc) in backend_services { - transports.entry(BackendTransportKey::from((&name, session_id))).and_modify(|e| e.service = svc); - } - drop(transports); + .collect(); let responses = responses .into_iter() @@ -290,37 +280,32 @@ where let backend_transports = session_manager.borrow_transports().await; info!("Borrowed transports {session_id:?} {backend_transports:?}"); - let (services, call_tool_tasks): (Vec<_>, Vec<_>) = backend_transports + let call_tool_tasks: Vec<_> = backend_transports .into_iter() - .map(|service_holder| { + .filter_map(|service_holder| { debug!( "call_tool: Finding backend for {} {service_holder:?} {backend_name} tool_name = {tool_name}", &request.name, ); - if service_holder.name == backend_name { + (service_holder.name == backend_name).then(|| { let mut request = request.clone(); request.name = tool_name.to_owned().into(); - ( - None, - Some(async move { + async move { if let Some(service) = service_holder.running_service { + let service = service.read().await; let response = service.call_tool(request).await; - (service_holder.name, Some(service), Some(response)) + (service_holder.name, Some(response)) } else { warn!("call_tool: trying to call a tool for which we have no backend {service_holder:?} {backend_name} tool_name = {tool_name}"); - (service_holder.name, None, None) + (service_holder.name, None) } - }), - ) - } else { - (Some(service_holder), None) - } - }) - .unzip(); + } + + }) + }).collect(); - let call_tool_tasks = call_tool_tasks.into_iter().flatten().collect::>(); if call_tool_tasks.len() > 1 { warn!("call_tool: More than one tool matching for tool name {}", request.name); @@ -333,18 +318,15 @@ where }); } - let call_tool_tasks_results: Vec<(String, Option<_>, Option<_>)> = - futures::future::join_all(call_tool_tasks).await; + let call_tool_tasks_results: Vec<(String, Option<_>)> = futures::future::join_all(call_tool_tasks).await; - let (backend_services, responses): (Vec<_>, Vec<_>) = call_tool_tasks_results + let responses: Vec<_> = call_tool_tasks_results .into_iter() - .map(|(name, service, response)| { + .map(|(name, response)| { info!("call_tool: backend {name} {response:?}"); - (ServiceHolder::new(name.clone(), service), (name, response)) + (name, response) }) - .unzip(); - - session_manager.return_transports(backend_services.into_iter().chain(services.into_iter().flatten())).await; + .collect(); let responses = responses .into_iter() @@ -377,31 +359,25 @@ where let request = request.clone(); async move { if let Some(service) = service_holder.running_service { + let service = service.read().await; let response = service.list_resources(request).await; - (service_holder.name, Some(service), Some(response)) + (service_holder.name, Some(response)) } else { - (service_holder.name, None, None) + (service_holder.name, None) } } }) .collect::>(); - let list_tools_tasks_results: Vec<(String, Option<_>, Option<_>)> = - futures::future::join_all(list_resources_tasks).await; + let list_tools_tasks_results: Vec<(String, Option<_>)> = futures::future::join_all(list_resources_tasks).await; - let (backend_services, responses): (Vec<_>, Vec<_>) = list_tools_tasks_results + let responses: Vec<_> = list_tools_tasks_results .into_iter() - .map(|(name, service, response)| { + .map(|(name, response)| { info!("list_resources: backend {name} {response:?}"); - ((name.clone(), service), (name, response)) + (name, response) }) - .unzip(); - - let mut transports = self.transports.lock().await; - for (name, svc) in backend_services { - transports.entry(BackendTransportKey::from((&name, session_id))).and_modify(|e| e.service = svc); - } - drop(transports); + .collect(); let responses = responses .into_iter() @@ -439,7 +415,7 @@ where let backend_transports = session_manager.borrow_transports().await; info!("Borrowed transports {session_id:?} {backend_transports:?}"); - let (services, call_tool_tasks): (Vec<_>, Vec<_>) = backend_transports + let call_tool_tasks: Vec<_> = backend_transports .into_iter() .map(|service_holder| { debug!( @@ -447,27 +423,22 @@ where &request.uri, ); - if service_holder.name == backend_name { + (service_holder.name == backend_name).then(|| { let mut request = request.clone(); request.uri = String::from(resource_uri); - ( - None, - Some(async move { + async move { if let Some(service) = service_holder.running_service { + let service = service.read().await; let response = service.read_resource(request).await; - (service_holder.name, Some(service), Some(response)) + (service_holder.name, Some(response)) } else { warn!("call_tool: trying to call a tool for which we have no backend {service_holder:?} {backend_name} resource_name = {resource_uri}"); - (service_holder.name, None, None) + (service_holder.name, None) } - }), - ) - } else { - (Some(service_holder), None) - } - }) - .unzip(); + } + }) + }).collect(); let call_tool_tasks = call_tool_tasks.into_iter().flatten().collect::>(); if call_tool_tasks.len() > 1 { @@ -482,18 +453,15 @@ where }); } - let call_tool_tasks_results: Vec<(String, Option<_>, Option<_>)> = - futures::future::join_all(call_tool_tasks).await; + let call_tool_tasks_results: Vec<_> = futures::future::join_all(call_tool_tasks).await; - let (backend_services, responses): (Vec<_>, Vec<_>) = call_tool_tasks_results + let responses: Vec<_> = call_tool_tasks_results .into_iter() - .map(|(name, service, response)| { + .map(|(name, response)| { info!("read_resource: backend {name} {response:?}"); - (ServiceHolder::new(name.clone(), service), (name, response)) + (name, response) }) - .unzip(); - - session_manager.return_transports(backend_services.into_iter().chain(services.into_iter().flatten())).await; + .collect(); let responses = responses .into_iter() 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 49e829ad..15f1b611 100644 --- a/crates/contextforge-gateway-rs-lib/src/gateway/session_manager.rs +++ b/crates/contextforge-gateway-rs-lib/src/gateway/session_manager.rs @@ -34,21 +34,21 @@ impl<'a> SessionManager<'a> { .filter_map(|name| { transports .get_mut(&BackendTransportKey::from((&name, self.session_id))) - .map(|b| ServiceHolder::new(name, b.service.take())) + .map(|b| ServiceHolder::new(name, b.service.clone())) }) .collect() } - 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; - for svc_holder in backend_transports { - transports - .entry(BackendTransportKey::from((&svc_holder.name, self.session_id))) - .and_modify(|e| e.service = svc_holder.running_service); - } - } + // 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; + // for svc_holder in backend_transports { + // transports + // .entry(BackendTransportKey::from((&svc_holder.name, self.session_id))) + // .and_modify(|e| e.service = svc_holder.running_service); + // } + // } pub async fn cleanup_backends(&self, reason: &'static str) { let names: Vec<_> = self.virtual_host.backends.keys().cloned().collect(); From 0c863dcc9cc1b99d3d582138dc69717fda3d3d05 Mon Sep 17 00:00:00 2001 From: Dawid Nowak Date: Mon, 18 May 2026 22:54:44 +0100 Subject: [PATCH 2/2] Removing of option from mcp_gateways/BackendTransportService.2 Signed-off-by: Dawid Nowak --- .../src/gateway/mcp_gateway.rs | 27 +++--- .../src/tests/gateway_end_to_end.rs | 93 ++++++++++++++++++- 2 files changed, 103 insertions(+), 17 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 ec7b5137..77418ad3 100644 --- a/crates/contextforge-gateway-rs-lib/src/gateway/mcp_gateway.rs +++ b/crates/contextforge-gateway-rs-lib/src/gateway/mcp_gateway.rs @@ -19,7 +19,7 @@ use rmcp::{ service::{RequestContext, RunningService}, transport::{StreamableHttpClientTransport, streamable_http_client::StreamableHttpClientTransportConfig}, }; -use tokio::sync::{Mutex, RwLock}; +use tokio::sync::Mutex; use tracing::{debug, info, warn}; use typed_builder::TypedBuilder; @@ -56,19 +56,16 @@ pub struct BackendTransportKey { session_id: String, } -pub type BackendService = RunningService; +type McpClientService = Arc>; #[derive(Debug)] pub struct ServiceHolder { pub name: String, - pub running_service: Option>>>, + pub running_service: Option, } impl ServiceHolder { - pub fn new( - name: String, - running_service: Option>>>, - ) -> ServiceHolder { + pub fn new(name: String, running_service: Option) -> ServiceHolder { Self { name, running_service } } } @@ -77,7 +74,7 @@ impl ServiceHolder { pub struct BackendTransportService { #[expect(dead_code, reason = "stored backend capabilities are kept with transport state for future routing")] capabilities: Option, - pub(crate) service: Option>>, + pub(crate) service: Option, } impl From<(&str, &str)> for BackendTransportKey { @@ -92,8 +89,8 @@ impl From<(&String, &SessionId)> for BackendTransportKey { } } -impl From<(Option, Option>>)> for BackendTransportService { - fn from((capabilities, service): (Option, Option>>)) -> Self { +impl From<(Option, Option)> for BackendTransportService { + fn from((capabilities, service): (Option, Option)) -> Self { Self { capabilities, service } } } @@ -179,7 +176,7 @@ where .map(|pi| pi.capabilities.clone())); ( (name.clone(), server_capabilities.clone()), - (name.clone(), BackendTransportService::from((server_capabilities, running_service.map(|rs| Arc::new(RwLock::new(rs)))))), + (name.clone(), BackendTransportService::from((server_capabilities, running_service.map(Arc::new)))), ) }) .unzip(); @@ -226,7 +223,7 @@ where let request = request.clone(); async move { if let Some(service) = service_holder.running_service { - let service = service.read().await; + //let service = service.read().await; let response = service.list_tools(request).await; (service_holder.name, Some(response)) } else { @@ -293,7 +290,7 @@ where request.name = tool_name.to_owned().into(); async move { if let Some(service) = service_holder.running_service { - let service = service.read().await; +// let service = service.read().await; let response = service.call_tool(request).await; (service_holder.name, Some(response)) @@ -359,7 +356,7 @@ where let request = request.clone(); async move { if let Some(service) = service_holder.running_service { - let service = service.read().await; + //let service = service.read().await; let response = service.list_resources(request).await; (service_holder.name, Some(response)) } else { @@ -428,7 +425,7 @@ where request.uri = String::from(resource_uri); async move { if let Some(service) = service_holder.running_service { - let service = service.read().await; + //let service = service.read().await; let response = service.read_resource(request).await; (service_holder.name, Some(response)) diff --git a/crates/contextforge-gateway-rs-lib/src/tests/gateway_end_to_end.rs b/crates/contextforge-gateway-rs-lib/src/tests/gateway_end_to_end.rs index 63db44db..f1e96aed 100644 --- a/crates/contextforge-gateway-rs-lib/src/tests/gateway_end_to_end.rs +++ b/crates/contextforge-gateway-rs-lib/src/tests/gateway_end_to_end.rs @@ -16,7 +16,7 @@ use http::{HeaderMap, HeaderValue}; use jsonwebtoken::{Algorithm, EncodingKey, Header, encode}; use rmcp::{ ServiceExt, - model::InitializeRequestParams, + model::{CallToolRequestParams, InitializeRequestParams}, transport::{ StreamableHttpClientTransport, StreamableHttpServerConfig, StreamableHttpService, streamable_http_client::StreamableHttpClientTransportConfig, @@ -161,7 +161,6 @@ async fn create_gateway_with_four_counters(user: &str, config: Config) -> crate: let gateway = Gateway::builder() .with_config(config.clone()) - //.with_user_config_store(Arc::new(mocked_user_config_store)) .with_session_manager(Arc::new(LocalSessionManager::default())) .with_user_config_store_type(crate::UserConfigStoreType::Test(Arc::new(mocked_user_config_store))) .build(); @@ -330,6 +329,96 @@ async fn plaintext_list_tools_end_to_end_test() -> crate::Result<()> { Ok(()) } +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +#[test_log::test] +async fn plaintext_overlapped_counter_tools_end_to_end_test() -> crate::Result<()> { + let gateway_port = create_ports(1)[0]; + + let config = Config { + address: Some(format!("127.0.0.1:{gateway_port}").parse().expect("This should work")), + token_verification_public_key: "../../assets/jwt.key.pub".into(), + upstream_connection_mode: Some(crate::common::UpstreamConnectionMode::PlainTextOrTls), + ..Default::default() + }; + + let user = "admin@example.com"; + + let Ok(TestSettings { handle, gateway_url, expected_tool_names }) = + create_gateway_with_four_counters(user, config).await + else { + panic!("Invalid configuration "); + }; + + let test_future: BoxFuture<'_, crate::Result<()>> = async { + tokio::time::sleep(Duration::from_millis(100)).await; + let mut default_headers = HeaderMap::new(); + let token = get_token(user.to_owned()); + default_headers.insert( + http::header::AUTHORIZATION, + HeaderValue::from_str(format!("Bearer {token}").as_str()).expect("This should work"), + ); + let client = reqwest::Client::builder().default_headers(default_headers).build().expect("This should work"); + + info!("Seding request to {gateway_url}"); + + let config = StreamableHttpClientTransportConfig::with_uri(gateway_url); + let transport = StreamableHttpClientTransport::with_client(client, config); + let request = InitializeRequestParams::default(); + + let maybe_service = request.serve(transport).await; + let Ok(running_service) = maybe_service else { + warn!("No Service {maybe_service:?}"); + return Err("Couldn't get a service".into()); + }; + + let increment_tool_name = + expected_tool_names.into_iter().find(|tn| tn.ends_with("increment")).expect("This should work"); + let get_value_tool_name = increment_tool_name.replace("increment", "get_value"); + + let call_tool = CallToolRequestParams::new(increment_tool_name); + + let tasks = (0..4).map(|_| async { running_service.call_tool(call_tool.clone()).await }).collect::>(); + + let results = futures::future::join_all(tasks).await; + for result in results { + let Ok(_) = result else { + let msg = format!("Call tools returned error {call_tool:?}"); + warn!(msg); + return Err(msg.into()); + }; + } + + let call_tool = CallToolRequestParams::new(get_value_tool_name); + + let Ok(get_value_result) = running_service.call_tool(call_tool).await else { + let msg = "Call tools get value returned error"; + warn!(msg); + return Err(msg.into()); + }; + info!("{get_value_result:?}"); + let actual_value: usize = get_value_result.into_typed().expect("This should work"); + if 4 != actual_value { + warn!("Wrong value for counters"); + return Err("Wrong value for counters".into()); + } + + Ok(()) + } + .boxed(); + + let maybe_passed = test_future.await; + + handle.abort(); + if maybe_passed.is_ok() { + info!("Test passed"); + } else { + info!("Test NOT passed {maybe_passed:?}"); + panic!() + } + + Ok(()) +} + #[tokio::test(flavor = "multi_thread", worker_threads = 1)] #[test_log::test] async fn tls_list_tools_end_to_end_test() -> crate::Result<()> {