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 4c5c6f37..0ae9c9dc 100644 --- a/crates/contextforge-gateway-rs-lib/src/gateway/mcp_gateway.rs +++ b/crates/contextforge-gateway-rs-lib/src/gateway/mcp_gateway.rs @@ -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, @@ -510,29 +510,45 @@ where request: CompleteRequestParams, cx: RequestContext, ) -> Result { - let maybe_parts = cx.extensions.get::(); - let maybe_session = maybe_parts.and_then(|parts| parts.extensions.get::()); - let maybe_user_config = maybe_parts.and_then(|parts| parts.extensions.get::()); - 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) } } @@ -629,7 +645,7 @@ async fn resolve_backend( } fn merge_capabilities(_server_capabilities: Vec<(String, Option)>) -> 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( diff --git a/crates/contextforge-gateway-rs-lib/tests/gateway_completions.rs b/crates/contextforge-gateway-rs-lib/tests/gateway_completions.rs new file mode 100644 index 00000000..4b41f589 --- /dev/null +++ b/crates/contextforge-gateway-rs-lib/tests/gateway_completions.rs @@ -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(()) +} diff --git a/crates/contextforge-gateway-rs-lib/tests/support/mock_counter.rs b/crates/contextforge-gateway-rs-lib/tests/support/mock_counter.rs index 42037bfa..b482132d 100644 --- a/crates/contextforge-gateway-rs-lib/tests/support/mock_counter.rs +++ b/crates/contextforge-gateway-rs-lib/tests/support/mock_counter.rs @@ -238,6 +238,27 @@ impl ServerHandler for Counter { } } + async fn complete( + &self, + request: CompleteRequestParams, + _: RequestContext, + ) -> Result { + // 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,