diff --git a/.changelog/reject-invalid-method-identifiers.md b/.changelog/reject-invalid-method-identifiers.md new file mode 100644 index 00000000..7beebcda --- /dev/null +++ b/.changelog/reject-invalid-method-identifiers.md @@ -0,0 +1,7 @@ +--- +"mpp": patch +--- + +Reject payment challenges whose method identifier contains characters other +than lowercase ASCII letters. Reject payment challenges reached through a +cross-origin redirect before a credential can be created or sent. diff --git a/src/client/error.rs b/src/client/error.rs index e2ceecd4..4ee1e784 100644 --- a/src/client/error.rs +++ b/src/client/error.rs @@ -21,6 +21,9 @@ pub enum HttpError { /// Request could not be cloned (required for retry) CloneFailed, + /// A redirect changed the request origin before returning a payment challenge + CrossOriginRedirect, + /// Payment provider error Payment(MppError), @@ -41,6 +44,9 @@ impl fmt::Display for HttpError { Self::InvalidChallenge(msg) => write!(f, "invalid challenge: {}", msg), Self::InvalidCredential(msg) => write!(f, "invalid credential: {}", msg), Self::CloneFailed => write!(f, "request could not be cloned for retry"), + Self::CrossOriginRedirect => { + write!(f, "Refusing to send payment credential across redirect") + } Self::Payment(e) => write!(f, "payment failed: {}", e), #[cfg(feature = "client")] Self::Request(e) => write!(f, "HTTP request failed: {}", e), diff --git a/src/client/fetch.rs b/src/client/fetch.rs index 86ddc4cc..0e937773 100644 --- a/src/client/fetch.rs +++ b/src/client/fetch.rs @@ -253,6 +253,25 @@ async fn send_with_payment( return Ok(resp); } + if url + .as_ref() + .is_some_and(|request_url| request_url.origin() != resp.url().origin()) + { + let err = HttpError::CrossOriginRedirect; + events + .emit(ClientEvent::PaymentFailed(PaymentFailedContext { + challenge: None, + error: err.to_string(), + reason: None, + })) + .await; + pending_payments + .rollback() + .await + .map_err(HttpError::Payment)?; + return Err(err); + } + let www_auth_values: Vec<&str> = resp .headers() .get_all(WWW_AUTHENTICATE) @@ -794,6 +813,76 @@ mod tests { assert_eq!(call_count.load(Ordering::SeqCst), 2); // initial 402 + retry } + #[tokio::test] + async fn cross_origin_redirect_before_402_is_rejected() { + let (_, www_auth) = test_challenge(); + let authorization_observed = Arc::new(AtomicU32::new(0)); + let observed = authorization_observed.clone(); + let target = Router::new().route( + "/paid", + get(move |req: axum::http::Request| { + let www_auth = www_auth.clone(); + let observed = observed.clone(); + async move { + if req.headers().contains_key("authorization") { + observed.fetch_add(1, Ordering::SeqCst); + } + ( + AxumStatusCode::PAYMENT_REQUIRED, + [(WWW_AUTH_NAME, www_auth)], + "pay up", + ) + } + }), + ); + let target_url = spawn_server(target).await; + let source = Router::new().route( + "/paid", + get(move || { + let target_url = target_url.clone(); + async move { + ( + AxumStatusCode::TEMPORARY_REDIRECT, + [(axum::http::header::LOCATION, format!("{target_url}/paid"))], + "redirect", + ) + } + }), + ); + let source_url = spawn_server(source).await; + let provider = MockProvider::new(); + let events = ClientEvents::default(); + let failed_count = Arc::new(AtomicU32::new(0)); + let _failed_sub = events.on_payment_failed({ + let failed_count = failed_count.clone(); + move |ctx| { + failed_count.fetch_add(1, Ordering::SeqCst); + async move { + assert!(ctx.challenge.is_none()); + assert_eq!( + ctx.error, + "Refusing to send payment credential across redirect" + ); + } + } + }); + + let err = reqwest::Client::new() + .get(format!("{source_url}/paid")) + .send_with_payment_options(&provider, &AcceptPaymentPolicy::Always, events) + .await + .unwrap_err(); + + assert!(matches!(err, HttpError::CrossOriginRedirect)); + assert_eq!( + err.to_string(), + "Refusing to send payment credential across redirect" + ); + assert_eq!(provider.call_count(), 0); + assert_eq!(authorization_observed.load(Ordering::SeqCst), 0); + assert_eq!(failed_count.load(Ordering::SeqCst), 1); + } + #[tokio::test] async fn dropped_paid_request_abandons_transient_provider_state() { let (_, www_auth) = test_challenge(); diff --git a/src/client/middleware.rs b/src/client/middleware.rs index a34323d5..5ee75513 100644 --- a/src/client/middleware.rs +++ b/src/client/middleware.rs @@ -17,6 +17,7 @@ use crate::client::events::{ CredentialCreatedContext, PaymentFailedContext, PaymentFailureReason, PaymentResponseContext, }; use crate::client::provider::{PaymentContext, PaymentProvider, PendingPayments}; +use crate::client::HttpError; use crate::client::DEFAULT_MAX_PAYMENT_RETRIES; use crate::protocol::core::accept_payment::ACCEPT_PAYMENT_HEADER; use crate::protocol::core::{ @@ -214,6 +215,21 @@ where return Ok(resp); } + if payment_context.url.origin() != resp.url().origin() { + let error = HttpError::CrossOriginRedirect.to_string(); + self.events + .emit(ClientEvent::PaymentFailed(PaymentFailedContext { + challenge: None, + error: error.clone(), + reason: None, + })) + .await; + rollback_middleware_payments(&mut pending_payments).await?; + return Err(reqwest_middleware::Error::Middleware(anyhow::anyhow!( + error + ))); + } + let www_auth_values: Vec<&str> = resp .headers() .get_all(WWW_AUTHENTICATE) @@ -644,6 +660,77 @@ mod tests { assert_eq!(call_count.load(Ordering::SeqCst), 2); } + #[tokio::test] + async fn cross_origin_redirect_before_402_is_rejected() { + let (_, www_auth) = test_challenge(); + let authorization_observed = Arc::new(AtomicU32::new(0)); + let observed = authorization_observed.clone(); + let target = Router::new().route( + "/paid", + get(move |req: axum::http::Request| { + let www_auth = www_auth.clone(); + let observed = observed.clone(); + async move { + if req.headers().contains_key("authorization") { + observed.fetch_add(1, Ordering::SeqCst); + } + ( + AxumStatusCode::PAYMENT_REQUIRED, + [(WWW_AUTH_NAME, www_auth)], + "pay up", + ) + } + }), + ); + let target_url = spawn_server(target).await; + let source = Router::new().route( + "/paid", + get(move || { + let target_url = target_url.clone(); + async move { + ( + AxumStatusCode::TEMPORARY_REDIRECT, + [(axum::http::header::LOCATION, format!("{target_url}/paid"))], + "redirect", + ) + } + }), + ); + let source_url = spawn_server(source).await; + let provider = TestProvider::new(); + let events = ClientEvents::default(); + let failed_count = Arc::new(AtomicU32::new(0)); + let _failed_sub = events.on_payment_failed({ + let failed_count = failed_count.clone(); + move |ctx| { + failed_count.fetch_add(1, Ordering::SeqCst); + async move { + assert!(ctx.challenge.is_none()); + assert_eq!( + ctx.error, + "Refusing to send payment credential across redirect" + ); + } + } + }); + let client = ClientBuilder::new(reqwest::Client::new()) + .with(PaymentMiddleware::new(provider.clone()).with_events(events)) + .build(); + + let err = client + .get(format!("{source_url}/paid")) + .send() + .await + .unwrap_err(); + + assert!(err + .to_string() + .contains("Refusing to send payment credential across redirect")); + assert_eq!(provider.call_count(), 0); + assert_eq!(authorization_observed.load(Ordering::SeqCst), 0); + assert_eq!(failed_count.load(Ordering::SeqCst), 1); + } + #[tokio::test] async fn test_middleware_passes_request_context_to_concurrent_payments() { #[derive(Clone)] diff --git a/src/client/ws.rs b/src/client/ws.rs index 7451acc9..314ac7c5 100644 --- a/src/client/ws.rs +++ b/src/client/ws.rs @@ -181,6 +181,23 @@ mod tests { assert!(matches!(parsed, WsServerMessage::Challenge { .. })); } + #[test] + fn test_ws_transport_rejects_invalid_method_name() { + let response = WsServerMessage::Challenge { + challenge: serde_json::json!({ + "id": "ch-1", + "realm": "test", + "method": "123", + "intent": "charge", + "request": "eyJ0ZXN0IjoidmFsdWUifQ" + }), + error: None, + }; + + let error = ws().get_challenge(&response).unwrap_err(); + assert!(error.to_string().contains("invalid method name")); + } + #[test] fn test_ws_server_message_need_voucher() { let json = r#"{"type":"needVoucher","channelId":"0xabc","requiredCumulative":"2000","acceptedCumulative":"1000","deposit":"5000"}"#; diff --git a/src/mcp.rs b/src/mcp.rs index 5890a4c5..474e6b2e 100644 --- a/src/mcp.rs +++ b/src/mcp.rs @@ -461,6 +461,22 @@ mod tests { assert_eq!(challenges[0].id, "ch_test_123"); } + #[test] + fn test_extract_challenges_rejects_invalid_method_name() { + let mut challenge = serde_json::to_value(test_challenge()).unwrap(); + challenge["method"] = json!("123"); + let error = json!({ + "code": PAYMENT_REQUIRED_CODE, + "message": "Payment Required", + "data": { + "httpStatus": 402, + "challenges": [challenge] + } + }); + + assert!(extract_challenges(&error).is_none()); + } + #[test] fn test_extract_challenges_no_data() { let error = json!({ diff --git a/src/protocol/core/headers.rs b/src/protocol/core/headers.rs index 15c8487e..86a55b6b 100644 --- a/src/protocol/core/headers.rs +++ b/src/protocol/core/headers.rs @@ -1022,6 +1022,17 @@ mod tests { assert!(err.to_string().contains("Invalid method")); } + #[test] + fn test_parse_www_authenticate_rejects_non_letter_method_names() { + for method in ["123", "*", "tempo!"] { + let header = format!( + r#"Payment id="abc", realm="api", method="{method}", intent="charge", request="e30""# + ); + let err = parse_www_authenticate(&header).unwrap_err(); + assert!(err.to_string().contains("Invalid method")); + } + } + #[test] fn test_parse_www_authenticate_rejects_mixed_case_method_name() { let header = diff --git a/src/protocol/core/types.rs b/src/protocol/core/types.rs index 400b23c0..e066fa4f 100644 --- a/src/protocol/core/types.rs +++ b/src/protocol/core/types.rs @@ -28,7 +28,7 @@ use crate::error::{MppError, Result}; /// let method2: MethodName = "TEMPO".into(); /// assert_eq!(method2.as_str(), "tempo"); /// ``` -#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)] +#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize)] #[serde(transparent)] pub struct MethodName(String); @@ -82,6 +82,22 @@ impl From for MethodName { } } +impl<'de> Deserialize<'de> for MethodName { + fn deserialize(deserializer: D) -> std::result::Result + where + D: Deserializer<'de>, + { + let method = Self(String::deserialize(deserializer)?); + if method.is_valid() { + Ok(method) + } else { + Err(serde::de::Error::custom( + "invalid method name: must contain only lowercase ASCII letters", + )) + } + } +} + /// Payment intent identifier (newtype over String). /// /// Represents a payment intent like "charge", "session", etc. @@ -404,6 +420,15 @@ mod tests { assert_eq!(parsed, method); } + #[test] + fn test_method_name_deserialization_rejects_invalid_names() { + for method in ["", "123", "*", "tempo!", "TEMPO"] { + let json = serde_json::to_string(method).unwrap(); + let error = serde_json::from_str::(&json).unwrap_err(); + assert!(error.to_string().contains("invalid method name")); + } + } + #[test] fn test_intent_name() { let intent: IntentName = "charge".into();