Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 7 additions & 0 deletions .changelog/reject-invalid-method-identifiers.md
Original file line number Diff line number Diff line change
@@ -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.
6 changes: 6 additions & 0 deletions src/client/error.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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),

Expand All @@ -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),
Expand Down
89 changes: 89 additions & 0 deletions src/client/fetch.rs
Original file line number Diff line number Diff line change
Expand Up @@ -253,6 +253,25 @@ async fn send_with_payment<P: PaymentProvider>(
return Ok(resp);
}

if url
.as_ref()
.is_some_and(|request_url| request_url.origin() != resp.url().origin())
{
Comment thread
brendanjryan marked this conversation as resolved.
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)
Expand Down Expand Up @@ -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<axum::body::Body>| {
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();
Expand Down
87 changes: 87 additions & 0 deletions src/client/middleware.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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::{
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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<axum::body::Body>| {
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)]
Expand Down
17 changes: 17 additions & 0 deletions src/client/ws.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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"}"#;
Expand Down
16 changes: 16 additions & 0 deletions src/mcp.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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!({
Expand Down
11 changes: 11 additions & 0 deletions src/protocol/core/headers.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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 =
Expand Down
27 changes: 26 additions & 1 deletion src/protocol/core/types.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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);

Expand Down Expand Up @@ -82,6 +82,22 @@ impl From<String> for MethodName {
}
}

impl<'de> Deserialize<'de> for MethodName {
fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
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.
Expand Down Expand Up @@ -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::<MethodName>(&json).unwrap_err();
assert!(error.to_string().contains("invalid method name"));
}
}

#[test]
fn test_intent_name() {
let intent: IntentName = "charge".into();
Expand Down
Loading