rama_http/layer/html_rewrite/
service.rs1use 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#[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 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 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 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
110fn is_html_content_type(content_type: &ContentType) -> bool {
112 content_type.mime().essence_str() == "text/html"
115}
116
117#[derive(Clone)]
121pub struct HtmlRewriteLayer<H> {
122 selectors: Box<[Selector]>,
123 handler: H,
124 policy: BodyRewritePolicy,
125}
126
127impl<H> HtmlRewriteLayer<H> {
128 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 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 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}