use reqwest::header::{HeaderMap, HeaderValue, AUTHORIZATION, WWW_AUTHENTICATE};
use reqwest::{RequestBuilder, Response, StatusCode};
use super::accept_payment_policy::AcceptPaymentPolicy;
use super::error::HttpError;
use super::events::{
ChallengeReceivedContext, ClientEvent, ClientEvents, CredentialCreatedContext,
PaymentFailedContext, PaymentFailureReason, PaymentResponseContext,
};
use super::provider::{PaymentContext, PaymentProvider, PendingPayments};
use super::DEFAULT_MAX_PAYMENT_RETRIES;
use crate::client::challenge_selection::{
expired_payment_error, select_supported_challenge, ChallengeSelectionError,
};
use crate::error::MppError;
use crate::protocol::core::accept_payment::ACCEPT_PAYMENT_HEADER;
use crate::protocol::core::{format_authorization, parse_www_authenticate_all};
pub trait PaymentExt: Sized {
fn send_with_payment<P: PaymentProvider>(
self,
provider: &P,
) -> impl std::future::Future<Output = Result<Response, HttpError>> + Send {
self.send_with_payment_policy(provider, &AcceptPaymentPolicy::Always)
}
fn send_with_payment_from_response<P: PaymentProvider>(
self,
provider: &P,
response: Response,
) -> impl std::future::Future<Output = Result<Response, HttpError>> + Send;
fn send_with_payment_policy<P: PaymentProvider>(
self,
provider: &P,
policy: &AcceptPaymentPolicy,
) -> impl std::future::Future<Output = Result<Response, HttpError>> + Send;
fn send_with_payment_max_retries<P: PaymentProvider>(
self,
provider: &P,
max_payment_retries: usize,
) -> impl std::future::Future<Output = Result<Response, HttpError>> + Send {
self.send_with_payment_policy_max_retries(
provider,
&AcceptPaymentPolicy::Always,
max_payment_retries,
)
}
fn send_with_payment_policy_max_retries<P: PaymentProvider>(
self,
provider: &P,
policy: &AcceptPaymentPolicy,
max_payment_retries: usize,
) -> impl std::future::Future<Output = Result<Response, HttpError>> + Send {
self.send_with_payment_options_max_retries(
provider,
policy,
ClientEvents::default(),
max_payment_retries,
)
}
fn send_with_payment_options<P: PaymentProvider>(
self,
provider: &P,
policy: &AcceptPaymentPolicy,
events: ClientEvents,
) -> impl std::future::Future<Output = Result<Response, HttpError>> + Send {
let _ = events;
self.send_with_payment_policy(provider, policy)
}
fn send_with_payment_options_max_retries<P: PaymentProvider>(
self,
provider: &P,
policy: &AcceptPaymentPolicy,
events: ClientEvents,
max_payment_retries: usize,
) -> impl std::future::Future<Output = Result<Response, HttpError>> + Send {
let _ = max_payment_retries;
self.send_with_payment_options(provider, policy, events)
}
}
impl PaymentExt for RequestBuilder {
async fn send_with_payment_from_response<P: PaymentProvider>(
self,
provider: &P,
response: Response,
) -> Result<Response, HttpError> {
send_with_payment(
self,
provider,
&AcceptPaymentPolicy::Always,
ClientEvents::default(),
DEFAULT_MAX_PAYMENT_RETRIES,
Some(response),
)
.await
}
async fn send_with_payment_policy<P: PaymentProvider>(
self,
provider: &P,
policy: &AcceptPaymentPolicy,
) -> Result<Response, HttpError> {
self.send_with_payment_options_max_retries(
provider,
policy,
ClientEvents::default(),
DEFAULT_MAX_PAYMENT_RETRIES,
)
.await
}
async fn send_with_payment_policy_max_retries<P: PaymentProvider>(
self,
provider: &P,
policy: &AcceptPaymentPolicy,
max_payment_retries: usize,
) -> Result<Response, HttpError> {
self.send_with_payment_options_max_retries(
provider,
policy,
ClientEvents::default(),
max_payment_retries,
)
.await
}
async fn send_with_payment_options<P: PaymentProvider>(
self,
provider: &P,
policy: &AcceptPaymentPolicy,
events: ClientEvents,
) -> Result<Response, HttpError> {
self.send_with_payment_options_max_retries(
provider,
policy,
events,
DEFAULT_MAX_PAYMENT_RETRIES,
)
.await
}
async fn send_with_payment_options_max_retries<P: PaymentProvider>(
self,
provider: &P,
policy: &AcceptPaymentPolicy,
events: ClientEvents,
max_payment_retries: usize,
) -> Result<Response, HttpError> {
send_with_payment(self, provider, policy, events, max_payment_retries, None).await
}
}
async fn send_with_payment<P: PaymentProvider>(
request: RequestBuilder,
provider: &P,
policy: &AcceptPaymentPolicy,
events: ClientEvents,
max_payment_retries: usize,
initial_response: Option<Response>,
) -> Result<Response, HttpError> {
let retry_builder = request.try_clone().ok_or(HttpError::CloneFailed)?;
let peek = retry_builder.try_clone().and_then(|b| b.build().ok());
let url = peek.as_ref().map(|r| r.url().clone());
let caller_accept = peek.as_ref().and_then(|r| {
r.headers()
.get(ACCEPT_PAYMENT_HEADER)
.and_then(|v| v.to_str().ok())
.map(String::from)
});
let provider_accept = provider.accept_payment_header();
let inject = caller_accept.is_none() && url.as_ref().is_some_and(|u| policy.allows(u));
let request = if inject {
if let Some(ref header) = provider_accept {
request.header(ACCEPT_PAYMENT_HEADER, header)
} else {
request
}
} else {
request
};
let ranking_accept = caller_accept.or(provider_accept);
let mut paid_challenge_ids = std::collections::HashSet::new();
let mut pending_payments = PendingPayments::new(provider.clone());
let mut retried_stale_session = false;
let mut refreshed_after_provider_setup = false;
let mut resp = match initial_response {
Some(response) => response,
None => request.send().await?,
};
let mut payment_attempt = 0;
while payment_attempt < max_payment_retries {
if resp.status() != StatusCode::PAYMENT_REQUIRED {
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)
.iter()
.filter_map(|v| v.to_str().ok())
.collect();
if www_auth_values.is_empty() {
if !paid_challenge_ids.is_empty() {
if !pending_payments.is_empty() {
if resp.headers().contains_key("payment-receipt") {
pending_payments
.commit()
.await
.map_err(HttpError::Payment)?;
} else {
pending_payments
.rollback()
.await
.map_err(HttpError::Payment)?;
}
}
return Ok(resp);
}
events
.emit(ClientEvent::PaymentFailed(PaymentFailedContext {
challenge: None,
error: HttpError::MissingChallenge.to_string(),
reason: None,
}))
.await;
pending_payments
.rollback()
.await
.map_err(HttpError::Payment)?;
return Err(HttpError::MissingChallenge);
}
let challenges: Vec<_> = parse_www_authenticate_all(www_auth_values)
.into_iter()
.filter_map(|r| r.ok())
.collect();
let challenge = match select_supported_challenge(
&challenges,
ranking_accept.as_deref(),
|challenge| provider.supports(challenge.method.as_str(), challenge.intent.as_str()),
|challenges| provider.select_challenge(challenges),
) {
Ok(challenge) => challenge.clone(),
Err(ChallengeSelectionError::Expired(challenge)) => {
let err = HttpError::Payment(expired_payment_error(&challenge));
let error = err.to_string();
let expires = challenge.expires.clone();
events
.emit(ClientEvent::PaymentFailed(PaymentFailedContext {
challenge: Some(*challenge),
error,
reason: Some(PaymentFailureReason::PreSigningExpired { expires }),
}))
.await;
pending_payments
.rollback()
.await
.map_err(HttpError::Payment)?;
return Err(err);
}
Err(ChallengeSelectionError::NoSupportedChallenge(message)) => {
let err = HttpError::NoSupportedChallenge(message);
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 Some(url) = url.clone() else {
pending_payments
.rollback()
.await
.map_err(HttpError::Payment)?;
return Err(HttpError::CloneFailed);
};
let payment_context = PaymentContext {
url,
headers: peek
.as_ref()
.map(|request| request.headers().clone())
.unwrap_or_default(),
};
let challenge = match provider
.prepare_http_payment_challenge(&challenge, payment_context.clone())
.await
{
Ok(Some(challenge)) => challenge,
Ok(None) => {
if refreshed_after_provider_setup {
return Err(HttpError::Payment(MppError::InvalidConfig(
"payment provider repeatedly requested a fresh HTTP challenge".to_owned(),
)));
}
refreshed_after_provider_setup = true;
pending_payments
.rollback()
.await
.map_err(HttpError::Payment)?;
paid_challenge_ids.clear();
resp = retry_builder
.try_clone()
.ok_or(HttpError::CloneFailed)?
.send()
.await
.map_err(HttpError::request)?;
continue;
}
Err(err) => {
let http_err = HttpError::Payment(err);
events
.emit(ClientEvent::PaymentFailed(PaymentFailedContext {
challenge: Some(challenge),
error: http_err.to_string(),
reason: None,
}))
.await;
pending_payments
.rollback()
.await
.map_err(HttpError::Payment)?;
return Err(http_err);
}
};
payment_attempt += 1;
if !paid_challenge_ids.insert(challenge.id.clone()) {
events
.emit(ClientEvent::PaymentFailed(PaymentFailedContext {
challenge: Some(challenge),
error: "payment retry returned a previously paid challenge".to_string(),
reason: None,
}))
.await;
pending_payments
.rollback()
.await
.map_err(HttpError::Payment)?;
return Ok(resp);
}
let override_credential = events
.emit_challenge_received(ChallengeReceivedContext {
challenge: challenge.clone(),
challenges: challenges.clone(),
})
.await;
let credential = match override_credential {
Some(credential) => credential,
None => match provider.pay_with_context(&challenge, payment_context).await {
Ok(credential) => credential,
Err(err) => {
let http_err = HttpError::Payment(err);
events
.emit(ClientEvent::PaymentFailed(PaymentFailedContext {
challenge: Some(challenge),
error: http_err.to_string(),
reason: None,
}))
.await;
pending_payments
.rollback()
.await
.map_err(HttpError::Payment)?;
return Err(http_err);
}
},
};
pending_payments.push((challenge.clone(), credential.clone()));
events
.emit(ClientEvent::CredentialCreated(CredentialCreatedContext {
challenge: challenge.clone(),
credential: credential.clone(),
}))
.await;
let auth_header = match format_authorization(&credential) {
Ok(auth_header) => auth_header,
Err(err) => {
let http_err = HttpError::InvalidCredential(err.to_string());
events
.emit(ClientEvent::PaymentFailed(PaymentFailedContext {
challenge: Some(challenge),
error: http_err.to_string(),
reason: None,
}))
.await;
pending_payments
.rollback()
.await
.map_err(HttpError::Payment)?;
return Err(http_err);
}
};
let auth_header = match HeaderValue::from_str(&auth_header) {
Ok(auth_header) => auth_header,
Err(err) => {
let http_err = HttpError::InvalidCredential(err.to_string());
events
.emit(ClientEvent::PaymentFailed(PaymentFailedContext {
challenge: Some(challenge),
error: http_err.to_string(),
reason: None,
}))
.await;
pending_payments
.rollback()
.await
.map_err(HttpError::Payment)?;
return Err(http_err);
}
};
let mut payment_headers = HeaderMap::new();
payment_headers.insert(AUTHORIZATION, auth_header);
let retry = match retry_builder.try_clone() {
Some(retry) => retry.headers(payment_headers),
None => {
pending_payments
.rollback()
.await
.map_err(HttpError::Payment)?;
return Err(HttpError::CloneFailed);
}
};
let (client, retry) = retry.build_split();
let mut retry = retry.map_err(HttpError::request)?;
*retry.url_mut() = resp.url().clone();
resp = match client.execute(retry).await {
Ok(resp) => resp,
Err(err) => {
let http_err = HttpError::request(err);
events
.emit(ClientEvent::PaymentFailed(PaymentFailedContext {
challenge: Some(challenge),
error: http_err.to_string(),
reason: None,
}))
.await;
pending_payments
.commit()
.await
.map_err(HttpError::Payment)?;
return Err(http_err);
}
};
let status = resp.status();
if status.is_success() {
events
.emit(ClientEvent::PaymentResponse(PaymentResponseContext {
challenge,
credential,
status,
}))
.await;
pending_payments
.commit()
.await
.map_err(HttpError::Payment)?;
return Ok(resp);
}
if !retried_stale_session
&& status == StatusCode::GONE
&& challenge.intent.as_str() == "session"
{
retried_stale_session = true;
pending_payments
.rollback()
.await
.map_err(HttpError::Payment)?;
paid_challenge_ids.clear();
resp = retry_builder
.try_clone()
.ok_or(HttpError::CloneFailed)?
.send()
.await
.map_err(HttpError::request)?;
continue;
}
if status != StatusCode::PAYMENT_REQUIRED || payment_attempt == max_payment_retries {
if resp.headers().contains_key("payment-receipt") {
pending_payments
.commit()
.await
.map_err(HttpError::Payment)?;
} else {
pending_payments
.rollback()
.await
.map_err(HttpError::Payment)?;
}
return Ok(resp);
}
if resp.headers().contains_key("payment-receipt") {
pending_payments
.commit()
.await
.map_err(HttpError::Payment)?;
} else {
pending_payments
.rollback()
.await
.map_err(HttpError::Payment)?;
}
}
Ok(resp)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_payment_ext_trait_exists() {
fn assert_payment_ext<T: PaymentExt>() {}
assert_payment_ext::<RequestBuilder>();
}
#[cfg(all(feature = "client", feature = "utils"))]
mod integration {
use super::*;
use crate::client::ClientEventKind;
use crate::error::MppError;
use crate::protocol::core::{
format_www_authenticate, Base64UrlJson, PaymentChallenge, PaymentCredential,
PaymentPayload,
};
use axum::http::header::WWW_AUTHENTICATE as WWW_AUTH_NAME;
use axum::http::StatusCode as AxumStatusCode;
use axum::response::IntoResponse;
use axum::routing::get;
use axum::Router;
use std::sync::atomic::{AtomicU32, Ordering};
use std::sync::{Arc, Mutex};
use tokio::net::TcpListener;
use tokio::sync::Notify;
#[derive(Clone)]
struct MockProvider {
pay_count: Arc<AtomicU32>,
commit_count: Arc<AtomicU32>,
rollback_count: Arc<AtomicU32>,
abandon_count: Arc<AtomicU32>,
challenge_ids: Arc<Mutex<Vec<String>>>,
fail: bool,
}
impl MockProvider {
fn new() -> Self {
Self {
pay_count: Arc::new(AtomicU32::new(0)),
commit_count: Arc::new(AtomicU32::new(0)),
rollback_count: Arc::new(AtomicU32::new(0)),
abandon_count: Arc::new(AtomicU32::new(0)),
challenge_ids: Arc::new(Mutex::new(Vec::new())),
fail: false,
}
}
fn failing() -> Self {
Self {
pay_count: Arc::new(AtomicU32::new(0)),
commit_count: Arc::new(AtomicU32::new(0)),
rollback_count: Arc::new(AtomicU32::new(0)),
abandon_count: Arc::new(AtomicU32::new(0)),
challenge_ids: Arc::new(Mutex::new(Vec::new())),
fail: true,
}
}
fn call_count(&self) -> u32 {
self.pay_count.load(Ordering::SeqCst)
}
fn challenge_ids(&self) -> Vec<String> {
self.challenge_ids.lock().unwrap().clone()
}
fn commit_count(&self) -> u32 {
self.commit_count.load(Ordering::SeqCst)
}
fn rollback_count(&self) -> u32 {
self.rollback_count.load(Ordering::SeqCst)
}
fn abandon_count(&self) -> u32 {
self.abandon_count.load(Ordering::SeqCst)
}
}
impl super::PaymentProvider for MockProvider {
fn supports(&self, _method: &str, _intent: &str) -> bool {
true
}
async fn pay(
&self,
challenge: &crate::protocol::core::PaymentChallenge,
) -> Result<PaymentCredential, MppError> {
self.pay_count.fetch_add(1, Ordering::SeqCst);
self.challenge_ids
.lock()
.unwrap()
.push(challenge.id.clone());
if self.fail {
return Err(MppError::Http("mock provider failure".into()));
}
let echo = challenge.to_echo();
Ok(PaymentCredential::new(
echo,
PaymentPayload::hash("0xmockhash"),
))
}
async fn commit_payment(
&self,
_challenge: &PaymentChallenge,
_credential: &PaymentCredential,
) -> Result<(), MppError> {
self.commit_count.fetch_add(1, Ordering::SeqCst);
Ok(())
}
async fn rollback_payment(
&self,
_challenge: &PaymentChallenge,
_credential: &PaymentCredential,
) -> Result<(), MppError> {
self.rollback_count.fetch_add(1, Ordering::SeqCst);
Ok(())
}
fn abandon_payment(
&self,
_challenge: &PaymentChallenge,
_credential: &PaymentCredential,
) {
self.abandon_count.fetch_add(1, Ordering::SeqCst);
}
}
fn test_challenge() -> (PaymentChallenge, String) {
test_challenge_with_id("test-id-123")
}
fn test_challenge_with_id(id: &str) -> (PaymentChallenge, String) {
let request =
Base64UrlJson::from_value(&serde_json::json!({"amount": "1000"})).unwrap();
let challenge =
PaymentChallenge::new(id, "test.example.com", "tempo", "charge", request);
let header = format_www_authenticate(&challenge).unwrap();
(challenge, header)
}
async fn spawn_server(app: Router) -> String {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
axum::serve(listener, app).await.unwrap();
});
format!("http://{}", addr)
}
#[tokio::test]
async fn test_happy_path_402_then_200() {
let (_, www_auth) = test_challenge();
let call_count = Arc::new(AtomicU32::new(0));
let counter = call_count.clone();
let app = Router::new().route(
"/paid",
get(move |req: axum::http::Request<axum::body::Body>| {
let www_auth = www_auth.clone();
let counter = counter.clone();
async move {
counter.fetch_add(1, Ordering::SeqCst);
if req.headers().get("authorization").is_some() {
(AxumStatusCode::OK, "ok").into_response()
} else {
(
AxumStatusCode::PAYMENT_REQUIRED,
[(WWW_AUTH_NAME, www_auth)],
"pay up",
)
.into_response()
}
}
}),
);
let base_url = spawn_server(app).await;
let provider = MockProvider::new();
let client = reqwest::Client::new();
let resp = client
.get(format!("{}/paid", base_url))
.send_with_payment(&provider)
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
assert_eq!(provider.call_count(), 1);
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 = 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 paid_retry_targets_final_same_origin_url() {
let (_, www_auth) = test_challenge();
let app = Router::new()
.route(
"/start",
get(|req: axum::http::Request<axum::body::Body>| async move {
if req.headers().contains_key("authorization") {
return AxumStatusCode::BAD_REQUEST.into_response();
}
(
AxumStatusCode::TEMPORARY_REDIRECT,
[(axum::http::header::LOCATION, "/paid")],
"redirect",
)
.into_response()
}),
)
.route(
"/paid",
get(move |req: axum::http::Request<axum::body::Body>| {
let www_auth = www_auth.clone();
async move {
if req.headers().contains_key("authorization") {
AxumStatusCode::OK.into_response()
} else {
(
AxumStatusCode::PAYMENT_REQUIRED,
[(WWW_AUTH_NAME, www_auth)],
"pay up",
)
.into_response()
}
}
}),
);
let base_url = spawn_server(app).await;
let provider = MockProvider::new();
let resp = reqwest::Client::new()
.get(format!("{base_url}/start"))
.send_with_payment(&provider)
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
assert_eq!(resp.url().path(), "/paid");
assert_eq!(provider.call_count(), 1);
}
#[tokio::test]
async fn dropped_paid_request_abandons_transient_provider_state() {
let (_, www_auth) = test_challenge();
let retry_started = Arc::new(Notify::new());
let handler_notify = retry_started.clone();
let app = Router::new().route(
"/paid",
get(move |req: axum::http::Request<axum::body::Body>| {
let www_auth = www_auth.clone();
let handler_notify = handler_notify.clone();
async move {
if req.headers().get("authorization").is_some() {
handler_notify.notify_one();
std::future::pending::<axum::response::Response>().await
} else {
(
AxumStatusCode::PAYMENT_REQUIRED,
[(WWW_AUTH_NAME, www_auth)],
"payment required",
)
.into_response()
}
}
}),
);
let url = spawn_server(app).await;
let provider = MockProvider::new();
let task_provider = provider.clone();
let task = tokio::spawn(async move {
reqwest::Client::new()
.get(format!("{url}/paid"))
.send_with_payment(&task_provider)
.await
});
retry_started.notified().await;
task.abort();
assert!(task.await.unwrap_err().is_cancelled());
assert_eq!(provider.call_count(), 1);
assert_eq!(provider.commit_count(), 0);
assert_eq!(provider.rollback_count(), 0);
assert_eq!(provider.abandon_count(), 1);
}
#[tokio::test]
async fn test_paid_retry_replaces_caller_authorization() {
let (_, www_auth) = test_challenge();
let app = Router::new().route(
"/paid",
get(move |req: axum::http::Request<axum::body::Body>| {
let www_auth = www_auth.clone();
async move {
let authorization = req
.headers()
.get_all("authorization")
.iter()
.filter_map(|value| value.to_str().ok())
.collect::<Vec<_>>();
match authorization.as_slice() {
[value] if value.starts_with("Payment ") => {
(AxumStatusCode::OK, "ok").into_response()
}
["Bearer upstream-token"] => (
AxumStatusCode::PAYMENT_REQUIRED,
[(WWW_AUTH_NAME, www_auth)],
"pay up",
)
.into_response(),
_ => (AxumStatusCode::BAD_REQUEST, "duplicate authorization")
.into_response(),
}
}
}),
);
let base_url = spawn_server(app).await;
let provider = MockProvider::new();
let client = reqwest::Client::new();
let resp = client
.get(format!("{}/paid", base_url))
.bearer_auth("upstream-token")
.send_with_payment(&provider)
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
assert_eq!(provider.call_count(), 1);
}
#[tokio::test]
async fn test_payment_events_fire_on_success() {
let (_, www_auth) = test_challenge();
let app = Router::new().route(
"/paid",
get(move |req: axum::http::Request<axum::body::Body>| {
let www_auth = www_auth.clone();
async move {
if req.headers().get("authorization").is_some() {
(AxumStatusCode::OK, "ok").into_response()
} else {
(
AxumStatusCode::PAYMENT_REQUIRED,
[(WWW_AUTH_NAME, www_auth)],
"pay up",
)
.into_response()
}
}
}),
);
let base_url = spawn_server(app).await;
let provider = MockProvider::new();
let client = reqwest::Client::new();
let events = ClientEvents::default();
let challenge_count = Arc::new(AtomicU32::new(0));
let credential_count = Arc::new(AtomicU32::new(0));
let response_count = Arc::new(AtomicU32::new(0));
let _challenge_sub = events.on_challenge_received({
let challenge_count = challenge_count.clone();
move |ctx| {
challenge_count.fetch_add(1, Ordering::SeqCst);
async move {
assert_eq!(ctx.challenge.method.as_str(), "tempo");
None
}
}
});
let _credential_sub = events.on_credential_created({
let credential_count = credential_count.clone();
move |ctx| {
credential_count.fetch_add(1, Ordering::SeqCst);
async move {
assert_eq!(ctx.credential.challenge.method.as_str(), "tempo");
}
}
});
let _response_sub = events.on_payment_response({
let response_count = response_count.clone();
move |ctx| {
response_count.fetch_add(1, Ordering::SeqCst);
async move {
assert_eq!(ctx.status, StatusCode::OK);
}
}
});
let resp = client
.get(format!("{}/paid", base_url))
.send_with_payment_options(&provider, &AcceptPaymentPolicy::Always, events)
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
assert_eq!(provider.call_count(), 1);
assert_eq!(provider.commit_count(), 1);
assert_eq!(provider.rollback_count(), 0);
assert_eq!(challenge_count.load(Ordering::SeqCst), 1);
assert_eq!(credential_count.load(Ordering::SeqCst), 1);
assert_eq!(response_count.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn test_unsuccessful_paid_retry_emits_no_payment_outcome() {
let (_, www_auth) = test_challenge();
let app = Router::new().route(
"/paid",
get(move |req: axum::http::Request<axum::body::Body>| {
let www_auth = www_auth.clone();
async move {
if req.headers().get("authorization").is_some() {
AxumStatusCode::FORBIDDEN.into_response()
} else {
(
AxumStatusCode::PAYMENT_REQUIRED,
[(WWW_AUTH_NAME, www_auth)],
"pay up",
)
.into_response()
}
}
}),
);
let base_url = spawn_server(app).await;
let provider = MockProvider::new();
let client = reqwest::Client::new();
let events = ClientEvents::default();
let response_count = Arc::new(AtomicU32::new(0));
let failed_count = Arc::new(AtomicU32::new(0));
let _response_sub = events.on_payment_response({
let response_count = response_count.clone();
move |_| {
response_count.fetch_add(1, Ordering::SeqCst);
async {}
}
});
let _failed_sub = events.on_payment_failed({
let failed_count = failed_count.clone();
move |_| {
failed_count.fetch_add(1, Ordering::SeqCst);
async {}
}
});
let resp = client
.get(format!("{}/paid", base_url))
.send_with_payment_options(&provider, &AcceptPaymentPolicy::Always, events)
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::FORBIDDEN);
assert_eq!(provider.call_count(), 1);
assert_eq!(provider.commit_count(), 0);
assert_eq!(provider.rollback_count(), 1);
assert_eq!(response_count.load(Ordering::SeqCst), 0);
assert_eq!(failed_count.load(Ordering::SeqCst), 0);
}
#[tokio::test]
async fn test_incremental_402_retries_stop_at_default_cap() {
let headers = Arc::new(
(0..DEFAULT_MAX_PAYMENT_RETRIES)
.map(|i| test_challenge_with_id(&format!("cap-{i}")).1)
.collect::<Vec<_>>(),
);
let request_count = Arc::new(AtomicU32::new(0));
let counter = request_count.clone();
let app = Router::new().route(
"/paid",
get(move || {
let headers = headers.clone();
let counter = counter.clone();
async move {
let index = counter.fetch_add(1, Ordering::SeqCst) as usize;
let www_auth = headers
.get(index)
.unwrap_or_else(|| headers.last().unwrap())
.clone();
(
AxumStatusCode::PAYMENT_REQUIRED,
[(WWW_AUTH_NAME, www_auth)],
"pay up",
)
}
}),
);
let base_url = spawn_server(app).await;
let provider = MockProvider::new();
let client = reqwest::Client::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 |_| {
failed_count.fetch_add(1, Ordering::SeqCst);
async {}
}
});
let resp = client
.get(format!("{}/paid", base_url))
.send_with_payment_options(&provider, &AcceptPaymentPolicy::Always, events)
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::PAYMENT_REQUIRED);
assert_eq!(provider.call_count(), DEFAULT_MAX_PAYMENT_RETRIES as u32);
assert_eq!(
request_count.load(Ordering::SeqCst),
DEFAULT_MAX_PAYMENT_RETRIES as u32 + 1
);
assert_eq!(failed_count.load(Ordering::SeqCst), 0);
}
#[tokio::test]
async fn test_incremental_402_retries_do_not_pay_repeated_challenge() {
let (_, www_auth) = test_challenge();
let request_count = Arc::new(AtomicU32::new(0));
let counter = request_count.clone();
let app = Router::new().route(
"/paid",
get(move || {
let www_auth = www_auth.clone();
let counter = counter.clone();
async move {
counter.fetch_add(1, Ordering::SeqCst);
(
AxumStatusCode::PAYMENT_REQUIRED,
[(WWW_AUTH_NAME, www_auth)],
"pay up",
)
}
}),
);
let base_url = spawn_server(app).await;
let provider = MockProvider::new();
let client = reqwest::Client::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_some());
assert!(ctx.error.contains("previously paid challenge"));
}
}
});
let resp = client
.get(format!("{}/paid", base_url))
.send_with_payment_options(&provider, &AcceptPaymentPolicy::Always, events)
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::PAYMENT_REQUIRED);
assert_eq!(provider.call_count(), 1);
assert_eq!(request_count.load(Ordering::SeqCst), 2);
assert_eq!(failed_count.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn test_incremental_402_retries_use_configured_cap() {
let (_, www_auth) = test_challenge();
let request_count = Arc::new(AtomicU32::new(0));
let counter = request_count.clone();
let app = Router::new().route(
"/paid",
get(move || {
let www_auth = www_auth.clone();
let counter = counter.clone();
async move {
counter.fetch_add(1, Ordering::SeqCst);
(
AxumStatusCode::PAYMENT_REQUIRED,
[(WWW_AUTH_NAME, www_auth)],
"pay up",
)
}
}),
);
let base_url = spawn_server(app).await;
let provider = MockProvider::new();
let client = reqwest::Client::new();
let resp = client
.get(format!("{}/paid", base_url))
.send_with_payment_max_retries(&provider, 1)
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::PAYMENT_REQUIRED);
assert_eq!(provider.call_count(), 1);
assert_eq!(request_count.load(Ordering::SeqCst), 2);
}
#[tokio::test]
async fn test_incremental_402_retries_pay_replacement_challenges() {
let (_, first_header) = test_challenge_with_id("first");
let (_, second_header) = test_challenge_with_id("second");
let (_, third_header) = test_challenge_with_id("third");
let headers = Arc::new([first_header, second_header, third_header]);
let request_count = Arc::new(AtomicU32::new(0));
let counter = request_count.clone();
let app = Router::new().route(
"/paid",
get(move || {
let headers = headers.clone();
let counter = counter.clone();
async move {
let index = counter.fetch_add(1, Ordering::SeqCst) as usize;
if let Some(www_auth) = headers.get(index) {
(
AxumStatusCode::PAYMENT_REQUIRED,
[(WWW_AUTH_NAME, www_auth.clone())],
"pay up",
)
.into_response()
} else {
(AxumStatusCode::OK, "ok").into_response()
}
}
}),
);
let base_url = spawn_server(app).await;
let provider = MockProvider::new();
let client = reqwest::Client::new();
let resp = client
.get(format!("{}/paid", base_url))
.send_with_payment(&provider)
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
assert_eq!(
provider.challenge_ids(),
vec![
"first".to_string(),
"second".to_string(),
"third".to_string()
]
);
assert_eq!(request_count.load(Ordering::SeqCst), 4);
}
#[tokio::test]
async fn test_challenge_received_can_override_credential() {
let (_, www_auth) = test_challenge();
let app = Router::new().route(
"/paid",
get(move |req: axum::http::Request<axum::body::Body>| {
let www_auth = www_auth.clone();
async move {
if req.headers().get("authorization").is_some() {
(AxumStatusCode::OK, "ok").into_response()
} else {
(
AxumStatusCode::PAYMENT_REQUIRED,
[(WWW_AUTH_NAME, www_auth)],
"pay up",
)
.into_response()
}
}
}),
);
let base_url = spawn_server(app).await;
let provider = MockProvider::new();
let client = reqwest::Client::new();
let events = ClientEvents::default();
let _sub = events.on(ClientEventKind::ChallengeReceived, |event| async move {
match event {
ClientEvent::ChallengeReceived(ctx) => Some(PaymentCredential::new(
ctx.challenge.to_echo(),
PaymentPayload::hash("0xoverride"),
)),
_ => None,
}
});
let resp = client
.get(format!("{}/paid", base_url))
.send_with_payment_options(&provider, &AcceptPaymentPolicy::Always, events)
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
assert_eq!(provider.call_count(), 0);
}
#[tokio::test]
async fn test_payment_event_panic_does_not_fail_request() {
let (_, www_auth) = test_challenge();
let app = Router::new().route(
"/paid",
get(move |req: axum::http::Request<axum::body::Body>| {
let www_auth = www_auth.clone();
async move {
if req.headers().get("authorization").is_some() {
(AxumStatusCode::OK, "ok").into_response()
} else {
(
AxumStatusCode::PAYMENT_REQUIRED,
[(WWW_AUTH_NAME, www_auth)],
"pay up",
)
.into_response()
}
}
}),
);
let base_url = spawn_server(app).await;
let provider = MockProvider::new();
let client = reqwest::Client::new();
let events = ClientEvents::default();
let _sub = events.on::<_, _, ()>(ClientEventKind::ChallengeReceived, |_| async move {
panic!("hook panic should be isolated");
#[allow(unreachable_code)]
()
});
let resp = client
.get(format!("{}/paid", base_url))
.send_with_payment_options(&provider, &AcceptPaymentPolicy::Always, events)
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
assert_eq!(provider.call_count(), 1);
}
#[tokio::test]
async fn test_non_402_passthrough() {
let app = Router::new().route("/free", get(|| async { "free content" }));
let base_url = spawn_server(app).await;
let provider = MockProvider::new();
let client = reqwest::Client::new();
let resp = client
.get(format!("{}/free", base_url))
.send_with_payment(&provider)
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
assert_eq!(provider.call_count(), 0);
}
#[tokio::test]
async fn test_402_missing_www_authenticate() {
let app = Router::new().route(
"/no-header",
get(|| async { AxumStatusCode::PAYMENT_REQUIRED }),
);
let base_url = spawn_server(app).await;
let provider = MockProvider::new();
let client = reqwest::Client::new();
let err = client
.get(format!("{}/no-header", base_url))
.send_with_payment(&provider)
.await
.unwrap_err();
assert!(matches!(err, HttpError::MissingChallenge));
}
#[tokio::test]
async fn test_authenticated_402_without_challenge_is_application_response() {
let (_, www_auth) = test_challenge();
let app = Router::new().route(
"/paid",
get(move |req: axum::http::Request<axum::body::Body>| {
let www_auth = www_auth.clone();
async move {
if req.headers().contains_key(AUTHORIZATION) {
(
AxumStatusCode::PAYMENT_REQUIRED,
r#"{"type":"https://paymentauth.org/problems/insufficient-balance"}"#,
)
.into_response()
} else {
(
AxumStatusCode::PAYMENT_REQUIRED,
[(WWW_AUTH_NAME, www_auth)],
"payment required",
)
.into_response()
}
}
}),
);
let base_url = spawn_server(app).await;
let provider = MockProvider::new();
let response = reqwest::Client::new()
.get(format!("{base_url}/paid"))
.send_with_payment_max_retries(&provider, 1)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::PAYMENT_REQUIRED);
assert!(response
.text()
.await
.unwrap()
.contains("insufficient-balance"));
assert_eq!(provider.commit_count(), 0);
assert_eq!(provider.rollback_count(), 1);
}
#[tokio::test]
async fn test_provider_setup_refreshes_challenge_before_payment() {
#[derive(Clone)]
struct RefreshingProvider {
preparation_count: Arc<AtomicU32>,
paid_challenges: Arc<Mutex<Vec<String>>>,
}
impl PaymentProvider for RefreshingProvider {
fn supports(&self, method: &str, intent: &str) -> bool {
method == "tempo" && intent == "charge"
}
async fn prepare_http_payment_challenge(
&self,
challenge: &PaymentChallenge,
_context: PaymentContext,
) -> Result<Option<PaymentChallenge>, MppError> {
if self.preparation_count.fetch_add(1, Ordering::SeqCst) == 0 {
Ok(None)
} else {
Ok(Some(challenge.clone()))
}
}
async fn pay(
&self,
challenge: &PaymentChallenge,
) -> Result<PaymentCredential, MppError> {
self.paid_challenges
.lock()
.unwrap()
.push(challenge.id.clone());
Ok(PaymentCredential::new(
challenge.to_echo(),
PaymentPayload::hash("0xfresh"),
))
}
}
let request_count = Arc::new(AtomicU32::new(0));
let observed = request_count.clone();
let app = Router::new().route(
"/paid",
get(move |req: axum::http::Request<axum::body::Body>| {
let count = observed.fetch_add(1, Ordering::SeqCst);
async move {
if req.headers().contains_key(AUTHORIZATION) {
(AxumStatusCode::OK, "ok").into_response()
} else {
let (_, challenge) = test_challenge_with_id(if count == 0 {
"before-setup"
} else {
"after-setup"
});
(
AxumStatusCode::PAYMENT_REQUIRED,
[(WWW_AUTH_NAME, challenge)],
"payment required",
)
.into_response()
}
}
}),
);
let provider = RefreshingProvider {
preparation_count: Arc::new(AtomicU32::new(0)),
paid_challenges: Arc::new(Mutex::new(Vec::new())),
};
let base_url = spawn_server(app).await;
let response = reqwest::Client::new()
.get(format!("{base_url}/paid"))
.send_with_payment(&provider)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(request_count.load(Ordering::SeqCst), 3);
assert_eq!(
provider.paid_challenges.lock().unwrap().as_slice(),
["after-setup"]
);
}
#[tokio::test]
async fn test_402_malformed_www_authenticate() {
let app = Router::new().route(
"/bad-header",
get(|| async {
(
AxumStatusCode::PAYMENT_REQUIRED,
[(WWW_AUTH_NAME, "garbage-not-a-valid-challenge")],
)
}),
);
let base_url = spawn_server(app).await;
let provider = MockProvider::new();
let client = reqwest::Client::new();
let err = client
.get(format!("{}/bad-header", base_url))
.send_with_payment(&provider)
.await
.unwrap_err();
assert!(matches!(err, HttpError::NoSupportedChallenge(_)));
}
#[tokio::test]
async fn test_provider_failure_bubbles_up() {
let (_, www_auth) = test_challenge();
let app = Router::new().route(
"/paid",
get(move || {
let www_auth = www_auth.clone();
async move {
(
AxumStatusCode::PAYMENT_REQUIRED,
[(WWW_AUTH_NAME, www_auth)],
)
}
}),
);
let base_url = spawn_server(app).await;
let provider = MockProvider::failing();
let client = reqwest::Client::new();
let err = client
.get(format!("{}/paid", base_url))
.send_with_payment(&provider)
.await
.unwrap_err();
assert!(matches!(err, HttpError::Payment(_)));
}
#[tokio::test]
async fn stale_session_channel_is_invalidated_and_reopened_once() {
let www_auth = challenge_header("session-1", "tempo", "session");
let requests = Arc::new(AtomicU32::new(0));
let app = Router::new().route(
"/paid",
get({
let requests = requests.clone();
move |request: axum::http::Request<axum::body::Body>| {
let www_auth = www_auth.clone();
let requests = requests.clone();
async move {
let attempt = requests.fetch_add(1, Ordering::SeqCst);
let paid = request.headers().contains_key("authorization");
match (attempt, paid) {
(0 | 2, false) => (
AxumStatusCode::PAYMENT_REQUIRED,
[(WWW_AUTH_NAME, www_auth)],
)
.into_response(),
(1, true) => AxumStatusCode::GONE.into_response(),
(3, true) => AxumStatusCode::OK.into_response(),
_ => AxumStatusCode::INTERNAL_SERVER_ERROR.into_response(),
}
}
}
}),
);
let base_url = spawn_server(app).await;
let provider = MockProvider::new();
let response = reqwest::Client::new()
.get(format!("{base_url}/paid"))
.send_with_payment(&provider)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(requests.load(Ordering::SeqCst), 4);
assert_eq!(provider.call_count(), 2);
assert_eq!(provider.rollback_count(), 1);
assert_eq!(provider.commit_count(), 1);
}
#[tokio::test]
async fn test_fetch_rejects_expired_challenge_with_pre_signing_reason() {
let www_auth = challenge_header_with_expires(
"exp",
"tempo",
"charge",
Some("2020-01-01T00:00:00Z"),
);
let app = Router::new().route(
"/paid",
get(move || {
let www_auth = www_auth.clone();
async move {
(
AxumStatusCode::PAYMENT_REQUIRED,
[(WWW_AUTH_NAME, www_auth)],
)
}
}),
);
let base_url = spawn_server(app).await;
let provider = MockProvider::new();
let events = ClientEvents::default();
let failed_count = Arc::new(AtomicU32::new(0));
let captured_reason: Arc<std::sync::Mutex<Option<PaymentFailureReason>>> =
Arc::new(Default::default());
let _failed_sub = events.on_payment_failed({
let failed_count = failed_count.clone();
let captured_reason = captured_reason.clone();
move |ctx| {
failed_count.fetch_add(1, Ordering::SeqCst);
*captured_reason.lock().unwrap() = ctx.reason.clone();
async {}
}
});
let err = reqwest::Client::new()
.get(format!("{}/paid", base_url))
.send_with_payment_options(&provider, &AcceptPaymentPolicy::Always, events)
.await
.unwrap_err();
assert!(matches!(
err,
HttpError::Payment(MppError::PaymentExpired(_))
));
assert_eq!(provider.call_count(), 0);
assert_eq!(failed_count.load(Ordering::SeqCst), 1);
assert_eq!(
captured_reason.lock().unwrap().clone(),
Some(PaymentFailureReason::PreSigningExpired {
expires: Some("2020-01-01T00:00:00Z".to_string()),
}),
);
}
#[derive(Clone)]
struct SelectiveProvider {
supported: Vec<(&'static str, &'static str)>,
pay_count: Arc<AtomicU32>,
}
impl SelectiveProvider {
fn new(supported: Vec<(&'static str, &'static str)>) -> Self {
Self {
supported,
pay_count: Arc::new(AtomicU32::new(0)),
}
}
fn call_count(&self) -> u32 {
self.pay_count.load(Ordering::SeqCst)
}
}
impl super::PaymentProvider for SelectiveProvider {
fn supports(&self, method: &str, intent: &str) -> bool {
self.supported
.iter()
.any(|(m, i)| *m == method && *i == intent)
}
async fn pay(
&self,
challenge: &PaymentChallenge,
) -> Result<PaymentCredential, MppError> {
self.pay_count.fetch_add(1, Ordering::SeqCst);
let echo = challenge.to_echo();
Ok(PaymentCredential::new(
echo,
PaymentPayload::hash("0xmockhash"),
))
}
}
fn challenge_header(id: &str, method: &str, intent: &str) -> String {
challenge_header_with_expires(id, method, intent, None)
}
fn challenge_header_with_expires(
id: &str,
method: &str,
intent: &str,
expires: Option<&str>,
) -> String {
let request =
Base64UrlJson::from_value(&serde_json::json!({"amount": "1000"})).unwrap();
let mut challenge =
PaymentChallenge::new(id, "test.example.com", method, intent, request);
if let Some(expires) = expires {
challenge = challenge.with_expires(expires);
}
format_www_authenticate(&challenge).unwrap()
}
#[tokio::test]
async fn test_multi_challenge_selects_supported_method() {
let stripe_header = challenge_header("s1", "stripe", "charge");
let tempo_header = challenge_header("t1", "tempo", "charge");
let combined = format!("{}, {}", stripe_header, tempo_header);
let app = Router::new().route(
"/paid",
get(move |req: axum::http::Request<axum::body::Body>| {
let combined = combined.clone();
async move {
if req.headers().get("authorization").is_some() {
(AxumStatusCode::OK, "ok").into_response()
} else {
(
AxumStatusCode::PAYMENT_REQUIRED,
[(WWW_AUTH_NAME, combined)],
"pay up",
)
.into_response()
}
}
}),
);
let base_url = spawn_server(app).await;
let provider = SelectiveProvider::new(vec![("tempo", "charge")]);
let client = reqwest::Client::new();
let resp = client
.get(format!("{}/paid", base_url))
.send_with_payment(&provider)
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
assert_eq!(provider.call_count(), 1);
}
#[tokio::test]
async fn test_multi_challenge_picks_first_supported() {
let tempo_header = challenge_header("t1", "tempo", "charge");
let stripe_header = challenge_header("s1", "stripe", "charge");
let combined = format!("{}, {}", tempo_header, stripe_header);
let app = Router::new().route(
"/paid",
get(move |req: axum::http::Request<axum::body::Body>| {
let combined = combined.clone();
async move {
if req.headers().get("authorization").is_some() {
(AxumStatusCode::OK, "ok").into_response()
} else {
(
AxumStatusCode::PAYMENT_REQUIRED,
[(WWW_AUTH_NAME, combined)],
"pay up",
)
.into_response()
}
}
}),
);
let base_url = spawn_server(app).await;
let provider = SelectiveProvider::new(vec![("tempo", "charge"), ("stripe", "charge")]);
let client = reqwest::Client::new();
let resp = client
.get(format!("{}/paid", base_url))
.send_with_payment(&provider)
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
assert_eq!(provider.call_count(), 1);
}
#[tokio::test]
async fn test_multi_challenge_skips_expired_preferred_supported() {
let tempo_header = challenge_header_with_expires(
"t1",
"tempo",
"charge",
Some("2020-01-01T00:00:00Z"),
);
let stripe_header = challenge_header("s1", "stripe", "charge");
let combined = format!("{}, {}", tempo_header, stripe_header);
let picked: Arc<std::sync::Mutex<Option<String>>> = Arc::new(Default::default());
let picked_clone = picked.clone();
let app = Router::new().route(
"/paid",
get(move |req: axum::http::Request<axum::body::Body>| {
let combined = combined.clone();
let picked = picked_clone.clone();
async move {
if let Some(auth) = req.headers().get("authorization") {
*picked.lock().unwrap() =
Some(auth.to_str().unwrap_or_default().to_string());
(AxumStatusCode::OK, "ok").into_response()
} else {
(
AxumStatusCode::PAYMENT_REQUIRED,
[(WWW_AUTH_NAME, combined)],
"pay up",
)
.into_response()
}
}
}),
);
let base_url = spawn_server(app).await;
let provider = SelectiveProvider::new(vec![("tempo", "charge"), ("stripe", "charge")]);
let client = reqwest::Client::new();
let resp = client
.get(format!("{}/paid", base_url))
.header("Accept-Payment", "tempo/charge, stripe/charge;q=0.5")
.send_with_payment(&provider)
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
assert_eq!(provider.call_count(), 1);
let used = picked.lock().unwrap().clone().unwrap_or_default();
let cred = crate::protocol::core::parse_authorization(&used).unwrap();
assert_eq!(cred.challenge.id, "s1");
}
#[tokio::test]
async fn test_multi_challenge_skips_malformed_expiry_first_supported() {
let bad_header =
challenge_header_with_expires("bad", "tempo", "charge", Some("not-a-date"));
let valid_header = challenge_header("valid", "tempo", "charge");
let combined = format!("{}, {}", bad_header, valid_header);
let picked: Arc<std::sync::Mutex<Option<String>>> = Arc::new(Default::default());
let picked_clone = picked.clone();
let app = Router::new().route(
"/paid",
get(move |req: axum::http::Request<axum::body::Body>| {
let combined = combined.clone();
let picked = picked_clone.clone();
async move {
if let Some(auth) = req.headers().get("authorization") {
*picked.lock().unwrap() =
Some(auth.to_str().unwrap_or_default().to_string());
(AxumStatusCode::OK, "ok").into_response()
} else {
(
AxumStatusCode::PAYMENT_REQUIRED,
[(WWW_AUTH_NAME, combined)],
"pay up",
)
.into_response()
}
}
}),
);
let base_url = spawn_server(app).await;
let provider = SelectiveProvider::new(vec![("tempo", "charge")]);
let client = reqwest::Client::new();
let resp = client
.get(format!("{}/paid", base_url))
.send_with_payment(&provider)
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
assert_eq!(provider.call_count(), 1);
let used = picked.lock().unwrap().clone().unwrap_or_default();
let cred = crate::protocol::core::parse_authorization(&used).unwrap();
assert_eq!(cred.challenge.id, "valid");
}
#[tokio::test]
async fn test_no_supported_challenge_error() {
let stripe_header = challenge_header("s1", "stripe", "charge");
let app = Router::new().route(
"/paid",
get(move || {
let stripe_header = stripe_header.clone();
async move {
(
AxumStatusCode::PAYMENT_REQUIRED,
[(WWW_AUTH_NAME, stripe_header)],
)
}
}),
);
let base_url = spawn_server(app).await;
let provider = SelectiveProvider::new(vec![("tempo", "charge")]);
let client = reqwest::Client::new();
let err = client
.get(format!("{}/paid", base_url))
.send_with_payment(&provider)
.await
.unwrap_err();
assert!(matches!(err, HttpError::NoSupportedChallenge(_)));
}
#[tokio::test]
async fn test_multiple_www_authenticate_headers() {
let stripe_header = challenge_header("s1", "stripe", "charge");
let tempo_header = challenge_header("t1", "tempo", "charge");
let app = Router::new().route(
"/paid",
get(move |req: axum::http::Request<axum::body::Body>| {
let stripe_header = stripe_header.clone();
let tempo_header = tempo_header.clone();
async move {
if req.headers().get("authorization").is_some() {
(AxumStatusCode::OK, "ok").into_response()
} else {
let mut resp =
(AxumStatusCode::PAYMENT_REQUIRED, "pay up").into_response();
let headers = resp.headers_mut();
headers.append(WWW_AUTH_NAME, stripe_header.parse().unwrap());
headers.append(WWW_AUTH_NAME, tempo_header.parse().unwrap());
resp
}
}
}),
);
let base_url = spawn_server(app).await;
let provider = SelectiveProvider::new(vec![("tempo", "charge")]);
let client = reqwest::Client::new();
let resp = client
.get(format!("{}/paid", base_url))
.send_with_payment(&provider)
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
assert_eq!(provider.call_count(), 1);
}
#[tokio::test]
async fn test_intent_matching() {
let session_header = challenge_header("t1", "tempo", "session");
let app = Router::new().route(
"/paid",
get(move || {
let session_header = session_header.clone();
async move {
(
AxumStatusCode::PAYMENT_REQUIRED,
[(WWW_AUTH_NAME, session_header)],
)
}
}),
);
let base_url = spawn_server(app).await;
let provider = SelectiveProvider::new(vec![("tempo", "charge")]);
let client = reqwest::Client::new();
let err = client
.get(format!("{}/paid", base_url))
.send_with_payment(&provider)
.await
.unwrap_err();
assert!(matches!(err, HttpError::NoSupportedChallenge(_)));
}
#[derive(Clone)]
struct AdvertisingProvider;
impl super::PaymentProvider for AdvertisingProvider {
fn supports(&self, _method: &str, _intent: &str) -> bool {
true
}
async fn pay(
&self,
_challenge: &PaymentChallenge,
) -> Result<PaymentCredential, MppError> {
unimplemented!("not used in policy test")
}
fn accept_payment_header(&self) -> Option<String> {
Some("tempo/charge".to_string())
}
}
async fn spawn_header_capture() -> (String, Arc<std::sync::Mutex<Option<String>>>) {
let captured: Arc<std::sync::Mutex<Option<String>>> = Arc::new(Default::default());
let captured_clone = captured.clone();
let app = Router::new().route(
"/probe",
get(move |req: axum::http::Request<axum::body::Body>| {
let captured = captured_clone.clone();
async move {
let v = req
.headers()
.get("accept-payment")
.and_then(|h| h.to_str().ok())
.map(|s| s.to_string());
*captured.lock().unwrap() = v;
AxumStatusCode::OK
}
}),
);
let url = spawn_server(app).await;
(url, captured)
}
#[tokio::test]
async fn test_send_with_payment_default_injects() {
let (base_url, captured) = spawn_header_capture().await;
reqwest::Client::new()
.get(format!("{}/probe", base_url))
.send_with_payment(&AdvertisingProvider)
.await
.unwrap();
assert_eq!(captured.lock().unwrap().as_deref(), Some("tempo/charge"));
}
#[tokio::test]
async fn test_send_with_payment_policy_never_blocks() {
let (base_url, captured) = spawn_header_capture().await;
reqwest::Client::new()
.get(format!("{}/probe", base_url))
.send_with_payment_policy(&AdvertisingProvider, &AcceptPaymentPolicy::Never)
.await
.unwrap();
assert_eq!(captured.lock().unwrap().as_deref(), None);
}
#[tokio::test]
async fn test_caller_header_not_overwritten() {
let (base_url, captured) = spawn_header_capture().await;
reqwest::Client::new()
.get(format!("{}/probe", base_url))
.header("Accept-Payment", "stripe/charge")
.send_with_payment(&AdvertisingProvider)
.await
.unwrap();
assert_eq!(captured.lock().unwrap().as_deref(), Some("stripe/charge"));
}
#[tokio::test]
async fn test_caller_header_drives_ranking() {
let tempo_header = challenge_header("t1", "tempo", "charge");
let stripe_header = challenge_header("s1", "stripe", "charge");
let combined = format!("{}, {}", tempo_header, stripe_header);
let picked: Arc<std::sync::Mutex<Option<String>>> = Arc::new(Default::default());
let picked_clone = picked.clone();
let app = Router::new().route(
"/paid",
get(move |req: axum::http::Request<axum::body::Body>| {
let combined = combined.clone();
let picked = picked_clone.clone();
async move {
if let Some(auth) = req.headers().get("authorization") {
let v = auth.to_str().unwrap_or("").to_string();
*picked.lock().unwrap() = Some(v);
(AxumStatusCode::OK, "ok").into_response()
} else {
(
AxumStatusCode::PAYMENT_REQUIRED,
[(WWW_AUTH_NAME, combined)],
"pay",
)
.into_response()
}
}
}),
);
let base_url = spawn_server(app).await;
let provider = SelectiveProvider::new(vec![("tempo", "charge"), ("stripe", "charge")]);
reqwest::Client::new()
.get(format!("{}/paid", base_url))
.header("Accept-Payment", "stripe/charge, tempo/charge;q=0.1")
.send_with_payment(&provider)
.await
.unwrap();
let used = picked.lock().unwrap().clone().unwrap_or_default();
let cred = crate::protocol::core::parse_authorization(&used).unwrap();
assert_eq!(
cred.challenge.id, "s1",
"expected stripe challenge (id s1) to be picked, got id: {}",
cred.challenge.id
);
}
#[tokio::test]
async fn test_send_with_payment_policy_same_origin_mismatch() {
let (base_url, captured) = spawn_header_capture().await;
reqwest::Client::new()
.get(format!("{}/probe", base_url))
.send_with_payment_policy(
&AdvertisingProvider,
&AcceptPaymentPolicy::SameOrigin {
same_origin: "https://app.example.com".to_string(),
},
)
.await
.unwrap();
assert_eq!(captured.lock().unwrap().as_deref(), None);
}
}
}