diff --git a/crates/contextforge-gateway-rs-apis/src/user_store.rs b/crates/contextforge-gateway-rs-apis/src/user_store.rs index 882e3d9..f4dc1d4 100644 --- a/crates/contextforge-gateway-rs-apis/src/user_store.rs +++ b/crates/contextforge-gateway-rs-apis/src/user_store.rs @@ -30,6 +30,8 @@ pub struct BackendMCPGateway { pub transport: Transport, pub passthrough_headers: Vec, pub allowed_tool_names: Vec, + #[serde(default)] + pub tool_name_aliases: HashMap, pub allowed_resource_names: Vec, pub allowed_prompt_names: Vec, } 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 f5a21f1..853469a 100644 --- a/crates/contextforge-gateway-rs-lib/src/gateway/mcp_gateway.rs +++ b/crates/contextforge-gateway-rs-lib/src/gateway/mcp_gateway.rs @@ -1,5 +1,6 @@ use std::{collections::HashMap, sync::Arc}; +use contextforge_gateway_rs_apis::user_store::VirtualHost; use contextforge_gateway_rs_cpex::{GatewayPluginRuntimeHandle, ToolPreCallResult}; use itertools::Itertools; use rmcp::{ @@ -256,7 +257,7 @@ where ) .await; - Ok(ListToolsResult { meta: None, tools: merge_tools(responses), next_cursor: None }) + Ok(ListToolsResult { meta: None, tools: merge_tools(responses, virtual_host), next_cursor: None }) } async fn call_tool( @@ -268,9 +269,19 @@ where 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, tool_name) = - route_prefixed_name(&session_manager, "call_tool", &request.name, "Routing problem... wrong tool name") - .await?; + 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? @@ -534,6 +545,40 @@ 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. +/// Older configs without aliases retain the legacy `{backend}-{tool}` route. +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(|| split_prefixed_name(name, backend_names)) +} + +/// Returns the control-plane name for an upstream tool, with the legacy +/// namespaced form as a fallback for configs published before aliases existed. +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(|| 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:?}"); @@ -668,7 +713,7 @@ fn log_list_backend_response( } } -fn merge_tools(tools: Vec<(String, ListToolsResult)>) -> Vec { +fn merge_tools(tools: Vec<(String, ListToolsResult)>, virtual_host: &VirtualHost) -> Vec { tools .into_iter() .flat_map(|(backend_name, result)| { @@ -676,7 +721,7 @@ fn merge_tools(tools: Vec<(String, ListToolsResult)>) -> Vec { .tools .into_iter() .map(|mut t| { - t.name = prefixed_name(&backend_name, &t.name).into(); + t.name = exposed_tool_name(virtual_host, &backend_name, &t.name).into(); t }) .collect::>() @@ -753,4 +798,64 @@ mod tests { let backend_names = vec!["counter_on", "counter_oneee", "counter_one"]; assert_eq!(Some(("counter_one", "get-value")), split_prefixed_name("counter_one-get-value", &backend_names)); } + + #[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 test_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": [] + } + } + }); + let virtual_host: VirtualHost = serde_json::from_value(config_json).expect("valid virtual host"); + let backend_names = vec!["compliance-reference"]; + + 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/tests/support/list_tools_gateway.rs b/crates/contextforge-gateway-rs-lib/tests/support/list_tools_gateway.rs index 1518295..556b5e3 100644 --- a/crates/contextforge-gateway-rs-lib/tests/support/list_tools_gateway.rs +++ b/crates/contextforge-gateway-rs-lib/tests/support/list_tools_gateway.rs @@ -209,15 +209,19 @@ fn create_backends(ports: &[u16], with_tls: bool) -> HashMap HashMap String { + format!("00000000-0000-0000-0000-{port:012}") +} + fn create_tool_names(ports: &[u16]) -> Vec { ports .iter() - .flat_map(|port| MOCK_COUNTER_TOOL_NAMES.iter().map(move |name| format!("backend-{port}-{name}"))) + .flat_map(|port| MOCK_COUNTER_TOOL_NAMES.iter().map(move |name| format!("backend-{port}.{name}"))) .collect() } fn create_prompt_names(ports: &[u16]) -> Vec { ports .iter() - .flat_map(|port| MOCK_COUNTER_PROMPT_NAMES.iter().map(move |name| format!("backend-{port}-{name}"))) + .flat_map(|port| { + let backend_id = backend_id(*port); + MOCK_COUNTER_PROMPT_NAMES.iter().map(move |name| format!("{backend_id}-{name}")) + }) .collect() } fn create_resource_template_names(ports: &[u16]) -> Vec { ports .iter() - .flat_map(|port| MOCK_COUNTER_RESOURCE_TEMPLATE_NAMES.iter().map(move |name| format!("backend-{port}-{name}"))) + .flat_map(|port| { + let backend_id = backend_id(*port); + MOCK_COUNTER_RESOURCE_TEMPLATE_NAMES.iter().map(move |name| format!("{backend_id}-{name}")) + }) .collect() } fn create_resource_template_uris(ports: &[u16]) -> Vec { ports .iter() - .flat_map(|port| MOCK_COUNTER_RESOURCE_TEMPLATE_URIS.iter().map(move |uri| format!("backend-{port}-{uri}"))) + .flat_map(|port| { + let backend_id = backend_id(*port); + MOCK_COUNTER_RESOURCE_TEMPLATE_URIS.iter().map(move |uri| format!("{backend_id}-{uri}")) + }) .collect() } diff --git a/crates/contextforge-gateway-rs-lib/tests/support/plugin_gateway.rs b/crates/contextforge-gateway-rs-lib/tests/support/plugin_gateway.rs index d330313..93c9731 100644 --- a/crates/contextforge-gateway-rs-lib/tests/support/plugin_gateway.rs +++ b/crates/contextforge-gateway-rs-lib/tests/support/plugin_gateway.rs @@ -260,6 +260,7 @@ async fn start_gateway_with_runtime( transport: Transport::default(), passthrough_headers: Vec::new(), allowed_tool_names: Vec::new(), + tool_name_aliases: HashMap::new(), allowed_resource_names: Vec::new(), allowed_prompt_names: Vec::new(), }, diff --git a/schemas/user_config.json b/schemas/user_config.json index 444d84c..d253640 100644 --- a/schemas/user_config.json +++ b/schemas/user_config.json @@ -53,6 +53,13 @@ "type": "string" } }, + "tool_name_aliases": { + "type": "object", + "additionalProperties": { + "type": "string" + }, + "default": {} + }, "allowed_resource_names": { "type": "array", "items": {