Skip to content
Draft
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
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,10 @@ use super::{

/// Preserves identifiers for a single backend. For multiple backends, splits a
/// `{backend}-{identifier}` namespace so duplicate identifiers remain routable.
fn route_identifier<'a, N: AsRef<str>>(identifier: &'a str, backend_names: &'a [N]) -> Option<(&'a str, &'a str)> {
pub(super) fn route_identifier<'a, N: AsRef<str>>(
identifier: &'a str,
backend_names: &'a [N],
) -> Option<(&'a str, &'a str)> {
if let [backend] = backend_names {
return Some((backend.as_ref(), identifier));
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@ impl<'a> AuthorizedCallValidator<'a> {
pub fn new(call_name: &'a str, ctx: &'a RequestContext<RoleServer>) -> Self {
Self { call_name, ctx }
}

pub fn validate(self) -> Result<(&'a VirtualHost, &'a SessionId, &'a ContextForgeClaims), ErrorData> {
let maybe_parts = self.ctx.extensions.get::<Parts>();
let maybe_session_id = maybe_parts.and_then(|parts| parts.extensions.get::<SessionId>());
Expand Down Expand Up @@ -82,31 +83,20 @@ impl<'a> AuthorizedCallValidator<'a> {

Ok((virtual_host, session_id, claims))
}
}

pub struct InitializeCallValidator<'a> {
ctx: &'a RequestContext<RoleServer>,
}

impl<'a> InitializeCallValidator<'a> {
pub fn new(ctx: &'a RequestContext<RoleServer>) -> Self {
Self { ctx }
}
pub fn validate(self) -> Result<(&'a VirtualHost, SessionId, &'a ContextForgeClaims), ErrorData> {
// once session_id is removed, validate_stateless should be the "validate"
pub fn validate_stateless(self) -> Result<(&'a VirtualHost, &'a ContextForgeClaims), ErrorData> {
let maybe_parts = self.ctx.extensions.get::<Parts>();

let downstream_session_id = SessionId::mock();
let maybe_user_config = maybe_parts.and_then(|parts| parts.extensions.get::<UserConfig>());
let maybe_virtual_host_id = maybe_parts.and_then(|parts| parts.extensions.get::<VirtualHostId>());
let maybe_claims = maybe_parts.and_then(|parts| parts.extensions.get::<ContextForgeClaims>());
let call_name = "initialize";
let maybe_virtual_host_id = maybe_parts.and_then(|parts| parts.extensions.get::<VirtualHostId>());
let call_name = self.call_name;
let has_user_config = maybe_user_config.is_some();
let virtual_hosts = maybe_user_config.map_or(0, |user_config| user_config.virtual_hosts.len());
let has_session_id = true;
let has_claims = maybe_claims.is_some();
let virtual_host_id = maybe_virtual_host_id.map_or("<missing>", |id| id.value().as_str());
debug!(
"InitializeCallValidator::validate - mcp call validation call_name = {call_name} has_user_config = {has_user_config} virtual_hosts = {virtual_hosts} has_session_id = {has_session_id} has_claims = {has_claims} virtual_host_id = {virtual_host_id}"
"AuthorizedCallValidator::validate - mcp call validation call_name = {call_name} has_user_config = {has_user_config} virtual_hosts = {virtual_hosts} has_claims = {has_claims} virtual_host_id = {virtual_host_id}"
);

let Some(user_config) = maybe_user_config else {
Expand All @@ -126,11 +116,11 @@ impl<'a> InitializeCallValidator<'a> {
};

let Some(virtual_host) = user_config.virtual_hosts.get(virtual_host_id.value()) else {
let call_name = "initialize";
let call_name = self.call_name;
let virtual_host_id = virtual_host_id.value();
let virtual_hosts = user_config.virtual_hosts.len();
debug!(
"InitializeCallValidator::validate - mcp virtual host config missing call_name = {call_name} virtual_host_id = {virtual_host_id} virtual_hosts = {virtual_hosts}"
"AuthorizedCallValidator::validate - mcp virtual host config missing call_name = {call_name} virtual_host_id = {virtual_host_id} virtual_hosts = {virtual_hosts}"
);
return Err(ErrorData {
code: ErrorCode::RESOURCE_NOT_FOUND,
Expand All @@ -147,6 +137,6 @@ impl<'a> InitializeCallValidator<'a> {
});
};

Ok((virtual_host, downstream_session_id, claims))
Ok((virtual_host, claims))
}
}
33 changes: 14 additions & 19 deletions crates/contextforge-data-plane-lib/src/gateway/mcp_service.rs
Original file line number Diff line number Diff line change
Expand Up @@ -8,45 +8,40 @@ 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,
CallToolRequestParams, CallToolResponse, CompleteRequestParams, CompleteResult, ErrorCode,
GetPromptRequestParams, GetPromptResponse, InitializeRequestParams, InitializeResult, ListPromptsResult,
ListResourceTemplatesResult, ListResourcesResult, ListToolsResult, PaginatedRequestParams,
ReadResourceRequestParams, ReadResourceResponse, SubscribeRequestParams, UnsubscribeRequestParams,
},
service::RequestContext,
};
use typed_builder::TypedBuilder;

use super::{backend_transports::BackendTransports, session_store::UserSessionStore};
use super::{backend_transports::BackendTransports};

#[derive(Clone, TypedBuilder)]
#[builder(field_defaults(setter(prefix = "with_")))]
pub struct McpService<T>
where
T: UserSessionStore,
pub struct McpService
{
#[builder(default = BackendTransports::default())]
transports: BackendTransports,
http_client: reqwest::Client,
user_session_store: T,
#[builder(default)]
plugin_runtime: Option<GatewayPluginRuntimeHandle>,
}

impl<T> ServerHandler for McpService<T>
where
T: UserSessionStore + Send + Sync + 'static,
impl ServerHandler for McpService
{
async fn ping(&self, _cx: RequestContext<RoleServer>) -> Result<(), ErrorData> {
Ok(())
}

async fn initialize(
&self,
request: InitializeRequestParams,
cx: RequestContext<RoleServer>,
_request: InitializeRequestParams,
_cx: RequestContext<RoleServer>,
) -> Result<InitializeResult, ErrorData> {
initialization::initialize(self, request, cx).await
}

async fn ping(&self, _cx: RequestContext<RoleServer>) -> Result<(), ErrorData> {
Ok(())
Err(ErrorData { code: ErrorCode::METHOD_NOT_FOUND, message: "Initialize not supported".into(), data: None })
}

async fn list_tools(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -10,16 +10,13 @@ use crate::gateway::{
identifier_routing::{backend_forward_error, route_identifier_to_backend},
mcp_call_validator::AuthorizedCallValidator,
session_manager::SessionManager,
session_store::UserSessionStore,
};

pub(super) async fn complete<T>(
mcp_service: &McpService<T>,
pub(super) async fn complete(
mcp_service: &McpService,
request: CompleteRequestParams,
cx: RequestContext<RoleServer>,
) -> Result<CompleteResult, ErrorData>
where
T: UserSessionStore + Send + Sync + 'static,
{
let mcp_call_validator = AuthorizedCallValidator::new("complete", &cx);
let (virtual_host, session_id, claims) = mcp_call_validator.validate()?;
Expand Down
Loading
Loading