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
64 changes: 40 additions & 24 deletions crates/contextforge-gateway-rs-lib/src/gateway/mcp_gateway.rs
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,7 @@ use itertools::Itertools;
use rmcp::{
ErrorData, RoleClient, RoleServer, ServerHandler, ServiceExt,
model::{
CallToolRequestParams, CallToolResult, CompleteRequestParams, CompleteResult, CompletionInfo, ErrorCode,
CallToolRequestParams, CallToolResult, CompleteRequestParams, CompleteResult, ErrorCode,
GetPromptRequestParams, GetPromptResult, Implementation, InitializeRequestParams, InitializeResult,
ListPromptsResult, ListResourceTemplatesResult, ListResourcesResult, ListToolsResult, PaginatedRequestParams,
Prompt, ReadResourceRequestParams, ReadResourceResult, Reference, Resource, ResourceTemplate,
Expand Down Expand Up @@ -510,29 +510,45 @@ where
request: CompleteRequestParams,
cx: RequestContext<RoleServer>,
) -> Result<CompleteResult, ErrorData> {
let maybe_parts = cx.extensions.get::<Parts>();
let maybe_session = maybe_parts.and_then(|parts| parts.extensions.get::<SessionId>());
let maybe_user_config = maybe_parts.and_then(|parts| parts.extensions.get::<UserConfig>());
info!("complete user_config = {maybe_user_config:#?} session_id = {maybe_session:#?}");
let values = match &request.r#ref {
Reference::Resource(_) => {
if request.argument.name == "id" {
vec!["1".into(), "2".into(), "3".into()]
} else {
vec![]
}
},
Reference::Prompt(prompt_ref) => {
if request.argument.name == "name" {
vec!["Alice".into(), "Bob".into(), "Charlie".into()]
} else if request.argument.name == "style" {
vec!["friendly".into(), "formal".into(), "casual".into()]
} else {
vec![prompt_ref.name.clone()]
}
},
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 backend_names = session_manager.get_backend_names();

// The reference carries a namespaced prompt name or resource URI; route on that.
let namespaced = match &request.r#ref {
Reference::Prompt(prompt) => prompt.name.as_str(),
Reference::Resource(resource) => resource.uri.as_str(),
};

let Some((backend_name, stripped)) = split_prefixed_name(namespaced, &backend_names) else {
return Err(ErrorData {
code: ErrorCode::INTERNAL_ERROR,
message: "Routing problem... wrong completion reference".into(),
data: None,
});
};
Ok(CompleteResult::new(CompletionInfo::new(values).map_err(|e| ErrorData::internal_error(e, None))?))
let backend_name = backend_name.to_owned();
let stripped = stripped.to_owned();

let (service_name, service) = resolve_backend(&session_manager, "complete", &backend_name).await?;

let mut routed_request = request;
match &mut routed_request.r#ref {
Reference::Prompt(prompt) => prompt.name = stripped,
Reference::Resource(resource) => resource.uri = stripped,
}
let response = service.complete(routed_request).await.map_err(|error| {
warn!("complete: backend {service_name} {error:?}");
ErrorData {
code: ErrorCode::INTERNAL_ERROR,
message: "Routing problem... got no responses from backends".into(),
data: None,
}
})?;
info!("complete: backend {service_name} returned {} values", response.completion.values.len());
Ok(response)
}
}

Expand Down Expand Up @@ -629,7 +645,7 @@ async fn resolve_backend(
}

fn merge_capabilities(_server_capabilities: Vec<(String, Option<ServerCapabilities>)>) -> ServerCapabilities {
ServerCapabilities::builder().enable_prompts().enable_resources().enable_tools().build()
ServerCapabilities::builder().enable_completions().enable_prompts().enable_resources().enable_tools().build()
}

