1use std::future::Future;
10use std::sync::Arc;
11
12use r402_core::resource_server::ResourceServer;
13use r402_core::wire::{Extensions, PaymentRequired, PaymentRequirements, ResourceInfo};
14use rmcp::model::{CallToolRequestParams, CallToolResult};
15
16use crate::encode::{
17 McpPaymentPayload, attach_settle_response, extract_payment_from_params,
18 payment_required_tool_result, settlement_failed_tool_result,
19};
20
21#[derive(Debug, Clone, Copy, thiserror::Error)]
23pub enum PaymentWrapperConfigError {
24 #[error("PaymentWrapperConfig.accepts must have at least one payment requirement")]
26 EmptyAccepts,
27}
28
29#[derive(Clone, Default)]
31pub struct PaymentWrapperHooks {
32 pub on_before_execution: Option<Arc<dyn Fn(ServerHookContext) -> bool + Send + Sync>>,
34 pub on_after_execution: Option<Arc<dyn Fn(AfterExecutionContext) + Send + Sync>>,
36 pub on_after_settlement: Option<Arc<dyn Fn(SettlementContext) + Send + Sync>>,
38}
39
40impl std::fmt::Debug for PaymentWrapperHooks {
41 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
42 f.debug_struct("PaymentWrapperHooks")
43 .field("on_before_execution", &self.on_before_execution.is_some())
44 .field("on_after_execution", &self.on_after_execution.is_some())
45 .field("on_after_settlement", &self.on_after_settlement.is_some())
46 .finish()
47 }
48}
49
50#[derive(Debug, Clone)]
52pub struct ServerHookContext {
53 pub tool_name: String,
55 pub arguments: serde_json::Map<String, serde_json::Value>,
57 pub payment_requirements: PaymentRequirements,
59 pub payment_payload: McpPaymentPayload,
61}
62
63#[derive(Debug, Clone)]
65pub struct AfterExecutionContext {
66 pub server: ServerHookContext,
68 pub result: CallToolResult,
70}
71
72#[derive(Debug, Clone)]
74pub struct SettlementContext {
75 pub server: ServerHookContext,
77 pub settlement: r402_core::wire::SettleResponse,
79}
80
81#[derive(Debug, Clone)]
83pub struct PaymentWrapperConfig {
84 pub accepts: Vec<PaymentRequirements>,
86 pub resource: Option<ResourceInfo>,
88 pub hooks: PaymentWrapperHooks,
90 pub extensions: Extensions,
92}
93
94impl PaymentWrapperConfig {
95 pub fn try_new(
101 accepts: Vec<PaymentRequirements>,
102 resource: Option<ResourceInfo>,
103 ) -> Result<Self, PaymentWrapperConfigError> {
104 if accepts.is_empty() {
105 return Err(PaymentWrapperConfigError::EmptyAccepts);
106 }
107 Ok(Self {
108 accepts,
109 resource,
110 hooks: PaymentWrapperHooks::default(),
111 extensions: Extensions::new(),
112 })
113 }
114
115 #[must_use]
117 pub fn with_hooks(mut self, hooks: PaymentWrapperHooks) -> Self {
118 self.hooks = hooks;
119 self
120 }
121
122 #[must_use]
124 pub fn with_extensions(mut self, extensions: Extensions) -> Self {
125 self.extensions = extensions;
126 self
127 }
128}
129
130#[derive(Debug, Clone)]
132pub struct PaymentWrapper {
133 server: ResourceServer,
134 config: PaymentWrapperConfig,
135}
136
137impl PaymentWrapper {
138 pub fn try_new(
144 server: ResourceServer,
145 config: PaymentWrapperConfig,
146 ) -> Result<Self, PaymentWrapperConfigError> {
147 if config.accepts.is_empty() {
148 return Err(PaymentWrapperConfigError::EmptyAccepts);
149 }
150 Ok(Self { server, config })
151 }
152
153 fn resource_or_default(&self) -> ResourceInfo {
155 self.config.resource.clone().unwrap_or_else(|| {
156 ResourceInfo::new("mcp://tool/unknown")
157 .with_description("Unknown tool")
158 .with_mime_type("application/json")
159 })
160 }
161
162 #[must_use]
164 pub fn payment_required_result(&self, error_msg: impl Into<String>) -> CallToolResult {
165 let mut required = PaymentRequired::new(self.resource_or_default())
166 .with_error(error_msg.into())
167 .with_accepts(self.config.accepts.clone());
168 if !self.config.extensions.is_empty() {
169 required = required.with_extensions(self.config.extensions.clone());
170 }
171 payment_required_tool_result(&required)
172 }
173
174 #[must_use]
176 pub fn settlement_failed_result(&self, error_msg: impl Into<String>) -> CallToolResult {
177 settlement_failed_tool_result(
178 &self.config.accepts,
179 &self.resource_or_default(),
180 &self.config.extensions,
181 error_msg,
182 )
183 }
184
185 pub async fn invoke<H, Fut>(&self, params: CallToolRequestParams, handler: H) -> CallToolResult
187 where
188 H: FnOnce(CallToolRequestParams) -> Fut,
189 Fut: Future<Output = CallToolResult>,
190 {
191 let Some(payload) = extract_payment_from_params(¶ms) else {
192 return self.payment_required_result("Payment Required");
193 };
194
195 let Some(requirements) = self
196 .server
197 .find_matching_requirements(&self.config.accepts, &payload)
198 .cloned()
199 else {
200 return self.payment_required_result("No matching payment requirements found");
201 };
202
203 match self.server.verify_payment(&payload, &requirements).await {
204 Ok(resp) if resp.is_valid() => {}
205 Ok(resp) => {
206 let reason = match resp {
207 r402_core::wire::VerifyResponse::Invalid {
208 reason, message, ..
209 } => message.map_or_else(|| reason.to_string(), |m| m.to_string()),
210 _ => "Payment verification failed".into(),
212 };
213 return self
214 .payment_required_result(format!("Payment verification failed: {reason}"));
215 }
216 Err(err) => {
217 return self.payment_required_result(format!("Payment verification error: {err}"));
218 }
219 }
220
221 let arguments = params.arguments.clone().unwrap_or_default();
222 let tool_name = params.name.to_string();
223 let hook_ctx = ServerHookContext {
224 tool_name,
225 arguments,
226 payment_requirements: requirements.clone(),
227 payment_payload: payload.clone(),
228 };
229
230 if let Some(ref before) = self.config.hooks.on_before_execution
231 && !before(hook_ctx.clone())
232 {
233 return self.payment_required_result("Execution aborted by OnBeforeExecution hook");
234 }
235
236 let result = handler(params).await;
237 if result.is_error.unwrap_or(false) {
238 return result;
239 }
240
241 if let Some(ref after) = self.config.hooks.on_after_execution {
242 after(AfterExecutionContext {
243 server: hook_ctx.clone(),
244 result: result.clone(),
245 });
246 }
247
248 let settle = match self.server.settle_payment(&payload, &requirements).await {
249 Ok(s) if s.is_success() => s,
250 Ok(s) => {
251 let reason = match s {
252 r402_core::wire::SettleResponse::Failure {
253 reason, message, ..
254 } => message.map_or_else(|| reason.to_string(), |m| m.to_string()),
255 _ => "Settlement failed".into(),
257 };
258 return self.settlement_failed_result(format!("Settlement failed: {reason}"));
259 }
260 Err(err) => {
261 return self.settlement_failed_result(format!("Settlement error: {err}"));
262 }
263 };
264
265 if let Some(ref after_settle) = self.config.hooks.on_after_settlement {
266 after_settle(SettlementContext {
267 server: hook_ctx,
268 settlement: settle.clone(),
269 });
270 }
271
272 attach_settle_response(result, &settle)
273 }
274
275 pub fn wrap<H, Fut>(
277 &self,
278 handler: H,
279 ) -> impl Fn(
280 CallToolRequestParams,
281 ) -> std::pin::Pin<Box<dyn Future<Output = CallToolResult> + Send>>
282 + Clone
283 + Send
284 + Sync
285 + 'static
286 where
287 H: Fn(CallToolRequestParams) -> Fut + Clone + Send + Sync + 'static,
288 Fut: Future<Output = CallToolResult> + Send + 'static,
289 {
290 let this = self.clone();
291 move |params| {
292 let this = this.clone();
293 let handler = handler.clone();
294 Box::pin(async move { this.invoke(params, handler).await })
295 }
296 }
297}
298
299#[cfg(test)]
300mod tests {
301 use std::future::Future;
302 use std::sync::Arc;
303 use std::sync::atomic::{AtomicUsize, Ordering};
304
305 use r402_core::FacilitatorError;
306 use r402_core::facilitator::Facilitator;
307 use r402_core::wire::{
308 PaymentRequirements, ResourceInfo, SettleRequest, SettleResponse, SupportedResponse,
309 VerifyRequest, VerifyResponse,
310 };
311 use rmcp::model::ContentBlock;
312 use serde_json::json;
313
314 use super::*;
315 use crate::encode::{McpPaymentPayload, attach_payment_to_params, extract_settle_response};
316
317 struct MockFacilitator {
318 verifies: AtomicUsize,
319 settles: AtomicUsize,
320 verify_ok: bool,
321 settle_ok: bool,
322 }
323
324 impl Facilitator for MockFacilitator {
325 fn verify(
326 &self,
327 _request: VerifyRequest,
328 ) -> impl Future<Output = Result<VerifyResponse, FacilitatorError>> + Send {
329 self.verifies.fetch_add(1, Ordering::SeqCst);
330 std::future::ready(if self.verify_ok {
331 Ok(VerifyResponse::valid("0xpayer"))
332 } else {
333 Ok(VerifyResponse::invalid(
334 None,
335 r402_core::ErrorReason::InvalidExactEvmPayloadAuthorizationValidAfter,
336 ))
337 })
338 }
339
340 fn settle(
341 &self,
342 _request: SettleRequest,
343 ) -> impl Future<Output = Result<SettleResponse, FacilitatorError>> + Send {
344 self.settles.fetch_add(1, Ordering::SeqCst);
345 std::future::ready(if self.settle_ok {
346 Ok(SettleResponse::Success {
347 payer: "0xpayer".into(),
348 transaction: "0xtx".into(),
349 network: "eip155:1".into(),
350 amount: Some("1".into()),
351 extensions: Extensions::new(),
352 })
353 } else {
354 Ok(SettleResponse::Failure {
355 reason: r402_core::ErrorReason::UnexpectedSettleError,
356 message: Some("boom".into()),
357 payer: None,
358 network: "eip155:1".into(),
359 extensions: Extensions::new(),
360 })
361 })
362 }
363
364 fn supported(
365 &self,
366 ) -> impl Future<Output = Result<SupportedResponse, FacilitatorError>> + Send {
367 std::future::ready(Ok(SupportedResponse::default()))
368 }
369 }
370
371 fn accepts() -> Vec<PaymentRequirements> {
372 vec![PaymentRequirements::new(
373 "exact".into(),
374 "eip155:1".parse().unwrap(),
375 "1".into(),
376 "0xa".into(),
377 "0xb".into(),
378 60,
379 )]
380 }
381
382 fn sample_payload() -> McpPaymentPayload {
383 let req = accepts()
384 .into_iter()
385 .next()
386 .expect("accepts fixture is non-empty");
387 McpPaymentPayload::new(req, json!({"sig": "0x"}))
388 }
389
390 fn wrapper(verify_ok: bool, settle_ok: bool) -> PaymentWrapper {
391 let fac = Arc::new(MockFacilitator {
392 verifies: AtomicUsize::new(0),
393 settles: AtomicUsize::new(0),
394 verify_ok,
395 settle_ok,
396 });
397 let server = ResourceServer::new(fac);
398 let config = PaymentWrapperConfig::try_new(
399 accepts(),
400 Some(ResourceInfo::new("mcp://tool/demo").with_description("Demo")),
401 )
402 .unwrap();
403 PaymentWrapper::try_new(server, config).unwrap()
404 }
405
406 #[tokio::test]
407 async fn missing_payment_returns_dual_format_402() {
408 let w = wrapper(true, true);
409 let params = CallToolRequestParams::new("demo");
410 let result = w
411 .invoke(params, |_| async {
412 CallToolResult::success(vec![ContentBlock::text("should not run")])
413 })
414 .await;
415 assert_eq!(result.is_error, Some(true));
416 assert!(result.structured_content.is_some());
417 let text = result
418 .content
419 .first()
420 .and_then(ContentBlock::as_text)
421 .unwrap();
422 assert!(text.text.contains("Payment Required"));
423 }
424
425 #[tokio::test]
426 async fn happy_path_settles_and_attaches_meta() {
427 let w = wrapper(true, true);
428 let params =
429 attach_payment_to_params(CallToolRequestParams::new("demo"), &sample_payload());
430 let result = w
431 .invoke(params, |_| async {
432 CallToolResult::success(vec![ContentBlock::text("ok")])
433 })
434 .await;
435 assert!(!result.is_error.unwrap_or(false));
436 let settle = extract_settle_response(&result).unwrap();
437 assert!(settle.is_success());
438 }
439
440 #[tokio::test]
441 async fn tool_error_skips_settlement() {
442 let fac = Arc::new(MockFacilitator {
443 verifies: AtomicUsize::new(0),
444 settles: AtomicUsize::new(0),
445 verify_ok: true,
446 settle_ok: true,
447 });
448 let settles = Arc::clone(&fac);
449 let server = ResourceServer::new(fac);
450 let config = PaymentWrapperConfig::try_new(accepts(), None).unwrap();
451 let w = PaymentWrapper::try_new(server, config).unwrap();
452 let params =
453 attach_payment_to_params(CallToolRequestParams::new("demo"), &sample_payload());
454 let result = w
455 .invoke(params, |_| async {
456 CallToolResult::structured_error(json!({"err": true}))
457 })
458 .await;
459 assert_eq!(result.is_error, Some(true));
460 assert_eq!(settles.settles.load(Ordering::SeqCst), 0);
461 assert_eq!(settles.verifies.load(Ordering::SeqCst), 1);
462 }
463
464 #[tokio::test]
465 async fn verify_failure_is_tool_error_not_transport() {
466 let w = wrapper(false, true);
467 let params =
468 attach_payment_to_params(CallToolRequestParams::new("demo"), &sample_payload());
469 let result = w
470 .invoke(params, |_| async {
471 CallToolResult::success(vec![ContentBlock::text("nope")])
472 })
473 .await;
474 assert_eq!(result.is_error, Some(true));
475 assert!(result.structured_content.is_some());
476 }
477
478 #[tokio::test]
479 async fn settle_failure_uses_payment_required_format() {
480 let w = wrapper(true, false);
481 let params =
482 attach_payment_to_params(CallToolRequestParams::new("demo"), &sample_payload());
483 let result = w
484 .invoke(params, |_| async {
485 CallToolResult::success(vec![ContentBlock::text("ok")])
486 })
487 .await;
488 assert_eq!(result.is_error, Some(true));
489 let text = result
490 .content
491 .first()
492 .and_then(ContentBlock::as_text)
493 .unwrap();
494 assert!(text.text.contains("Settlement failed"));
495 assert!(extract_settle_response(&result).is_none());
496 }
497
498 #[test]
499 fn empty_accepts_rejected() {
500 let err = PaymentWrapperConfig::try_new(vec![], None).unwrap_err();
501 assert!(matches!(err, PaymentWrapperConfigError::EmptyAccepts));
502 }
503}