mod interaction;
#[cfg(test)]
mod tests;
use std::sync::Arc;
use std::sync::atomic::AtomicU64;
use asupersync::Cx;
use fastmcp_client::http_auth::discovery::OAuthDiscoveryError;
use fastmcp_client::http_auth::discovery::client_credentials::rpc::catalog::{
ClientCredentialsCatalogClient, ClientCredentialsCatalogError, ClientCredentialsCatalogLimits,
};
use fastmcp_client::http_auth::discovery::client_credentials::rpc::{
ClientCredentialsCoreCall, ClientCredentialsCoreError,
};
use fastmcp_client::http_auth::discovery::client_credentials::{
ClientCredentialsClient, ClientCredentialsError,
};
use fastmcp_client::http_auth::rpc::catalog::ManagedCatalogError;
use fastmcp_client::http_auth::rpc::{ManagedCoreEvent, ManagedCoreLimits};
use fastmcp_core::{McpContext, McpError, McpResult};
use fastmcp_protocol::{CoreRequest, CoreResult, FinalCoreResult, RequestId};
use serde_json::json;
use super::super::{
BoxFuture, CoreBackend, Forwarder, ManagedOAuthPrompt, ManagedOAuthResource, ManagedOAuthTool,
UNEXPECTED_RESULT, allocate_request_id, build_prompts, build_resources, build_tools, check_cx,
core_request, forward_notification, upstream_error,
};
use super::{ManagedOAuthResourceTemplate, build_templates};
use interaction::{InteractiveMachineResponse, MachineInputs, interaction_error};
const MACHINE_FAILURE: &str = "Machine-authenticated upstream request failed";
const CATALOG_FAILURE: &str = "Machine-authenticated upstream catalog acquisition failed";
#[derive(Clone)]
pub struct ClientCredentialsProvider {
source: Arc<dyn MachineSource>,
forwarder: Arc<Forwarder>,
catalog_limits: ClientCredentialsCatalogLimits,
namespace: Option<String>,
}
impl ClientCredentialsProvider {
pub fn new(client: ClientCredentialsClient) -> Self {
Self::from_source(Arc::new(NativeMachineSource(client)))
}
fn from_source(source: Arc<dyn MachineSource>) -> Self {
let next_id = Arc::new(AtomicU64::new(1));
Self {
forwarder: Arc::new(Forwarder {
backend: Arc::new(MachineBackend {
source: Arc::clone(&source),
next_id: Arc::clone(&next_id),
inputs: None,
}),
next_id,
limits: ManagedCoreLimits::default(),
}),
source,
catalog_limits: ClientCredentialsCatalogLimits::default(),
namespace: None,
}
}
pub fn with_limits(
mut self,
calls: ManagedCoreLimits,
catalogs: ClientCredentialsCatalogLimits,
) -> Self {
self.forwarder = Arc::new(Forwarder {
backend: Arc::clone(&self.forwarder.backend),
next_id: Arc::clone(&self.forwarder.next_id),
limits: calls,
});
self.catalog_limits = catalogs;
self
}
pub fn with_namespace(mut self, namespace: impl Into<String>) -> McpResult<Self> {
let namespace = namespace.into();
if namespace.is_empty()
|| namespace.len() > 64
|| !namespace
.bytes()
.all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'_' | b'-' | b'.'))
{
return Err(McpError::invalid_params(
"Invalid machine provider namespace",
));
}
self.namespace = Some(namespace);
Ok(self)
}
pub async fn tools(&self, cx: &Cx) -> McpResult<Vec<ManagedOAuthTool>> {
let mut entries = Vec::new();
for page in self.catalog(cx, "tools/list").await? {
let CoreResult::Final(FinalCoreResult::ToolsList { result, .. }) = page else {
return Err(McpError::invalid_request(UNEXPECTED_RESULT));
};
entries.extend(result.payload.tools);
}
build_tools(
Arc::clone(&self.forwarder),
entries,
self.namespace.as_deref(),
)
}
pub async fn resources(&self, cx: &Cx) -> McpResult<Vec<ManagedOAuthResource>> {
let mut entries = Vec::new();
for page in self.catalog(cx, "resources/list").await? {
let CoreResult::Final(FinalCoreResult::ResourcesList { result, .. }) = page else {
return Err(McpError::invalid_request(UNEXPECTED_RESULT));
};
entries.extend(result.payload.resources);
}
build_resources(Arc::clone(&self.forwarder), entries)
}
pub async fn prompts(&self, cx: &Cx) -> McpResult<Vec<ManagedOAuthPrompt>> {
let mut entries = Vec::new();
for page in self.catalog(cx, "prompts/list").await? {
let CoreResult::Final(FinalCoreResult::PromptsList { result, .. }) = page else {
return Err(McpError::invalid_request(UNEXPECTED_RESULT));
};
entries.extend(result.payload.prompts);
}
build_prompts(
Arc::clone(&self.forwarder),
entries,
self.namespace.as_deref(),
)
}
pub async fn resource_templates(
&self,
cx: &Cx,
) -> McpResult<Vec<ManagedOAuthResourceTemplate>> {
let mut entries = Vec::new();
for page in self.catalog(cx, "resources/templates/list").await? {
let CoreResult::Final(FinalCoreResult::ResourceTemplatesList { result, .. }) = page
else {
return Err(McpError::invalid_request(UNEXPECTED_RESULT));
};
entries.extend(result.payload.resource_templates);
}
build_templates(Arc::clone(&self.forwarder), entries)
}
async fn catalog(&self, cx: &Cx, method: &'static str) -> McpResult<Vec<CoreResult>> {
check_cx(cx)?;
let result = self
.source
.collect(cx, method, &self.forwarder.next_id, self.catalog_limits)
.await;
check_cx(cx)?;
result
}
}
struct MachineCall {
request: CoreRequest,
discovery_id: RequestId,
request_id: RequestId,
limits: ManagedCoreLimits,
inputs: Option<Arc<MachineInputs>>,
}
trait MachineSource: Send + Sync {
fn collect<'a>(
&'a self,
cx: &'a Cx,
method: &'static str,
ids: &'a AtomicU64,
limits: ClientCredentialsCatalogLimits,
) -> BoxFuture<'a, McpResult<Vec<CoreResult>>>;
fn start<'a>(
&'a self,
ctx: &'a McpContext,
cx: &'a Cx,
call: MachineCall,
) -> BoxFuture<'a, McpResult<Box<dyn MachineResponse>>>;
}
trait MachineResponse: Send {
fn next_event<'a>(
&'a mut self,
cx: &'a Cx,
) -> BoxFuture<'a, McpResult<Option<ManagedCoreEvent>>>;
}
struct NativeMachineSource(ClientCredentialsClient);
impl MachineSource for NativeMachineSource {
fn collect<'a>(
&'a self,
cx: &'a Cx,
method: &'static str,
ids: &'a AtomicU64,
limits: ClientCredentialsCatalogLimits,
) -> BoxFuture<'a, McpResult<Vec<CoreResult>>> {
Box::pin(async move {
let request = core_request(method, json!({}), None)?;
let collector = ClientCredentialsCatalogClient::new(self.0.clone(), limits);
let pages = collector
.collect(
cx,
request,
|| {
next_pair(ids).map_err(|_| {
ClientCredentialsCatalogError::Catalog(
ManagedCatalogError::AbortedByHost,
)
})
},
|_| {
Err(ClientCredentialsCatalogError::Catalog(
ManagedCatalogError::AbortedByHost,
))
},
)
.await
.map_err(catalog_error)?;
Ok(pages.into_pages())
})
}
fn start<'a>(
&'a self,
ctx: &'a McpContext,
cx: &'a Cx,
call: MachineCall,
) -> BoxFuture<'a, McpResult<Box<dyn MachineResponse>>> {
Box::pin(async move {
let MachineCall {
request,
discovery_id,
request_id,
limits,
inputs,
} = call;
let cancellation = ctx.request_cancellation();
if let Some(inputs) = inputs {
let operation = self
.0
.start_core_interaction_with_cancellation(
cx,
&cancellation,
request,
discovery_id,
request_id,
inputs.policy.limits(limits)?,
)
.await
.map_err(interaction_error)?;
return Ok(Box::new(InteractiveMachineResponse::new(
operation,
ctx.clone().with_request_cx(cx.clone()),
inputs,
)) as Box<dyn MachineResponse>);
}
let call = self
.0
.request_core_with_cancellation(
cx,
&cancellation,
request,
discovery_id,
request_id,
limits,
)
.await
.map_err(machine_error)?;
Ok(Box::new(NativeMachineResponse(call)) as Box<dyn MachineResponse>)
})
}
}
struct NativeMachineResponse(ClientCredentialsCoreCall);
impl MachineResponse for NativeMachineResponse {
fn next_event<'a>(
&'a mut self,
cx: &'a Cx,
) -> BoxFuture<'a, McpResult<Option<ManagedCoreEvent>>> {
Box::pin(async move { self.0.next_event(cx).await.map_err(machine_error) })
}
}
struct MachineBackend {
source: Arc<dyn MachineSource>,
next_id: Arc<AtomicU64>,
inputs: Option<Arc<MachineInputs>>,
}
impl CoreBackend for MachineBackend {
fn execute<'a>(
&'a self,
ctx: &'a McpContext,
cx: &'a Cx,
request: CoreRequest,
id: RequestId,
limits: ManagedCoreLimits,
) -> BoxFuture<'a, McpResult<FinalCoreResult>> {
Box::pin(async move {
ctx.checkpoint()?;
check_cx(cx)?;
let (request, inputs) = match &self.inputs {
Some(inputs) => match inputs.policy.select_request(&request)? {
Some(selected) => (selected, Some(Arc::clone(inputs))),
None => (request, None),
},
None => (request, None),
};
let discovery_id = allocate_request_id(&self.next_id)?;
let mut response = self
.source
.start(
ctx,
cx,
MachineCall {
request,
discovery_id,
request_id: id,
limits,
inputs,
},
)
.await?;
loop {
ctx.checkpoint()?;
check_cx(cx)?;
let event = response.next_event(cx).await?;
ctx.checkpoint()?;
check_cx(cx)?;
match event {
Some(ManagedCoreEvent::Notification(notification)) => {
forward_notification(ctx, *notification)?;
}
Some(ManagedCoreEvent::Result(result)) => match *result {
CoreResult::Final(result) => return Ok(result),
_ => return Err(McpError::invalid_request(UNEXPECTED_RESULT)),
},
None => return Err(McpError::invalid_request(MACHINE_FAILURE)),
}
}
})
}
}
fn next_pair(ids: &AtomicU64) -> McpResult<(RequestId, RequestId)> {
Ok((allocate_request_id(ids)?, allocate_request_id(ids)?))
}
fn machine_error(error: ClientCredentialsCoreError) -> McpError {
match error {
ClientCredentialsCoreError::Protocol(error) => upstream_error(error),
ClientCredentialsCoreError::Authentication(ClientCredentialsError::Discovery(
OAuthDiscoveryError::Cancelled | OAuthDiscoveryError::TimedOut,
)) => McpError::request_cancelled(),
ClientCredentialsCoreError::Authentication(_) => McpError::invalid_request(MACHINE_FAILURE),
}
}
fn catalog_error(error: ClientCredentialsCatalogError) -> McpError {
match error {
ClientCredentialsCatalogError::Core(error) => machine_error(error),
ClientCredentialsCatalogError::Catalog(_) => McpError::invalid_request(CATALOG_FAILURE),
}
}