Skip to main content

mithril_aggregator_discovery/
capabilities_discoverer.rs

1use std::sync::Arc;
2
3use mithril_common::{
4    AggregateSignatureType, StdResult,
5    entities::{MithrilNetwork, SignedEntityTypeDiscriminants},
6    messages::AggregatorCapabilities,
7};
8
9use crate::{AggregatorDiscoverer, AggregatorEndpoint, model::AggregatorEndpointWithCapabilities};
10
11/// Required capabilities for an aggregator.
12#[derive(Clone, PartialEq, Eq, Debug)]
13pub enum RequiredAggregatorCapabilities {
14    /// All
15    All,
16    /// Signed entity type.
17    SignedEntityType(SignedEntityTypeDiscriminants),
18    /// Aggregate signature type.
19    AggregateSignatureType(AggregateSignatureType),
20    /// Logical OR of required capabilities.
21    Or(Vec<RequiredAggregatorCapabilities>),
22    /// Logical AND of required capabilities.
23    And(Vec<RequiredAggregatorCapabilities>),
24}
25
26impl From<SignedEntityTypeDiscriminants> for RequiredAggregatorCapabilities {
27    fn from(signed_entity_type: SignedEntityTypeDiscriminants) -> Self {
28        RequiredAggregatorCapabilities::SignedEntityType(signed_entity_type)
29    }
30}
31
32impl RequiredAggregatorCapabilities {
33    /// Check if the available capabilities match the required capabilities.
34    fn matches(&self, available: &AggregatorCapabilities) -> bool {
35        match self {
36            RequiredAggregatorCapabilities::All => true,
37            RequiredAggregatorCapabilities::SignedEntityType(required_signed_entity_type) => {
38                available
39                    .signed_entity_types
40                    .iter()
41                    .any(|req| req == required_signed_entity_type)
42            }
43            RequiredAggregatorCapabilities::AggregateSignatureType(
44                required_aggregate_signature_types,
45            ) => *required_aggregate_signature_types == available.aggregate_signature_type,
46            RequiredAggregatorCapabilities::Or(requirements) => {
47                requirements.iter().any(|req| req.matches(available))
48            }
49            RequiredAggregatorCapabilities::And(requirements) => {
50                requirements.iter().all(|req| req.matches(available))
51            }
52        }
53    }
54}
55
56/// An aggregator discoverer for specific capabilities.
57pub struct CapableAggregatorDiscoverer {
58    required_capabilities: RequiredAggregatorCapabilities,
59    inner_discoverer: Arc<dyn AggregatorDiscoverer<AggregatorEndpoint>>,
60}
61
62impl CapableAggregatorDiscoverer {
63    /// Creates a new `CapableAggregatorDiscoverer` instance with the provided capabilities.
64    pub fn new(
65        capabilities: RequiredAggregatorCapabilities,
66        inner_discoverer: Arc<dyn AggregatorDiscoverer<AggregatorEndpoint>>,
67    ) -> Self {
68        Self {
69            required_capabilities: capabilities,
70            inner_discoverer,
71        }
72    }
73}
74
75#[async_trait::async_trait]
76impl AggregatorDiscoverer<AggregatorEndpointWithCapabilities> for CapableAggregatorDiscoverer {
77    async fn get_available_aggregators(
78        &self,
79        network: MithrilNetwork,
80    ) -> StdResult<Box<dyn Iterator<Item = AggregatorEndpointWithCapabilities>>> {
81        let aggregator_endpoints = self.inner_discoverer.get_available_aggregators(network).await?;
82
83        Ok(Box::new(CapableAggregatorDiscovererIterator {
84            required_capabilities: self.required_capabilities.clone(),
85            inner_iterator: aggregator_endpoints,
86        }))
87    }
88}
89
90/// An iterator over aggregator endpoints filtered by capabilities.
91struct CapableAggregatorDiscovererIterator {
92    required_capabilities: RequiredAggregatorCapabilities,
93    inner_iterator: Box<dyn Iterator<Item = AggregatorEndpoint>>,
94}
95
96impl Iterator for CapableAggregatorDiscovererIterator {
97    type Item = AggregatorEndpointWithCapabilities;
98
99    fn next(&mut self) -> Option<Self::Item> {
100        for aggregator_endpoint in self.inner_iterator.by_ref() {
101            if let Ok(aggregator_with_capabilities) =
102                AggregatorEndpointWithCapabilities::try_from(aggregator_endpoint)
103                && self
104                    .required_capabilities
105                    .matches(aggregator_with_capabilities.capabilities())
106            {
107                return Some(aggregator_with_capabilities);
108            }
109        }
110
111        None
112    }
113}
114
115#[cfg(test)]
116mod tests {
117    use std::collections::BTreeSet;
118
119    use httpmock::MockServer;
120    use serde_json::json;
121
122    use mithril_common::{
123        AggregateSignatureType::Concatenation,
124        entities::SignedEntityTypeDiscriminants::{
125            CardanoDatabase, CardanoStakeDistribution, CardanoTransactions,
126            MithrilStakeDistribution,
127        },
128        messages::AggregatorFeaturesMessage,
129    };
130
131    use super::*;
132
133    mod required_capabilities {
134        use super::*;
135
136        #[test]
137        fn required_capabilities_match_all_success() {
138            let required = RequiredAggregatorCapabilities::All;
139            let available = AggregatorCapabilities {
140                aggregate_signature_type: Concatenation,
141                signed_entity_types: BTreeSet::from([]),
142                cardano_transactions_prover: None,
143            };
144
145            assert!(required.matches(&available));
146        }
147
148        #[test]
149        fn required_capabilities_match_signed_entity_types_success() {
150            let required =
151                RequiredAggregatorCapabilities::SignedEntityType(CardanoStakeDistribution);
152            let available = AggregatorCapabilities {
153                aggregate_signature_type: Concatenation,
154                signed_entity_types: BTreeSet::from([
155                    CardanoTransactions,
156                    CardanoStakeDistribution,
157                    CardanoDatabase,
158                ]),
159                cardano_transactions_prover: None,
160            };
161
162            assert!(required.matches(&available));
163        }
164
165        #[test]
166        fn required_capabilities_match_signed_entity_types_failure() {
167            let required =
168                RequiredAggregatorCapabilities::SignedEntityType(MithrilStakeDistribution);
169            let available = AggregatorCapabilities {
170                aggregate_signature_type: Concatenation,
171                signed_entity_types: BTreeSet::from([
172                    CardanoTransactions,
173                    CardanoStakeDistribution,
174                    CardanoDatabase,
175                ]),
176                cardano_transactions_prover: None,
177            };
178
179            assert!(!required.matches(&available));
180        }
181
182        #[test]
183        fn required_capabilities_match_signed_aggregate_signature_type_success() {
184            let required = RequiredAggregatorCapabilities::AggregateSignatureType(Concatenation);
185            let available = AggregatorCapabilities {
186                aggregate_signature_type: Concatenation,
187                signed_entity_types: BTreeSet::from([
188                    CardanoTransactions,
189                    CardanoStakeDistribution,
190                    CardanoDatabase,
191                ]),
192                cardano_transactions_prover: None,
193            };
194
195            assert!(required.matches(&available));
196        }
197
198        #[test]
199        fn required_capabilities_match_or_success() {
200            let required = RequiredAggregatorCapabilities::Or(vec![
201                RequiredAggregatorCapabilities::SignedEntityType(MithrilStakeDistribution),
202                RequiredAggregatorCapabilities::AggregateSignatureType(Concatenation),
203            ]);
204            let available = AggregatorCapabilities {
205                aggregate_signature_type: Concatenation,
206                signed_entity_types: BTreeSet::from([
207                    CardanoTransactions,
208                    CardanoStakeDistribution,
209                    CardanoDatabase,
210                ]),
211                cardano_transactions_prover: None,
212            };
213
214            assert!(required.matches(&available));
215        }
216
217        #[test]
218        fn required_capabilities_match_and_success() {
219            let required = RequiredAggregatorCapabilities::And(vec![
220                RequiredAggregatorCapabilities::SignedEntityType(CardanoTransactions),
221                RequiredAggregatorCapabilities::SignedEntityType(CardanoStakeDistribution),
222                RequiredAggregatorCapabilities::AggregateSignatureType(Concatenation),
223            ]);
224            let available = AggregatorCapabilities {
225                aggregate_signature_type: Concatenation,
226                signed_entity_types: BTreeSet::from([
227                    CardanoTransactions,
228                    CardanoStakeDistribution,
229                    CardanoDatabase,
230                ]),
231                cardano_transactions_prover: None,
232            };
233
234            assert!(required.matches(&available));
235        }
236
237        #[test]
238        fn required_capabilities_match_and_failure() {
239            let required = RequiredAggregatorCapabilities::And(vec![
240                RequiredAggregatorCapabilities::SignedEntityType(CardanoTransactions),
241                RequiredAggregatorCapabilities::SignedEntityType(CardanoStakeDistribution),
242                RequiredAggregatorCapabilities::AggregateSignatureType(Concatenation),
243            ]);
244            let available = AggregatorCapabilities {
245                aggregate_signature_type: Concatenation,
246                signed_entity_types: BTreeSet::from([CardanoTransactions]),
247                cardano_transactions_prover: None,
248            };
249
250            assert!(!required.matches(&available));
251        }
252    }
253
254    mod capable_discoverer {
255        use super::*;
256
257        fn create_aggregator_features_message(
258            capabilities: AggregatorCapabilities,
259        ) -> AggregatorFeaturesMessage {
260            AggregatorFeaturesMessage {
261                open_api_version: "1.0.0".to_string(),
262                documentation_url: "https://docs".to_string(),
263                capabilities,
264            }
265        }
266
267        #[tokio::test(flavor = "multi_thread")]
268        async fn get_available_aggregators_success() {
269            let capabilities = AggregatorCapabilities {
270                aggregate_signature_type: Concatenation,
271                signed_entity_types: BTreeSet::from([
272                    CardanoStakeDistribution,
273                    CardanoTransactions,
274                ]),
275                cardano_transactions_prover: None,
276            };
277            let aggregator_server = MockServer::start();
278            let aggregator_server_mock = aggregator_server.mock(|when, then| {
279                when.path("/");
280                then.status(200)
281                    .body(json!(create_aggregator_features_message(capabilities)).to_string());
282            });
283            let discoverer = CapableAggregatorDiscoverer::new(
284                RequiredAggregatorCapabilities::And(vec![
285                    RequiredAggregatorCapabilities::SignedEntityType(CardanoTransactions),
286                    RequiredAggregatorCapabilities::AggregateSignatureType(Concatenation),
287                ]),
288                Arc::new(crate::test::double::AggregatorDiscovererFake::new(vec![
289                    Ok(vec![AggregatorEndpoint::new(aggregator_server.url("/"))]),
290                ])),
291            );
292
293            let mut aggregators = discoverer
294                .get_available_aggregators(MithrilNetwork::new("release-devnet".into()))
295                .await
296                .unwrap();
297
298            let next_aggregator = aggregators.next().map(Into::into);
299            aggregator_server_mock.assert();
300            assert_eq!(
301                Some(AggregatorEndpoint::new(aggregator_server.url("/"))),
302                next_aggregator
303            );
304        }
305
306        #[tokio::test(flavor = "multi_thread")]
307        async fn get_available_aggregators_succeeds_when_aggregator_capabilities_do_not_match() {
308            let capabilities = AggregatorCapabilities {
309                aggregate_signature_type: Concatenation,
310                signed_entity_types: BTreeSet::from([CardanoTransactions]),
311                cardano_transactions_prover: None,
312            };
313            let aggregator_server = MockServer::start();
314            let aggregator_server_mock = aggregator_server.mock(|when, then| {
315                when.path("/");
316                then.status(200)
317                    .body(json!(create_aggregator_features_message(capabilities)).to_string());
318            });
319            let discoverer = CapableAggregatorDiscoverer::new(
320                RequiredAggregatorCapabilities::And(vec![
321                    RequiredAggregatorCapabilities::SignedEntityType(CardanoDatabase),
322                    RequiredAggregatorCapabilities::AggregateSignatureType(Concatenation),
323                ]),
324                Arc::new(crate::test::double::AggregatorDiscovererFake::new(vec![
325                    Ok(vec![AggregatorEndpoint::new(aggregator_server.url("/"))]),
326                ])),
327            );
328
329            let mut aggregators = discoverer
330                .get_available_aggregators(MithrilNetwork::new("release-devnet".into()))
331                .await
332                .unwrap();
333
334            let next_aggregator = aggregators.next();
335            aggregator_server_mock.assert();
336            assert!(next_aggregator.is_none());
337        }
338
339        #[tokio::test(flavor = "multi_thread")]
340        async fn get_available_aggregators_succeeds_when_one_aggregator_returns_an_error() {
341            let aggregator_server_1 = MockServer::start();
342            let aggregator_server_mock_1 = aggregator_server_1.mock(|when, then| {
343                when.path("/");
344                then.status(500);
345            });
346            let capabilities_2 = AggregatorCapabilities {
347                aggregate_signature_type: Concatenation,
348                signed_entity_types: BTreeSet::from([CardanoStakeDistribution, CardanoDatabase]),
349                cardano_transactions_prover: None,
350            };
351            let aggregator_server_2 = MockServer::start();
352            let aggregator_server_mock_2 = aggregator_server_2.mock(|when, then| {
353                when.path("/");
354                then.status(200)
355                    .body(json!(create_aggregator_features_message(capabilities_2)).to_string());
356            });
357            let discoverer = CapableAggregatorDiscoverer::new(
358                RequiredAggregatorCapabilities::And(vec![
359                    RequiredAggregatorCapabilities::SignedEntityType(CardanoDatabase),
360                    RequiredAggregatorCapabilities::AggregateSignatureType(Concatenation),
361                ]),
362                Arc::new(crate::test::double::AggregatorDiscovererFake::new(vec![
363                    Ok(vec![
364                        AggregatorEndpoint::new(aggregator_server_1.url("/")),
365                        AggregatorEndpoint::new(aggregator_server_2.url("/")),
366                    ]),
367                ])),
368            );
369
370            let mut aggregators = discoverer
371                .get_available_aggregators(MithrilNetwork::new("release-devnet".into()))
372                .await
373                .unwrap();
374
375            let next_aggregator = aggregators.next().map(|endpoint| endpoint.into());
376            aggregator_server_mock_1.assert();
377            aggregator_server_mock_2.assert();
378            assert_eq!(
379                Some(AggregatorEndpoint::new(aggregator_server_2.url("/"))),
380                next_aggregator
381            );
382        }
383
384        #[tokio::test(flavor = "multi_thread")]
385        async fn get_available_aggregators_succeeds_and_makes_minimum_calls_to_aggregators() {
386            let aggregator_server_1 = MockServer::start();
387            let aggregator_server_mock_1 = aggregator_server_1.mock(|when, then| {
388                when.path("/");
389                then.status(500);
390            });
391            let capabilities_2 = AggregatorCapabilities {
392                aggregate_signature_type: Concatenation,
393                signed_entity_types: BTreeSet::from([CardanoStakeDistribution]),
394                cardano_transactions_prover: None,
395            };
396            let aggregator_server_2 = MockServer::start();
397            let aggregator_server_mock_2 = aggregator_server_2.mock(|when, then| {
398                when.path("/");
399                then.status(200)
400                    .body(json!(create_aggregator_features_message(capabilities_2)).to_string());
401            });
402            let capabilities_3 = AggregatorCapabilities {
403                aggregate_signature_type: Concatenation,
404                signed_entity_types: BTreeSet::from([CardanoDatabase]),
405                cardano_transactions_prover: None,
406            };
407            let aggregator_server_3 = MockServer::start();
408            let aggregator_server_mock_3 = aggregator_server_3.mock(|when, then| {
409                when.path("/");
410                then.status(200)
411                    .body(json!(create_aggregator_features_message(capabilities_3)).to_string());
412            });
413            let capabilities_4 = AggregatorCapabilities {
414                aggregate_signature_type: Concatenation,
415                signed_entity_types: BTreeSet::from([CardanoDatabase]),
416                cardano_transactions_prover: None,
417            };
418            let aggregator_server_4 = MockServer::start();
419            let aggregator_server_mock_4 = aggregator_server_4.mock(|when, then| {
420                when.path("/");
421                then.status(200)
422                    .body(json!(create_aggregator_features_message(capabilities_4)).to_string());
423            });
424            let discoverer = CapableAggregatorDiscoverer::new(
425                RequiredAggregatorCapabilities::And(vec![
426                    RequiredAggregatorCapabilities::SignedEntityType(CardanoDatabase),
427                    RequiredAggregatorCapabilities::AggregateSignatureType(Concatenation),
428                ]),
429                Arc::new(crate::test::double::AggregatorDiscovererFake::new(vec![
430                    Ok(vec![
431                        AggregatorEndpoint::new(aggregator_server_1.url("/")),
432                        AggregatorEndpoint::new(aggregator_server_2.url("/")),
433                        AggregatorEndpoint::new(aggregator_server_3.url("/")),
434                        AggregatorEndpoint::new(aggregator_server_4.url("/")),
435                    ]),
436                ])),
437            );
438
439            let mut aggregators = discoverer
440                .get_available_aggregators(MithrilNetwork::new("release-devnet".into()))
441                .await
442                .unwrap();
443
444            let next_aggregator = aggregators.next().map(|endpoint| endpoint.into());
445            aggregator_server_mock_1.assert();
446            aggregator_server_mock_2.assert();
447            aggregator_server_mock_3.assert();
448            assert_eq!(0, aggregator_server_mock_4.calls());
449            assert_eq!(
450                Some(AggregatorEndpoint::new(aggregator_server_3.url("/"))),
451                next_aggregator
452            );
453        }
454    }
455}