Skip to main content

scematica_protocol/
middleware.rs

1use axum::{
2    body::Body,
3    extract::State,
4    http::{Request, StatusCode},
5    middleware::Next,
6    Json,
7};
8use base64::Engine;
9use std::sync::Arc;
10use tracing::{debug, warn};
11
12use crate::{
13    facilitator::Facilitator,
14    types::{PaymentPayload, PaymentRequired, PaymentRequirements, ResourceInfo, X402_VERSION},
15};
16
17pub const PAYMENT_HEADER: &str = "X-Payment";
18pub const PAYMENT_RESPONSE_HEADER: &str = "X-Payment-Response";
19
20/// Shared state injected into the middleware via `State`.
21#[derive(Clone)]
22pub struct PaymentGate {
23    pub facilitator: Arc<Facilitator>,
24    /// All accepted payment options — served in 402 responses.
25    pub requirements: Vec<PaymentRequirements>,
26}
27
28/// axum middleware: returns 402 if no valid payment header, otherwise passes through.
29pub async fn payment_middleware(
30    State(gate): State<Arc<PaymentGate>>,
31    req: Request<Body>,
32    next: Next<Body>,
33) -> Result<axum::response::Response, (StatusCode, Json<serde_json::Value>)> {
34    let payment_value = req
35        .headers()
36        .get(PAYMENT_HEADER)
37        .and_then(|v| v.to_str().ok())
38        .map(|s| s.to_owned());
39
40    if let Some(encoded) = payment_value {
41        if let Some(payload) = decode_payment_payload(&encoded) {
42            if let Some(req_spec) = gate
43                .requirements
44                .iter()
45                .find(|r| r.network == payload.network && r.scheme == payload.scheme)
46            {
47                let verify = gate.facilitator.verify(&payload, req_spec);
48                if verify.is_valid {
49                    debug!(payer = ?verify.payer, "Payment verified — proceeding");
50
51                    // Pass through to the actual handler
52                    let mut response = next.run(req).await;
53
54                    // Settle asynchronously so the API response is not delayed
55                    let fac = gate.facilitator.clone();
56                    let p = payload.clone();
57                    let r = req_spec.clone();
58                    tokio::spawn(async move {
59                        let result = fac.settle(&p, &r).await;
60                        if !result.success {
61                            warn!("Async payment settlement failed: {:?}", result.error);
62                        }
63                        // Encode settlement result for potential header attachment
64                        if let Ok(json) = serde_json::to_string(&result) {
65                            let b64 = base64::engine::general_purpose::STANDARD.encode(json);
66                            drop(b64); // header can't be set after response sent — log only
67                        }
68                    });
69
70                    response
71                        .headers_mut()
72                        .insert(PAYMENT_RESPONSE_HEADER, "payment-verified".parse().unwrap());
73                    return Ok(response);
74                }
75                warn!(reason = ?verify.invalid_reason, "Payment verification failed");
76            }
77        }
78    }
79
80    // No valid payment → 402 with what we accept
81    let url = req.uri().to_string();
82    let method = req.method().to_string();
83    let body = PaymentRequired {
84        x402_version: X402_VERSION,
85        resource: ResourceInfo {
86            url,
87            method,
88            description: "Scematica Protocol — pay per API call".into(),
89        },
90        accepts: gate.requirements.clone(),
91    };
92    Err((
93        StatusCode::PAYMENT_REQUIRED,
94        Json(serde_json::to_value(&body).unwrap_or_default()),
95    ))
96}
97
98fn decode_payment_payload(encoded: &str) -> Option<PaymentPayload> {
99    let bytes = base64::engine::general_purpose::STANDARD
100        .decode(encoded)
101        .ok()?;
102    serde_json::from_slice(&bytes).ok()
103}