use r402::facilitator::BoxFuture;
use r402::proto;
use r402::scheme::{
ClientError, FirstMatch, PaymentCandidate, PaymentPolicy, PaymentSelector, SchemeClient,
};
use crate::PAYMENT_META_KEY;
use crate::error::{McpPaymentError, PaymentRequiredError};
use crate::extract;
use crate::types::{
AfterPaymentContext, BeforePaymentContext, CallToolParams, CallToolResult, ClientHooks,
ClientOptions, NoClientHooks, PaidToolCallResult, PaymentRequiredContext,
};
pub trait McpCaller: Send + Sync {
fn call_tool(
&self,
params: CallToolParams,
) -> BoxFuture<'_, Result<CallToolResult, McpPaymentError>>;
}
pub struct X402McpClient {
caller: Box<dyn McpCaller>,
scheme_clients: Vec<Box<dyn SchemeClient>>,
selector: Box<dyn PaymentSelector>,
policies: Vec<Box<dyn PaymentPolicy>>,
options: ClientOptions,
hooks: Box<dyn ClientHooks>,
}
impl std::fmt::Debug for X402McpClient {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("X402McpClient")
.field("scheme_clients", &self.scheme_clients.len())
.field("options", &self.options)
.finish_non_exhaustive()
}
}
impl X402McpClient {
pub fn builder(caller: impl McpCaller + 'static) -> X402McpClientBuilder {
X402McpClientBuilder {
caller: Box::new(caller),
scheme_clients: Vec::new(),
selector: None,
policies: Vec::new(),
options: ClientOptions::default(),
hooks: None,
}
}
#[must_use]
pub fn caller(&self) -> &dyn McpCaller {
&*self.caller
}
pub async fn call_tool(
&self,
name: &str,
arguments: serde_json::Map<String, serde_json::Value>,
) -> Result<PaidToolCallResult, McpPaymentError> {
let params = CallToolParams {
name: name.to_owned(),
arguments: arguments.clone(),
meta: None,
};
let result = self.caller.call_tool(params).await?;
if !result.is_error {
return Ok(build_paid_result(result, false));
}
let payment_required = match extract::extract_payment_required_from_result(&result) {
Some(pr) if !pr.accepts.is_empty() => pr,
_ => return Ok(build_paid_result(result, false)),
};
let pr_ctx = PaymentRequiredContext {
tool_name: name.to_owned(),
arguments: arguments.clone(),
payment_required: payment_required.clone(),
};
let custom_payment = self.hooks.on_payment_required(&pr_ctx).await?;
if let Some(payload) = custom_payment {
return self.call_tool_with_payload(name, arguments, payload).await;
}
if !self.options.auto_payment {
return Err(McpPaymentError::PaymentRequired(Box::new(
PaymentRequiredError::new("Payment required", payment_required),
)));
}
let approved = self.hooks.on_payment_requested(&pr_ctx).await?;
if !approved {
return Err(McpPaymentError::PaymentRequired(Box::new(
PaymentRequiredError::new("Payment denied by hook", payment_required),
)));
}
let before_ctx = BeforePaymentContext {
tool_name: name.to_owned(),
payment_required: payment_required.clone(),
};
self.hooks.on_before_payment(&before_ctx).await?;
let payload = self.create_payment(&payment_required).await?;
self.call_tool_with_payload(name, arguments, payload).await
}
pub async fn call_tool_with_payment(
&self,
name: &str,
arguments: serde_json::Map<String, serde_json::Value>,
payload: serde_json::Value,
) -> Result<PaidToolCallResult, McpPaymentError> {
self.call_tool_with_payload(name, arguments, payload).await
}
pub async fn get_tool_payment_requirements(
&self,
name: &str,
arguments: serde_json::Map<String, serde_json::Value>,
) -> Result<Option<proto::PaymentRequired>, McpPaymentError> {
let params = CallToolParams {
name: name.to_owned(),
arguments,
meta: None,
};
let result = self.caller.call_tool(params).await?;
Ok(extract::extract_payment_required_from_result(&result))
}
async fn call_tool_with_payload(
&self,
name: &str,
arguments: serde_json::Map<String, serde_json::Value>,
payload: serde_json::Value,
) -> Result<PaidToolCallResult, McpPaymentError> {
let mut meta = serde_json::Map::new();
meta.insert(PAYMENT_META_KEY.to_owned(), payload.clone());
let params = CallToolParams {
name: name.to_owned(),
arguments,
meta: Some(meta),
};
let result = self.caller.call_tool(params).await?;
let settle_response = result
.meta
.as_ref()
.and_then(extract::extract_payment_response_from_meta);
let after_ctx = AfterPaymentContext {
tool_name: name.to_owned(),
payment_payload: payload,
result: result.clone(),
settle_response: settle_response.clone(),
};
let _ = self.hooks.on_after_payment(&after_ctx).await;
Ok(build_paid_result(result, true))
}
async fn create_payment(
&self,
payment_required: &proto::PaymentRequired,
) -> Result<serde_json::Value, McpPaymentError> {
let mut candidates: Vec<PaymentCandidate> = Vec::new();
for client in &self.scheme_clients {
candidates.extend(client.accept(payment_required));
}
if candidates.is_empty() {
return Err(McpPaymentError::NoMatchingPaymentOption);
}
let mut refs: Vec<&PaymentCandidate> = candidates.iter().collect();
for policy in &self.policies {
refs = policy.apply(refs);
}
let selected = self
.selector
.select(&refs)
.ok_or(McpPaymentError::NoMatchingPaymentOption)?;
let signed_json = selected.sign().await.map_err(|e| match e {
ClientError::SigningError(msg) => McpPaymentError::SigningFailed(msg),
other => McpPaymentError::PaymentCreationFailed(other.to_string()),
})?;
let payload: serde_json::Value = serde_json::from_str(&signed_json)?;
Ok(payload)
}
}
pub struct X402McpClientBuilder {
caller: Box<dyn McpCaller>,
scheme_clients: Vec<Box<dyn SchemeClient>>,
selector: Option<Box<dyn PaymentSelector>>,
policies: Vec<Box<dyn PaymentPolicy>>,
options: ClientOptions,
hooks: Option<Box<dyn ClientHooks>>,
}
impl std::fmt::Debug for X402McpClientBuilder {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("X402McpClientBuilder")
.field("scheme_clients", &self.scheme_clients.len())
.field("options", &self.options)
.finish_non_exhaustive()
}
}
impl X402McpClientBuilder {
#[must_use]
pub fn scheme_client(mut self, client: Box<dyn SchemeClient>) -> Self {
self.scheme_clients.push(client);
self
}
#[must_use]
pub fn selector(mut self, selector: Box<dyn PaymentSelector>) -> Self {
self.selector = Some(selector);
self
}
#[must_use]
pub fn policy(mut self, policy: Box<dyn PaymentPolicy>) -> Self {
self.policies.push(policy);
self
}
#[must_use]
pub const fn options(mut self, options: ClientOptions) -> Self {
self.options = options;
self
}
#[must_use]
pub const fn auto_payment(mut self, enabled: bool) -> Self {
self.options.auto_payment = enabled;
self
}
#[must_use]
pub fn hooks(mut self, hooks: Box<dyn ClientHooks>) -> Self {
self.hooks = Some(hooks);
self
}
#[must_use]
pub fn build(self) -> X402McpClient {
assert!(
!self.scheme_clients.is_empty(),
"at least one scheme client must be registered"
);
X402McpClient {
caller: self.caller,
scheme_clients: self.scheme_clients,
selector: self.selector.unwrap_or_else(|| Box::new(FirstMatch)),
policies: self.policies,
options: self.options,
hooks: self.hooks.unwrap_or_else(|| Box::new(NoClientHooks)),
}
}
}
pub async fn call_paid_tool(
caller: &dyn McpCaller,
scheme_clients: &[&dyn SchemeClient],
name: &str,
arguments: serde_json::Map<String, serde_json::Value>,
) -> Result<PaidToolCallResult, McpPaymentError> {
let params = CallToolParams {
name: name.to_owned(),
arguments: arguments.clone(),
meta: None,
};
let result = caller.call_tool(params).await?;
if !result.is_error {
return Ok(build_paid_result(result, false));
}
let payment_required = match extract::extract_payment_required_from_result(&result) {
Some(pr) if !pr.accepts.is_empty() => pr,
_ => return Ok(build_paid_result(result, false)),
};
let mut candidates: Vec<PaymentCandidate> = Vec::new();
for client in scheme_clients {
candidates.extend(client.accept(&payment_required));
}
let refs: Vec<&PaymentCandidate> = candidates.iter().collect();
let selected = FirstMatch
.select(&refs)
.ok_or(McpPaymentError::NoMatchingPaymentOption)?;
let signed_json = selected.sign().await.map_err(|e| match e {
ClientError::SigningError(msg) => McpPaymentError::SigningFailed(msg),
other => McpPaymentError::PaymentCreationFailed(other.to_string()),
})?;
let payload: serde_json::Value = serde_json::from_str(&signed_json)?;
let mut meta = serde_json::Map::new();
meta.insert(PAYMENT_META_KEY.to_owned(), payload);
let params = CallToolParams {
name: name.to_owned(),
arguments,
meta: Some(meta),
};
let result = caller.call_tool(params).await?;
Ok(build_paid_result(result, true))
}
fn build_paid_result(result: CallToolResult, payment_made: bool) -> PaidToolCallResult {
let payment_response = result
.meta
.as_ref()
.and_then(extract::extract_payment_response_from_meta);
PaidToolCallResult {
content: result.content.clone(),
is_error: result.is_error,
payment_response,
payment_made,
raw_result: result,
}
}