1use std::future::Future;
15use std::pin::Pin;
16use std::sync::Arc;
17use std::task::{Context, Poll};
18
19use tower::Service;
20
21use camel_api::body::Body;
22use camel_api::{BoxValueFuture, CamelError, ClaimCheckRepository, Exchange, Message, Value};
23
24#[allow(clippy::type_complexity)]
34#[derive(Clone)]
35pub enum ClaimKeySource {
36 Sync(Arc<dyn Fn(&Exchange) -> Result<String, CamelError> + Send + Sync>),
38 Async(Arc<dyn Fn(&Exchange) -> BoxValueFuture + Send + Sync>),
40}
41
42impl ClaimKeySource {
43 pub async fn key(&self, exchange: &Exchange) -> Result<String, CamelError> {
45 match self {
46 Self::Sync(f) => f(exchange),
47 Self::Async(f) => {
48 let value = f(exchange).await?;
49 match value {
50 Value::Null => Err(claim_key_null_or_empty()),
51 Value::String(s) if s.is_empty() => Err(claim_key_null_or_empty()),
52 Value::String(s) => Ok(s),
53 Value::Array(_) | Value::Object(_) => {
54 Err(CamelError::ProcessorError(
55 "claim_check key expression returned a non-scalar value (array/object); expected a string key".into(),
56 ))
57 }
58 other => Ok(other.to_string()),
59 }
60 }
61 }
62 }
63}
64
65fn claim_key_null_or_empty() -> CamelError {
66 CamelError::ValidationError("claim_check key expression evaluated to null or empty".into())
67}
68
69#[derive(Clone, Debug, PartialEq, Eq)]
71pub enum ClaimCheckOp {
72 Set,
74 Get,
76 GetAndRemove,
78 Push,
80 Pop,
82}
83
84#[derive(Clone)]
91pub struct ClaimCheckService {
92 repository: Arc<dyn ClaimCheckRepository>,
93 operation: ClaimCheckOp,
94 key_expression: ClaimKeySource,
95 filter: Option<ClaimCheckFilter>,
96}
97
98impl std::fmt::Debug for ClaimCheckService {
99 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
100 f.debug_struct("ClaimCheckService")
101 .field("repository", &self.repository.name())
102 .field("operation", &self.operation)
103 .field("filter", &self.filter)
104 .finish()
105 }
106}
107
108impl ClaimCheckService {
109 pub fn new(
111 repository: Arc<dyn ClaimCheckRepository>,
112 operation: ClaimCheckOp,
113 key_expression: ClaimKeySource,
114 ) -> Self {
115 Self {
116 repository,
117 operation,
118 key_expression,
119 filter: None,
120 }
121 }
122
123 pub fn with_filter(mut self, filter: ClaimCheckFilter) -> Self {
125 self.filter = Some(filter);
126 self
127 }
128}
129
130fn merge_stashed(current: &mut Message, stashed: &Message, filter: &ClaimCheckFilter) {
134 match filter.body {
135 FilterAction::Include => current.body = stashed.body.clone(),
136 FilterAction::Exclude => {}
137 FilterAction::Remove => current.body = Body::Empty,
138 }
139
140 match &filter.headers_action {
141 HeadersAction::All(action) => match action {
142 FilterAction::Include => {
143 for (k, v) in &stashed.headers {
144 current.headers.insert(k.clone(), v.clone());
145 }
146 }
147 FilterAction::Exclude => {}
148 FilterAction::Remove => current.headers.clear(),
149 },
150 HeadersAction::ByPattern {
151 include,
152 exclude,
153 remove,
154 } => {
155 if !include.is_empty() {
156 for (k, v) in &stashed.headers {
157 if include.iter().any(|p| p.matches(k)) {
158 current.headers.insert(k.clone(), v.clone());
159 }
160 }
161 }
162 if !exclude.is_empty() {
163 for (k, v) in &stashed.headers {
164 if !exclude.iter().any(|p| p.matches(k)) {
165 current.headers.insert(k.clone(), v.clone());
166 }
167 }
168 }
169 if !remove.is_empty() {
170 current
171 .headers
172 .retain(|k, _| !remove.iter().any(|p| p.matches(k)));
173 }
174 }
175 }
176}
177
178impl Service<Exchange> for ClaimCheckService {
179 type Response = Exchange;
180 type Error = CamelError;
181 type Future = Pin<Box<dyn Future<Output = Result<Exchange, CamelError>> + Send>>;
182
183 fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
184 Poll::Ready(Ok(()))
185 }
186
187 fn call(&mut self, mut exchange: Exchange) -> Self::Future {
188 let repository = self.repository.clone();
189 let operation = self.operation.clone();
190 let key_source = self.key_expression.clone();
191 let filter = self.filter.clone();
192
193 Box::pin(async move {
194 let key = key_source.key(&exchange).await?;
197 match operation {
198 ClaimCheckOp::Set => {
199 let stashed = exchange.input.clone();
200 repository.set(&key, stashed).await?;
201 exchange.input.body = Body::Text(key);
202 Ok(exchange)
203 }
204 ClaimCheckOp::Get => {
205 let stashed = repository.get(&key).await?;
206 if let Some(ref f) = filter {
207 merge_stashed(&mut exchange.input, &stashed, f);
208 } else {
209 exchange.input.body = stashed.body;
210 }
211 Ok(exchange)
212 }
213 ClaimCheckOp::GetAndRemove => {
214 let stashed = repository.get_and_remove(&key).await?;
215 if let Some(ref f) = filter {
216 merge_stashed(&mut exchange.input, &stashed, f);
217 } else {
218 exchange.input.body = stashed.body;
219 }
220 Ok(exchange)
221 }
222 ClaimCheckOp::Push => {
223 let stashed = exchange.input.clone();
224 repository.push(&key, stashed).await?;
225 exchange.input.body = Body::Text(key);
226 Ok(exchange)
227 }
228 ClaimCheckOp::Pop => {
229 let stashed = repository.pop(&key).await?;
230 if let Some(ref f) = filter {
231 merge_stashed(&mut exchange.input, &stashed, f);
232 } else {
233 exchange.input.body = stashed.body;
234 }
235 Ok(exchange)
236 }
237 }
238 })
239 }
240}
241
242fn has_regex_metachars(s: &str) -> bool {
243 s.contains(['^', '$', '(', ')', '[', ']', '{', '}', '|', '+', '.', '\\'])
244}
245
246#[derive(Debug, Clone)]
247pub enum HeaderPattern {
248 All,
249 Prefix(String),
250 Exact(String),
251 Regex(regex::Regex),
252}
253
254impl PartialEq for HeaderPattern {
255 fn eq(&self, other: &Self) -> bool {
256 match (self, other) {
257 (Self::All, Self::All) => true,
258 (Self::Prefix(a), Self::Prefix(b)) => a == b,
259 (Self::Exact(a), Self::Exact(b)) => a == b,
260 (Self::Regex(a), Self::Regex(b)) => a.as_str() == b.as_str(),
261 _ => false,
262 }
263 }
264}
265
266impl HeaderPattern {
267 fn compile(pattern: &str) -> Result<Self, FilterParseError> {
268 if pattern == "*" {
269 return Ok(Self::All);
270 }
271 if let Some(prefix) = pattern.strip_suffix('*') {
272 return Ok(Self::Prefix(prefix.to_string()));
273 }
274 if has_regex_metachars(pattern) {
275 match regex::Regex::new(pattern) {
276 Ok(re) => Ok(Self::Regex(re)),
277 Err(_) => Err(FilterParseError::InvalidPattern(pattern.to_string())),
278 }
279 } else {
280 Ok(Self::Exact(pattern.to_string()))
281 }
282 }
283
284 fn matches(&self, header_key: &str) -> bool {
285 match self {
286 Self::All => true,
287 Self::Prefix(prefix) => header_key.starts_with(prefix),
288 Self::Exact(exact) => header_key == exact,
289 Self::Regex(re) => re.is_match(header_key),
290 }
291 }
292}
293
294#[derive(Debug, Clone, PartialEq)]
295pub struct ClaimCheckFilter {
296 pub body: FilterAction,
297 pub headers_action: HeadersAction,
298}
299
300#[derive(Debug, Clone, Copy, PartialEq, Eq)]
301pub enum FilterAction {
302 Include,
303 Exclude,
304 Remove,
305}
306
307#[derive(Debug, Clone, PartialEq)]
308pub enum HeadersAction {
309 All(FilterAction),
310 ByPattern {
311 include: Vec<HeaderPattern>,
312 exclude: Vec<HeaderPattern>,
313 remove: Vec<HeaderPattern>,
314 },
315}
316
317#[derive(Debug)]
318pub enum FilterParseError {
319 InvalidToken(String),
320 InvalidPattern(String),
321 MixedIncludeExclude,
322}
323
324impl std::fmt::Display for FilterParseError {
325 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
326 match self {
327 Self::InvalidToken(tok) => write!(f, "invalid filter segment '{tok}'"),
328 Self::InvalidPattern(pat) => write!(f, "invalid header pattern '{pat}'"),
329 Self::MixedIncludeExclude => {
330 write!(f, "cannot mix include (+) and exclude (-) header patterns")
331 }
332 }
333 }
334}
335
336impl ClaimCheckFilter {
337 pub fn parse(input: &str) -> Result<Self, FilterParseError> {
340 let mut body_action: Option<FilterAction> = None;
341 let mut headers_action: Option<HeadersAction> = None;
342 let mut has_positive_include = false;
343
344 for token in input.split(',') {
345 let token = token.trim();
346 if token.is_empty() {
347 continue;
348 }
349
350 let (prefix, rule) = if let Some(rest) = token.strip_prefix("--") {
351 (FilterAction::Remove, rest)
352 } else if let Some(rest) = token.strip_prefix('-') {
353 (FilterAction::Exclude, rest)
354 } else {
355 let rest = token.strip_prefix('+').unwrap_or(token);
356 (FilterAction::Include, rest)
357 };
358
359 match rule {
360 "body" => {
361 if prefix == FilterAction::Include {
362 has_positive_include = true;
363 }
364 body_action = Some(prefix);
365 }
366 "headers" | "header" => {
367 if prefix == FilterAction::Include {
368 has_positive_include = true;
369 }
370 headers_action = Some(HeadersAction::All(prefix));
371 }
372 "attachments" | "attachment" => {
373 }
375 r if r.starts_with("header:") || r.starts_with("headers:") => {
376 let (_, pattern_str) = r
377 .split_once(':')
378 .ok_or_else(|| FilterParseError::InvalidToken(r.to_string()))?;
379 let pattern = HeaderPattern::compile(pattern_str)?;
380
381 if prefix == FilterAction::Include {
382 has_positive_include = true;
383 }
384
385 let (mut include, mut exclude, mut remove) = match headers_action.take() {
386 Some(HeadersAction::ByPattern {
387 include,
388 exclude,
389 remove,
390 }) => (include, exclude, remove),
391 _ => (vec![], vec![], vec![]),
392 };
393 match prefix {
394 FilterAction::Include => {
395 if !exclude.is_empty() {
396 return Err(FilterParseError::MixedIncludeExclude);
397 }
398 include.push(pattern);
399 }
400 FilterAction::Exclude => {
401 if !include.is_empty() {
402 return Err(FilterParseError::MixedIncludeExclude);
403 }
404 exclude.push(pattern);
405 }
406 FilterAction::Remove => remove.push(pattern),
407 }
408 headers_action = Some(HeadersAction::ByPattern {
409 include,
410 exclude,
411 remove,
412 });
413 }
414 other => return Err(FilterParseError::InvalidToken(other.to_string())),
415 }
416 }
417
418 let default_if_omitted = if has_positive_include {
419 FilterAction::Exclude
420 } else {
421 FilterAction::Include
422 };
423
424 Ok(ClaimCheckFilter {
425 body: body_action.unwrap_or(default_if_omitted),
426 headers_action: headers_action.unwrap_or(HeadersAction::All(default_if_omitted)),
427 })
428 }
429}
430
431#[cfg(test)]
432#[path = "claim_check_tests.rs"]
433mod tests;