fn log_list_backend_response<T, E: std::fmt::Debug>(
Expand Down
125 changes: 125 additions & 0 deletions crates/contextforge-gateway-rs-lib/tests/gateway_completions.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,125 @@
mod support;

use contextforge_gateway_rs_lib::{Config, Result, UpstreamConnectionMode};
use tracing::info;

use support::{
ListToolsGatewaySettings, connect_client, create_client, create_gateway_with_four_counters, create_ports,
};

fn plaintext_config(gateway_port: u16) -> Config {
Config {
address: Some(format!("127.0.0.1:{gateway_port}").parse().expect("This should work")),
token_verification_public_key: Some("../../assets/jwt.key.pub".into()),
upstream_connection_mode: Some(UpstreamConnectionMode::PlainTextOrTls),
..Default::default()
}
}

#[tokio::test(flavor = "multi_thread", worker_threads = 1)]
#[test_log::test]
async fn plaintext_completes_prompt_argument_through_prefixed_backend() -> Result<()> {
let gateway_port = create_ports(1)[0];
let user = "admin@example.com";
let Ok(ListToolsGatewaySettings { handle, gateway_url, .. }) =
create_gateway_with_four_counters(user, plaintext_config(gateway_port)).await
else {
panic!("Invalid configuration ");
};

let client = create_client(user);
let maybe_passed = assert_prompt_completion(gateway_url, client).await;

handle.abort();
maybe_passed
}

#[tokio::test(flavor = "multi_thread", worker_threads = 1)]
#[test_log::test]
async fn plaintext_completes_resource_argument_through_prefixed_backend() -> Result<()> {
let gateway_port = create_ports(1)[0];
let user = "admin@example.com";
let Ok(ListToolsGatewaySettings { handle, gateway_url, .. }) =
create_gateway_with_four_counters(user, plaintext_config(gateway_port)).await
else {
panic!("Invalid configuration ");
};

let client = create_client(user);
let maybe_passed = assert_resource_completion(gateway_url, client).await;

handle.abort();
maybe_passed
}

#[tokio::test(flavor = "multi_thread", worker_threads = 1)]
#[test_log::test]
async fn plaintext_complete_for_unrouted_reference_errors() -> Result<()> {
let gateway_port = create_ports(1)[0];
let user = "admin@example.com";
let Ok(ListToolsGatewaySettings { handle, gateway_url, .. }) =
create_gateway_with_four_counters(user, plaintext_config(gateway_port)).await
else {
panic!("Invalid configuration ");
};

let client = create_client(user);
let maybe_passed = assert_unrouted_completion_errors(gateway_url, client).await;

handle.abort();
maybe_passed
}

async fn assert_prompt_completion(gateway_url: String, client: reqwest::Client) -> Result<()> {
info!("Sending request to {gateway_url}");
let running_service = connect_client(gateway_url, client).await?;

// Spec-compliant clients only issue completion/complete when the server advertises the
// capability, so the gateway must declare it for the proxying below to be reachable.
if running_service.peer_info().and_then(|info| info.capabilities.completions.clone()).is_none() {
return Err("gateway must advertise the completions capability".into());
}

let prompts = running_service.list_prompts(None).await?;
let prompt_name = prompts
.prompts
.iter()
.find(|p| p.name.ends_with("-example_prompt"))
.map(|p| p.name.clone())
.ok_or("expected a federated example_prompt")?;

// The backend only knows the un-prefixed prompt name, so a non-empty result proves the gateway
// routed to a single backend and stripped the namespace prefix before forwarding.
let values = running_service.complete_prompt_simple(prompt_name, "message", "h").await?;
if !values.contains(&"hello".to_owned()) {
return Err(format!("expected backend prompt completions, got: {values:?}").into());
}

Ok(())
}

async fn assert_resource_completion(gateway_url: String, client: reqwest::Client) -> Result<()> {
let running_service = connect_client(gateway_url, client).await?;

let resources = running_service.list_resources(None).await?;
let uri = resources.resources.first().ok_or("expected at least one federated resource")?.uri.clone();

// The mock echoes back the URI it received; it must be the stripped backend-local URI.
let values = running_service.complete_resource_simple(uri, "path", "").await?;
match values.first() {
Some(value) if !value.starts_with("backend-") => Ok(()),
other => Err(format!("expected stripped backend URI in completion, got: {other:?}").into()),
}
}

async fn assert_unrouted_completion_errors(gateway_url: String, client: reqwest::Client) -> Result<()> {
let running_service = connect_client(gateway_url, client).await?;

// No backend namespace prefix => no route, so the gateway must reject it.
let result = running_service.complete_prompt_simple("unrouted_prompt", "message", "h").await;
if result.is_ok() {
return Err("expected a routing error for an unrouted completion reference".into());
}

Ok(())
}
21 changes: 21 additions & 0 deletions crates/contextforge-gateway-rs-lib/tests/support/mock_counter.rs
Original file line number Diff line number Diff line change
Expand Up @@ -238,6 +238,27 @@ impl ServerHandler for Counter {
}
}

async fn complete(
&self,
request: CompleteRequestParams,
_: RequestContext<RoleServer>,
) -> Result<CompleteResult, McpError> {
// Only backend-local references are known here; a still-namespaced name/URI won't match,
// proving the gateway stripped the prefix before forwarding.
let values = match &request.r#ref {
Reference::Prompt(prompt) if prompt.name == "example_prompt" && request.argument.name == "message" => {
vec!["hello".to_owned(), "hola".to_owned()]
},
Reference::Resource(resource)
if matches!(resource.uri.as_str(), "str:////Users/to/some/path/" | "memo://insights") =>
{
vec![resource.uri.clone()]
},
_ => return Err(McpError::invalid_params("unknown completion reference", None)),
};
Ok(CompleteResult::new(CompletionInfo::new(values).map_err(|e| McpError::internal_error(e, None))?))
}

async fn list_resource_templates(
&self,
_request: Option<PaginatedRequestParams>,
Expand Down