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
2 changes: 2 additions & 0 deletions crates/contextforge-gateway-rs-apis/src/user_store.rs
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,8 @@ pub struct BackendMCPGateway {
pub transport: Transport,
pub passthrough_headers: Vec<String>,
pub allowed_tool_names: Vec<String>,
#[serde(default)]
pub tool_name_aliases: HashMap<String, String>,
pub allowed_resource_names: Vec<String>,
pub allowed_prompt_names: Vec<String>,
}
Expand Down
117 changes: 111 additions & 6 deletions crates/contextforge-gateway-rs-lib/src/gateway/mcp_gateway.rs
Original file line number Diff line number Diff line change
@@ -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::{
Expand Down Expand Up @@ -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(
Expand All @@ -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?
Expand Down Expand Up @@ -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<str>>(
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:?}");
Expand Down Expand Up @@ -668,15 +713,15 @@ fn log_list_backend_response<T, E: std::fmt::Debug>(
}
}

fn merge_tools(tools: Vec<(String, ListToolsResult)>) -> Vec<Tool> {
fn merge_tools(tools: Vec<(String, ListToolsResult)>, virtual_host: &VirtualHost) -> Vec<Tool> {
tools
.into_iter()
.flat_map(|(backend_name, result)| {
result
.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::<Vec<_>>()
Expand Down Expand Up @@ -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)
);
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -209,15 +209,19 @@ fn create_backends(ports: &[u16], with_tls: bool) -> HashMap<String, BackendMCPG
format!("http://127.0.0.1:{port}/mcp").parse().expect("This should work")
};

let name = format!("backend-{port}");
let backend_id = backend_id(*port);
(
name.clone(),
backend_id,
BackendMCPGateway {
name,
name: format!("backend-{port}"),
url,
transport: Transport::default(),
passthrough_headers: Vec::new(),
allowed_tool_names: Vec::new(),
tool_name_aliases: MOCK_COUNTER_TOOL_NAMES
.iter()
.map(|tool_name| (format!("backend-{port}.{tool_name}"), (*tool_name).to_owned()))
.collect(),
allowed_resource_names: Vec::new(),
allowed_prompt_names: Vec::new(),
},
Expand All @@ -226,31 +230,44 @@ fn create_backends(ports: &[u16], with_tls: bool) -> HashMap<String, BackendMCPG
.collect()
}

fn backend_id(port: u16) -> String {
format!("00000000-0000-0000-0000-{port:012}")
}

fn create_tool_names(ports: &[u16]) -> Vec<String> {
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<String> {
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<String> {
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<String> {
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()
}

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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(),
},
Expand Down
7 changes: 7 additions & 0 deletions schemas/user_config.json
Original file line number Diff line number Diff line change
Expand Up @@ -53,6 +53,13 @@
"type": "string"
}
},
"tool_name_aliases": {
"type": "object",
"additionalProperties": {
"type": "string"
},
"default": {}
},
"allowed_resource_names": {
"type": "array",
"items": {
Expand Down