Skip to main content

rama_http/layer/json_rewrite/
service.rs

1//! [`Service`]s that rewrite JSON request or response bodies.
2
3use std::fmt;
4
5use rama_core::error::BoxError;
6use rama_core::extensions::Extensions;
7use rama_core::{Layer, Service};
8use rama_json::path::JsonPath;
9use rama_json::rewrite::JsonValueHandler;
10use rama_json::tokenizer::DEFAULT_MAX_BUFFERED_BYTES;
11use rama_utils::macros::{define_inner_service_accessors, generate_set_and_with};
12
13use super::JsonRewriteBody;
14use crate::headers::ContentType;
15use crate::layer::remove_header::{
16    remove_cache_validation_response_headers, remove_payload_metadata_headers,
17};
18use crate::layer::util::rewrite_policy::BodyRewritePolicy;
19use crate::{HeaderMap, Request, Response, StreamingBody};
20
21/// Rewrites JSON response bodies of the underlying service, using rama's
22/// streaming [`JsonRewriter`](rama_json::rewrite::JsonRewriter).
23///
24/// See the [module docs](crate::layer::json_rewrite) for details. Construct it
25/// directly with [`new`](Self::new) or via [`JsonRewriteLayer`].
26#[derive(Clone)]
27pub struct JsonRewrite<S, H> {
28    pub(crate) inner: S,
29    pub(crate) selectors: Box<[JsonPath]>,
30    pub(crate) handler: H,
31    policy: BodyRewritePolicy,
32    max_buffered_bytes: usize,
33}
34
35impl<S, H> JsonRewrite<S, H> {
36    /// Creates a new [`JsonRewrite`] service.
37    ///
38    /// The `selector` index passed to the handler is the index into
39    /// `selectors`; keep the two aligned.
40    pub fn new(inner: S, selectors: impl IntoIterator<Item = JsonPath>, handler: H) -> Self {
41        Self {
42            inner,
43            selectors: selectors.into_iter().collect(),
44            handler,
45            policy: BodyRewritePolicy::unencoded_content_type(is_json_content_type),
46            max_buffered_bytes: DEFAULT_MAX_BUFFERED_BYTES,
47        }
48    }
49
50    generate_set_and_with! {
51        /// Sets a custom response rewrite policy.
52        ///
53        /// The predicate receives the response headers and extensions and can
54        /// narrow rewriting beyond the built-in `Content-Encoding` guard.
55        pub fn rewrite_policy(
56            mut self,
57            policy: impl Fn(&HeaderMap, &Extensions) -> bool + Send + Sync + 'static,
58        ) -> Self {
59            self.policy = BodyRewritePolicy::custom(policy);
60            self
61        }
62    }
63
64    generate_set_and_with! {
65        /// Sets the tokenizer buffered-input limit for each rewritten body.
66        pub fn max_buffered_bytes(mut self, max_buffered_bytes: usize) -> Self {
67            self.max_buffered_bytes = max_buffered_bytes;
68            self
69        }
70    }
71
72    define_inner_service_accessors!();
73}
74
75impl<S: fmt::Debug, H> fmt::Debug for JsonRewrite<S, H> {
76    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
77        f.debug_struct("JsonRewrite")
78            .field("inner", &self.inner)
79            .field("selectors", &self.selectors)
80            .field("handler", &std::any::type_name::<H>())
81            .field("policy", &self.policy)
82            .field("max_buffered_bytes", &self.max_buffered_bytes)
83            .finish()
84    }
85}
86
87impl<S, H, ReqBody, ResBody> Service<Request<ReqBody>> for JsonRewrite<S, H>
88where
89    S: Service<Request<ReqBody>, Output = Response<ResBody>>,
90    ResBody: StreamingBody<Data: Send + 'static, Error: Into<BoxError> + Send + 'static>
91        + Send
92        + 'static,
93    H: JsonValueHandler + Clone + Send + Sync + 'static,
94    ReqBody: Send + 'static,
95{
96    type Output = Response<JsonRewriteBody<ResBody, H>>;
97    type Error = S::Error;
98
99    async fn serve(&self, req: Request<ReqBody>) -> Result<Self::Output, Self::Error> {
100        let res = self.inner.serve(req).await?;
101        let (mut parts, body) = res.into_parts();
102        let rewrite = !self.selectors.is_empty()
103            && self
104                .policy
105                .should_rewrite(&parts.headers, &parts.extensions);
106        let body = if rewrite {
107            // Rewriting changes the body length and invalidates range support,
108            // so drop the now-stale payload metadata (Content-Length,
109            // Transfer-Encoding, Accept-Ranges, ...) and representation
110            // validators (ETag, Last-Modified, ...); the response becomes
111            // chunked / unknown-length.
112            remove_payload_metadata_headers(&mut parts.headers);
113            remove_cache_validation_response_headers(&mut parts.headers);
114            JsonRewriteBody::with_max_buffered_bytes(
115                body,
116                self.selectors.iter().cloned(),
117                self.handler.clone(),
118                self.max_buffered_bytes,
119            )
120        } else {
121            JsonRewriteBody::passthrough(body)
122        };
123        Ok(Response::from_parts(parts, body))
124    }
125}
126
127/// Rewrites JSON request bodies before they reach the underlying service,
128/// using rama's streaming [`JsonRewriter`](rama_json::rewrite::JsonRewriter).
129///
130/// See the [module docs](crate::layer::json_rewrite) for details. Construct it
131/// directly with [`new`](Self::new) or via [`JsonRequestRewriteLayer`].
132#[derive(Clone)]
133pub struct JsonRequestRewrite<S, H> {
134    pub(crate) inner: S,
135    pub(crate) selectors: Box<[JsonPath]>,
136    pub(crate) handler: H,
137    policy: BodyRewritePolicy,
138    max_buffered_bytes: usize,
139}
140
141impl<S, H> JsonRequestRewrite<S, H> {
142    /// Creates a new [`JsonRequestRewrite`] service.
143    ///
144    /// The `selector` index passed to the handler is the index into
145    /// `selectors`; keep the two aligned.
146    pub fn new(inner: S, selectors: impl IntoIterator<Item = JsonPath>, handler: H) -> Self {
147        Self {
148            inner,
149            selectors: selectors.into_iter().collect(),
150            handler,
151            policy: BodyRewritePolicy::unencoded_content_type(is_json_content_type),
152            max_buffered_bytes: DEFAULT_MAX_BUFFERED_BYTES,
153        }
154    }
155
156    generate_set_and_with! {
157        /// Sets a custom request rewrite policy.
158        ///
159        /// The predicate receives the request headers and extensions and can
160        /// narrow rewriting beyond the built-in `Content-Encoding` guard.
161        pub fn rewrite_policy(
162            mut self,
163            policy: impl Fn(&HeaderMap, &Extensions) -> bool + Send + Sync + 'static,
164        ) -> Self {
165            self.policy = BodyRewritePolicy::custom(policy);
166            self
167        }
168    }
169
170    generate_set_and_with! {
171        /// Sets the tokenizer buffered-input limit for each rewritten body.
172        pub fn max_buffered_bytes(mut self, max_buffered_bytes: usize) -> Self {
173            self.max_buffered_bytes = max_buffered_bytes;
174            self
175        }
176    }
177
178    define_inner_service_accessors!();
179}
180
181impl<S: fmt::Debug, H> fmt::Debug for JsonRequestRewrite<S, H> {
182    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
183        f.debug_struct("JsonRequestRewrite")
184            .field("inner", &self.inner)
185            .field("selectors", &self.selectors)
186            .field("handler", &std::any::type_name::<H>())
187            .field("policy", &self.policy)
188            .field("max_buffered_bytes", &self.max_buffered_bytes)
189            .finish()
190    }
191}
192
193impl<S, H, ReqBody> Service<Request<ReqBody>> for JsonRequestRewrite<S, H>
194where
195    S: Service<Request<JsonRewriteBody<ReqBody, H>>>,
196    ReqBody: StreamingBody<Data: Send + 'static, Error: Into<BoxError> + Send + 'static>
197        + Send
198        + 'static,
199    H: JsonValueHandler + Clone + Send + Sync + 'static,
200{
201    type Output = S::Output;
202    type Error = S::Error;
203
204    async fn serve(&self, req: Request<ReqBody>) -> Result<Self::Output, Self::Error> {
205        let (mut parts, body) = req.into_parts();
206        let rewrite = !self.selectors.is_empty()
207            && self
208                .policy
209                .should_rewrite(&parts.headers, &parts.extensions);
210        let body = if rewrite {
211            // Rewriting changes the request body length and transfer shape, so
212            // drop stale payload metadata before forwarding upstream.
213            remove_payload_metadata_headers(&mut parts.headers);
214            JsonRewriteBody::with_max_buffered_bytes(
215                body,
216                self.selectors.iter().cloned(),
217                self.handler.clone(),
218                self.max_buffered_bytes,
219            )
220        } else {
221            JsonRewriteBody::passthrough(body)
222        };
223        self.inner.serve(Request::from_parts(parts, body)).await
224    }
225}
226
227/// Whether this content type is JSON that can be rewritten.
228fn is_json_content_type(content_type: &ContentType) -> bool {
229    let mime = content_type.mime();
230    mime.type_() == "application"
231        && (mime.subtype() == "json" || mime.suffix().is_some_and(|name| name == "json"))
232}
233
234/// Layer that applies [`JsonRewrite`] to the responses of the wrapped service.
235///
236/// See the [module docs](crate::layer::json_rewrite).
237#[derive(Clone)]
238pub struct JsonRewriteLayer<H> {
239    selectors: Box<[JsonPath]>,
240    handler: H,
241    policy: BodyRewritePolicy,
242    max_buffered_bytes: usize,
243}
244
245/// Layer that applies [`JsonRequestRewrite`] to requests before they reach the
246/// wrapped service.
247///
248/// See the [module docs](crate::layer::json_rewrite).
249#[derive(Clone)]
250pub struct JsonRequestRewriteLayer<H> {
251    selectors: Box<[JsonPath]>,
252    handler: H,
253    policy: BodyRewritePolicy,
254    max_buffered_bytes: usize,
255}
256
257impl<H> JsonRequestRewriteLayer<H> {
258    /// Creates a new [`JsonRequestRewriteLayer`] that rewrites values matching
259    /// `selectors` with `handler` (the handler is cloned per request, so it
260    /// starts fresh for each one).
261    pub fn new(selectors: impl IntoIterator<Item = JsonPath>, handler: H) -> Self {
262        Self {
263            selectors: selectors.into_iter().collect(),
264            handler,
265            policy: BodyRewritePolicy::unencoded_content_type(is_json_content_type),
266            max_buffered_bytes: DEFAULT_MAX_BUFFERED_BYTES,
267        }
268    }
269
270    generate_set_and_with! {
271        /// Sets a custom request rewrite policy.
272        ///
273        /// The predicate receives the request headers and extensions and can
274        /// narrow rewriting beyond the built-in `Content-Encoding` guard.
275        pub fn rewrite_policy(
276            mut self,
277            policy: impl Fn(&HeaderMap, &Extensions) -> bool + Send + Sync + 'static,
278        ) -> Self {
279            self.policy = BodyRewritePolicy::custom(policy);
280            self
281        }
282    }
283
284    generate_set_and_with! {
285        /// Sets the tokenizer buffered-input limit for each rewritten body.
286        pub fn max_buffered_bytes(mut self, max_buffered_bytes: usize) -> Self {
287            self.max_buffered_bytes = max_buffered_bytes;
288            self
289        }
290    }
291
292    /// Wraps a body directly using this layer's selector set and handler.
293    ///
294    /// This is useful for services that need request-specific gating before
295    /// deciding whether a single request body should be rewritten, while
296    /// still sharing the same layer configuration.
297    pub fn rewrite_body<B>(&self, body: B) -> JsonRewriteBody<B, H>
298    where
299        H: JsonValueHandler + Clone,
300    {
301        JsonRewriteBody::with_max_buffered_bytes(
302            body,
303            self.selectors.iter().cloned(),
304            self.handler.clone(),
305            self.max_buffered_bytes,
306        )
307    }
308}
309
310impl<H: fmt::Debug> fmt::Debug for JsonRequestRewriteLayer<H> {
311    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
312        f.debug_struct("JsonRequestRewriteLayer")
313            .field("selectors", &self.selectors)
314            .field("handler", &self.handler)
315            .field("policy", &self.policy)
316            .field("max_buffered_bytes", &self.max_buffered_bytes)
317            .finish()
318    }
319}
320
321impl<S, H: Clone> Layer<S> for JsonRequestRewriteLayer<H> {
322    type Service = JsonRequestRewrite<S, H>;
323
324    fn layer(&self, inner: S) -> Self::Service {
325        JsonRequestRewrite {
326            inner,
327            selectors: self.selectors.clone(),
328            handler: self.handler.clone(),
329            policy: self.policy.clone(),
330            max_buffered_bytes: self.max_buffered_bytes,
331        }
332    }
333
334    fn into_layer(self, inner: S) -> Self::Service {
335        JsonRequestRewrite {
336            inner,
337            selectors: self.selectors,
338            handler: self.handler,
339            policy: self.policy,
340            max_buffered_bytes: self.max_buffered_bytes,
341        }
342    }
343}
344
345impl<H> JsonRewriteLayer<H> {
346    /// Creates a new [`JsonRewriteLayer`] that rewrites values matching
347    /// `selectors` with `handler` (the handler is cloned per response, so it
348    /// starts fresh for each one).
349    pub fn new(selectors: impl IntoIterator<Item = JsonPath>, handler: H) -> Self {
350        Self {
351            selectors: selectors.into_iter().collect(),
352            handler,
353            policy: BodyRewritePolicy::unencoded_content_type(is_json_content_type),
354            max_buffered_bytes: DEFAULT_MAX_BUFFERED_BYTES,
355        }
356    }
357
358    generate_set_and_with! {
359        /// Sets a custom response rewrite policy.
360        ///
361        /// The predicate receives the response headers and extensions and can
362        /// narrow rewriting beyond the built-in `Content-Encoding` guard.
363        pub fn rewrite_policy(
364            mut self,
365            policy: impl Fn(&HeaderMap, &Extensions) -> bool + Send + Sync + 'static,
366        ) -> Self {
367            self.policy = BodyRewritePolicy::custom(policy);
368            self
369        }
370    }
371
372    generate_set_and_with! {
373        /// Sets the tokenizer buffered-input limit for each rewritten body.
374        pub fn max_buffered_bytes(mut self, max_buffered_bytes: usize) -> Self {
375            self.max_buffered_bytes = max_buffered_bytes;
376            self
377        }
378    }
379
380    /// Wraps a body directly using this layer's selector set and handler.
381    ///
382    /// This is useful for services that need request-specific gating before
383    /// deciding whether a single response body should be rewritten, while
384    /// still sharing the same layer configuration.
385    pub fn rewrite_body<B>(&self, body: B) -> JsonRewriteBody<B, H>
386    where
387        H: JsonValueHandler + Clone,
388    {
389        JsonRewriteBody::with_max_buffered_bytes(
390            body,
391            self.selectors.iter().cloned(),
392            self.handler.clone(),
393            self.max_buffered_bytes,
394        )
395    }
396}
397
398impl<H: fmt::Debug> fmt::Debug for JsonRewriteLayer<H> {
399    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
400        f.debug_struct("JsonRewriteLayer")
401            .field("selectors", &self.selectors)
402            .field("handler", &self.handler)
403            .field("policy", &self.policy)
404            .field("max_buffered_bytes", &self.max_buffered_bytes)
405            .finish()
406    }
407}
408
409impl<S, H: Clone> Layer<S> for JsonRewriteLayer<H> {
410    type Service = JsonRewrite<S, H>;
411
412    fn layer(&self, inner: S) -> Self::Service {
413        JsonRewrite {
414            inner,
415            selectors: self.selectors.clone(),
416            handler: self.handler.clone(),
417            policy: self.policy.clone(),
418            max_buffered_bytes: self.max_buffered_bytes,
419        }
420    }
421
422    fn into_layer(self, inner: S) -> Self::Service {
423        JsonRewrite {
424            inner,
425            selectors: self.selectors,
426            handler: self.handler,
427            policy: self.policy,
428            max_buffered_bytes: self.max_buffered_bytes,
429        }
430    }
431}
432
433#[cfg(test)]
434mod tests {
435    use super::*;
436    use crate::headers::HeaderMapExt;
437    use crate::{HeaderMap, header};
438
439    #[test]
440    fn rewrite_content_type_policy() {
441        let cases = [
442            ("application/json", true),
443            ("application/json; charset=utf-8", true),
444            ("application/problem+json", true),
445            ("text/json", false),
446            ("text/plain", false),
447        ];
448
449        for (content_type, expected) in cases {
450            let mut headers = HeaderMap::new();
451            headers.insert(
452                header::CONTENT_TYPE,
453                content_type.parse().expect("valid header"),
454            );
455            let content_type = headers.typed_get::<ContentType>().expect("content type");
456            assert_eq!(
457                is_json_content_type(&content_type),
458                expected,
459                "{content_type}"
460            );
461        }
462    }
463
464    #[test]
465    fn rewrite_policy_skips_content_encoded_json() {
466        let mut headers = HeaderMap::new();
467        headers.insert(header::CONTENT_TYPE, "application/json".parse().unwrap());
468        headers.insert(header::CONTENT_ENCODING, "gzip".parse().unwrap());
469        let policy = BodyRewritePolicy::unencoded_content_type(is_json_content_type);
470        assert!(!policy.should_rewrite(&headers, &Extensions::new()));
471    }
472}