use crate::escape_html;
pub const ORIGIN_TRIAL_HEADER: &str = "Origin-Trial";
pub fn origin_trial_header(token: &str) -> (&'static str, String) {
(ORIGIN_TRIAL_HEADER, token.to_owned())
}
pub fn apply_origin_trial_headers(headers: &mut Vec<(String, String)>, token: &str) {
if !token.is_empty()
&& !headers
.iter()
.any(|(name, _)| name.eq_ignore_ascii_case(ORIGIN_TRIAL_HEADER))
{
headers.push((ORIGIN_TRIAL_HEADER.into(), token.into()));
}
}
#[cfg(not(feature = "http"))]
pub use apply_origin_trial_headers as apply_origin_trial;
#[cfg(feature = "http")]
pub fn apply_origin_trial(
headers: &mut http::HeaderMap,
token: &str,
) -> Result<(), http::header::InvalidHeaderValue> {
if !token.is_empty() && !headers.contains_key(ORIGIN_TRIAL_HEADER) {
headers.insert(
http::header::HeaderName::from_static("origin-trial"),
http::HeaderValue::from_str(token)?,
);
}
Ok(())
}
pub fn origin_trial_meta_tag(token: &str) -> String {
if token.is_empty() {
String::new()
} else {
format!(
"<meta http-equiv=\"origin-trial\" content=\"{}\">",
escape_html(token)
)
}
}
#[cfg(feature = "tower")]
mod middleware {
use http::{HeaderValue, Response};
use std::{
future::Future,
pin::Pin,
sync::{
atomic::{AtomicBool, Ordering},
Arc,
},
task::{Context, Poll},
};
use tower::{Layer, Service};
const WARNING: &str = "WebMCP: Origin-Agent-Cluster: ?0 may cause SecurityError in older Origin-Trial builds; header left unchanged";
#[derive(Clone, Debug)]
pub struct OriginTrialLayer {
token: Option<HeaderValue>,
warn_on_oac_opt_out: bool,
warned: Arc<AtomicBool>,
}
impl OriginTrialLayer {
pub fn new(token: &str) -> Result<Self, http::header::InvalidHeaderValue> {
Ok(Self {
token: if token.is_empty() {
None
} else {
Some(HeaderValue::from_str(token)?)
},
warn_on_oac_opt_out: false,
warned: Arc::new(AtomicBool::new(false)),
})
}
pub fn warn_on_oac_opt_out(mut self, enabled: bool) -> Self {
self.warn_on_oac_opt_out = enabled;
self
}
fn apply<B>(&self, response: &mut Response<B>) {
if let Some(token) = &self.token {
let headers = response.headers_mut();
if !headers.contains_key("origin-trial") {
headers.insert("origin-trial", token.clone());
}
if self.warn_on_oac_opt_out
&& headers
.get_all("origin-agent-cluster")
.iter()
.any(|value| value == "?0")
&& !self.warned.swap(true, Ordering::Relaxed)
{
eprintln!("{WARNING}");
}
}
}
}
impl<S> Layer<S> for OriginTrialLayer {
type Service = OriginTrialService<S>;
fn layer(&self, inner: S) -> Self::Service {
OriginTrialService {
inner,
layer: self.clone(),
}
}
}
#[derive(Clone, Debug)]
pub struct OriginTrialService<S> {
inner: S,
layer: OriginTrialLayer,
}
impl<S, Request, B> Service<Request> for OriginTrialService<S>
where
S: Service<Request, Response = Response<B>>,
{
type Response = Response<B>;
type Error = S::Error;
type Future = OriginTrialFuture<S::Future>;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, request: Request) -> Self::Future {
OriginTrialFuture {
inner: Box::pin(self.inner.call(request)),
layer: self.layer.clone(),
}
}
}
#[derive(Debug)]
pub struct OriginTrialFuture<F> {
inner: Pin<Box<F>>,
layer: OriginTrialLayer,
}
impl<F, B, E> Future for OriginTrialFuture<F>
where
F: Future<Output = Result<Response<B>, E>>,
{
type Output = Result<Response<B>, E>;
fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
match self.inner.as_mut().poll(cx) {
Poll::Ready(Ok(mut response)) => {
self.layer.apply(&mut response);
Poll::Ready(Ok(response))
}
other => other,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn warning_is_opt_in_once_shared_and_never_rewrites() {
let mut response = Response::new(());
response
.headers_mut()
.insert("origin-agent-cluster", HeaderValue::from_static("?0"));
let quiet = OriginTrialLayer::new("token").unwrap();
quiet.apply(&mut response);
assert!(!quiet.warned.load(Ordering::Relaxed));
let layer = quiet.warn_on_oac_opt_out(true);
layer.clone().apply(&mut response);
assert!(layer.warned.load(Ordering::Relaxed));
layer.apply(&mut response);
assert_eq!(response.headers()["origin-agent-cluster"], "?0");
let empty = OriginTrialLayer::new("").unwrap().warn_on_oac_opt_out(true);
empty.apply(&mut response);
assert!(!empty.warned.load(Ordering::Relaxed));
}
}
}
#[cfg(feature = "tower")]
pub use middleware::{OriginTrialFuture, OriginTrialLayer, OriginTrialService};