1use std::future::Future;
2use std::pin::Pin;
3use std::task::{Context, Poll};
4
5use tower::Service;
6
7use camel_api::{CamelError, Exchange, ValueSource};
8
9#[derive(Clone)]
14pub struct DynamicSetHeader<P> {
15 inner: P,
16 key: String,
17 source: ValueSource,
18}
19
20impl<P> DynamicSetHeader<P> {
21 pub fn new(inner: P, key: impl Into<String>, source: impl Into<ValueSource>) -> Self {
22 Self {
23 inner,
24 key: key.into(),
25 source: source.into(),
26 }
27 }
28}
29
30#[derive(Clone)]
32pub struct DynamicSetHeaderLayer {
33 key: String,
34 source: ValueSource,
35}
36
37impl DynamicSetHeaderLayer {
38 pub fn new(key: impl Into<String>, source: impl Into<ValueSource>) -> Self {
39 Self {
40 key: key.into(),
41 source: source.into(),
42 }
43 }
44}
45
46impl<S> tower::Layer<S> for DynamicSetHeaderLayer {
47 type Service = DynamicSetHeader<S>;
48
49 fn layer(&self, inner: S) -> Self::Service {
50 DynamicSetHeader {
51 inner,
52 key: self.key.clone(),
53 source: self.source.clone(),
54 }
55 }
56}
57
58impl<P> Service<Exchange> for DynamicSetHeader<P>
59where
60 P: Service<Exchange, Response = Exchange, Error = CamelError> + Clone + Send + 'static,
61 P::Future: Send,
62{
63 type Response = Exchange;
64 type Error = CamelError;
65 type Future = Pin<Box<dyn Future<Output = Result<Exchange, CamelError>> + Send>>;
66
67 fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
68 self.inner.poll_ready(cx)
69 }
70
71 fn call(&mut self, mut exchange: Exchange) -> Self::Future {
72 let source = self.source.clone();
73 let key = self.key.clone();
74 let clone = self.inner.clone();
77 let mut inner = std::mem::replace(&mut self.inner, clone);
78 Box::pin(async move {
79 let value = source.evaluate(&exchange).await?;
80 exchange.input.headers.insert(key, value);
81 inner.call(exchange).await
82 })
83 }
84}
85
86#[derive(Clone)]
91pub struct DynamicSetHeaderIfAbsent<P> {
92 inner: P,
93 key: String,
94 source: ValueSource,
95}
96
97impl<P> DynamicSetHeaderIfAbsent<P> {
98 pub fn new(inner: P, key: impl Into<String>, source: impl Into<ValueSource>) -> Self {
99 Self {
100 inner,
101 key: key.into(),
102 source: source.into(),
103 }
104 }
105}
106
107impl<P> Service<Exchange> for DynamicSetHeaderIfAbsent<P>
108where
109 P: Service<Exchange, Response = Exchange, Error = CamelError> + Clone + Send + 'static,
110 P::Future: Send,
111{
112 type Response = Exchange;
113 type Error = CamelError;
114 type Future = Pin<Box<dyn Future<Output = Result<Exchange, CamelError>> + Send>>;
115
116 fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
117 self.inner.poll_ready(cx)
118 }
119
120 fn call(&mut self, mut exchange: Exchange) -> Self::Future {
121 if exchange.input.headers.contains_key(&self.key) {
124 let clone = self.inner.clone();
125 let mut inner = std::mem::replace(&mut self.inner, clone);
126 return Box::pin(inner.call(exchange));
127 }
128
129 let source = self.source.clone();
130 let key = self.key.clone();
131 let clone = self.inner.clone();
133 let mut inner = std::mem::replace(&mut self.inner, clone);
134 Box::pin(async move {
135 let value = source.evaluate(&exchange).await?;
136 exchange.input.headers.insert(key, value);
137 inner.call(exchange).await
138 })
139 }
140}
141
142#[cfg(test)]
143mod tests {
144 use std::sync::Arc;
145 use std::sync::atomic::{AtomicUsize, Ordering};
146
147 use camel_api::{BoxValueFuture, Exchange, IdentityProcessor, Message, Value, ValueSource};
148 use tower::ServiceExt;
149
150 use super::*;
151
152 fn sync_source<F>(f: F) -> ValueSource
153 where
154 F: Fn(&Exchange) -> Value + Send + Sync + 'static,
155 {
156 ValueSource::Sync(Arc::new(f))
157 }
158
159 fn failing_source() -> ValueSource {
160 ValueSource::Async(Arc::new(|_: &Exchange| {
161 Box::pin(async { Err(CamelError::ProcessorError("expr boom".into())) })
162 as BoxValueFuture
163 }))
164 }
165
166 #[derive(Clone)]
168 struct CountingInner {
169 called: Arc<AtomicUsize>,
170 }
171
172 impl Service<Exchange> for CountingInner {
173 type Response = Exchange;
174 type Error = CamelError;
175 type Future = Pin<Box<dyn Future<Output = Result<Exchange, CamelError>> + Send>>;
176
177 fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
178 Poll::Ready(Ok(()))
179 }
180
181 fn call(&mut self, exchange: Exchange) -> Self::Future {
182 self.called.fetch_add(1, Ordering::SeqCst);
183 Box::pin(async { Ok(exchange) })
184 }
185 }
186
187 #[tokio::test]
188 async fn dynamic_set_header_error_propagates() {
189 let called = Arc::new(AtomicUsize::new(0));
190 let svc = DynamicSetHeader::new(
191 CountingInner {
192 called: Arc::clone(&called),
193 },
194 "greeting",
195 failing_source(),
196 );
197
198 let result = svc.oneshot(Exchange::new(Message::new("world"))).await;
199
200 assert!(result.is_err(), "failed evaluation must fail the step");
201 assert_eq!(
202 called.load(Ordering::SeqCst),
203 0,
204 "inner must NOT run when evaluation fails"
205 );
206 }
207
208 #[tokio::test]
209 async fn dynamic_set_header_if_absent_error_propagates() {
210 let called = Arc::new(AtomicUsize::new(0));
211 let svc = DynamicSetHeaderIfAbsent::new(
212 CountingInner {
213 called: Arc::clone(&called),
214 },
215 "greeting",
216 failing_source(),
217 );
218
219 let result = svc.oneshot(Exchange::new(Message::new("world"))).await;
221
222 assert!(result.is_err(), "failed evaluation must fail the step");
223 assert_eq!(
224 called.load(Ordering::SeqCst),
225 0,
226 "inner must NOT run when evaluation fails"
227 );
228 }
229
230 #[tokio::test]
231 async fn dynamic_set_header_if_absent_present_header_skips_evaluation() {
232 let evals = Arc::new(AtomicUsize::new(0));
233 let evals_clone = evals.clone();
234 let source = ValueSource::Async(Arc::new(move |_: &Exchange| {
235 evals_clone.fetch_add(1, Ordering::SeqCst);
236 Box::pin(async { Ok(Value::Bool(true)) }) as BoxValueFuture
237 }));
238
239 let mut msg = Message::new("computed");
240 msg.set_header("key", Value::String("original".into()));
241
242 let svc = DynamicSetHeaderIfAbsent::new(IdentityProcessor, "key", source);
243 let result = svc.oneshot(Exchange::new(msg)).await.unwrap();
244
245 assert_eq!(
246 result.input.header("key"),
247 Some(&Value::String("original".into()))
248 );
249 assert_eq!(
250 evals.load(Ordering::SeqCst),
251 0,
252 "expression must NOT be evaluated when the header is already present"
253 );
254 }
255
256 #[tokio::test]
257 async fn test_dynamic_set_header_from_body() {
258 let exchange = Exchange::new(Message::new("world"));
259
260 let svc = DynamicSetHeader::new(
261 IdentityProcessor,
262 "greeting",
263 sync_source(|ex: &Exchange| {
264 Value::String(format!("hello {}", ex.input.body.as_text().unwrap_or("")))
265 }),
266 );
267
268 let result = svc.oneshot(exchange).await.unwrap();
269 assert_eq!(
270 result.input.header("greeting"),
271 Some(&Value::String("hello world".into()))
272 );
273 }
274
275 #[tokio::test]
276 async fn test_dynamic_set_header_overwrites_existing() {
277 let mut msg = Message::new("new");
278 msg.set_header("key", Value::String("old".into()));
279 let exchange = Exchange::new(msg);
280
281 let svc = DynamicSetHeader::new(
282 IdentityProcessor,
283 "key",
284 sync_source(|ex: &Exchange| {
285 Value::String(ex.input.body.as_text().unwrap_or("").into())
286 }),
287 );
288
289 let result = svc.oneshot(exchange).await.unwrap();
290 assert_eq!(
291 result.input.header("key"),
292 Some(&Value::String("new".into()))
293 );
294 }
295
296 #[tokio::test]
297 async fn test_dynamic_set_header_preserves_body() {
298 let exchange = Exchange::new(Message::new("body content"));
299
300 let svc = DynamicSetHeader::new(
301 IdentityProcessor,
302 "len",
303 sync_source(|ex: &Exchange| {
304 let len = ex.input.body.as_text().map(|t| t.len() as i64).unwrap_or(0);
305 Value::Number(len.into())
306 }),
307 );
308
309 let result = svc.oneshot(exchange).await.unwrap();
310 assert_eq!(result.input.body.as_text(), Some("body content"));
311 assert_eq!(result.input.header("len"), Some(&Value::Number(12.into())));
312 }
313
314 #[tokio::test]
315 async fn test_dynamic_set_header_layer_composes() {
316 use tower::ServiceBuilder;
317
318 let svc = ServiceBuilder::new()
319 .layer(DynamicSetHeaderLayer::new(
320 "computed",
321 sync_source(|_ex: &Exchange| Value::Bool(true)),
322 ))
323 .service(IdentityProcessor);
324
325 let exchange = Exchange::new(Message::default());
326 let result = svc.oneshot(exchange).await.unwrap();
327 assert_eq!(result.input.header("computed"), Some(&Value::Bool(true)));
328 }
329
330 #[tokio::test]
333 async fn test_dynamic_set_header_if_absent_adds_when_missing() {
334 let exchange = Exchange::new(Message::new("world"));
335
336 let svc = DynamicSetHeaderIfAbsent::new(
337 IdentityProcessor,
338 "greeting",
339 sync_source(|ex: &Exchange| {
340 Value::String(format!("hello {}", ex.input.body.as_text().unwrap_or("")))
341 }),
342 );
343
344 let result = svc.oneshot(exchange).await.unwrap();
345 assert_eq!(
346 result.input.header("greeting"),
347 Some(&Value::String("hello world".into()))
348 );
349 }
350
351 #[tokio::test]
352 async fn test_dynamic_set_header_if_absent_preserves_existing() {
353 use std::sync::Arc;
354 use std::sync::atomic::{AtomicUsize, Ordering};
355
356 let mut msg = Message::new("computed");
357 msg.set_header("key", Value::String("original".into()));
358 let exchange = Exchange::new(msg);
359
360 let call_count = Arc::new(AtomicUsize::new(0));
363 let cc = call_count.clone();
364
365 let svc = DynamicSetHeaderIfAbsent::new(
366 IdentityProcessor,
367 "key",
368 sync_source(move |ex: &Exchange| {
369 cc.fetch_add(1, Ordering::SeqCst);
370 Value::String(ex.input.body.as_text().unwrap_or("").into())
371 }),
372 );
373
374 let result = svc.oneshot(exchange).await.unwrap();
375 assert_eq!(
376 result.input.header("key"),
377 Some(&Value::String("original".into()))
378 );
379 assert_eq!(
380 call_count.load(Ordering::SeqCst),
381 0,
382 "expression must NOT be evaluated when header is present"
383 );
384 }
385
386 #[tokio::test]
387 async fn test_dynamic_set_header_if_absent_preserves_body() {
388 let exchange = Exchange::new(Message::new("body content"));
389
390 let svc = DynamicSetHeaderIfAbsent::new(
391 IdentityProcessor,
392 "len",
393 sync_source(|ex: &Exchange| {
394 let len = ex.input.body.as_text().map(|t| t.len() as i64).unwrap_or(0);
395 Value::Number(len.into())
396 }),
397 );
398
399 let result = svc.oneshot(exchange).await.unwrap();
400 assert_eq!(result.input.body.as_text(), Some("body content"));
401 assert_eq!(result.input.header("len"), Some(&Value::Number(12.into())));
402 }
403}