1use 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#[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 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 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 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 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#[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 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 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 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 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
227fn 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#[derive(Clone)]
238pub struct JsonRewriteLayer<H> {
239 selectors: Box<[JsonPath]>,
240 handler: H,
241 policy: BodyRewritePolicy,
242 max_buffered_bytes: usize,
243}
244
245#[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 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 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 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 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 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 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 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 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}