Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
143 changes: 54 additions & 89 deletions crates/contextforge-gateway-rs-lib/src/gateway/mcp_gateway.rs
Original file line number Diff line number Diff line change
Expand Up @@ -56,19 +56,16 @@ pub struct BackendTransportKey {
session_id: String,
}

pub type BackendService = RunningService<RoleClient, InitializeRequestParams>;
type McpClientService = Arc<RunningService<RoleClient, InitializeRequestParams>>;

#[derive(Debug)]
pub struct ServiceHolder {
pub name: String,
pub running_service: Option<RunningService<RoleClient, InitializeRequestParams>>,
Comment thread
dawid-nowak marked this conversation as resolved.
pub running_service: Option<McpClientService>,
}

impl ServiceHolder {
pub fn new(
name: String,
running_service: Option<RunningService<RoleClient, InitializeRequestParams>>,
) -> ServiceHolder {
pub fn new(name: String, running_service: Option<McpClientService>) -> ServiceHolder {
Self { name, running_service }
}
}
Expand All @@ -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<ServerCapabilities>,
pub(crate) service: Option<BackendService>,
pub(crate) service: Option<McpClientService>,
}

impl From<(&str, &str)> for BackendTransportKey {
Expand All @@ -92,13 +89,9 @@ impl From<(&String, &SessionId)> for BackendTransportKey {
}
}

impl From<(Option<ServerCapabilities>, Option<BackendService>)> for BackendTransportService {
fn from((capabilities, service): (Option<ServerCapabilities>, Option<BackendService>)) -> Self {
if let Some(service) = service {
Self { capabilities, service: Some(service) }
} else {
Self { capabilities, service: None }
}
impl From<(Option<ServerCapabilities>, Option<McpClientService>)> for BackendTransportService {
fn from((capabilities, service): (Option<ServerCapabilities>, Option<McpClientService>)) -> Self {
Self { capabilities, service }
}
}

Expand Down Expand Up @@ -183,7 +176,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(Arc::new)))),
)
})
.unzip();
Expand Down Expand Up @@ -230,31 +223,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::<Vec<_>>();

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()
Expand Down Expand Up @@ -290,37 +277,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}",
Comment thread
dawid-nowak marked this conversation as resolved.
&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::<Vec<_>>();
if call_tool_tasks.len() > 1 {
warn!("call_tool: More than one tool matching for tool name {}", request.name);

Expand All @@ -333,18 +315,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()
Expand Down Expand Up @@ -377,31 +356,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::<Vec<_>>();

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()
Expand Down Expand Up @@ -439,35 +412,30 @@ 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!(
"read_resource: Finding backend for {} {service_holder:?} {backend_name} read_resource = {resource_uri}",
&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::<Vec<_>>();
if call_tool_tasks.len() > 1 {
Expand All @@ -482,18 +450,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()
Expand Down
22 changes: 11 additions & 11 deletions crates/contextforge-gateway-rs-lib/src/gateway/session_manager.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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()))
Comment thread
dawid-nowak marked this conversation as resolved.
.map(|b| ServiceHolder::new(name, b.service.clone()))
})
.collect()
}

pub async fn return_transports(&self, backend_transports: impl Iterator<Item = ServiceHolder>) {
let backend_transports = backend_transports.collect::<Vec<_>>();
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<Item = ServiceHolder>) {
// let backend_transports = backend_transports.collect::<Vec<_>>();
// 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();
Expand Down
Loading