diff --git a/_context/wiki/architecture.md b/_context/wiki/architecture.md index b9d79b1e..be839b60 100644 --- a/_context/wiki/architecture.md +++ b/_context/wiki/architecture.md @@ -84,8 +84,10 @@ Order is invariant: auth/config before backend selection; request plugins before | JWT decoders | `ContextForgeDataPlaneAppState` | Process | | User config | `RedisUserConfigStore` (LRU + Redis) | Request-path consumed; control-plane authored | | Request identity / VirtualHostId | Request extensions | One HTTP request | +| Gateway request context | `virtual_host_config_layer` task-local | One HTTP request; copied into the RMCP service factory for context-aware local methods | | Downstream session id | RMCP + `SessionId` extension | MCP session | | Backend RMCP services | `BackendTransports` map | Local process, per principal/backend/session | +| Downstream subscription sinks | `DownstreamSubscriptionRegistry` | Listen-stream lifetime, keyed by principal, virtual host, subscription id, registration id, and notification kind | | Local user session mapping | `LocalUserSessionStore` | Local LRU, 50k entries, 1 hour | | Plugin manager | `CpexRuntimeRegistry` | Process, reloadable | @@ -107,6 +109,7 @@ In multi-runtime mode, the first thread initializes the optional CPEX plugin run | State | Lock | Contention profile | | --- | --- | --- | | `BackendTransports` map | `Arc>>` | Locked briefly on initialize insert, per-call borrow, and cleanup. Borrowing clones `Arc` handles so the lock is not held across backend calls. | +| `DownstreamSubscriptionRegistry` map | `Arc>>` | Locked only for short insert/remove operations. Cleanup runs from `Drop`, so it cannot await. No registry lock is held across stream waits or backend I/O. | | Subscription set | `Arc>>` | Local `subscribe`/`unsubscribe` only. | | User config LRU cache | `Arc>` inside `RedisUserConfigStore` | One lock per config lookup on the hot path; misses add a Redis round trip. | | User session LRU cache | Same pattern in `LocalUserSessionStore` | Initialize and delete paths. | @@ -128,6 +131,7 @@ The binary sets `tikv_jemallocator` as the global allocator. jemalloc holds up b - List methods fan out to all connected backends concurrently and merge. - Targeted calls resolve exactly one backend service handle. - `call_tool` watches the downstream cancellation token and forwards a cancel to the backend if the client gives up first; backend progress notifications are forwarded downstream while the call is in flight. +- `subscriptions/listen` registers downstream sinks, parks on RMCP cancellation, and removes those sinks when the listen stream closes. ## Startup And Response Flow diff --git a/_context/wiki/routing.md b/_context/wiki/routing.md index 4a681a85..d214179c 100644 --- a/_context/wiki/routing.md +++ b/_context/wiki/routing.md @@ -60,6 +60,57 @@ This is **local process state only**. Implications: - Gateway restart → all sessions lost → clients must re-run `initialize`. - Multi-runtime mode (`--single-runtime false`): each runtime thread has its own `BackendTransports` with no cross-thread affinity. +## Modern Discovery And Subscription Admission + +Modern downstream clients use `server/discover` with MCP `2026-07-28`. +`McpService::supported_protocol_versions()` advertises only `2026-07-28`, so +requests that declare older protocol versions fail protocol-version validation. +Full removal of RMCP legacy session behavior belongs to the stateless routing +migration tracked separately. + +For normal requests, `virtual_host_config_layer` carries the authenticated +principal, selected virtual host id, and selected `VirtualHost` into the RMCP +service factory. The per-request `McpService` then uses that context for local +modern methods: + +```text +MCP client + -> /servers/{virtual_host_id}/mcp + -> auth + user config + virtual host check + -> RMCP service factory receives principal + virtual host + -> server/discover and subscriptions/listen use that context +``` + +`server/discover` reports capabilities derived from the selected virtual host. +Today this is vhost-accurate, not backend-live-accurate: a non-empty virtual +host advertises list-change and resource subscription support, while backend +capability-cache refinement belongs to the stateless routing work. + +`subscriptions/listen` uses RMCP's subscription machinery. The gateway narrows +the requested filter to supported notification kinds and routable resource +subscription URIs, registers the accepted `SubscriptionSink`s in +`DownstreamSubscriptionRegistry`, and removes them when the listen stream is +cancelled or closed. + +```text +MCP client + | + | subscriptions/listen + v +Gateway + | + | narrow filter against selected virtual host + | register accepted sinks + | wait for listen cancellation + | remove sinks on close + v +DownstreamSubscriptionRegistry +``` + +Notification delivery is intentionally separate follow-up work: +`*/list_changed` relay, `resources/updated` relay over `subscriptions/listen`, +and upstream `subscriptions/listen` management. + ```mermaid sequenceDiagram @@ -131,4 +182,6 @@ If RMCP rejects the delete, local state is untouched. | `get_prompt` | Targeted | Single-backend: name unchanged. Multi-backend: strips prefix. | | `complete` | Targeted | Routes on prompt name or resource URI inside `ref`. | | `ping` | Local | Returns success; no backend fanout. | +| `server/discover` | Local | Reports `2026-07-28` support and capabilities derived from the authenticated user's selected virtual host. | +| `subscriptions/listen` | Local | Narrows the requested subscription filter, registers downstream sinks, and cleans them up when the listen stream closes. Backend notification delivery is follow-up work. | | `DELETE` | Session | RMCP handles first; on success `session_id_layer` removes local session + backend transports. | diff --git a/crates/contextforge-data-plane-lib/src/gateway/downstream_subscriptions.rs b/crates/contextforge-data-plane-lib/src/gateway/downstream_subscriptions.rs new file mode 100644 index 00000000..3f512b6f --- /dev/null +++ b/crates/contextforge-data-plane-lib/src/gateway/downstream_subscriptions.rs @@ -0,0 +1,207 @@ +use std::{ + collections::HashMap, + sync::{ + Arc, Mutex, + atomic::{AtomicU64, Ordering}, + }, +}; + +use rmcp::{ + model::{RequestId, SubscriptionFilter}, + service::SubscriptionSink, +}; + +use crate::layers::request_context::GatewayRequestContext; + +#[derive(Clone, Default)] +pub(crate) struct DownstreamSubscriptionRegistry { + inner: Arc>>, + next_registration_id: Arc, +} + +impl DownstreamSubscriptionRegistry { + pub(crate) fn register( + &self, + context: &GatewayRequestContext, + filter: &SubscriptionFilter, + sink: &SubscriptionSink, + ) -> DownstreamSubscriptionGuard { + let registration_id = self.next_registration_id.fetch_add(1, Ordering::Relaxed); + let keys = subscription_keys(context, filter, sink.id(), registration_id); + let mut subscriptions = self.inner.lock().expect("downstream subscription registry lock poisoned"); + for key in &keys { + subscriptions.insert(key.clone(), sink.clone()); + } + DownstreamSubscriptionGuard { registry: self.clone(), keys } + } + + fn remove_all(&self, keys: &[DownstreamSubscriptionKey]) { + let mut subscriptions = self.inner.lock().expect("downstream subscription registry lock poisoned"); + for key in keys { + subscriptions.remove(key); + } + } +} + +pub(crate) struct DownstreamSubscriptionGuard { + registry: DownstreamSubscriptionRegistry, + keys: Vec, +} + +impl Drop for DownstreamSubscriptionGuard { + fn drop(&mut self) { + self.registry.remove_all(&self.keys); + } +} + +#[derive(Clone, Debug, Eq, Hash, PartialEq)] +pub(crate) struct DownstreamSubscriptionKey { + principal: String, + virtual_host_id: String, + subscription_id: RequestId, + registration_id: u64, + notification: DownstreamSubscriptionNotification, +} + +#[derive(Clone, Debug, Eq, Hash, PartialEq)] +pub(crate) enum DownstreamSubscriptionNotification { + ToolsListChanged, + PromptsListChanged, + ResourcesListChanged, + ResourceUpdated { uri: String }, +} + +pub(super) fn subscription_keys( + context: &GatewayRequestContext, + filter: &SubscriptionFilter, + subscription_id: &RequestId, + registration_id: u64, +) -> Vec { + let mut keys = Vec::new(); + if filter.tools_list_changed == Some(true) { + keys.push(subscription_key( + context, + subscription_id, + registration_id, + DownstreamSubscriptionNotification::ToolsListChanged, + )); + } + if filter.prompts_list_changed == Some(true) { + keys.push(subscription_key( + context, + subscription_id, + registration_id, + DownstreamSubscriptionNotification::PromptsListChanged, + )); + } + if filter.resources_list_changed == Some(true) { + keys.push(subscription_key( + context, + subscription_id, + registration_id, + DownstreamSubscriptionNotification::ResourcesListChanged, + )); + } + if let Some(uris) = &filter.resource_subscriptions { + keys.extend(uris.iter().map(|uri| { + subscription_key( + context, + subscription_id, + registration_id, + DownstreamSubscriptionNotification::ResourceUpdated { uri: uri.clone() }, + ) + })); + } + keys +} + +fn subscription_key( + context: &GatewayRequestContext, + subscription_id: &RequestId, + registration_id: u64, + notification: DownstreamSubscriptionNotification, +) -> DownstreamSubscriptionKey { + DownstreamSubscriptionKey { + principal: context.principal().to_owned(), + virtual_host_id: context.virtual_host_id().to_owned(), + subscription_id: subscription_id.clone(), + registration_id, + notification, + } +} + +#[cfg(test)] +mod tests { + use std::collections::HashMap; + + use contextforge_data_plane_apis::user_store::{BackendMCPGateway, VirtualHost}; + use rmcp::model::RequestId; + + use super::*; + + #[test] + fn keys_include_subscription_id_and_notification_kind() { + let gateway_context = GatewayRequestContext::new(&test_claims(), &test_virtual_host_id(), &test_virtual_host()); + let filter = SubscriptionFilter::builder().tools_list_changed().resource_subscription("memo://known").build(); + + let keys = subscription_keys(&gateway_context, &filter, &RequestId::Number(7), 9); + + assert_eq!(2, keys.len()); + assert!(keys.iter().all(|key| key.subscription_id == RequestId::Number(7))); + assert!(keys.iter().all(|key| key.registration_id == 9)); + } + + #[test] + fn registration_id_is_part_of_key_identity() { + let gateway_context = GatewayRequestContext::new(&test_claims(), &test_virtual_host_id(), &test_virtual_host()); + let filter = SubscriptionFilter::builder().tools_list_changed().build(); + + let first = subscription_keys(&gateway_context, &filter, &RequestId::Number(7), 0); + let second = subscription_keys(&gateway_context, &filter, &RequestId::Number(7), 1); + + assert_ne!(first, second); + } + + fn test_virtual_host() -> VirtualHost { + VirtualHost { + backends: HashMap::from([( + "backend-one".to_owned(), + BackendMCPGateway { + name: "backend-one".to_owned(), + url: "http://127.0.0.1:9999/mcp".parse().expect("valid URL"), + passthrough_headers: Vec::new(), + add_headers: HashMap::new(), + remove_headers: Vec::new(), + allowed_tool_names: Vec::new(), + tool_name_aliases: HashMap::new(), + allowed_resource_names: Vec::new(), + allowed_prompt_names: Vec::new(), + }, + )]), + } + } + + fn test_claims() -> crate::common::ContextForgeClaims { + crate::common::ContextForgeClaims { + sub: "test-principal".to_owned(), + jti: "test-jti".to_owned(), + token_use: None, + iat: None, + iss: "test-issuer".to_owned(), + aud: "test-audience".to_owned(), + exp: 1, + teams: None, + user: crate::common::User::builder() + .email("test@example.com".to_owned()) + .full_name(None) + .is_admin(false) + .auth_provider("test".to_owned()) + .build(), + scopes: None, + } + } + + fn test_virtual_host_id() -> crate::layers::virtual_host_id::VirtualHostId { + crate::layers::virtual_host_id::VirtualHostId::new("test-vhost".to_owned()) + } +} diff --git a/crates/contextforge-data-plane-lib/src/gateway/mcp_service.rs b/crates/contextforge-data-plane-lib/src/gateway/mcp_service.rs index 452a4669..9f8cd464 100644 --- a/crates/contextforge-data-plane-lib/src/gateway/mcp_service.rs +++ b/crates/contextforge-data-plane-lib/src/gateway/mcp_service.rs @@ -4,23 +4,35 @@ mod prompts; mod resources; mod tools; +use std::{borrow::Cow, collections::HashSet}; + +use contextforge_data_plane_apis::user_store::VirtualHost; use contextforge_data_plane_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, + GetPromptResponse, Implementation, InitializeRequestParams, InitializeResult, ListPromptsResult, + ListResourceTemplatesResult, ListResourcesResult, ListToolsResult, PaginatedRequestParams, ProtocolVersion, + ReadResourceRequestParams, ReadResourceResponse, ServerCapabilities, ServerInfo, SubscribeRequestParams, + SubscriptionFilter, UnsubscribeRequestParams, }, - service::RequestContext, + service::{RequestContext, SubscriptionContext}, }; use typed_builder::TypedBuilder; -use super::{backend_transports::BackendTransports, session_store::UserSessionStore}; +use super::{DownstreamSubscriptionRegistry, backend_transports::BackendTransports, session_store::UserSessionStore}; +use crate::layers::request_context::GatewayRequestContext; + +const SUPPORTED_PROTOCOL_VERSIONS: &[ProtocolVersion] = &[ProtocolVersion::V_2026_07_28]; #[derive(Clone, TypedBuilder)] -#[builder(field_defaults(setter(prefix = "with_")))] +#[builder( + field_defaults(setter(prefix = "with_")), + builder_method(vis = "pub(crate)"), + builder_type(vis = "pub(crate)"), + build_method(vis = "pub(crate)") +)] pub struct McpService where T: UserSessionStore, @@ -31,12 +43,79 @@ where user_session_store: T, #[builder(default)] plugin_runtime: Option, + #[builder(default, setter(skip))] + gateway_context: Option, + #[builder(default, setter(skip))] + capabilities: ServerCapabilities, + #[builder(default, setter(skip))] + downstream_subscriptions: DownstreamSubscriptionRegistry, +} + +impl McpService +where + T: UserSessionStore, +{ + pub(crate) fn with_gateway_request_context(mut self, gateway_context: Option) -> Self { + self.capabilities = gateway_context + .as_ref() + .map_or_else(ServerCapabilities::default, |context| capabilities_for_virtual_host(context.virtual_host())); + self.gateway_context = gateway_context; + self + } + + pub(crate) fn with_downstream_subscription_registry(mut self, registry: DownstreamSubscriptionRegistry) -> Self { + self.downstream_subscriptions = registry; + self + } + + #[cfg(test)] + fn with_test_capabilities(mut self, capabilities: ServerCapabilities) -> Self { + self.capabilities = capabilities; + self + } } impl ServerHandler for McpService where T: UserSessionStore + Send + Sync + 'static, { + fn supported_protocol_versions(&self) -> Cow<'static, [ProtocolVersion]> { + Cow::Borrowed(SUPPORTED_PROTOCOL_VERSIONS) + } + + fn get_info(&self) -> ServerInfo { + ServerInfo::new(self.capabilities.clone()) + .with_server_info(gateway_server_implementation()) + .with_protocol_version(ProtocolVersion::V_2026_07_28) + } + + fn accepted_subscription_filter(&self, requested: &SubscriptionFilter) -> Option { + let gateway_context = self.gateway_context.as_ref()?; + let mut accepted = requested.supported_by(&self.capabilities); + accepted.resource_subscriptions = accepted.resource_subscriptions.take().and_then(|uris| { + let mut seen = HashSet::new(); + let uris = uris + .into_iter() + .filter(|uri| { + seen.insert(uri.clone()) && is_routable_resource_subscription(gateway_context.virtual_host(), uri) + }) + .collect::>(); + (!uris.is_empty()).then_some(uris) + }); + Some(accepted) + } + + async fn listen(&self, context: SubscriptionContext) -> Result<(), ErrorData> { + let Some(gateway_context) = self.gateway_context.as_ref() else { + return Err(ErrorData::internal_error("subscriptions/listen missing gateway request context", None)); + }; + let _subscription_guard = + self.downstream_subscriptions.register(gateway_context, context.accepted(), context.sink()); + + context.cancelled().await; + Ok(()) + } + async fn initialize( &self, request: InitializeRequestParams, @@ -129,3 +208,186 @@ where completion::complete(self, request, cx).await } } + +pub(crate) fn capabilities_for_virtual_host(virtual_host: &VirtualHost) -> ServerCapabilities { + if virtual_host.backends.is_empty() { + return ServerCapabilities::default(); + } + + ServerCapabilities::builder() + .enable_completions() + .enable_prompts() + .enable_prompts_list_changed() + .enable_resources() + .enable_resources_subscribe() + .enable_resources_list_changed() + .enable_tools() + .enable_tool_list_changed() + .build() +} + +fn gateway_server_implementation() -> Implementation { + Implementation::new("rust-conformance-server", "0.1.0") +} + +fn is_routable_resource_subscription(virtual_host: &VirtualHost, uri: &str) -> bool { + if virtual_host.backends.len() <= 1 { + return !virtual_host.backends.is_empty(); + } + + virtual_host + .backends + .keys() + .any(|backend_name| uri.strip_prefix(backend_name).and_then(|rest| rest.strip_prefix('-')).is_some()) +} + +#[cfg(test)] +mod tests { + use std::collections::HashMap; + + use contextforge_data_plane_apis::user_store::BackendMCPGateway; + use rmcp::model::RequestId; + + use super::*; + use crate::gateway::session_store::LocalUserSessionStore; + + #[test] + fn capabilities_for_non_empty_virtual_host_include_subscription_capabilities() { + let virtual_host = virtual_host_with_backends(&["backend-one"]); + + let capabilities = capabilities_for_virtual_host(&virtual_host); + + assert!(capabilities.completions.is_some()); + assert_eq!(Some(true), capabilities.tools.and_then(|tools| tools.list_changed)); + assert_eq!(Some(true), capabilities.prompts.and_then(|prompts| prompts.list_changed)); + let resources = capabilities.resources.expect("resources are enabled"); + assert_eq!(Some(true), resources.subscribe); + assert_eq!(Some(true), resources.list_changed); + } + + #[test] + fn accepted_subscription_filter_keeps_only_routable_resource_uris() { + let virtual_host = virtual_host_with_backends(&["backend-one", "backend-two"]); + let gateway_context = GatewayRequestContext::new(&test_claims(), &test_virtual_host_id(), &virtual_host); + let service = test_service(Some(gateway_context), capabilities_for_virtual_host(&virtual_host)); + let requested = SubscriptionFilter::builder() + .tools_list_changed() + .prompts_list_changed() + .resources_list_changed() + .resource_subscription("backend-one-memo://known") + .resource_subscription("missing-memo://unknown") + .resource_subscription("backend-one-memo://known") + .build(); + + let accepted = service.accepted_subscription_filter(&requested).expect("subscriptions/listen is implemented"); + + assert_eq!(Some(true), accepted.tools_list_changed); + assert_eq!(Some(true), accepted.prompts_list_changed); + assert_eq!(Some(true), accepted.resources_list_changed); + assert_eq!(Some(vec!["backend-one-memo://known".to_owned()]), accepted.resource_subscriptions); + } + + #[test] + fn accepted_subscription_filter_is_not_implemented_without_gateway_context() { + let service = + test_service(None, ServerCapabilities::builder().enable_tools().enable_tool_list_changed().build()); + let requested = SubscriptionFilter::builder().tools_list_changed().build(); + + assert!(service.accepted_subscription_filter(&requested).is_none()); + } + + #[test] + fn subscription_filter_for_empty_vhost_accepts_no_notifications() { + let virtual_host = VirtualHost { backends: HashMap::new() }; + let gateway_context = GatewayRequestContext::new(&test_claims(), &test_virtual_host_id(), &virtual_host); + let service = test_service(Some(gateway_context), capabilities_for_virtual_host(&virtual_host)); + let requested = SubscriptionFilter::builder() + .tools_list_changed() + .prompts_list_changed() + .resources_list_changed() + .resource_subscription("memo://known") + .build(); + + let accepted = service.accepted_subscription_filter(&requested).expect("context exists"); + + assert_eq!(None, accepted.tools_list_changed); + assert_eq!(None, accepted.prompts_list_changed); + assert_eq!(None, accepted.resources_list_changed); + assert_eq!(None, accepted.resource_subscriptions); + } + + fn test_service( + gateway_context: Option, + capabilities: ServerCapabilities, + ) -> McpService { + McpService::builder() + .with_user_session_store(LocalUserSessionStore::new()) + .with_http_client(reqwest::Client::new()) + .build() + .with_gateway_request_context(gateway_context) + .with_test_capabilities(capabilities) + } + + fn virtual_host_with_backends(names: &[&str]) -> VirtualHost { + VirtualHost { + backends: names + .iter() + .map(|name| { + ( + (*name).to_owned(), + BackendMCPGateway { + name: (*name).to_owned(), + url: "http://127.0.0.1:9999/mcp".parse().expect("valid URL"), + passthrough_headers: Vec::new(), + add_headers: HashMap::new(), + remove_headers: Vec::new(), + allowed_tool_names: Vec::new(), + tool_name_aliases: HashMap::new(), + allowed_resource_names: Vec::new(), + allowed_prompt_names: Vec::new(), + }, + ) + }) + .collect(), + } + } + + fn test_claims() -> crate::common::ContextForgeClaims { + crate::common::ContextForgeClaims { + sub: "test-principal".to_owned(), + jti: "test-jti".to_owned(), + token_use: None, + iat: None, + iss: "test-issuer".to_owned(), + aud: "test-audience".to_owned(), + exp: 1, + teams: None, + user: crate::common::User::builder() + .email("test@example.com".to_owned()) + .full_name(None) + .is_admin(false) + .auth_provider("test".to_owned()) + .build(), + scopes: None, + } + } + + fn test_virtual_host_id() -> crate::layers::virtual_host_id::VirtualHostId { + crate::layers::virtual_host_id::VirtualHostId::new("test-vhost".to_owned()) + } + + #[test] + fn registry_keys_include_subscription_id_and_notification_kind() { + let virtual_host = virtual_host_with_backends(&["backend-one"]); + let gateway_context = GatewayRequestContext::new(&test_claims(), &test_virtual_host_id(), &virtual_host); + let filter = SubscriptionFilter::builder().tools_list_changed().resource_subscription("memo://known").build(); + let keys = super::super::downstream_subscriptions::subscription_keys( + &gateway_context, + &filter, + &RequestId::Number(7), + 0, + ); + + assert_eq!(2, keys.len()); + } +} diff --git a/crates/contextforge-data-plane-lib/src/gateway/mod.rs b/crates/contextforge-data-plane-lib/src/gateway/mod.rs index 8bf5f23c..0ff6b6f7 100644 --- a/crates/contextforge-data-plane-lib/src/gateway/mod.rs +++ b/crates/contextforge-data-plane-lib/src/gateway/mod.rs @@ -1,5 +1,6 @@ mod backend_client; mod backend_transports; +mod downstream_subscriptions; mod identifier_routing; mod list_aggregation; mod mcp_call_validator; @@ -8,5 +9,6 @@ mod session_manager; mod session_store; pub use backend_transports::BackendTransports; +pub(crate) use downstream_subscriptions::DownstreamSubscriptionRegistry; pub use mcp_service::McpService; pub use session_store::{LocalUserSessionStore, UserSession, UserSessionStore}; diff --git a/crates/contextforge-data-plane-lib/src/layers/mod.rs b/crates/contextforge-data-plane-lib/src/layers/mod.rs index e4b52754..a34dd18a 100644 --- a/crates/contextforge-data-plane-lib/src/layers/mod.rs +++ b/crates/contextforge-data-plane-lib/src/layers/mod.rs @@ -1,5 +1,6 @@ pub mod claims_id; pub mod mcp_origin; +pub(crate) mod request_context; pub mod session_id; pub mod user_config_store; pub mod virtual_host_config; diff --git a/crates/contextforge-data-plane-lib/src/layers/request_context.rs b/crates/contextforge-data-plane-lib/src/layers/request_context.rs new file mode 100644 index 00000000..e3e9d366 --- /dev/null +++ b/crates/contextforge-data-plane-lib/src/layers/request_context.rs @@ -0,0 +1,53 @@ +use std::future::Future; + +use contextforge_data_plane_apis::user_store::VirtualHost; + +use crate::{common::ContextForgeClaims, layers::virtual_host_id::VirtualHostId}; + +tokio::task_local! { + static GATEWAY_REQUEST_CONTEXT: GatewayRequestContext; +} + +#[derive(Clone, Debug)] +pub(crate) struct GatewayRequestContext { + principal: String, + virtual_host_id: String, + virtual_host: VirtualHost, +} + +impl GatewayRequestContext { + pub(crate) fn new( + claims: &ContextForgeClaims, + virtual_host_id: &VirtualHostId, + virtual_host: &VirtualHost, + ) -> Self { + Self { + principal: claims.sub.clone(), + virtual_host_id: virtual_host_id.value().clone(), + virtual_host: virtual_host.clone(), + } + } + + pub(crate) fn principal(&self) -> &str { + &self.principal + } + + pub(crate) fn virtual_host_id(&self) -> &str { + &self.virtual_host_id + } + + pub(crate) fn virtual_host(&self) -> &VirtualHost { + &self.virtual_host + } +} + +pub(crate) async fn scope_gateway_request_context(context: GatewayRequestContext, future: F) -> R +where + F: Future, +{ + GATEWAY_REQUEST_CONTEXT.scope(context, future).await +} + +pub(crate) fn current_gateway_request_context() -> Option { + GATEWAY_REQUEST_CONTEXT.try_with(Clone::clone).ok() +} diff --git a/crates/contextforge-data-plane-lib/src/layers/virtual_host_config.rs b/crates/contextforge-data-plane-lib/src/layers/virtual_host_config.rs index 7634e58d..1b69cb4b 100644 --- a/crates/contextforge-data-plane-lib/src/layers/virtual_host_config.rs +++ b/crates/contextforge-data-plane-lib/src/layers/virtual_host_config.rs @@ -3,7 +3,13 @@ use contextforge_data_plane_apis::user_store::UserConfig; use http::{StatusCode, header}; use tracing::debug; -use crate::layers::virtual_host_id::VirtualHostId; +use crate::{ + common::ContextForgeClaims, + layers::{ + request_context::{GatewayRequestContext, scope_gateway_request_context}, + virtual_host_id::VirtualHostId, + }, +}; const SERVER_NOT_FOUND_BODY: &str = r#"{"detail":"Server not found"}"#; @@ -11,24 +17,28 @@ pub async fn virtual_host_config_layer(request: http::Request, let virtual_host_id = request.extensions().get::(); let user_config = request.extensions().get::(); - if let (Some(virtual_host_id), Some(user_config)) = (virtual_host_id, user_config) - && !has_virtual_host(user_config, virtual_host_id) - { + if let (Some(virtual_host_id), Some(user_config)) = (virtual_host_id, user_config) { + let Some(virtual_host) = user_config.virtual_hosts.get(virtual_host_id.value()) else { + let virtual_host_id = virtual_host_id.value(); + let virtual_hosts = user_config.virtual_hosts.len(); + debug!( + "virtual_host_config_layer - virtual host config missing virtual_host_id = {virtual_host_id} virtual_hosts = {virtual_hosts}" + ); + return server_not_found_response(); + }; + + if let Some(claims) = request.extensions().get::() { + let gateway_context = GatewayRequestContext::new(claims, virtual_host_id, virtual_host); + return scope_gateway_request_context(gateway_context, next.run(request)).await; + } + let virtual_host_id = virtual_host_id.value(); - let virtual_hosts = user_config.virtual_hosts.len(); - debug!( - "virtual_host_config_layer - virtual host config missing virtual_host_id = {virtual_host_id} virtual_hosts = {virtual_hosts}" - ); - return server_not_found_response(); + debug!("virtual_host_config_layer - claims missing virtual_host_id = {virtual_host_id}"); } next.run(request).await } -fn has_virtual_host(user_config: &UserConfig, virtual_host_id: &VirtualHostId) -> bool { - user_config.virtual_hosts.contains_key(virtual_host_id.value()) -} - fn server_not_found_response() -> Response { Response::builder() .status(StatusCode::NOT_FOUND) diff --git a/crates/contextforge-data-plane-lib/src/lib.rs b/crates/contextforge-data-plane-lib/src/lib.rs index 4a7f9b32..69b99d01 100644 --- a/crates/contextforge-data-plane-lib/src/lib.rs +++ b/crates/contextforge-data-plane-lib/src/lib.rs @@ -22,7 +22,7 @@ mod tools; mod user_config_store; pub use common::{RedisClient, RedisConfig, UpstreamConnectionMode}; -use gateway::{BackendTransports, McpService}; +use gateway::{BackendTransports, DownstreamSubscriptionRegistry, McpService}; use layers::session_id::SessionId; use tower_http::cors::{Any, CorsLayer}; use tower_http::trace::TraceLayer; @@ -77,6 +77,7 @@ impl Gateway { let user_session_store = LocalUserSessionStore::new(); let backend_transports = BackendTransports::default(); + let downstream_subscriptions = DownstreamSubscriptionRegistry::default(); let session_id_state = SessionIdState { user_session_store: Arc::new(user_session_store.clone()), backend_transports: backend_transports.clone(), @@ -99,12 +100,15 @@ impl Gateway { let mcp_service: StreamableHttpService, LocalSessionManager> = StreamableHttpService::new( move || { + let gateway_context = layers::request_context::current_gateway_request_context(); Ok(McpService::builder() .with_user_session_store(user_session_store.clone()) .with_http_client(reqwest_backend_client.clone()) .with_transports(backend_transports.clone()) .with_plugin_runtime(mcp_plugin_runtime.clone()) - .build()) + .build() + .with_gateway_request_context(gateway_context) + .with_downstream_subscription_registry(downstream_subscriptions.clone())) }, session_manager, streamable_config, diff --git a/crates/contextforge-data-plane-lib/tests/gateway_modern_subscriptions.rs b/crates/contextforge-data-plane-lib/tests/gateway_modern_subscriptions.rs new file mode 100644 index 00000000..59829d51 --- /dev/null +++ b/crates/contextforge-data-plane-lib/tests/gateway_modern_subscriptions.rs @@ -0,0 +1,222 @@ +mod support; + +use std::time::{Duration, Instant}; + +use contextforge_data_plane_lib::Result; +use rmcp::{ + ClientLifecycleMode, ClientServiceExt, + model::{ClientInfo, ErrorCode, ProtocolVersion, SubscriptionFilter}, + transport::{StreamableHttpClientTransport, streamable_http_client::StreamableHttpClientTransportConfig}, +}; +use serde_json::{Value, json}; + +use support::{ + ListToolsGatewaySettings, TEST_USER_ID, create_client, create_gateway_with_four_counters, create_ports, + plaintext_config, +}; + +const MODERN_PROTOCOL_VERSION: &str = "2026-07-28"; +const LEGACY_PROTOCOL_VERSION: &str = "2025-11-25"; + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +#[test_log::test] +async fn modern_discover_reports_context_capabilities_and_listen_acknowledges_filter() -> Result<()> { + let gateway_port = create_ports(1)[0]; + let user = TEST_USER_ID; + let Ok(ListToolsGatewaySettings { handle, gateway_url, .. }) = + create_gateway_with_four_counters(user, plaintext_config(gateway_port)).await + else { + panic!("invalid test gateway configuration"); + }; + wait_for_gateway_port(gateway_port).await; + + let maybe_passed = assert_modern_discover_and_listen(gateway_url, user).await; + + handle.abort(); + maybe_passed +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +#[test_log::test] +async fn modern_endpoint_rejects_legacy_protocol_and_legacy_methods() -> Result<()> { + let gateway_port = create_ports(1)[0]; + let user = TEST_USER_ID; + let Ok(ListToolsGatewaySettings { handle, gateway_url, .. }) = + create_gateway_with_four_counters(user, plaintext_config(gateway_port)).await + else { + panic!("invalid test gateway configuration"); + }; + wait_for_gateway_port(gateway_port).await; + + let maybe_passed = assert_rmcp_subscription_method_gating(&gateway_url, user).await; + + handle.abort(); + maybe_passed +} + +async fn assert_modern_discover_and_listen(gateway_url: String, user: &str) -> Result<()> { + let client = create_client(user); + let discover = post_raw_mcp( + &client, + &gateway_url, + MODERN_PROTOCOL_VERSION, + json!({ + "jsonrpc": "2.0", + "id": 1, + "method": "server/discover", + "params": { + "_meta": request_meta(MODERN_PROTOCOL_VERSION) + } + }), + ) + .await?; + assert_eq!(json!([MODERN_PROTOCOL_VERSION]), discover["result"]["supportedVersions"]); + + let transport = + StreamableHttpClientTransport::with_client(client, StreamableHttpClientTransportConfig::with_uri(gateway_url)); + let running_service = ClientInfo::default() + .serve_with_lifecycle( + transport, + ClientLifecycleMode::Discover { preferred_versions: vec![ProtocolVersion::V_2026_07_28] }, + ) + .await?; + let peer_info = running_service.peer_info().expect("discover lifecycle sets peer info"); + assert_eq!(ProtocolVersion::V_2026_07_28, peer_info.protocol_version); + assert_eq!(Some(true), peer_info.capabilities.tools.as_ref().and_then(|tools| tools.list_changed)); + assert_eq!(Some(true), peer_info.capabilities.prompts.as_ref().and_then(|prompts| prompts.list_changed)); + let resources = peer_info.capabilities.resources.as_ref().expect("resources are advertised"); + assert_eq!(Some(true), resources.subscribe); + assert_eq!(Some(true), resources.list_changed); + + let mut subscription = running_service + .listen( + SubscriptionFilter::builder() + .tools_list_changed() + .resources_list_changed() + .resource_subscription("unroutable://resource") + .build(), + ) + .await?; + let acknowledged = subscription.acknowledged(); + assert_eq!(Some(true), acknowledged.tools_list_changed); + assert_eq!(Some(true), acknowledged.resources_list_changed); + assert_eq!(None, acknowledged.resource_subscriptions); + + subscription.cancel().await?; + running_service.cancel().await?; + Ok(()) +} + +async fn assert_rmcp_subscription_method_gating(gateway_url: &str, user: &str) -> Result<()> { + let client = create_client(user); + + let modern_subscribe = post_raw_mcp( + &client, + gateway_url, + MODERN_PROTOCOL_VERSION, + json!({ + "jsonrpc": "2.0", + "id": 1, + "method": "resources/subscribe", + "params": { + "uri": "memo://known", + "_meta": request_meta(MODERN_PROTOCOL_VERSION) + } + }), + ) + .await?; + assert_jsonrpc_error_code(&modern_subscribe, i64::from(ErrorCode::METHOD_NOT_FOUND.0)); + + let legacy_listen = post_raw_mcp( + &client, + gateway_url, + LEGACY_PROTOCOL_VERSION, + json!({ + "jsonrpc": "2.0", + "id": 2, + "method": "subscriptions/listen", + "params": { + "notifications": { + "toolsListChanged": true + }, + "_meta": request_meta(LEGACY_PROTOCOL_VERSION) + } + }), + ) + .await?; + assert_jsonrpc_error_code(&legacy_listen, i64::from(ErrorCode::UNSUPPORTED_PROTOCOL_VERSION.0)); + + Ok(()) +} + +async fn post_raw_mcp( + client: &reqwest::Client, + gateway_url: &str, + protocol_version: &str, + body: Value, +) -> Result { + let method = body.get("method").and_then(Value::as_str).expect("test request has method"); + let mut request = client + .post(gateway_url) + .header(http::header::CONTENT_TYPE, "application/json") + .header(http::header::ACCEPT, "application/json, text/event-stream") + .header("MCP-Protocol-Version", protocol_version) + .header("Mcp-Method", method) + .json(&body); + if let Some(name) = + body.get("params").and_then(|params| params.get("name").or_else(|| params.get("uri"))).and_then(Value::as_str) + { + request = request.header("Mcp-Name", name); + } + let response = request.send().await?; + let status = response.status(); + let body = response.text().await?; + let values = response_values(&body)?; + assert!(!values.is_empty(), "expected JSON-RPC response body for status {status}, got empty body"); + Ok(values.into_iter().next().expect("non-empty response values")) +} + +fn request_meta(protocol_version: &str) -> Value { + json!({ + "io.modelcontextprotocol/protocolVersion": protocol_version, + "io.modelcontextprotocol/clientInfo": { + "name": "modern-subscription-test-client", + "version": "0.1.0" + }, + "io.modelcontextprotocol/clientCapabilities": {} + }) +} + +fn response_values(body: &str) -> Result> { + let values = body + .lines() + .filter_map(|line| line.strip_prefix("data:")) + .map(str::trim) + .filter(|data| !data.is_empty()) + .map(serde_json::from_str) + .collect::, _>>()?; + if !values.is_empty() { + return Ok(values); + } + + Ok(vec![serde_json::from_str(body.trim())?]) +} + +fn assert_jsonrpc_error_code(response: &Value, expected_code: i64) { + assert_eq!( + Some(expected_code), + response.pointer("/error/code").and_then(Value::as_i64), + "unexpected JSON-RPC response: {response}" + ); +} + +async fn wait_for_gateway_port(port: u16) { + let deadline = Instant::now() + Duration::from_secs(5); + loop { + if tokio::net::TcpStream::connect(("127.0.0.1", port)).await.is_ok() { + return; + } + assert!(Instant::now() < deadline, "gateway did not start on port {port}"); + tokio::time::sleep(Duration::from_millis(20)).await; + } +}