1use std::collections::HashMap;
34use std::sync::Arc;
35
36use async_trait::async_trait;
37
38use apl_core::attributes::AttributeBag;
39use apl_core::step::{PdpCall, PdpDecision, PdpDialect, PdpError, PdpResolver};
40
41#[derive(Clone)]
50pub struct PdpRouter {
51 resolvers: HashMap<PdpDialect, Arc<dyn PdpResolver>>,
52}
53
54impl PdpRouter {
55 pub fn new() -> Self {
56 Self {
57 resolvers: HashMap::new(),
58 }
59 }
60
61 pub fn register(&mut self, resolver: Arc<dyn PdpResolver>) -> &mut Self {
66 let dialect = resolver.dialect();
67 if self.resolvers.contains_key(&dialect) {
68 tracing::warn!(
69 dialect = ?dialect,
70 "PdpRouter: resolver for dialect already registered — keeping existing",
71 );
72 return self;
73 }
74 self.resolvers.insert(dialect, resolver);
75 self
76 }
77
78 pub fn replace(&mut self, resolver: Arc<dyn PdpResolver>) -> &mut Self {
82 let dialect = resolver.dialect();
83 self.resolvers.insert(dialect, resolver);
84 self
85 }
86
87 pub fn len(&self) -> usize {
89 self.resolvers.len()
90 }
91
92 pub fn is_empty(&self) -> bool {
93 self.resolvers.is_empty()
94 }
95}
96
97impl Default for PdpRouter {
98 fn default() -> Self {
99 Self::new()
100 }
101}
102
103#[async_trait]
104impl PdpResolver for PdpRouter {
105 fn dialect(&self) -> PdpDialect {
106 PdpDialect::Custom("router".to_string())
111 }
112
113 async fn evaluate(&self, call: &PdpCall, bag: &AttributeBag) -> Result<PdpDecision, PdpError> {
114 let resolver = self
115 .resolvers
116 .get(&call.dialect)
117 .ok_or_else(|| PdpError::NoResolver(call.dialect.clone()))?;
118 resolver.evaluate(call, bag).await
119 }
120}
121
122#[cfg(test)]
123mod tests {
124 use super::*;
125 use apl_core::evaluator::Decision;
126
127 struct FakePdp {
128 dialect: PdpDialect,
129 decision: Decision,
130 }
131
132 #[async_trait]
133 impl PdpResolver for FakePdp {
134 fn dialect(&self) -> PdpDialect {
135 self.dialect.clone()
136 }
137
138 async fn evaluate(
139 &self,
140 _call: &PdpCall,
141 _bag: &AttributeBag,
142 ) -> Result<PdpDecision, PdpError> {
143 Ok(PdpDecision {
144 decision: self.decision.clone(),
145 diagnostics: Vec::new(),
146 })
147 }
148 }
149
150 #[tokio::test]
151 async fn routes_by_dialect() {
152 let mut router = PdpRouter::new();
153 router.register(Arc::new(FakePdp {
154 dialect: PdpDialect::Cedar,
155 decision: Decision::Allow,
156 }));
157 router.register(Arc::new(FakePdp {
158 dialect: PdpDialect::Opa,
159 decision: Decision::Deny {
160 reason: Some("opa says no".into()),
161 rule_source: "opa".into(),
162 },
163 }));
164
165 let bag = AttributeBag::default();
166 let cedar_call = PdpCall {
167 dialect: PdpDialect::Cedar,
168 args: serde_yaml::Value::Null,
169 };
170 let opa_call = PdpCall {
171 dialect: PdpDialect::Opa,
172 args: serde_yaml::Value::Null,
173 };
174
175 let cedar_res = router.evaluate(&cedar_call, &bag).await.unwrap();
176 assert!(matches!(cedar_res.decision, Decision::Allow));
177
178 let opa_res = router.evaluate(&opa_call, &bag).await.unwrap();
179 assert!(matches!(opa_res.decision, Decision::Deny { .. }));
180 }
181
182 #[tokio::test]
183 async fn missing_dialect_returns_no_resolver() {
184 let router = PdpRouter::new();
185 let bag = AttributeBag::default();
186 let call = PdpCall {
187 dialect: PdpDialect::Cedar,
188 args: serde_yaml::Value::Null,
189 };
190 let err = router.evaluate(&call, &bag).await.unwrap_err();
191 assert!(matches!(err, PdpError::NoResolver(_)));
192 }
193
194 #[tokio::test]
195 async fn duplicate_register_keeps_first() {
196 let mut router = PdpRouter::new();
197 router.register(Arc::new(FakePdp {
198 dialect: PdpDialect::Cedar,
199 decision: Decision::Allow,
200 }));
201 router.register(Arc::new(FakePdp {
202 dialect: PdpDialect::Cedar,
203 decision: Decision::Deny {
204 reason: Some("shouldn't fire".into()),
205 rule_source: "test".into(),
206 },
207 }));
208 let call = PdpCall {
209 dialect: PdpDialect::Cedar,
210 args: serde_yaml::Value::Null,
211 };
212 let res = router
213 .evaluate(&call, &AttributeBag::default())
214 .await
215 .unwrap();
216 assert!(matches!(res.decision, Decision::Allow));
217 }
218}