diff --git a/crates/contextforge-gateway-rs-lib/src/gateway/backend_client.rs b/crates/contextforge-gateway-rs-lib/src/gateway/backend_client.rs index b73d6be..7d9b2eb 100644 --- a/crates/contextforge-gateway-rs-lib/src/gateway/backend_client.rs +++ b/crates/contextforge-gateway-rs-lib/src/gateway/backend_client.rs @@ -14,7 +14,7 @@ use tokio::sync::{Mutex, RwLock}; use tokio_util::sync::CancellationToken; use tracing::{debug, warn}; -use super::mcp_gateway::prefixed_name; +use super::identifier_routing::prefixed_name; #[derive(Clone)] pub(crate) struct GatewayBackendClient { diff --git a/crates/contextforge-gateway-rs-lib/src/gateway/backend_transports.rs b/crates/contextforge-gateway-rs-lib/src/gateway/backend_transports.rs new file mode 100644 index 0000000..40b387b --- /dev/null +++ b/crates/contextforge-gateway-rs-lib/src/gateway/backend_transports.rs @@ -0,0 +1,75 @@ +use std::{collections::HashMap, sync::Arc}; + +use rmcp::{RoleClient, model::ServerCapabilities, service::RunningService}; +use tokio::sync::Mutex; + +use super::backend_client::GatewayBackendClient; +use crate::SessionId; + +pub(crate) type McpClientService = Arc>; + +#[derive(Clone, Default)] +pub struct BackendTransports(Arc>>); + +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(crate) fn inner(&self) -> &Arc>> { + &self.0 + } +} + +#[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)] +pub(crate) struct BackendTransportKey { + principal: String, + backend_name: String, + session_id: String, +} + +#[derive(Debug)] +pub(crate) struct ServiceHolder { + pub(crate) name: String, + pub(crate) running_service: Option, +} + +impl ServiceHolder { + pub(crate) fn new(name: String, running_service: Option) -> Self { + Self { name, running_service } + } +} + +#[derive(Debug)] +pub(crate) struct BackendTransportService { + #[expect(dead_code, reason = "stored backend capabilities are kept with transport state for future routing")] + capabilities: Option, + pub(crate) service: Option, +} + +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, &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(), + } + } +} + +impl From<(Option, Option)> for BackendTransportService { + fn from((capabilities, service): (Option, Option)) -> Self { + Self { capabilities, service } + } +} diff --git a/crates/contextforge-gateway-rs-lib/src/gateway/identifier_routing.rs b/crates/contextforge-gateway-rs-lib/src/gateway/identifier_routing.rs new file mode 100644 index 0000000..ff36274 --- /dev/null +++ b/crates/contextforge-gateway-rs-lib/src/gateway/identifier_routing.rs @@ -0,0 +1,243 @@ +use contextforge_gateway_rs_apis::user_store::VirtualHost; +use rmcp::{ErrorData, model::ErrorCode}; +use tracing::{debug, warn}; + +use super::{ + backend_transports::{McpClientService, ServiceHolder}, + session_manager::SessionManager, +}; + +/// Preserves identifiers for a single backend. For multiple backends, splits a +/// `{backend}-{identifier}` namespace so duplicate identifiers remain routable. +fn route_identifier<'a, N: AsRef>(identifier: &'a str, backend_names: &'a [N]) -> Option<(&'a str, &'a str)> { + if let [backend] = backend_names { + return Some((backend.as_ref(), identifier)); + } + + backend_names.iter().find_map(|backend| { + let backend = backend.as_ref(); + identifier.strip_prefix(backend)?.strip_prefix('-').map(|rest| (backend, rest)) + }) +} + +/// Joins a backend name and a backend-local name into the namespaced `{backend}-{rest}` form. +pub(crate) fn prefixed_name(backend_name: &str, rest: &str) -> String { + format!("{backend_name}-{rest}") +} + +/// Resolves an exact control-plane alias to its backend and upstream name. Without an alias, +/// single-backend hosts preserve the upstream name and multi-backend hosts use the legacy prefix. +pub(super) fn resolve_tool_route<'a, N: AsRef>( + virtual_host: &'a VirtualHost, + name: &'a str, + backend_names: &'a [N], +) -> Option<(&'a str, &'a str)> { + let mut aliases = backend_names.iter().filter_map(|backend_name| { + let backend_name = backend_name.as_ref(); + let original_name = virtual_host.backends.get(backend_name)?.tool_name_aliases.get(name)?; + Some((backend_name, original_name.as_str())) + }); + let alias = aliases.next(); + if aliases.next().is_some() { + return None; + } + alias.or_else(|| route_identifier(name, backend_names)) +} + +/// Returns the control-plane alias for an upstream tool when configured. Without an alias, +/// single-backend hosts preserve the upstream name and multi-backend hosts use the legacy prefix. +pub(super) fn exposed_tool_name(virtual_host: &VirtualHost, backend_name: &str, original_name: &str) -> String { + virtual_host + .backends + .get(backend_name) + .and_then(|backend| { + backend + .tool_name_aliases + .iter() + .find_map(|(alias, original)| (original == original_name).then(|| alias.clone())) + }) + .unwrap_or_else(|| { + if virtual_host.backends.len() == 1 { + original_name.to_owned() + } else { + prefixed_name(backend_name, original_name) + } + }) +} + +/// Logs a backend forwarding failure and maps it to the routing error every handler returns. +pub(super) fn backend_forward_error(op: &str, backend_name: &str, error: &impl std::fmt::Debug) -> ErrorData { + warn!("{op}: backend {backend_name} {error:?}"); + ErrorData { + code: ErrorCode::INTERNAL_ERROR, + message: "Routing problem... got no responses from backends".into(), + data: None, + } +} + +/// Routes an identifier to its backend, preserving it for a single backend and splitting the +/// namespace for multiple backends. Returns `(backend_name, service, backend_local_identifier)`. +pub(super) async fn route_identifier_to_backend( + session_manager: &SessionManager<'_>, + op: &str, + identifier: &str, + no_route_message: &'static str, +) -> Result<(String, McpClientService, String), ErrorData> { + let backend_names = session_manager.get_backend_names(); + let Some((backend_name, routed_identifier)) = route_identifier(identifier, &backend_names) else { + return Err(ErrorData { code: ErrorCode::INTERNAL_ERROR, message: no_route_message.into(), data: None }); + }; + let routed_identifier = routed_identifier.to_owned(); + let (backend_name, service) = resolve_backend(session_manager, op, backend_name).await?; + Ok((backend_name, service, routed_identifier)) +} + +/// Resolves the single connected backend named `backend_name` and takes its running service. +/// Shared by tool, resource, and prompt routing so they reject duplicate or missing backends +/// the same way; a duplicate match means the session is invalid, so it is cleaned up. +pub(super) async fn resolve_backend( + session_manager: &SessionManager<'_>, + op: &str, + backend_name: &str, +) -> Result<(String, McpClientService), ErrorData> { + let backend_transports = session_manager.borrow_transports().await; + debug!("{op}: resolving backend {backend_name} from {backend_transports:?}"); + + let mut target = None; + for service_holder in backend_transports { + if service_holder.name == backend_name { + if target.is_some() { + warn!("{op}: more than one backend matching {backend_name}"); + session_manager.cleanup_backends("invalid session.. duplicate backends detected").await; + return Err(ErrorData { + code: ErrorCode::INVALID_REQUEST, + message: "Routing problem... multiple matching backends".into(), + data: None, + }); + } + target = Some(service_holder); + } + } + + let Some(ServiceHolder { name, running_service }) = target else { + return Err(ErrorData { + code: ErrorCode::INTERNAL_ERROR, + message: "Routing problem... got no responses from backends".into(), + data: None, + }); + }; + let Some(service) = running_service else { + warn!("{op}: no running backend for {backend_name}"); + return Err(ErrorData { + code: ErrorCode::INTERNAL_ERROR, + message: "Routing problem... got no responses from backends".into(), + data: None, + }); + }; + Ok((name, service)) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn multi_backend_route_requires_exact_backend_prefix() { + let backend_names = vec!["counter-on", "counter-oneee", "counter-one"]; + assert_eq!(Some(("counter-one", "increment")), route_identifier("counter-one-increment", &backend_names)); + assert_eq!(None, route_identifier("counter-oneincrement", &backend_names)); + assert_eq!(None, route_identifier("counteroneincrement", &backend_names)); + assert_eq!(Some(("counter-one", "get-value")), route_identifier("counter-one-get-value", &backend_names)); + + // Tool, resource, and prompt routing all share this splitter. + assert_eq!( + Some(("counter-one", "example-prompt")), + route_identifier("counter-one-example-prompt", &backend_names) + ); + assert_eq!(None, route_identifier("counter-oneexample-prompt", &backend_names)); + + let backend_names = vec!["counter_on", "counter_oneee", "counter_one"]; + assert_eq!(Some(("counter_one", "get-value")), route_identifier("counter_one-get-value", &backend_names)); + } + + #[test] + fn single_backend_routes_unprefixed_identifier_unchanged() { + let backend_names = vec!["backend-id"]; + + assert_eq!(Some(("backend-id", "test_simple_text")), route_identifier("test_simple_text", &backend_names)); + assert_eq!(Some(("backend-id", "backend-id-tool")), route_identifier("backend-id-tool", &backend_names)); + assert_eq!( + Some(("backend-id", "test://template/123/data")), + route_identifier("test://template/123/data", &backend_names) + ); + } + + #[test] + fn control_plane_alias_is_advertised_and_routes_to_original_name() { + let config_json = serde_json::json!({ + "backends": { + "79fabb70-2188-4de8-95ed-dc1e976e14d4": { + "name": "compliance_reference", + "url": "http://upstream:9000/mcp", + "transport": "STREAMABLEHTTP", + "passthrough_headers": [], + "allowed_tool_names": ["get_stats", "echo"], + "tool_name_aliases": { + "Public.Tool": "get_stats", + "Echo_Tool": "echo" + }, + "allowed_resource_names": [], + "allowed_prompt_names": [] + } + } + }); + let virtual_host: VirtualHost = serde_json::from_value(config_json).expect("valid virtual host"); + let backend_ids = vec!["79fabb70-2188-4de8-95ed-dc1e976e14d4"]; + + assert_eq!( + "Public.Tool", + exposed_tool_name(&virtual_host, "79fabb70-2188-4de8-95ed-dc1e976e14d4", "get_stats") + ); + assert_eq!( + Some(("79fabb70-2188-4de8-95ed-dc1e976e14d4", "get_stats")), + resolve_tool_route(&virtual_host, "Public.Tool", &backend_ids) + ); + } + + #[test] + fn multi_backend_tool_routing_falls_back_to_legacy_prefixed_names() { + let config_json = serde_json::json!({ + "backends": { + "compliance-reference": { + "name": "compliance_reference", + "url": "http://upstream:9000/mcp", + "transport": "STREAMABLEHTTP", + "passthrough_headers": [], + "allowed_tool_names": ["get_stats"], + "allowed_resource_names": [], + "allowed_prompt_names": [] + }, + "other": { + "name": "other", + "url": "http://other:9000/mcp", + "transport": "STREAMABLEHTTP", + "passthrough_headers": [], + "allowed_tool_names": [], + "allowed_resource_names": [], + "allowed_prompt_names": [] + } + } + }); + let virtual_host: VirtualHost = serde_json::from_value(config_json).expect("valid virtual host"); + let backend_names = vec!["compliance-reference", "other"]; + + assert_eq!( + "compliance-reference-get_stats", + exposed_tool_name(&virtual_host, "compliance-reference", "get_stats") + ); + assert_eq!( + Some(("compliance-reference", "get_stats")), + resolve_tool_route(&virtual_host, "compliance-reference-get_stats", &backend_names) + ); + } +} diff --git a/crates/contextforge-gateway-rs-lib/src/gateway/list_aggregation.rs b/crates/contextforge-gateway-rs-lib/src/gateway/list_aggregation.rs new file mode 100644 index 0000000..2612f9e --- /dev/null +++ b/crates/contextforge-gateway-rs-lib/src/gateway/list_aggregation.rs @@ -0,0 +1,201 @@ +use contextforge_gateway_rs_apis::user_store::VirtualHost; +use rmcp::model::{ + ListPromptsResult, ListResourceTemplatesResult, ListResourcesResult, ListToolsResult, Prompt, Resource, + ResourceTemplate, Tool, +}; +use tracing::{info, warn}; + +use super::{ + backend_transports::{McpClientService, ServiceHolder}, + identifier_routing::{exposed_tool_name, prefixed_name}, +}; + +/// Fans a paginated list request out to every connected backend concurrently, logs each response, +/// and returns the `(backend_name, result)` pairs that succeeded. +pub(super) async fn fan_out_list( + backends: Vec, + op: &str, + item_count: C, + call: F, +) -> Vec<(String, R)> +where + F: Fn(McpClientService) -> Fut, + Fut: std::future::Future>, + C: Fn(&R) -> usize, + E: std::fmt::Debug, +{ + let tasks = backends.into_iter().map(|service_holder| { + let call = &call; + async move { + let response = match service_holder.running_service { + Some(service) => Some(call(service).await), + None => None, + }; + (service_holder.name, response) + } + }); + + futures::future::join_all(tasks) + .await + .into_iter() + .filter_map(|(name, response)| { + log_backend_response(op, &name, response.as_ref(), &item_count); + match response { + Some(Ok(response)) => Some((name, response)), + _ => None, + } + }) + .collect() +} + +fn log_backend_response( + kind: &str, + name: &str, + response: Option<&Result>, + item_count: impl Fn(&T) -> usize, +) { + match response { + Some(Ok(response)) => info!("{kind}: backend {name} completed ({} items)", item_count(response)), + Some(Err(error)) => warn!("{kind}: backend {name} {error:?}"), + None => info!("{kind}: backend {name} unavailable"), + } +} + +pub(super) fn merge_tools(tools: Vec<(String, ListToolsResult)>, virtual_host: &VirtualHost) -> Vec { + let mut tools = tools + .into_iter() + .flat_map(|(backend_name, result)| { + result + .tools + .into_iter() + .map(|mut tool| { + tool.name = exposed_tool_name(virtual_host, &backend_name, &tool.name).into(); + tool + }) + .collect::>() + }) + .collect::>(); + tools.sort_unstable_by(|tool, other| tool.name.cmp(&other.name)); + tools +} + +pub(super) fn merge_resources( + resources: Vec<(String, ListResourcesResult)>, + namespace_identifiers: bool, +) -> Vec { + let mut resources = resources + .into_iter() + .flat_map(|(backend_name, result)| { + result + .resources + .into_iter() + .map(|mut resource| { + if namespace_identifiers { + resource.name = prefixed_name(&backend_name, &resource.name); + resource.uri = prefixed_name(&backend_name, &resource.uri); + } + resource + }) + .collect::>() + }) + .collect::>(); + resources.sort_unstable_by(|resource, other| resource.name.cmp(&other.name)); + resources +} + +pub(super) fn merge_resource_templates( + templates: Vec<(String, ListResourceTemplatesResult)>, + namespace_identifiers: bool, +) -> Vec { + let mut templates = templates + .into_iter() + .flat_map(|(backend_name, result)| { + result.resource_templates.into_iter().map(move |mut template| { + if namespace_identifiers { + template.name = prefixed_name(&backend_name, &template.name); + template.uri_template = prefixed_name(&backend_name, &template.uri_template); + } + template + }) + }) + .collect::>(); + templates.sort_unstable_by(|template, other| template.name.cmp(&other.name)); + templates +} + +pub(super) fn merge_prompts(prompts: Vec<(String, ListPromptsResult)>, namespace_identifiers: bool) -> Vec { + let mut prompts = prompts + .into_iter() + .flat_map(|(backend_name, result)| { + result.prompts.into_iter().map(move |mut prompt| { + if namespace_identifiers { + prompt.name = prefixed_name(&backend_name, &prompt.name); + } + prompt + }) + }) + .collect::>(); + prompts.sort_unstable_by(|prompt, other| prompt.name.cmp(&other.name)); + prompts +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn single_backend_listings_preserve_identifiers() { + let config_json = serde_json::json!({ + "backends": { + "backend-id": { + "name": "backend", + "url": "http://upstream:9000/mcp", + "transport": "STREAMABLEHTTP", + "passthrough_headers": [], + "allowed_tool_names": ["test_simple_text"], + "allowed_resource_names": [], + "allowed_prompt_names": [] + } + } + }); + let virtual_host: VirtualHost = serde_json::from_value(config_json).expect("valid virtual host"); + let tools = merge_tools( + vec![( + "backend-id".to_owned(), + ListToolsResult::with_all_items(vec![Tool::new("test_simple_text", "", serde_json::Map::new())]), + )], + &virtual_host, + ); + let prompts = merge_prompts( + vec![( + "backend-id".to_owned(), + ListPromptsResult::with_all_items(vec![Prompt::new("test_prompt", None::, None)]), + )], + false, + ); + let resources = merge_resources( + vec![( + "backend-id".to_owned(), + ListResourcesResult::with_all_items(vec![Resource::new("test://resource", "test_resource")]), + )], + false, + ); + let templates = merge_resource_templates( + vec![( + "backend-id".to_owned(), + ListResourceTemplatesResult::with_all_items(vec![ResourceTemplate::new( + "test://template/{id}/data", + "test_template", + )]), + )], + false, + ); + + assert_eq!("test_simple_text", tools[0].name); + assert_eq!("test_prompt", prompts[0].name); + assert_eq!("test_resource", resources[0].name); + assert_eq!("test://resource", resources[0].uri); + assert_eq!("test_template", templates[0].name); + assert_eq!("test://template/{id}/data", templates[0].uri_template); + } +} diff --git a/crates/contextforge-gateway-rs-lib/src/gateway/mcp_gateway.rs b/crates/contextforge-gateway-rs-lib/src/gateway/mcp_gateway.rs deleted file mode 100644 index 377ff22..0000000 --- a/crates/contextforge-gateway-rs-lib/src/gateway/mcp_gateway.rs +++ /dev/null @@ -1,979 +0,0 @@ -use std::{collections::HashMap, sync::Arc}; - -use contextforge_gateway_rs_apis::user_store::VirtualHost; -use contextforge_gateway_rs_cpex::{GatewayPluginRuntimeHandle, ToolPreCallResult}; -use rmcp::{ - ErrorData, RoleClient, RoleServer, ServerHandler, ServiceExt, - model::{ - CallToolRequestParams, CallToolResponse, CompleteRequestParams, CompleteResult, ErrorCode, - GetPromptRequestParams, GetPromptResponse, Implementation, InitializeRequestParams, InitializeResult, - ListPromptsResult, ListResourceTemplatesResult, ListResourcesResult, ListToolsResult, PaginatedRequestParams, - Prompt, ReadResourceRequestParams, ReadResourceResponse, Reference, Resource, ResourceTemplate, - ServerCapabilities, SubscribeRequestParams, Tool, UnsubscribeRequestParams, - }, - service::{RequestContext, RunningService}, - transport::{StreamableHttpClientTransport, streamable_http_client::StreamableHttpClientTransportConfig}, -}; -use tokio::sync::Mutex; -use tracing::{debug, info, warn}; -use typed_builder::TypedBuilder; - -use super::{ - backend_client::{GatewayBackendClient, call_backend_tool}, - mcp_call_validator::AuthorizedCallValidator, -}; -pub use crate::gateway::session_store::LocalUserSessionStore; -use crate::{ - SessionId, - gateway::{ - mcp_call_validator::InitializeCallValidator, - session_manager::SessionManager, - session_store::{UserSession, UserSessionStore}, - }, -}; - -#[derive(Clone, Default)] -pub struct BackendTransports(Arc>>); - -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 fn inner(&self) -> &Arc>> { - &self.0 - } -} - -#[derive(Clone, TypedBuilder)] -#[builder(field_defaults(setter(prefix = "with_")))] -pub struct McpService -where - T: UserSessionStore, -{ - #[builder(default = BackendTransports::default())] - transports: BackendTransports, - http_client: reqwest::Client, - user_session_store: T, - #[builder(default)] - plugin_runtime: Option, -} - -#[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)] -pub struct BackendTransportKey { - principal: String, - backend_name: String, - session_id: String, -} - -type McpClientService = Arc>; - -#[derive(Debug)] -pub struct ServiceHolder { - pub name: String, - pub running_service: Option, -} - -impl ServiceHolder { - pub fn new(name: String, running_service: Option) -> ServiceHolder { - Self { name, running_service } - } -} - -#[derive(Debug)] -pub struct BackendTransportService { - #[expect(dead_code, reason = "stored backend capabilities are kept with transport state for future routing")] - capabilities: Option, - pub(crate) service: Option, -} - -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, &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(), - } - } -} - -impl From<(Option, Option)> for BackendTransportService { - fn from((capabilities, service): (Option, Option)) -> Self { - Self { capabilities, service } - } -} - -impl ServerHandler for McpService -where - T: UserSessionStore + Send + Sync + 'static, -{ - async fn initialize( - &self, - request: InitializeRequestParams, - cx: RequestContext, - ) -> Result { - let call_validator = InitializeCallValidator::new(&cx); - 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(claims.sub.clone(), Arc::clone(&downstream_session_id.session_id))) - .await - { - maybe_session_mapping.unwrap_or_default() - } else { - return Err(ErrorData { - code: ErrorCode::INTERNAL_ERROR, - message: "Internal problem... session store can't be accessed".into(), - data: None, - }); - }; - - let namespace_identifiers = virtual_host.backends.len() > 1; - let tasks: Vec<_> = virtual_host - .backends - .iter() - .map(|(name, backend)| { - let client = self.http_client.clone(); - let backend_client = GatewayBackendClient::new( - name.clone(), - namespace_identifiers, - request.clone(), - self.plugin_runtime.clone(), - ); - let backend_url = backend.url.clone(); - let downstream_session_id = downstream_session_id.clone(); - - Box::pin(async move { - let mut headers = HashMap::new(); - if let Some(host) = backend_url.host_str() && backend_url.scheme() == "https"{ - let host = if let Some(port) = backend_url.port(){ - format!("{host}:{port}") - }else{ - host.to_owned() - }; - - if let Ok(value) = http::HeaderValue::from_str(&host){ - headers.insert(http::header::HOST, value); - }else{ - warn!("Really can't set the host header for {:?}",backend_url.host_str()); - } - } - - let config = StreamableHttpClientTransportConfig::with_uri(backend_url.to_string()) - .custom_headers(headers); - let transport = StreamableHttpClientTransport::with_client(client, config); - let maybe_running_service = backend_client.serve(transport).await; - if let Ok(running_service) = maybe_running_service { - info!("initialize: intialized for {downstream_session_id:?} {name:?}"); - (name, Some(running_service)) - } else { - warn!("initialize: Unable to initialize for {downstream_session_id:?} {name:?} {maybe_running_service:?}",); - (name, None) - } - }) - }).collect(); - - let initialization_results: Vec<(&String, Option>)> = - futures::future::join_all(tasks).await; - - let (capabilities, backend_services): (Vec<_>, Vec<_>) = initialization_results - .into_iter() - .map(|(name, running_service):(_,_)| { - info!("initialize: Adding transport: session_id {downstream_session_id:#?} backend {name} {running_service:?}"); - - let server_capabilities = - running_service.as_ref() - .and_then(|rs| - rs.peer() - .peer_info() - .as_ref() - .map(|pi| pi.capabilities.clone())); - ( - (name.clone(), server_capabilities.clone()), - (name.clone(), BackendTransportService::from((server_capabilities, running_service.map(Arc::new)))), - ) - }) - .unzip(); - - if self - .user_session_store - .set_session( - &UserSession::new(claims.sub.clone(), Arc::clone(&downstream_session_id.session_id)), - &session_mapping, - ) - .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.inner().lock().await; - for (name, svc) in backend_services { - transports - .entry(BackendTransportKey::from((name.as_str(), downstream_session_id.value(), claims.sub.as_str()))) - .insert_entry(svc); - } - drop(transports); - - Ok(InitializeResult::new(merge_capabilities(capabilities)) - .with_server_info(Implementation::new("rust-conformance-server", "0.1.0")) - .with_instructions("Rust MCP conformance test server")) - } - - async fn ping(&self, _cx: RequestContext) -> Result<(), ErrorData> { - Ok(()) - } - - async fn list_tools( - &self, - request: Option, - cx: RequestContext, - ) -> Result { - let mcp_call_validator = AuthorizedCallValidator::new("list_tools", &cx); - 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_transports: Vec<_> = session_manager.borrow_transports().await; - - let responses = fan_out_list( - backend_transports, - "list_tools", - |response: &ListToolsResult| response.tools.len(), - |service| { - let request = request.clone(); - async move { service.list_tools(request).await } - }, - ) - .await; - - Ok(ListToolsResult::with_all_items(merge_tools(responses, virtual_host))) - } - - async fn call_tool( - &self, - request: CallToolRequestParams, - cx: RequestContext, - ) -> Result { - let mcp_call_validator = AuthorizedCallValidator::new("call_tool", &cx); - 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(); - - let Some((backend_name, tool_name)) = resolve_tool_route(virtual_host, &request.name, &backend_names) else { - return Err(ErrorData { - code: ErrorCode::INTERNAL_ERROR, - message: "Routing problem... wrong tool name".into(), - data: None, - }); - }; - let backend_name = backend_name.to_owned(); - let tool_name = tool_name.to_owned(); - - let (service_name, service) = resolve_backend(&session_manager, "call_tool", &backend_name).await?; - - let pre_result = if let Some(plugin_runtime) = &self.plugin_runtime { - plugin_runtime.before_tool_call(&request, &tool_name, &service_name).await? - } else { - ToolPreCallResult::unchanged() - }; - let post_state = pre_result.state; - let mut routed_request = request; - pre_result.arguments.apply_to_request(&mut routed_request, &tool_name); - - let progress_token = cx.meta.get_progress_token(); - let handle = service - .service() - .start_tool_call( - service.peer(), - routed_request, - progress_token, - tool_name.clone(), - cx.peer.clone(), - post_state.clone(), - ) - .await - .map_err(|error| backend_forward_error("call_tool", &service_name, &error))?; - let backend_progress_token = handle.progress_token.clone(); - let response = call_backend_tool(handle, cx.ct.clone()).await; - service.service().stop_tracking_tool_call(&backend_progress_token).await; - - let response = response.map_err(|error| backend_forward_error("call_tool", &service_name, &error))?; - let response = match (&self.plugin_runtime, post_state) { - (Some(plugin_runtime), Some(post_state)) => { - plugin_runtime.after_tool_call(&tool_name, response, Some(post_state)).await? - }, - _ => response, - }; - info!("call_tool: backend {service_name} completed"); - Ok(response.into()) - } - - async fn list_resources( - &self, - request: Option, - cx: RequestContext, - ) -> Result { - let mcp_call_validator = AuthorizedCallValidator::new("list_resources", &cx); - let (virtual_host, session_id, claims) = mcp_call_validator.validate()?; - let namespace_identifiers = virtual_host.backends.len() > 1; - - 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 responses = fan_out_list( - backend_transports, - "list_resources", - |response: &ListResourcesResult| response.resources.len(), - |service| { - let request = request.clone(); - async move { service.list_resources(request).await } - }, - ) - .await; - - Ok(ListResourcesResult::with_all_items(merge_resources(responses, namespace_identifiers))) - } - - async fn read_resource( - &self, - request: ReadResourceRequestParams, - cx: RequestContext, - ) -> Result { - let mcp_call_validator = AuthorizedCallValidator::new("read_resource", &cx); - 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 (service_name, service, resource_uri) = route_identifier_to_backend( - &session_manager, - "read_resource", - &request.uri, - "Routing problem... wrong resource name", - ) - .await?; - - let mut routed_request = request; - routed_request.uri = resource_uri; - let response = service - .read_resource(routed_request) - .await - .map_err(|error| backend_forward_error("read_resource", &service_name, &error))?; - info!("read_resource: backend {service_name} returned {} contents", response.contents.len()); - Ok(response.into()) - } - - async fn list_resource_templates( - &self, - request: Option, - cx: RequestContext, - ) -> Result { - let mcp_call_validator = AuthorizedCallValidator::new("list_resource_templates", &cx); - let (virtual_host, session_id, claims) = mcp_call_validator.validate()?; - let namespace_identifiers = virtual_host.backends.len() > 1; - - 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 responses = fan_out_list( - backend_transports, - "list_resource_templates", - |response: &ListResourceTemplatesResult| response.resource_templates.len(), - |service| { - let request = request.clone(); - async move { service.list_resource_templates(request).await } - }, - ) - .await; - - Ok(ListResourceTemplatesResult::with_all_items(merge_resource_templates(responses, namespace_identifiers))) - } - - async fn subscribe( - &self, - request: SubscribeRequestParams, - cx: RequestContext, - ) -> Result<(), ErrorData> { - let mcp_call_validator = AuthorizedCallValidator::new("subscribe", &cx); - 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 (service_name, service, resource_uri) = route_identifier_to_backend( - &session_manager, - "subscribe", - &request.uri, - "Routing problem... wrong resource name", - ) - .await?; - - let mut routed_request = request; - routed_request.uri = resource_uri.clone(); - service.service().track_resource_subscription(&resource_uri, cx.peer.clone()).await; - - if let Err(error) = service.subscribe(routed_request).await { - service.service().stop_tracking_resource_subscription(&resource_uri).await; - return Err(backend_forward_error("subscribe", &service_name, &error)); - } - info!("subscribe: backend {service_name} completed"); - Ok(()) - } - - async fn unsubscribe( - &self, - request: UnsubscribeRequestParams, - cx: RequestContext, - ) -> Result<(), ErrorData> { - let mcp_call_validator = AuthorizedCallValidator::new("unsubscribe", &cx); - 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 (service_name, service, resource_uri) = route_identifier_to_backend( - &session_manager, - "unsubscribe", - &request.uri, - "Routing problem... wrong resource name", - ) - .await?; - - let mut routed_request = request; - routed_request.uri = resource_uri.clone(); - service - .unsubscribe(routed_request) - .await - .map_err(|error| backend_forward_error("unsubscribe", &service_name, &error))?; - service.service().stop_tracking_resource_subscription(&resource_uri).await; - info!("unsubscribe: backend {service_name} completed"); - Ok(()) - } - - async fn list_prompts( - &self, - request: Option, - cx: RequestContext, - ) -> Result { - let mcp_call_validator = AuthorizedCallValidator::new("list_prompts", &cx); - let (virtual_host, session_id, claims) = mcp_call_validator.validate()?; - let namespace_identifiers = virtual_host.backends.len() > 1; - - 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 responses = fan_out_list( - backend_transports, - "list_prompts", - |response: &ListPromptsResult| response.prompts.len(), - |service| { - let request = request.clone(); - async move { service.list_prompts(request).await } - }, - ) - .await; - - Ok(ListPromptsResult::with_all_items(merge_prompts(responses, namespace_identifiers))) - } - - async fn get_prompt( - &self, - request: GetPromptRequestParams, - cx: RequestContext, - ) -> Result { - let mcp_call_validator = AuthorizedCallValidator::new("get_prompt", &cx); - 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 (service_name, service, prompt_name) = route_identifier_to_backend( - &session_manager, - "get_prompt", - &request.name, - "Routing problem... wrong prompt name", - ) - .await?; - - let mut routed_request = request; - routed_request.name = prompt_name; - let response = service - .get_prompt(routed_request) - .await - .map_err(|error| backend_forward_error("get_prompt", &service_name, &error))?; - info!("get_prompt: backend {service_name} returned {} messages", response.messages.len()); - Ok(response.into()) - } - - async fn complete( - &self, - request: CompleteRequestParams, - cx: RequestContext, - ) -> Result { - let mcp_call_validator = AuthorizedCallValidator::new("complete", &cx); - 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 identifier = match &request.r#ref { - Reference::Prompt(prompt) => prompt.name.as_str(), - Reference::Resource(resource) => resource.uri.as_str(), - _ => return Err(ErrorData::invalid_params("Unsupported completion reference", None)), - }; - - let (service_name, service, routed_identifier) = route_identifier_to_backend( - &session_manager, - "complete", - identifier, - "Routing problem... wrong completion reference", - ) - .await?; - - let mut routed_request = request; - match &mut routed_request.r#ref { - Reference::Prompt(prompt) => prompt.name = routed_identifier, - Reference::Resource(resource) => resource.uri = routed_identifier, - _ => return Err(ErrorData::invalid_params("Unsupported completion reference", None)), - } - let response = service - .complete(routed_request) - .await - .map_err(|error| backend_forward_error("complete", &service_name, &error))?; - info!("complete: backend {service_name} returned {} values", response.completion.values.len()); - Ok(response) - } -} - -/// Preserves identifiers for a single backend. For multiple backends, splits a -/// `{backend}-{identifier}` namespace so duplicate identifiers remain routable. -fn route_identifier<'a, N: AsRef>(identifier: &'a str, backend_names: &'a [N]) -> Option<(&'a str, &'a str)> { - if let [backend] = backend_names { - return Some((backend.as_ref(), identifier)); - } - - backend_names.iter().find_map(|backend| { - let backend = backend.as_ref(); - identifier.strip_prefix(backend)?.strip_prefix('-').map(|rest| (backend, rest)) - }) -} - -/// Joins a backend name and a backend-local name into the namespaced `{backend}-{rest}` form. -pub(crate) fn prefixed_name(backend_name: &str, rest: &str) -> String { - format!("{backend_name}-{rest}") -} - -/// Resolves an exact control-plane alias to its backend and upstream name. Without an alias, -/// single-backend hosts preserve the upstream name and multi-backend hosts use the legacy prefix. -fn resolve_tool_route<'a, N: AsRef>( - virtual_host: &'a VirtualHost, - name: &'a str, - backend_names: &'a [N], -) -> Option<(&'a str, &'a str)> { - let mut aliases = backend_names.iter().filter_map(|backend_name| { - let backend_name = backend_name.as_ref(); - let original_name = virtual_host.backends.get(backend_name)?.tool_name_aliases.get(name)?; - Some((backend_name, original_name.as_str())) - }); - let alias = aliases.next(); - if aliases.next().is_some() { - return None; - } - alias.or_else(|| route_identifier(name, backend_names)) -} - -/// Returns the control-plane alias for an upstream tool when configured. Without an alias, -/// single-backend hosts preserve the upstream name and multi-backend hosts use the legacy prefix. -fn exposed_tool_name(virtual_host: &VirtualHost, backend_name: &str, original_name: &str) -> String { - virtual_host - .backends - .get(backend_name) - .and_then(|backend| { - backend - .tool_name_aliases - .iter() - .find_map(|(alias, original)| (original == original_name).then(|| alias.clone())) - }) - .unwrap_or_else(|| { - if virtual_host.backends.len() == 1 { - original_name.to_owned() - } else { - prefixed_name(backend_name, original_name) - } - }) -} - -/// Logs a backend forwarding failure and maps it to the routing error every handler returns. -fn backend_forward_error(op: &str, backend_name: &str, error: &impl std::fmt::Debug) -> ErrorData { - warn!("{op}: backend {backend_name} {error:?}"); - ErrorData { - code: ErrorCode::INTERNAL_ERROR, - message: "Routing problem... got no responses from backends".into(), - data: None, - } -} - -/// Fans a paginated list request out to every connected backend concurrently, logs each response, -/// and returns the `(backend_name, result)` pairs that succeeded. -async fn fan_out_list( - backends: Vec, - op: &str, - item_count: C, - call: F, -) -> Vec<(String, R)> -where - F: Fn(McpClientService) -> Fut, - Fut: std::future::Future>, - C: Fn(&R) -> usize, - E: std::fmt::Debug, -{ - let tasks = backends.into_iter().map(|service_holder| { - let call = &call; - async move { - let response = match service_holder.running_service { - Some(service) => Some(call(service).await), - None => None, - }; - (service_holder.name, response) - } - }); - - futures::future::join_all(tasks) - .await - .into_iter() - .filter_map(|(name, response)| { - log_list_backend_response(op, &name, response.as_ref(), &item_count); - match response { - Some(Ok(response)) => Some((name, response)), - _ => None, - } - }) - .collect() -} - -/// Routes an identifier to its backend, preserving it for a single backend and splitting the -/// namespace for multiple backends. Returns `(backend_name, service, backend_local_identifier)`. -async fn route_identifier_to_backend( - session_manager: &SessionManager<'_>, - op: &str, - identifier: &str, - no_route_message: &'static str, -) -> Result<(String, McpClientService, String), ErrorData> { - let backend_names = session_manager.get_backend_names(); - let Some((backend_name, routed_identifier)) = route_identifier(identifier, &backend_names) else { - return Err(ErrorData { code: ErrorCode::INTERNAL_ERROR, message: no_route_message.into(), data: None }); - }; - let routed_identifier = routed_identifier.to_owned(); - let (backend_name, service) = resolve_backend(session_manager, op, backend_name).await?; - Ok((backend_name, service, routed_identifier)) -} - -/// Resolves the single connected backend named `backend_name` and takes its running service. -/// Shared by tool, resource, and prompt routing so they reject duplicate or missing backends -/// the same way; a duplicate match means the session is invalid, so it is cleaned up. -async fn resolve_backend( - session_manager: &SessionManager<'_>, - op: &str, - backend_name: &str, -) -> Result<(String, McpClientService), ErrorData> { - let backend_transports = session_manager.borrow_transports().await; - debug!("{op}: resolving backend {backend_name} from {backend_transports:?}"); - - let mut target = None; - for service_holder in backend_transports { - if service_holder.name == backend_name { - if target.is_some() { - warn!("{op}: more than one backend matching {backend_name}"); - session_manager.cleanup_backends("invalid session.. duplicate backends detected").await; - return Err(ErrorData { - code: ErrorCode::INVALID_REQUEST, - message: "Routing problem... multiple matching backends".into(), - data: None, - }); - } - target = Some(service_holder); - } - } - - let Some(ServiceHolder { name, running_service }) = target else { - return Err(ErrorData { - code: ErrorCode::INTERNAL_ERROR, - message: "Routing problem... got no responses from backends".into(), - data: None, - }); - }; - let Some(service) = running_service else { - warn!("{op}: no running backend for {backend_name}"); - return Err(ErrorData { - code: ErrorCode::INTERNAL_ERROR, - message: "Routing problem... got no responses from backends".into(), - data: None, - }); - }; - Ok((name, service)) -} - -fn merge_capabilities(_server_capabilities: Vec<(String, Option)>) -> ServerCapabilities { - ServerCapabilities::builder() - .enable_completions() - .enable_prompts() - .enable_resources() - .enable_resources_subscribe() - .enable_tools() - .build() -} - -fn log_list_backend_response( - kind: &str, - name: &str, - response: Option<&Result>, - item_count: impl Fn(&T) -> usize, -) { - match response { - Some(Ok(response)) => info!("{kind}: backend {name} completed ({} items)", item_count(response)), - Some(Err(error)) => warn!("{kind}: backend {name} {error:?}"), - None => info!("{kind}: backend {name} unavailable"), - } -} - -fn merge_tools(tools: Vec<(String, ListToolsResult)>, virtual_host: &VirtualHost) -> Vec { - let mut tools = tools - .into_iter() - .flat_map(|(backend_name, result)| { - result - .tools - .into_iter() - .map(|mut t| { - t.name = exposed_tool_name(virtual_host, &backend_name, &t.name).into(); - t - }) - .collect::>() - }) - .collect::>(); - tools.sort_unstable_by(|tool, other| tool.name.cmp(&other.name)); - tools -} - -fn merge_resources(resources: Vec<(String, ListResourcesResult)>, namespace_identifiers: bool) -> Vec { - let mut resources = resources - .into_iter() - .flat_map(|(backend_name, result)| { - result - .resources - .into_iter() - .map(|mut t| { - if namespace_identifiers { - t.name = prefixed_name(&backend_name, &t.name); - t.uri = prefixed_name(&backend_name, &t.uri); - } - t - }) - .collect::>() - }) - .collect::>(); - resources.sort_unstable_by(|resource, other| resource.name.cmp(&other.name)); - resources -} - -fn merge_resource_templates( - templates: Vec<(String, ListResourceTemplatesResult)>, - namespace_identifiers: bool, -) -> Vec { - let mut templates = templates - .into_iter() - .flat_map(|(backend_name, result)| { - result.resource_templates.into_iter().map(move |mut template| { - if namespace_identifiers { - template.name = prefixed_name(&backend_name, &template.name); - template.uri_template = prefixed_name(&backend_name, &template.uri_template); - } - template - }) - }) - .collect::>(); - templates.sort_unstable_by(|template, other| template.name.cmp(&other.name)); - templates -} - -fn merge_prompts(prompts: Vec<(String, ListPromptsResult)>, namespace_identifiers: bool) -> Vec { - let mut prompts = prompts - .into_iter() - .flat_map(|(backend_name, result)| { - result.prompts.into_iter().map(move |mut prompt| { - if namespace_identifiers { - prompt.name = prefixed_name(&backend_name, &prompt.name); - } - prompt - }) - }) - .collect::>(); - prompts.sort_unstable_by(|prompt, other| prompt.name.cmp(&other.name)); - prompts -} - -#[cfg(test)] -mod tests { - // Note this useful idiom: importing names from outer (for mod tests) scope. - use super::*; - - #[test] - fn test_splitting() { - let backend_names = vec!["counter-on", "counter-oneee", "counter-one"]; - assert_eq!(Some(("counter-one", "increment")), route_identifier("counter-one-increment", &backend_names)); - assert_eq!(None, route_identifier("counter-oneincrement", &backend_names)); - assert_eq!(None, route_identifier("counteroneincrement", &backend_names)); - assert_eq!(Some(("counter-one", "get-value")), route_identifier("counter-one-get-value", &backend_names)); - - // Tool, resource, and prompt routing all share this splitter. - assert_eq!( - Some(("counter-one", "example-prompt")), - route_identifier("counter-one-example-prompt", &backend_names) - ); - assert_eq!(None, route_identifier("counter-oneexample-prompt", &backend_names)); - - let backend_names = vec!["counter_on", "counter_oneee", "counter_one"]; - assert_eq!(Some(("counter_one", "get-value")), route_identifier("counter_one-get-value", &backend_names)); - } - - #[test] - fn single_backend_routes_unprefixed_identifier_unchanged() { - let backend_names = vec!["backend-id"]; - - assert_eq!(Some(("backend-id", "test_simple_text")), route_identifier("test_simple_text", &backend_names)); - assert_eq!(Some(("backend-id", "backend-id-tool")), route_identifier("backend-id-tool", &backend_names)); - assert_eq!( - Some(("backend-id", "test://template/123/data")), - route_identifier("test://template/123/data", &backend_names) - ); - } - - #[test] - fn single_backend_listings_preserve_identifiers() { - let config_json = serde_json::json!({ - "backends": { - "backend-id": { - "name": "backend", - "url": "http://upstream:9000/mcp", - "transport": "STREAMABLEHTTP", - "passthrough_headers": [], - "allowed_tool_names": ["test_simple_text"], - "allowed_resource_names": [], - "allowed_prompt_names": [] - } - } - }); - let virtual_host: VirtualHost = serde_json::from_value(config_json).expect("valid virtual host"); - let tools = merge_tools( - vec![( - "backend-id".to_owned(), - ListToolsResult::with_all_items(vec![Tool::new("test_simple_text", "", serde_json::Map::new())]), - )], - &virtual_host, - ); - let prompts = merge_prompts( - vec![( - "backend-id".to_owned(), - ListPromptsResult::with_all_items(vec![Prompt::new("test_prompt", None::, None)]), - )], - false, - ); - let resources = merge_resources( - vec![( - "backend-id".to_owned(), - ListResourcesResult::with_all_items(vec![Resource::new("test://resource", "test_resource")]), - )], - false, - ); - let templates = merge_resource_templates( - vec![( - "backend-id".to_owned(), - ListResourceTemplatesResult::with_all_items(vec![ResourceTemplate::new( - "test://template/{id}/data", - "test_template", - )]), - )], - false, - ); - - assert_eq!("test_simple_text", tools[0].name); - assert_eq!("test_prompt", prompts[0].name); - assert_eq!("test_resource", resources[0].name); - assert_eq!("test://resource", resources[0].uri); - assert_eq!("test_template", templates[0].name); - assert_eq!("test://template/{id}/data", templates[0].uri_template); - } - - #[test] - fn test_control_plane_alias_is_advertised_and_routes_to_original_name() { - let config_json = serde_json::json!({ - "backends": { - "79fabb70-2188-4de8-95ed-dc1e976e14d4": { - "name": "compliance_reference", - "url": "http://upstream:9000/mcp", - "transport": "STREAMABLEHTTP", - "passthrough_headers": [], - "allowed_tool_names": ["get_stats", "echo"], - "tool_name_aliases": { - "Public.Tool": "get_stats", - "Echo_Tool": "echo" - }, - "allowed_resource_names": [], - "allowed_prompt_names": [] - } - } - }); - let virtual_host: VirtualHost = serde_json::from_value(config_json).expect("valid virtual host"); - let backend_ids = vec!["79fabb70-2188-4de8-95ed-dc1e976e14d4"]; - - assert_eq!( - "Public.Tool", - exposed_tool_name(&virtual_host, "79fabb70-2188-4de8-95ed-dc1e976e14d4", "get_stats") - ); - assert_eq!( - Some(("79fabb70-2188-4de8-95ed-dc1e976e14d4", "get_stats")), - resolve_tool_route(&virtual_host, "Public.Tool", &backend_ids) - ); - } - - #[test] - fn multi_backend_tool_routing_falls_back_to_legacy_prefixed_names() { - let config_json = serde_json::json!({ - "backends": { - "compliance-reference": { - "name": "compliance_reference", - "url": "http://upstream:9000/mcp", - "transport": "STREAMABLEHTTP", - "passthrough_headers": [], - "allowed_tool_names": ["get_stats"], - "allowed_resource_names": [], - "allowed_prompt_names": [] - }, - "other": { - "name": "other", - "url": "http://other:9000/mcp", - "transport": "STREAMABLEHTTP", - "passthrough_headers": [], - "allowed_tool_names": [], - "allowed_resource_names": [], - "allowed_prompt_names": [] - } - } - }); - let virtual_host: VirtualHost = serde_json::from_value(config_json).expect("valid virtual host"); - let backend_names = vec!["compliance-reference", "other"]; - - assert_eq!( - "compliance-reference-get_stats", - exposed_tool_name(&virtual_host, "compliance-reference", "get_stats") - ); - assert_eq!( - Some(("compliance-reference", "get_stats")), - resolve_tool_route(&virtual_host, "compliance-reference-get_stats", &backend_names) - ); - } -} diff --git a/crates/contextforge-gateway-rs-lib/src/gateway/mcp_service.rs b/crates/contextforge-gateway-rs-lib/src/gateway/mcp_service.rs new file mode 100644 index 0000000..564da26 --- /dev/null +++ b/crates/contextforge-gateway-rs-lib/src/gateway/mcp_service.rs @@ -0,0 +1,131 @@ +mod completion; +mod initialization; +mod prompts; +mod resources; +mod tools; + +use contextforge_gateway_rs_cpex::GatewayPluginRuntimeHandle; +use rmcp::{ + ErrorData, RoleServer, ServerHandler, + model::{ + CallToolRequestParams, CallToolResponse, CompleteRequestParams, CompleteResult, GetPromptRequestParams, + GetPromptResponse, InitializeRequestParams, InitializeResult, ListPromptsResult, ListResourceTemplatesResult, + ListResourcesResult, ListToolsResult, PaginatedRequestParams, ReadResourceRequestParams, ReadResourceResponse, + SubscribeRequestParams, UnsubscribeRequestParams, + }, + service::RequestContext, +}; +use typed_builder::TypedBuilder; + +use super::{backend_transports::BackendTransports, session_store::UserSessionStore}; + +#[derive(Clone, TypedBuilder)] +#[builder(field_defaults(setter(prefix = "with_")))] +pub struct McpService +where + T: UserSessionStore, +{ + #[builder(default = BackendTransports::default())] + transports: BackendTransports, + http_client: reqwest::Client, + user_session_store: T, + #[builder(default)] + plugin_runtime: Option, +} + +impl ServerHandler for McpService +where + T: UserSessionStore + Send + Sync + 'static, +{ + async fn initialize( + &self, + request: InitializeRequestParams, + cx: RequestContext, + ) -> Result { + initialization::initialize(self, request, cx).await + } + + async fn ping(&self, _cx: RequestContext) -> Result<(), ErrorData> { + Ok(()) + } + + async fn list_tools( + &self, + request: Option, + cx: RequestContext, + ) -> Result { + tools::list_tools(self, request, cx).await + } + + async fn call_tool( + &self, + request: CallToolRequestParams, + cx: RequestContext, + ) -> Result { + tools::call_tool(self, request, cx).await + } + + async fn list_resources( + &self, + request: Option, + cx: RequestContext, + ) -> Result { + resources::list_resources(self, request, cx).await + } + + async fn read_resource( + &self, + request: ReadResourceRequestParams, + cx: RequestContext, + ) -> Result { + resources::read_resource(self, request, cx).await + } + + async fn list_resource_templates( + &self, + request: Option, + cx: RequestContext, + ) -> Result { + resources::list_resource_templates(self, request, cx).await + } + + async fn subscribe( + &self, + request: SubscribeRequestParams, + cx: RequestContext, + ) -> Result<(), ErrorData> { + resources::subscribe(self, request, cx).await + } + + async fn unsubscribe( + &self, + request: UnsubscribeRequestParams, + cx: RequestContext, + ) -> Result<(), ErrorData> { + resources::unsubscribe(self, request, cx).await + } + + async fn list_prompts( + &self, + request: Option, + cx: RequestContext, + ) -> Result { + prompts::list_prompts(self, request, cx).await + } + + async fn get_prompt( + &self, + request: GetPromptRequestParams, + cx: RequestContext, + ) -> Result { + prompts::get_prompt(self, request, cx).await + } + + async fn complete( + &self, + request: CompleteRequestParams, + cx: RequestContext, + ) -> Result { + completion::complete(self, request, cx).await + } +} diff --git a/crates/contextforge-gateway-rs-lib/src/gateway/mcp_service/completion.rs b/crates/contextforge-gateway-rs-lib/src/gateway/mcp_service/completion.rs new file mode 100644 index 0000000..7fe5211 --- /dev/null +++ b/crates/contextforge-gateway-rs-lib/src/gateway/mcp_service/completion.rs @@ -0,0 +1,54 @@ +use rmcp::{ + ErrorData, RoleServer, + model::{CompleteRequestParams, CompleteResult, Reference}, + service::RequestContext, +}; +use tracing::info; + +use super::McpService; +use crate::gateway::{ + identifier_routing::{backend_forward_error, route_identifier_to_backend}, + mcp_call_validator::AuthorizedCallValidator, + session_manager::SessionManager, + session_store::UserSessionStore, +}; + +pub(super) async fn complete( + mcp_service: &McpService, + request: CompleteRequestParams, + cx: RequestContext, +) -> Result +where + T: UserSessionStore + Send + Sync + 'static, +{ + let mcp_call_validator = AuthorizedCallValidator::new("complete", &cx); + let (virtual_host, session_id, claims) = mcp_call_validator.validate()?; + let session_manager = SessionManager::new(virtual_host, session_id, claims.sub.as_str(), &mcp_service.transports); + + let identifier = match &request.r#ref { + Reference::Prompt(prompt) => prompt.name.as_str(), + Reference::Resource(resource) => resource.uri.as_str(), + _ => return Err(ErrorData::invalid_params("Unsupported completion reference", None)), + }; + + let (service_name, service, routed_identifier) = route_identifier_to_backend( + &session_manager, + "complete", + identifier, + "Routing problem... wrong completion reference", + ) + .await?; + + let mut routed_request = request; + match &mut routed_request.r#ref { + Reference::Prompt(prompt) => prompt.name = routed_identifier, + Reference::Resource(resource) => resource.uri = routed_identifier, + _ => return Err(ErrorData::invalid_params("Unsupported completion reference", None)), + } + let response = service + .complete(routed_request) + .await + .map_err(|error| backend_forward_error("complete", &service_name, &error))?; + info!("complete: backend {service_name} returned {} values", response.completion.values.len()); + Ok(response) +} diff --git a/crates/contextforge-gateway-rs-lib/src/gateway/mcp_service/initialization.rs b/crates/contextforge-gateway-rs-lib/src/gateway/mcp_service/initialization.rs new file mode 100644 index 0000000..792392f --- /dev/null +++ b/crates/contextforge-gateway-rs-lib/src/gateway/mcp_service/initialization.rs @@ -0,0 +1,150 @@ +use std::{collections::HashMap, sync::Arc}; + +use rmcp::{ + ErrorData, RoleClient, RoleServer, ServiceExt, + model::{ErrorCode, Implementation, InitializeRequestParams, InitializeResult, ServerCapabilities}, + service::{RequestContext, RunningService}, + transport::{StreamableHttpClientTransport, streamable_http_client::StreamableHttpClientTransportConfig}, +}; +use tracing::{info, warn}; + +use super::McpService; +use crate::gateway::{ + backend_client::GatewayBackendClient, + backend_transports::{BackendTransportKey, BackendTransportService}, + mcp_call_validator::InitializeCallValidator, + session_store::{UserSession, UserSessionStore}, +}; + +pub(super) async fn initialize( + mcp_service: &McpService, + request: InitializeRequestParams, + cx: RequestContext, +) -> Result +where + T: UserSessionStore + Send + Sync + 'static, +{ + let call_validator = InitializeCallValidator::new(&cx); + let (virtual_host, downstream_session_id, claims) = call_validator.validate()?; + let session_mapping = if let Ok(maybe_session_mapping) = mcp_service + .user_session_store + .get_session(&UserSession::new(claims.sub.clone(), Arc::clone(&downstream_session_id.session_id))) + .await + { + maybe_session_mapping.unwrap_or_default() + } else { + return Err(ErrorData { + code: ErrorCode::INTERNAL_ERROR, + message: "Internal problem... session store can't be accessed".into(), + data: None, + }); + }; + + let namespace_identifiers = virtual_host.backends.len() > 1; + let tasks: Vec<_> = virtual_host + .backends + .iter() + .map(|(name, backend)| { + let client = mcp_service.http_client.clone(); + let backend_client = GatewayBackendClient::new( + name.clone(), + namespace_identifiers, + request.clone(), + mcp_service.plugin_runtime.clone(), + ); + let backend_url = backend.url.clone(); + let downstream_session_id = downstream_session_id.clone(); + + Box::pin(async move { + let mut headers = HashMap::new(); + if let Some(host) = backend_url.host_str() + && backend_url.scheme() == "https" + { + let host = if let Some(port) = backend_url.port() { + format!("{host}:{port}") + } else { + host.to_owned() + }; + + if let Ok(value) = http::HeaderValue::from_str(&host) { + headers.insert(http::header::HOST, value); + } else { + warn!("Really can't set the host header for {:?}", backend_url.host_str()); + } + } + + let config = + StreamableHttpClientTransportConfig::with_uri(backend_url.to_string()).custom_headers(headers); + let transport = StreamableHttpClientTransport::with_client(client, config); + let maybe_running_service = backend_client.serve(transport).await; + if let Ok(running_service) = maybe_running_service { + info!("initialize: intialized for {downstream_session_id:?} {name:?}"); + (name, Some(running_service)) + } else { + warn!( + "initialize: Unable to initialize for {downstream_session_id:?} {name:?} {maybe_running_service:?}", + ); + (name, None) + } + }) + }) + .collect(); + + let initialization_results: Vec<(&String, Option>)> = + futures::future::join_all(tasks).await; + + let (capabilities, backend_services): (Vec<_>, Vec<_>) = initialization_results + .into_iter() + .map(|(name, running_service): (_, _)| { + info!( + "initialize: Adding transport: session_id {downstream_session_id:#?} backend {name} {running_service:?}" + ); + + let server_capabilities = running_service + .as_ref() + .and_then(|rs| rs.peer().peer_info().as_ref().map(|pi| pi.capabilities.clone())); + ( + (name.clone(), server_capabilities.clone()), + (name.clone(), BackendTransportService::from((server_capabilities, running_service.map(Arc::new)))), + ) + }) + .unzip(); + + if mcp_service + .user_session_store + .set_session( + &UserSession::new(claims.sub.clone(), Arc::clone(&downstream_session_id.session_id)), + &session_mapping, + ) + .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 = mcp_service.transports.inner().lock().await; + for (name, service) in backend_services { + transports + .entry(BackendTransportKey::from((name.as_str(), downstream_session_id.value(), claims.sub.as_str()))) + .insert_entry(service); + } + drop(transports); + + Ok(InitializeResult::new(merge_capabilities(capabilities)) + .with_server_info(Implementation::new("rust-conformance-server", "0.1.0")) + .with_instructions("Rust MCP conformance test server")) +} + +fn merge_capabilities(_server_capabilities: Vec<(String, Option)>) -> ServerCapabilities { + ServerCapabilities::builder() + .enable_completions() + .enable_prompts() + .enable_resources() + .enable_resources_subscribe() + .enable_tools() + .build() +} diff --git a/crates/contextforge-gateway-rs-lib/src/gateway/mcp_service/prompts.rs b/crates/contextforge-gateway-rs-lib/src/gateway/mcp_service/prompts.rs new file mode 100644 index 0000000..2e620d9 --- /dev/null +++ b/crates/contextforge-gateway-rs-lib/src/gateway/mcp_service/prompts.rs @@ -0,0 +1,74 @@ +use rmcp::{ + ErrorData, RoleServer, + model::{GetPromptRequestParams, GetPromptResponse, ListPromptsResult, PaginatedRequestParams}, + service::RequestContext, +}; +use tracing::info; + +use super::McpService; +use crate::gateway::{ + identifier_routing::{backend_forward_error, route_identifier_to_backend}, + list_aggregation::{fan_out_list, merge_prompts}, + mcp_call_validator::AuthorizedCallValidator, + session_manager::SessionManager, + session_store::UserSessionStore, +}; + +pub(super) async fn list_prompts( + mcp_service: &McpService, + request: Option, + cx: RequestContext, +) -> Result +where + T: UserSessionStore + Send + Sync + 'static, +{ + let mcp_call_validator = AuthorizedCallValidator::new("list_prompts", &cx); + let (virtual_host, session_id, claims) = mcp_call_validator.validate()?; + let namespace_identifiers = virtual_host.backends.len() > 1; + + let session_manager = SessionManager::new(virtual_host, session_id, claims.sub.as_str(), &mcp_service.transports); + let backend_transports: Vec<_> = session_manager.borrow_transports().await; + + let responses = fan_out_list( + backend_transports, + "list_prompts", + |response: &ListPromptsResult| response.prompts.len(), + |service| { + let request = request.clone(); + async move { service.list_prompts(request).await } + }, + ) + .await; + + Ok(ListPromptsResult::with_all_items(merge_prompts(responses, namespace_identifiers))) +} + +pub(super) async fn get_prompt( + mcp_service: &McpService, + request: GetPromptRequestParams, + cx: RequestContext, +) -> Result +where + T: UserSessionStore + Send + Sync + 'static, +{ + let mcp_call_validator = AuthorizedCallValidator::new("get_prompt", &cx); + let (virtual_host, session_id, claims) = mcp_call_validator.validate()?; + let session_manager = SessionManager::new(virtual_host, session_id, claims.sub.as_str(), &mcp_service.transports); + + let (service_name, service, prompt_name) = route_identifier_to_backend( + &session_manager, + "get_prompt", + &request.name, + "Routing problem... wrong prompt name", + ) + .await?; + + let mut routed_request = request; + routed_request.name = prompt_name; + let response = service + .get_prompt(routed_request) + .await + .map_err(|error| backend_forward_error("get_prompt", &service_name, &error))?; + info!("get_prompt: backend {service_name} returned {} messages", response.messages.len()); + Ok(response.into()) +} diff --git a/crates/contextforge-gateway-rs-lib/src/gateway/mcp_service/resources.rs b/crates/contextforge-gateway-rs-lib/src/gateway/mcp_service/resources.rs new file mode 100644 index 0000000..ae37b3a --- /dev/null +++ b/crates/contextforge-gateway-rs-lib/src/gateway/mcp_service/resources.rs @@ -0,0 +1,169 @@ +use rmcp::{ + ErrorData, RoleServer, + model::{ + ListResourceTemplatesResult, ListResourcesResult, PaginatedRequestParams, ReadResourceRequestParams, + ReadResourceResponse, SubscribeRequestParams, UnsubscribeRequestParams, + }, + service::RequestContext, +}; +use tracing::info; + +use super::McpService; +use crate::gateway::{ + identifier_routing::{backend_forward_error, route_identifier_to_backend}, + list_aggregation::{fan_out_list, merge_resource_templates, merge_resources}, + mcp_call_validator::AuthorizedCallValidator, + session_manager::SessionManager, + session_store::UserSessionStore, +}; + +pub(super) async fn list_resources( + mcp_service: &McpService, + request: Option, + cx: RequestContext, +) -> Result +where + T: UserSessionStore + Send + Sync + 'static, +{ + let mcp_call_validator = AuthorizedCallValidator::new("list_resources", &cx); + let (virtual_host, session_id, claims) = mcp_call_validator.validate()?; + let namespace_identifiers = virtual_host.backends.len() > 1; + + let session_manager = SessionManager::new(virtual_host, session_id, claims.sub.as_str(), &mcp_service.transports); + let backend_transports: Vec<_> = session_manager.borrow_transports().await; + + let responses = fan_out_list( + backend_transports, + "list_resources", + |response: &ListResourcesResult| response.resources.len(), + |service| { + let request = request.clone(); + async move { service.list_resources(request).await } + }, + ) + .await; + + Ok(ListResourcesResult::with_all_items(merge_resources(responses, namespace_identifiers))) +} + +pub(super) async fn read_resource( + mcp_service: &McpService, + request: ReadResourceRequestParams, + cx: RequestContext, +) -> Result +where + T: UserSessionStore + Send + Sync + 'static, +{ + let mcp_call_validator = AuthorizedCallValidator::new("read_resource", &cx); + let (virtual_host, session_id, claims) = mcp_call_validator.validate()?; + let session_manager = SessionManager::new(virtual_host, session_id, claims.sub.as_str(), &mcp_service.transports); + + let (service_name, service, resource_uri) = route_identifier_to_backend( + &session_manager, + "read_resource", + &request.uri, + "Routing problem... wrong resource name", + ) + .await?; + + let mut routed_request = request; + routed_request.uri = resource_uri; + let response = service + .read_resource(routed_request) + .await + .map_err(|error| backend_forward_error("read_resource", &service_name, &error))?; + info!("read_resource: backend {service_name} returned {} contents", response.contents.len()); + Ok(response.into()) +} + +pub(super) async fn list_resource_templates( + mcp_service: &McpService, + request: Option, + cx: RequestContext, +) -> Result +where + T: UserSessionStore + Send + Sync + 'static, +{ + let mcp_call_validator = AuthorizedCallValidator::new("list_resource_templates", &cx); + let (virtual_host, session_id, claims) = mcp_call_validator.validate()?; + let namespace_identifiers = virtual_host.backends.len() > 1; + + let session_manager = SessionManager::new(virtual_host, session_id, claims.sub.as_str(), &mcp_service.transports); + let backend_transports: Vec<_> = session_manager.borrow_transports().await; + + let responses = fan_out_list( + backend_transports, + "list_resource_templates", + |response: &ListResourceTemplatesResult| response.resource_templates.len(), + |service| { + let request = request.clone(); + async move { service.list_resource_templates(request).await } + }, + ) + .await; + + Ok(ListResourceTemplatesResult::with_all_items(merge_resource_templates(responses, namespace_identifiers))) +} + +pub(super) async fn subscribe( + mcp_service: &McpService, + request: SubscribeRequestParams, + cx: RequestContext, +) -> Result<(), ErrorData> +where + T: UserSessionStore + Send + Sync + 'static, +{ + let mcp_call_validator = AuthorizedCallValidator::new("subscribe", &cx); + let (virtual_host, session_id, claims) = mcp_call_validator.validate()?; + let session_manager = SessionManager::new(virtual_host, session_id, claims.sub.as_str(), &mcp_service.transports); + + let (service_name, service, resource_uri) = route_identifier_to_backend( + &session_manager, + "subscribe", + &request.uri, + "Routing problem... wrong resource name", + ) + .await?; + + let mut routed_request = request; + routed_request.uri = resource_uri.clone(); + service.service().track_resource_subscription(&resource_uri, cx.peer.clone()).await; + + if let Err(error) = service.subscribe(routed_request).await { + service.service().stop_tracking_resource_subscription(&resource_uri).await; + return Err(backend_forward_error("subscribe", &service_name, &error)); + } + info!("subscribe: backend {service_name} completed"); + Ok(()) +} + +pub(super) async fn unsubscribe( + mcp_service: &McpService, + request: UnsubscribeRequestParams, + cx: RequestContext, +) -> Result<(), ErrorData> +where + T: UserSessionStore + Send + Sync + 'static, +{ + let mcp_call_validator = AuthorizedCallValidator::new("unsubscribe", &cx); + let (virtual_host, session_id, claims) = mcp_call_validator.validate()?; + let session_manager = SessionManager::new(virtual_host, session_id, claims.sub.as_str(), &mcp_service.transports); + + let (service_name, service, resource_uri) = route_identifier_to_backend( + &session_manager, + "unsubscribe", + &request.uri, + "Routing problem... wrong resource name", + ) + .await?; + + let mut routed_request = request; + routed_request.uri = resource_uri.clone(); + service + .unsubscribe(routed_request) + .await + .map_err(|error| backend_forward_error("unsubscribe", &service_name, &error))?; + service.service().stop_tracking_resource_subscription(&resource_uri).await; + info!("unsubscribe: backend {service_name} completed"); + Ok(()) +} diff --git a/crates/contextforge-gateway-rs-lib/src/gateway/mcp_service/tools.rs b/crates/contextforge-gateway-rs-lib/src/gateway/mcp_service/tools.rs new file mode 100644 index 0000000..05b2e38 --- /dev/null +++ b/crates/contextforge-gateway-rs-lib/src/gateway/mcp_service/tools.rs @@ -0,0 +1,107 @@ +use contextforge_gateway_rs_cpex::ToolPreCallResult; +use rmcp::{ + ErrorData, RoleServer, + model::{CallToolRequestParams, CallToolResponse, ErrorCode, ListToolsResult, PaginatedRequestParams}, + service::RequestContext, +}; +use tracing::info; + +use super::McpService; +use crate::gateway::{ + backend_client::call_backend_tool, + identifier_routing::{backend_forward_error, resolve_backend, resolve_tool_route}, + list_aggregation::{fan_out_list, merge_tools}, + mcp_call_validator::AuthorizedCallValidator, + session_manager::SessionManager, + session_store::UserSessionStore, +}; + +pub(super) async fn list_tools( + mcp_service: &McpService, + request: Option, + cx: RequestContext, +) -> Result +where + T: UserSessionStore + Send + Sync + 'static, +{ + let mcp_call_validator = AuthorizedCallValidator::new("list_tools", &cx); + let (virtual_host, session_id, claims) = mcp_call_validator.validate()?; + let session_manager = SessionManager::new(virtual_host, session_id, claims.sub.as_str(), &mcp_service.transports); + let backend_transports: Vec<_> = session_manager.borrow_transports().await; + + let responses = fan_out_list( + backend_transports, + "list_tools", + |response: &ListToolsResult| response.tools.len(), + |service| { + let request = request.clone(); + async move { service.list_tools(request).await } + }, + ) + .await; + + Ok(ListToolsResult::with_all_items(merge_tools(responses, virtual_host))) +} + +pub(super) async fn call_tool( + mcp_service: &McpService, + request: CallToolRequestParams, + cx: RequestContext, +) -> Result +where + T: UserSessionStore + Send + Sync + 'static, +{ + let mcp_call_validator = AuthorizedCallValidator::new("call_tool", &cx); + let (virtual_host, session_id, claims) = mcp_call_validator.validate()?; + let session_manager = SessionManager::new(virtual_host, session_id, claims.sub.as_str(), &mcp_service.transports); + + let backend_names = session_manager.get_backend_names(); + + let Some((backend_name, tool_name)) = resolve_tool_route(virtual_host, &request.name, &backend_names) else { + return Err(ErrorData { + code: ErrorCode::INTERNAL_ERROR, + message: "Routing problem... wrong tool name".into(), + data: None, + }); + }; + let backend_name = backend_name.to_owned(); + let tool_name = tool_name.to_owned(); + + let (service_name, backend_service) = resolve_backend(&session_manager, "call_tool", &backend_name).await?; + + let pre_result = if let Some(plugin_runtime) = &mcp_service.plugin_runtime { + plugin_runtime.before_tool_call(&request, &tool_name, &service_name).await? + } else { + ToolPreCallResult::unchanged() + }; + let post_state = pre_result.state; + let mut routed_request = request; + pre_result.arguments.apply_to_request(&mut routed_request, &tool_name); + + let progress_token = cx.meta.get_progress_token(); + let handle = backend_service + .service() + .start_tool_call( + backend_service.peer(), + routed_request, + progress_token, + tool_name.clone(), + cx.peer.clone(), + post_state.clone(), + ) + .await + .map_err(|error| backend_forward_error("call_tool", &service_name, &error))?; + let backend_progress_token = handle.progress_token.clone(); + let response = call_backend_tool(handle, cx.ct.clone()).await; + backend_service.service().stop_tracking_tool_call(&backend_progress_token).await; + + let response = response.map_err(|error| backend_forward_error("call_tool", &service_name, &error))?; + let response = match (&mcp_service.plugin_runtime, post_state) { + (Some(plugin_runtime), Some(post_state)) => { + plugin_runtime.after_tool_call(&tool_name, response, Some(post_state)).await? + }, + _ => response, + }; + info!("call_tool: backend {service_name} completed"); + Ok(response.into()) +} diff --git a/crates/contextforge-gateway-rs-lib/src/gateway/mod.rs b/crates/contextforge-gateway-rs-lib/src/gateway/mod.rs index 978229a..8bf5f23 100644 --- a/crates/contextforge-gateway-rs-lib/src/gateway/mod.rs +++ b/crates/contextforge-gateway-rs-lib/src/gateway/mod.rs @@ -1,8 +1,12 @@ mod backend_client; +mod backend_transports; +mod identifier_routing; +mod list_aggregation; mod mcp_call_validator; -pub(crate) mod mcp_gateway; +mod mcp_service; mod session_manager; mod session_store; -pub use mcp_gateway::{BackendTransports, LocalUserSessionStore, McpService}; -pub use session_store::{UserSession, UserSessionStore}; +pub use backend_transports::BackendTransports; +pub use mcp_service::McpService; +pub use session_store::{LocalUserSessionStore, 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 fdd69fe..af361e7 100644 --- a/crates/contextforge-gateway-rs-lib/src/gateway/session_manager.rs +++ b/crates/contextforge-gateway-rs-lib/src/gateway/session_manager.rs @@ -1,7 +1,7 @@ use contextforge_gateway_rs_apis::user_store::VirtualHost; use tracing::{debug, info}; -use super::mcp_gateway::{BackendTransportKey, BackendTransports, ServiceHolder}; +use super::backend_transports::{BackendTransportKey, BackendTransports, ServiceHolder}; use crate::layers::session_id::SessionId; pub struct SessionManager<'a> { diff --git a/docs/book/src/backend-connections-and-transports.md b/docs/book/src/backend-connections-and-transports.md index 8f30bd9..3328068 100644 --- a/docs/book/src/backend-connections-and-transports.md +++ b/docs/book/src/backend-connections-and-transports.md @@ -15,7 +15,7 @@ architecture roles. | Transport class | Current implementation | Main owner | Purpose | | --- | --- | --- | --- | | Downstream listener | Axum/Hyper over TCP and optional Rustls TLS. | `transports/` and `Gateway::run_gateway`. | Accept MCP streamable HTTP traffic from clients or the front door. | -| Upstream backend | Shared `reqwest::Client` plus RMCP `StreamableHttpClientTransport`. | `common.rs` and `gateway/mcp_gateway.rs`. | Open MCP client sessions to configured backend MCP servers. | +| Upstream backend | Shared `reqwest::Client` plus RMCP `StreamableHttpClientTransport`. | `common.rs`, `gateway/mcp_service/initialization.rs`, and `gateway/backend_transports.rs`. | Open MCP client sessions to configured backend MCP servers. | | Config store | Redis plain, TLS, or mTLS connection manager. | `common.rs` and `user_config_store/`. | Load `UserConfig` and plugin runtime config from control-plane authored storage. | The current MCP dataplane only uses streamable HTTP for backend MCP traffic.