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
363 changes: 284 additions & 79 deletions crates/contextforge-gateway-rs-lib/src/gateway/list_aggregation.rs

Large diffs are not rendered by default.

Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@ use tracing::info;
use super::McpService;
use crate::gateway::{
identifier_routing::{backend_forward_error, route_identifier_to_backend},
list_aggregation::{fan_out_list, merge_prompts},
list_aggregation::{decode_gateway_cursor, fan_out_list, merge_prompts},
mcp_call_validator::AuthorizedCallValidator,
session_manager::SessionManager,
session_store::UserSessionStore,
Expand All @@ -27,20 +27,41 @@ where
let namespace_identifiers = virtual_host.backends.len() > 1;

let session_manager = SessionManager::new(virtual_host, session_id, claims.sub.as_str(), &mcp_service.transports);
let backend_transports: Vec<_> = session_manager.borrow_transports().await;
let all_transports: Vec<_> = session_manager.borrow_transports().await;

let gateway_cursor = decode_gateway_cursor(request.as_ref().and_then(|r| r.cursor.as_deref()), "list_prompts")?;
let backend_transports: Vec<_> = if request.as_ref().and_then(|r| r.cursor.as_ref()).is_some() {
all_transports.into_iter().filter(|b| gateway_cursor.backends.contains_key(&b.name)).collect()
} else {
all_transports
};

let responses = fan_out_list(
backend_transports,
"list_prompts",
|response: &ListPromptsResult| response.prompts.len(),
|service| {
let request = request.clone();
async move { service.list_prompts(request).await }
|name, service| {
let cursor = gateway_cursor.backends.get(&name).cloned();
let req = request.clone();
async move {
let backend_req = match cursor {
Some(c) => {
let mut r = req.unwrap_or_default();
r.cursor = Some(c);
Some(r)
},
None => req,
};
service.list_prompts(backend_req).await
}
},
)
.await;

Ok(ListPromptsResult::with_all_items(merge_prompts(responses, namespace_identifiers)))
let (prompts, next_cursor) = merge_prompts(responses, namespace_identifiers, &gateway_cursor, "list_prompts");
let mut result = ListPromptsResult::with_all_items(prompts);
result.next_cursor = next_cursor;
Ok(result)
}

pub(super) async fn get_prompt<T>(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,7 @@ use tracing::info;
use super::McpService;
use crate::gateway::{
identifier_routing::{backend_forward_error, route_identifier_to_backend},
list_aggregation::{fan_out_list, merge_resource_templates, merge_resources},
list_aggregation::{decode_gateway_cursor, fan_out_list, merge_resource_templates, merge_resources},
mcp_call_validator::AuthorizedCallValidator,
session_manager::SessionManager,
session_store::UserSessionStore,
Expand All @@ -30,20 +30,41 @@ where
let namespace_identifiers = virtual_host.backends.len() > 1;

let session_manager = SessionManager::new(virtual_host, session_id, claims.sub.as_str(), &mcp_service.transports);
let backend_transports: Vec<_> = session_manager.borrow_transports().await;
let all_transports: Vec<_> = session_manager.borrow_transports().await;

let gateway_cursor = decode_gateway_cursor(request.as_ref().and_then(|r| r.cursor.as_deref()), "list_resources")?;
let backend_transports: Vec<_> = if request.as_ref().and_then(|r| r.cursor.as_ref()).is_some() {
all_transports.into_iter().filter(|b| gateway_cursor.backends.contains_key(&b.name)).collect()
} else {
all_transports
};

let responses = fan_out_list(
backend_transports,
"list_resources",
|response: &ListResourcesResult| response.resources.len(),
|service| {
let request = request.clone();
async move { service.list_resources(request).await }
|name, service| {
let cursor = gateway_cursor.backends.get(&name).cloned();
let req = request.clone();
async move {
let backend_req = match cursor {
Some(c) => {
let mut r = req.unwrap_or_default();
r.cursor = Some(c);
Some(r)
},
None => req,
};
service.list_resources(backend_req).await
}
},
)
.await;

Ok(ListResourcesResult::with_all_items(merge_resources(responses, namespace_identifiers)))
let (resources, next_cursor) = merge_resources(responses, namespace_identifiers, &gateway_cursor, "list_resources");
let mut result = ListResourcesResult::with_all_items(resources);
result.next_cursor = next_cursor;
Ok(result)
}

pub(super) async fn read_resource<T>(
Expand Down Expand Up @@ -89,20 +110,43 @@ where
let namespace_identifiers = virtual_host.backends.len() > 1;

let session_manager = SessionManager::new(virtual_host, session_id, claims.sub.as_str(), &mcp_service.transports);
let backend_transports: Vec<_> = session_manager.borrow_transports().await;
let all_transports: Vec<_> = session_manager.borrow_transports().await;

let gateway_cursor =
decode_gateway_cursor(request.as_ref().and_then(|r| r.cursor.as_deref()), "list_resource_templates")?;
let backend_transports: Vec<_> = if request.as_ref().and_then(|r| r.cursor.as_ref()).is_some() {
all_transports.into_iter().filter(|b| gateway_cursor.backends.contains_key(&b.name)).collect()
} else {
all_transports
};

let responses = fan_out_list(
backend_transports,
"list_resource_templates",
|response: &ListResourceTemplatesResult| response.resource_templates.len(),
|service| {
let request = request.clone();
async move { service.list_resource_templates(request).await }
|name, service| {
let cursor = gateway_cursor.backends.get(&name).cloned();
let req = request.clone();
async move {
let backend_req = match cursor {
Some(c) => {
let mut r = req.unwrap_or_default();
r.cursor = Some(c);
Some(r)
},
None => req,
};
service.list_resource_templates(backend_req).await
}
},
)
.await;

Ok(ListResourceTemplatesResult::with_all_items(merge_resource_templates(responses, namespace_identifiers)))
let (resource_templates, next_cursor) =
merge_resource_templates(responses, namespace_identifiers, &gateway_cursor, "list_resource_templates");
let mut result = ListResourceTemplatesResult::with_all_items(resource_templates);
result.next_cursor = next_cursor;
Ok(result)
}

#[expect(deprecated, reason = "temporary RMCP v3 compatibility; subscriptions/listen migration is deferred")]
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,7 @@ use super::McpService;
use crate::gateway::{
backend_client::call_backend_tool,
identifier_routing::{backend_forward_error, resolve_backend, resolve_tool_route},
list_aggregation::{fan_out_list, merge_tools},
list_aggregation::{decode_gateway_cursor, fan_out_list, merge_tools},
mcp_call_validator::AuthorizedCallValidator,
session_manager::SessionManager,
session_store::UserSessionStore,
Expand All @@ -27,20 +27,44 @@ where
let mcp_call_validator = AuthorizedCallValidator::new("list_tools", &cx);
let (virtual_host, session_id, claims) = mcp_call_validator.validate()?;
let session_manager = SessionManager::new(virtual_host, session_id, claims.sub.as_str(), &mcp_service.transports);
let backend_transports: Vec<_> = session_manager.borrow_transports().await;
let all_transports: Vec<_> = session_manager.borrow_transports().await;

let gateway_cursor = decode_gateway_cursor(request.as_ref().and_then(|r| r.cursor.as_deref()), "list_tools")?;
// On resume, skip backends already exhausted in the prior page.
// ponytail: topology changes between pages silently drop removed backends;
// add a cursor version field if reconfiguration stability matters.
let backend_transports: Vec<_> = if request.as_ref().and_then(|r| r.cursor.as_ref()).is_some() {
all_transports.into_iter().filter(|b| gateway_cursor.backends.contains_key(&b.name)).collect()
} else {
all_transports
};

let responses = fan_out_list(
backend_transports,
"list_tools",
|response: &ListToolsResult| response.tools.len(),
|service| {
let request = request.clone();
async move { service.list_tools(request).await }
|name, service| {
let cursor = gateway_cursor.backends.get(&name).cloned();
let req = request.clone();
async move {
let backend_req = match cursor {
Some(c) => {
let mut r = req.unwrap_or_default();
r.cursor = Some(c);
Some(r)
},
None => req,
};
service.list_tools(backend_req).await
}
},
)
.await;

Ok(ListToolsResult::with_all_items(merge_tools(responses, virtual_host)))
let (tools, next_cursor) = merge_tools(responses, virtual_host, &gateway_cursor, "list_tools");
let mut result = ListToolsResult::with_all_items(tools);
result.next_cursor = next_cursor;
Ok(result)
}

pub(super) async fn call_tool<T>(
Expand Down
174 changes: 174 additions & 0 deletions crates/contextforge-gateway-rs-lib/tests/gateway_pagination.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,174 @@
mod support;

use std::{collections::HashMap, sync::Arc};

use contextforge_gateway_rs_apis::{
User,
user_store::{BackendMCPGateway, Transport, UserConfig, VirtualHost},
};
use contextforge_gateway_rs_lib::{Config, Gateway, Result, UserConfigStore, UserConfigStoreType};
use rmcp::{
model::PaginatedRequestParams,
transport::{
StreamableHttpServerConfig, StreamableHttpService, streamable_http_server::session::local::LocalSessionManager,
},
};
use tracing::warn;

use support::{
MemoryUserConfigStore, TEST_USER_ID, connect_client, create_client, create_ports, paginating_mock, plaintext_config,
};

/// Build a single-backend `BackendMCPGateway` pointed at `port`.
fn paginating_backend(port: u16) -> BackendMCPGateway {
BackendMCPGateway {
name: format!("backend-{port}"),
url: format!("http://127.0.0.1:{port}/mcp").parse().expect("valid url"),
transport: Transport::default(),
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 backend_id(port: u16) -> String {
format!("00000000-0000-0000-0000-{port:012}")
}

/// Bind the TCP port for a backend; returns the ready listener.
/// Call this *before* `tokio::spawn` so the port is reserved before the test proceeds.
async fn bind_backend_port(port: u16) -> tokio::net::TcpListener {
tokio::net::TcpListener::bind(format!("127.0.0.1:{port}")).await.expect("bind backend")
}

/// Start an axum MCP server on an already-bound listener serving a `PaginatingServer`.
async fn serve_paginating_backend(listener: tokio::net::TcpListener) {
let service = StreamableHttpService::new(
|| Ok(paginating_mock::PaginatingServer),
LocalSessionManager::default().into(),
StreamableHttpServerConfig::default(),
);
let router = axum::Router::new().route_service("/mcp", service);
axum::serve(listener, router).await.expect("backend server");
}

/// Boot the gateway with the given config and user config; return the gateway URL.
async fn start_gateway(config: Config, virtual_host_id: &str, user_config: UserConfig) -> String {
let store = MemoryUserConfigStore::default();
store.set_config(&User::new(TEST_USER_ID), &user_config).await.expect("set config");

let address = config.address.expect("address required");
let gateway_url = format!("http://{address}/contextforge-rs/servers/{virtual_host_id}/mcp");

let gateway = Gateway::builder()
.with_config(config)
.with_session_manager(Arc::new(LocalSessionManager::default()))
.with_user_config_store_type(UserConfigStoreType::Test(Arc::new(store)))
.build();

tokio::spawn(async move {
let res = gateway.run_gateway().await;
warn!("Gateway exited {res:?}");
});

gateway_url
}

/// A paginating backend returns tools across two pages; the gateway must expose
/// all of them to the client without any items being silently dropped.
#[tokio::test(flavor = "multi_thread", worker_threads = 1)]
#[test_log::test]
async fn single_backend_pagination_all_tools_reachable() -> Result<()> {
let ports = create_ports(2);
let (backend_port, gateway_port) = (ports[0], ports[1]);
let config = plaintext_config(gateway_port);

let virtual_host_id = "22222222-2222-2222-2222-222222222222";
let backends = HashMap::from([(backend_id(backend_port), paginating_backend(backend_port))]);
let user_config =
UserConfig { virtual_hosts: HashMap::from([(virtual_host_id.to_owned(), VirtualHost { backends })]) };

let backend_listener = bind_backend_port(backend_port).await;
tokio::spawn(serve_paginating_backend(backend_listener));
let gateway_url = start_gateway(config, virtual_host_id, user_config).await;

let svc = connect_client(gateway_url, create_client(TEST_USER_ID)).await?;

// Page 1
let page1 = svc.list_tools(None).await.expect("page 1");
let page1_names: Vec<&str> = page1.tools.iter().map(|t| t.name.as_ref()).collect();
assert!(page1.next_cursor.is_some(), "page 1 must carry a next_cursor");
assert_eq!(page1_names, ["tool_alpha", "tool_beta"]);

// Page 2
let cursor = page1.next_cursor.map(|c| PaginatedRequestParams::default().with_cursor(Some(c)));
let page2 = svc.list_tools(cursor).await.expect("page 2");
let page2_names: Vec<&str> = page2.tools.iter().map(|t| t.name.as_ref()).collect();
assert!(page2.next_cursor.is_none(), "page 2 must be the final page");
assert_eq!(page2_names, ["tool_gamma"]);

// All tools reachable with no duplication
let mut all_names = page1_names.clone();
all_names.extend_from_slice(&page2_names);
all_names.sort_unstable();
assert_eq!(all_names, paginating_mock::PaginatingServer::all_tool_names());

Ok(())
}

/// When one backend exhausts its pages, it must be excluded from the resume
/// request. Without the filter, the exhausted backend would be re-queried and
/// its tools would appear in every subsequent page as duplicates.
#[tokio::test(flavor = "multi_thread", worker_threads = 1)]
#[test_log::test]
async fn multi_backend_exhausted_backend_not_requeried() -> Result<()> {
// Backend A: PaginatingServer (2 pages: 2 tools + 1 tool)
// Backend B: another PaginatingServer (same 2 pages, different backend ID)
//
// With 2 backends, tool names get the backend-ID prefix.
// Page 1: 2 tools from A page 1 + 2 tools from B page 1 = 4 total
// Page 2: 1 tool from A page 2 + 1 tool from B page 2 = 2 total
// If either backend were re-queried, its page-1 tools would reappear.
let ports = create_ports(3);
let (port_a, port_b, gateway_port) = (ports[0], ports[1], ports[2]);
let config = plaintext_config(gateway_port);

let virtual_host_id = "33333333-3333-3333-3333-333333333333";
let backends = HashMap::from([
(backend_id(port_a), paginating_backend(port_a)),
(backend_id(port_b), paginating_backend(port_b)),
]);
let user_config =
UserConfig { virtual_hosts: HashMap::from([(virtual_host_id.to_owned(), VirtualHost { backends })]) };

let listener_a = bind_backend_port(port_a).await;
let listener_b = bind_backend_port(port_b).await;
tokio::spawn(serve_paginating_backend(listener_a));
tokio::spawn(serve_paginating_backend(listener_b));
let gateway_url = start_gateway(config, virtual_host_id, user_config).await;

let svc = connect_client(gateway_url, create_client(TEST_USER_ID)).await?;

// Page 1: both backends contribute their first page (2 tools each)
let page1 = svc.list_tools(None).await.expect("page 1");
assert!(page1.next_cursor.is_some(), "page 1 must carry a next_cursor");
assert_eq!(page1.tools.len(), 4, "page 1 should have 2 tools from each backend");

// Page 2: both backends contribute their second page (1 tool each)
let cursor = page1.next_cursor.map(|c| PaginatedRequestParams::default().with_cursor(Some(c)));
let page2 = svc.list_tools(cursor).await.expect("page 2");
assert!(page2.next_cursor.is_none(), "page 2 must be the final page");
assert_eq!(page2.tools.len(), 2, "page 2 should have 1 tool from each backend");

// Union has 6 unique tools, no duplicates
let mut all_names: Vec<_> = page1.tools.iter().chain(page2.tools.iter()).map(|t| t.name.clone()).collect();
all_names.sort_unstable();
all_names.dedup();
assert_eq!(all_names.len(), page1.tools.len() + page2.tools.len(), "no duplicate tools across pages");

Ok(())
}
Loading
Loading