Skip to main content

macp_auth/auth/
chain.rs

1use super::resolver::AuthResolver;
2use crate::security::AuthIdentity;
3use macp_core::error::MacpError;
4use tonic::metadata::MetadataMap;
5
6pub struct AuthResolverChain {
7    resolvers: Vec<Box<dyn AuthResolver>>,
8}
9
10impl AuthResolverChain {
11    pub fn new(resolvers: Vec<Box<dyn AuthResolver>>) -> Self {
12        let names: Vec<&str> = resolvers.iter().map(|r| r.name()).collect();
13        tracing::info!(chain = ?names, "auth resolver chain initialized");
14        Self { resolvers }
15    }
16
17    pub async fn authenticate(&self, metadata: &MetadataMap) -> Result<AuthIdentity, MacpError> {
18        for resolver in &self.resolvers {
19            match resolver.resolve(metadata).await {
20                Ok(Some(identity)) => {
21                    tracing::debug!(
22                        resolver = resolver.name(),
23                        sender = %identity.sender,
24                        "authenticated"
25                    );
26                    return Ok(identity.into());
27                }
28                Ok(None) => continue,
29                Err(e) => {
30                    tracing::warn!(
31                        resolver = resolver.name(),
32                        error = %e,
33                        "auth resolver rejected credential"
34                    );
35                    return Err(MacpError::Unauthenticated);
36                }
37            }
38        }
39        Err(MacpError::Unauthenticated)
40    }
41}
42
43#[cfg(test)]
44mod tests {
45    use super::*;
46    use crate::auth::resolver::{AuthError, ResolvedIdentity};
47    use std::sync::atomic::{AtomicUsize, Ordering};
48    use std::sync::Arc;
49
50    enum Outcome {
51        Claim(&'static str),
52        Decline,
53        Fail,
54    }
55
56    struct StubResolver {
57        name: &'static str,
58        outcome: Outcome,
59        calls: Arc<AtomicUsize>,
60    }
61
62    impl StubResolver {
63        fn boxed(
64            name: &'static str,
65            outcome: Outcome,
66        ) -> (Box<dyn AuthResolver>, Arc<AtomicUsize>) {
67            let calls = Arc::new(AtomicUsize::new(0));
68            (
69                Box::new(Self {
70                    name,
71                    outcome,
72                    calls: calls.clone(),
73                }),
74                calls,
75            )
76        }
77
78        fn identity(&self, sender: &str) -> ResolvedIdentity {
79            ResolvedIdentity {
80                sender: sender.to_string(),
81                allowed_modes: None,
82                can_start_sessions: true,
83                max_open_sessions: None,
84                can_manage_mode_registry: false,
85                is_observer: false,
86                resolver: self.name.to_string(),
87            }
88        }
89    }
90
91    #[async_trait::async_trait]
92    impl AuthResolver for StubResolver {
93        fn name(&self) -> &str {
94            self.name
95        }
96
97        async fn resolve(
98            &self,
99            _metadata: &MetadataMap,
100        ) -> Result<Option<ResolvedIdentity>, AuthError> {
101            self.calls.fetch_add(1, Ordering::SeqCst);
102            match &self.outcome {
103                Outcome::Claim(sender) => Ok(Some(self.identity(sender))),
104                Outcome::Decline => Ok(None),
105                Outcome::Fail => Err(AuthError::InvalidCredential("stub rejection".to_string())),
106            }
107        }
108    }
109
110    #[tokio::test]
111    async fn first_resolver_that_claims_the_credential_wins() {
112        let (first, first_calls) = StubResolver::boxed("first", Outcome::Claim("agent://first"));
113        let (second, second_calls) =
114            StubResolver::boxed("second", Outcome::Claim("agent://second"));
115        let chain = AuthResolverChain::new(vec![first, second]);
116
117        let identity = chain.authenticate(&MetadataMap::new()).await.expect("ok");
118        assert_eq!(identity.sender, "agent://first");
119        assert_eq!(first_calls.load(Ordering::SeqCst), 1);
120        assert_eq!(
121            second_calls.load(Ordering::SeqCst),
122            0,
123            "chain must stop at the first positive verification"
124        );
125    }
126
127    #[tokio::test]
128    async fn declining_resolver_passes_to_the_next() {
129        let (first, first_calls) = StubResolver::boxed("first", Outcome::Decline);
130        let (second, second_calls) =
131            StubResolver::boxed("second", Outcome::Claim("agent://second"));
132        let chain = AuthResolverChain::new(vec![first, second]);
133
134        let identity = chain.authenticate(&MetadataMap::new()).await.expect("ok");
135        assert_eq!(identity.sender, "agent://second");
136        assert_eq!(first_calls.load(Ordering::SeqCst), 1);
137        assert_eq!(second_calls.load(Ordering::SeqCst), 1);
138    }
139
140    #[tokio::test]
141    async fn all_resolvers_declining_yields_unauthenticated() {
142        let (first, first_calls) = StubResolver::boxed("first", Outcome::Decline);
143        let (second, second_calls) = StubResolver::boxed("second", Outcome::Decline);
144        let chain = AuthResolverChain::new(vec![first, second]);
145
146        let err = chain.authenticate(&MetadataMap::new()).await.unwrap_err();
147        assert!(matches!(err, MacpError::Unauthenticated), "got {err:?}");
148        assert_eq!(first_calls.load(Ordering::SeqCst), 1);
149        assert_eq!(second_calls.load(Ordering::SeqCst), 1);
150    }
151
152    /// A resolver that claims the credential type but finds it invalid stops
153    /// the chain with Unauthenticated — later resolvers must not get a second
154    /// chance at a credential that positively failed verification.
155    #[tokio::test]
156    async fn resolver_error_stops_the_chain_as_unauthenticated() {
157        let (first, first_calls) = StubResolver::boxed("first", Outcome::Fail);
158        let (second, second_calls) =
159            StubResolver::boxed("second", Outcome::Claim("agent://second"));
160        let chain = AuthResolverChain::new(vec![first, second]);
161
162        let err = chain.authenticate(&MetadataMap::new()).await.unwrap_err();
163        assert!(matches!(err, MacpError::Unauthenticated), "got {err:?}");
164        assert_eq!(first_calls.load(Ordering::SeqCst), 1);
165        assert_eq!(
166            second_calls.load(Ordering::SeqCst),
167            0,
168            "a hard resolver error must not fall through to later resolvers"
169        );
170    }
171
172    #[tokio::test]
173    async fn empty_chain_is_unauthenticated() {
174        let chain = AuthResolverChain::new(vec![]);
175        let err = chain.authenticate(&MetadataMap::new()).await.unwrap_err();
176        assert!(matches!(err, MacpError::Unauthenticated), "got {err:?}");
177    }
178}