Skip to main content

rama_http/layer/html_rewrite/
service.rs

1//! [`Service`] that rewrites `text/html` response bodies.
2
3use std::fmt;
4
5use rama_core::error::BoxError;
6use rama_core::extensions::Extensions;
7use rama_core::{Layer, Service};
8use rama_utils::macros::{define_inner_service_accessors, generate_set_and_with};
9
10use super::HtmlRewriteBody;
11use crate::headers::ContentType;
12use crate::layer::remove_header::{
13    remove_cache_validation_response_headers, remove_payload_metadata_headers,
14};
15use crate::layer::util::rewrite_policy::BodyRewritePolicy;
16use crate::protocols::html::rewrite::ElementContentHandler;
17use crate::protocols::html::selector::Selector;
18use crate::{HeaderMap, Request, Response, StreamingBody};
19
20/// Rewrites the `text/html` response bodies of the underlying service, using
21/// rama's streaming [`HtmlRewriter`](crate::protocols::html::rewrite::HtmlRewriter).
22///
23/// See the [module docs](crate::layer::html_rewrite) for details. Construct it
24/// directly with [`new`](Self::new) or via [`HtmlRewriteLayer`].
25#[derive(Clone)]
26pub struct HtmlRewrite<S, H> {
27    pub(crate) inner: S,
28    pub(crate) selectors: Box<[Selector]>,
29    pub(crate) handler: H,
30    policy: BodyRewritePolicy,
31}
32
33impl<S, H> HtmlRewrite<S, H> {
34    /// Creates a new [`HtmlRewrite`] service.
35    ///
36    /// The `selector` index passed to the handler is the index into
37    /// `selectors`; keep the two aligned.
38    pub fn new(inner: S, selectors: impl IntoIterator<Item = Selector>, handler: H) -> Self {
39        Self {
40            inner,
41            selectors: selectors.into_iter().collect(),
42            handler,
43            policy: BodyRewritePolicy::unencoded_content_type(is_html_content_type),
44        }
45    }
46
47    generate_set_and_with! {
48        /// Sets a custom response rewrite policy.
49        ///
50        /// The predicate receives the response headers and extensions and can
51        /// narrow rewriting beyond the built-in `Content-Encoding` guard.
52        pub fn rewrite_policy(
53            mut self,
54            policy: impl Fn(&HeaderMap, &Extensions) -> bool + Send + Sync + 'static,
55        ) -> Self {
56            self.policy = BodyRewritePolicy::custom(policy);
57            self
58        }
59    }
60
61    define_inner_service_accessors!();
62}
63
64impl<S: fmt::Debug, H> fmt::Debug for HtmlRewrite<S, H> {
65    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
66        f.debug_struct("HtmlRewrite")
67            .field("inner", &self.inner)
68            .field("selectors", &self.selectors)
69            .field("handler", &std::any::type_name::<H>())
70            .field("policy", &self.policy)
71            .finish()
72    }
73}
74
75impl<S, H, ReqBody, ResBody> Service<Request<ReqBody>> for HtmlRewrite<S, H>
76where
77    S: Service<Request<ReqBody>, Output = Response<ResBody>>,
78    ResBody: StreamingBody<Data: Send + 'static, Error: Into<BoxError> + Send + 'static>
79        + Send
80        + 'static,
81    H: ElementContentHandler + Clone + Send + Sync + 'static,
82    ReqBody: Send + 'static,
83{
84    type Output = Response<HtmlRewriteBody<ResBody, H>>;
85    type Error = S::Error;
86
87    async fn serve(&self, req: Request<ReqBody>) -> Result<Self::Output, Self::Error> {
88        let res = self.inner.serve(req).await?;
89        let (mut parts, body) = res.into_parts();
90        let rewrite = !self.selectors.is_empty()
91            && self
92                .policy
93                .should_rewrite(&parts.headers, &parts.extensions);
94        let body = if rewrite {
95            // Rewriting changes the body length and invalidates range support,
96            // so drop the now-stale payload metadata (Content-Length,
97            // Transfer-Encoding, Accept-Ranges, …) and representation
98            // validators (ETag, Last-Modified, …); the response becomes
99            // chunked / unknown-length.
100            remove_payload_metadata_headers(&mut parts.headers);
101            remove_cache_validation_response_headers(&mut parts.headers);
102            HtmlRewriteBody::new(body, &self.selectors, self.handler.clone())
103        } else {
104            HtmlRewriteBody::passthrough(body)
105        };
106        Ok(Response::from_parts(parts, body))
107    }
108}
109
110/// Whether this content type is an HTML document that can be rewritten.
111fn is_html_content_type(content_type: &ContentType) -> bool {
112    // `essence_str` is the `type/subtype` without parameters (e.g. drops
113    // `; charset=...`) and is already lower-cased by the mime parser.
114    content_type.mime().essence_str() == "text/html"
115}
116
117/// Layer that applies [`HtmlRewrite`] to the responses of the wrapped service.
118///
119/// See the [module docs](crate::layer::html_rewrite).
120#[derive(Clone)]
121pub struct HtmlRewriteLayer<H> {
122    selectors: Box<[Selector]>,
123    handler: H,
124    policy: BodyRewritePolicy,
125}
126
127impl<H> HtmlRewriteLayer<H> {
128    /// Creates a new [`HtmlRewriteLayer`] that rewrites elements matching
129    /// `selectors` with `handler` (the handler is cloned per response, so it
130    /// starts fresh for each one).
131    pub fn new(selectors: impl IntoIterator<Item = Selector>, handler: H) -> Self {
132        Self {
133            selectors: selectors.into_iter().collect(),
134            handler,
135            policy: BodyRewritePolicy::unencoded_content_type(is_html_content_type),
136        }
137    }
138
139    generate_set_and_with! {
140        /// Sets a custom response rewrite policy.
141        ///
142        /// The predicate receives the response headers and extensions and can
143        /// narrow rewriting beyond the built-in `Content-Encoding` guard.
144        pub fn rewrite_policy(
145            mut self,
146            policy: impl Fn(&HeaderMap, &Extensions) -> bool + Send + Sync + 'static,
147        ) -> Self {
148            self.policy = BodyRewritePolicy::custom(policy);
149            self
150        }
151    }
152
153    /// Wraps a body directly using this layer's selector set and handler.
154    ///
155    /// This is useful for services that need request-specific gating before
156    /// deciding whether a single response body should be rewritten, while
157    /// still sharing the same layer configuration.
158    pub fn rewrite_body<B>(&self, body: B) -> HtmlRewriteBody<B, H>
159    where
160        H: ElementContentHandler + Clone,
161    {
162        HtmlRewriteBody::new(body, &self.selectors, self.handler.clone())
163    }
164}
165
166impl<H: fmt::Debug> fmt::Debug for HtmlRewriteLayer<H> {
167    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
168        f.debug_struct("HtmlRewriteLayer")
169            .field("selectors", &self.selectors)
170            .field("handler", &self.handler)
171            .field("policy", &self.policy)
172            .finish()
173    }
174}
175
176impl<S, H: Clone> Layer<S> for HtmlRewriteLayer<H> {
177    type Service = HtmlRewrite<S, H>;
178
179    fn layer(&self, inner: S) -> Self::Service {
180        HtmlRewrite {
181            inner,
182            selectors: self.selectors.clone(),
183            handler: self.handler.clone(),
184            policy: self.policy.clone(),
185        }
186    }
187
188    fn into_layer(self, inner: S) -> Self::Service {
189        HtmlRewrite {
190            inner,
191            selectors: self.selectors,
192            handler: self.handler,
193            policy: self.policy,
194        }
195    }
196}
197
198#[cfg(test)]
199mod tests {
200    use super::*;
201    use crate::headers::HeaderMapExt;
202    use crate::{HeaderMap, header};
203
204    #[test]
205    fn rewrite_content_type_policy() {
206        let cases = [
207            ("text/html", true),
208            ("text/html; charset=utf-8", true),
209            ("application/xhtml+xml", false),
210            ("application/json", false),
211        ];
212
213        for (content_type, expected) in cases {
214            let mut headers = HeaderMap::new();
215            headers.insert(
216                header::CONTENT_TYPE,
217                content_type.parse().expect("valid header"),
218            );
219            let content_type = headers.typed_get::<ContentType>().expect("content type");
220            assert_eq!(
221                is_html_content_type(&content_type),
222                expected,
223                "{content_type}"
224            );
225        }
226    }
227
228    #[test]
229    fn rewrite_policy_skips_content_encoded_html() {
230        let mut headers = HeaderMap::new();
231        headers.insert(header::CONTENT_TYPE, "text/html".parse().unwrap());
232        headers.insert(header::CONTENT_ENCODING, "gzip".parse().unwrap());
233        let policy = BodyRewritePolicy::unencoded_content_type(is_html_content_type);
234        assert!(!policy.should_rewrite(&headers, &Extensions::new()));
235    }
236}