#![cfg(feature = "http")]
use http::{HeaderMap, HeaderValue};
use webmcp::apply_origin_trial;
#[test]
fn http_headers_fill_only_when_missing_and_invalid_tokens_do_not_mutate() {
let mut headers = HeaderMap::new();
apply_origin_trial(&mut headers, "").unwrap();
assert!(headers.is_empty());
assert!(apply_origin_trial(&mut headers, "bad\r\nvalue").is_err());
assert!(headers.is_empty());
apply_origin_trial(&mut headers, "abc").unwrap();
apply_origin_trial(&mut headers, "def").unwrap();
assert_eq!(headers["Origin-Trial"], "abc");
headers.insert("origin-trial", HeaderValue::from_static(""));
apply_origin_trial(&mut headers, "new").unwrap();
assert_eq!(headers["origin-trial"], "");
apply_origin_trial(&mut headers, "ignored\ninvalid").unwrap();
}
#[cfg(feature = "tower")]
mod tower_tests {
use http::{HeaderValue, Response};
use std::{
future::Future,
pin::Pin,
sync::Arc,
task::{Context, Poll, Wake, Waker},
};
use tower::{Layer, Service};
use webmcp::OriginTrialLayer;
struct Noop;
impl Wake for Noop {
fn wake(self: Arc<Self>) {}
}
struct Inner {
ready: bool,
fail: bool,
existing: Option<&'static str>,
}
struct ResponseFuture {
pending: bool,
response: Option<Result<Response<String>, &'static str>>,
}
impl Future for ResponseFuture {
type Output = Result<Response<String>, &'static str>;
fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
if self.pending {
self.pending = false;
cx.waker().wake_by_ref();
Poll::Pending
} else {
Poll::Ready(self.response.take().unwrap())
}
}
}
impl Service<()> for Inner {
type Response = Response<String>;
type Error = &'static str;
type Future = ResponseFuture;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
if self.ready {
Poll::Ready(Ok(()))
} else {
self.ready = true;
cx.waker().wake_by_ref();
Poll::Pending
}
}
fn call(&mut self, _: ()) -> Self::Future {
let response = if self.fail {
Err("inner error")
} else {
let mut response = Response::builder()
.status(201)
.body("body".to_owned())
.unwrap();
response
.headers_mut()
.insert("origin-agent-cluster", HeaderValue::from_static("?0"));
if let Some(value) = self.existing {
response
.headers_mut()
.insert("origin-trial", HeaderValue::from_static(value));
}
Ok(response)
};
ResponseFuture {
pending: true,
response: Some(response),
}
}
}
#[test]
fn layer_preserves_readiness_pending_responses_errors_status_and_body() {
let waker = Waker::from(Arc::new(Noop));
let mut cx = Context::from_waker(&waker);
for fail in [false, true] {
let layer = OriginTrialLayer::new("token").unwrap();
let mut service = layer.layer(Inner {
ready: false,
fail,
existing: None,
});
assert!(service.poll_ready(&mut cx).is_pending());
assert_eq!(service.poll_ready(&mut cx), Poll::Ready(Ok(())));
let mut future = Box::pin(service.call(()));
assert!(future.as_mut().poll(&mut cx).is_pending());
let Poll::Ready(result) = future.as_mut().poll(&mut cx) else {
panic!("expected ready");
};
if fail {
assert_eq!(result.unwrap_err(), "inner error");
} else {
let response = result.unwrap();
assert_eq!(response.status(), 201);
assert_eq!(response.body(), "body");
assert_eq!(response.headers()["origin-trial"], "token");
assert_eq!(response.headers()["origin-agent-cluster"], "?0");
}
}
}
#[test]
fn layer_preserves_existing_headers_and_empty_token_is_passthrough() {
let waker = Waker::from(Arc::new(Noop));
let mut cx = Context::from_waker(&waker);
for token in ["", "token"] {
for existing in [None, Some(""), Some("previous")] {
let mut service = OriginTrialLayer::new(token)
.unwrap()
.warn_on_oac_opt_out(true)
.layer(Inner {
ready: true,
fail: false,
existing,
});
let mut future = Box::pin(service.call(()));
assert!(future.as_mut().poll(&mut cx).is_pending());
let Poll::Ready(Ok(response)) = future.as_mut().poll(&mut cx) else {
panic!("expected response");
};
let expected = existing.or(if token.is_empty() { None } else { Some(token) });
assert_eq!(
response
.headers()
.get("origin-trial")
.map(|h| h.to_str().unwrap()),
expected
);
assert_eq!(response.headers()["origin-agent-cluster"], "?0");
}
}
assert!(OriginTrialLayer::new("invalid\n").is_err());
}
}