use crate::{
extensions::{Extensions, ExtensionsRef},
matcher::Matcher,
};
use crate::std::vec::Vec;
use super::{Policy, PolicyOutput, PolicyResult};
impl<M, P, Input> Policy<Input> for Vec<(M, P)>
where
M: Matcher<Input>,
P: Policy<Input>,
Input: Send + ExtensionsRef + 'static,
{
type Guard = Option<P::Guard>;
type Error = P::Error;
async fn check(&self, input: Input) -> PolicyResult<Input, Self::Guard, Self::Error> {
for (matcher, policy) in self.iter() {
let ext = Extensions::new();
if matcher.matches(Some(&ext), &input) {
input.extensions().extend(&ext);
let result = policy.check(input).await;
return match result.output {
PolicyOutput::Ready(guard) => {
let guard = Some(guard);
PolicyResult {
input: result.input,
output: PolicyOutput::Ready(guard),
}
}
PolicyOutput::Abort(err) => PolicyResult {
input: result.input,
output: PolicyOutput::Abort(err),
},
PolicyOutput::Retry => PolicyResult {
input: result.input,
output: PolicyOutput::Retry,
},
};
}
}
PolicyResult {
input,
output: PolicyOutput::Ready(None),
}
}
}
impl<M, P, Input> Policy<Input> for (Vec<(M, P)>, P)
where
M: Matcher<Input>,
P: Policy<Input>,
Input: Send + ExtensionsRef + 'static,
{
type Guard = P::Guard;
type Error = P::Error;
async fn check(&self, input: Input) -> PolicyResult<Input, Self::Guard, Self::Error> {
let (matchers, default_policy) = self;
for (matcher, policy) in matchers.iter() {
let ext = Extensions::new();
if matcher.matches(Some(&ext), &input) {
input.extensions().extend(&ext);
return policy.check(input).await;
}
}
default_policy.check(input).await
}
}
#[cfg(all(test, feature = "std"))]
mod tests {
use crate::std::sync::Arc;
use crate::{
ServiceInput,
extensions::Extensions,
layer::limit::policy::{ConcurrentCounter, ConcurrentPolicy},
};
use super::*;
fn assert_ready<R, G, E>(result: PolicyResult<R, G, E>) -> Option<G> {
let guard = match result.output {
PolicyOutput::Ready(guard) => Some(guard),
PolicyOutput::Abort(_) | PolicyOutput::Retry => None,
};
assert!(guard.is_some(), "unexpected output, expected ready");
guard
}
fn assert_abort<R, G, E>(result: &PolicyResult<R, G, E>) {
assert!(
matches!(result.output, PolicyOutput::Abort(_)),
"unexpected output, expected abort"
);
}
type NumberedInput = ServiceInput<usize>;
#[tokio::test]
async fn matcher_policy_empty() {
let policy = Vec::<(bool, ConcurrentPolicy<(), ConcurrentCounter>)>::new();
for i in 0..10 {
drop(assert_ready(policy.check(NumberedInput::new(i)).await));
}
}
#[tokio::test]
async fn matcher_policy_always() {
let concurrency_policy = ConcurrentPolicy::max(2);
let policy = Arc::new(vec![(true, concurrency_policy)]);
let guard_1 = assert_ready(policy.check(Extensions::new()).await);
let guard_2 = assert_ready(policy.check(Extensions::new()).await);
assert_abort(&policy.check(Extensions::new()).await);
drop(guard_1);
let _guard_3 = assert_ready(policy.check(Extensions::new()).await);
assert_abort(&policy.check(Extensions::new()).await);
drop(guard_2);
drop(assert_ready(policy.check(Extensions::new()).await));
}
#[derive(Debug, Clone)]
enum TestMatchers {
Const(usize),
Odd,
}
impl Matcher<NumberedInput> for TestMatchers {
fn matches(&self, _ext: Option<&Extensions>, req: &NumberedInput) -> bool {
match self {
Self::Const(n) => *n == req.input,
Self::Odd => req.input % 2 == 1,
}
}
}
#[tokio::test]
async fn matcher_policy_scoped_limits() {
let policy = vec![
(TestMatchers::Odd, ConcurrentPolicy::max(2)),
(TestMatchers::Const(42), ConcurrentPolicy::max(1)),
];
for i in 1..10 {
drop(assert_ready(policy.check(NumberedInput::new(i * 2)).await));
}
let odd_guard_1 = assert_ready(policy.check(NumberedInput::new(1)).await);
let const_guard_1 = assert_ready(policy.check(NumberedInput::new(42)).await);
let odd_guard_2 = assert_ready(policy.check(NumberedInput::new(3)).await);
assert_abort(&policy.check(NumberedInput::new(5)).await);
assert_abort(&policy.check(NumberedInput::new(42)).await);
for i in 1..10 {
drop(assert_ready(policy.check(NumberedInput::new(i * 2)).await));
}
drop(odd_guard_1);
let _odd_guard_3 = assert_ready(policy.check(NumberedInput::new(9)).await);
assert_abort(&policy.check(NumberedInput::new(42)).await);
drop(const_guard_1);
drop(assert_ready(policy.check(NumberedInput::new(42)).await));
assert_abort(&policy.check(NumberedInput::new(11)).await);
drop(odd_guard_2);
drop(assert_ready(policy.check(NumberedInput::new(13)).await));
for i in 1..10 {
drop(assert_ready(policy.check(NumberedInput::new(i * 2)).await));
}
}
}