Skip to main content

snix_castore/directoryservice/combinators/
race.rs

1use std::sync::Arc;
2
3use futures::{StreamExt, TryStreamExt, stream::BoxStream};
4use tonic::async_trait;
5use tracing::instrument;
6
7use crate::{
8    B3Digest, Directory, combinators,
9    composition::{CompositionContext, ServiceBuilder},
10    directoryservice::{self, DirectoryPutter, DirectoryService, FailingPutter},
11};
12
13/// Holds references to multiple directory services.
14/// Read requests try services in parallel.
15/// The first positive response is returned (Ok(None) does only bubble up if all backends return this)
16/// Write requests are not implemented.
17pub struct Race<DS> {
18    instance_name: String,
19    services: Vec<DS>,
20}
21
22impl<DS> Race<DS> {
23    /// Construct from an iterator of services.
24    pub fn new<I: IntoIterator<Item = DS>>(instance_name: String, iter: I) -> Race<DS> {
25        Self {
26            instance_name,
27            services: Vec::from_iter(iter),
28        }
29    }
30
31    /// Add another sevice to the list.
32    pub fn add(&mut self, svc: DS) {
33        self.services.push(svc);
34    }
35}
36
37#[async_trait]
38impl<DS> DirectoryService for Race<DS>
39where
40    DS: DirectoryService,
41{
42    #[instrument(skip(self, digest), fields(directory.digest = %digest, instance_name = %self.instance_name))]
43    async fn get(&self, digest: &B3Digest) -> Result<Option<Directory>, directoryservice::Error> {
44        Ok(combinators::race::race_unary(&self.services, |svc| async {
45            // Skip over `Ok(None)` by returning None,
46            // but keep the Option<Directory> in the returned Ok() value.
47            Some(svc.get(digest).await.transpose()?.map(Some))
48        })
49        .await
50        .map_err(Error::Racing)?)
51    }
52
53    #[instrument(skip_all, fields(directory.digest = %root_directory_digest, instance_name = %self.instance_name))]
54    fn get_recursive(
55        &self,
56        root_directory_digest: &B3Digest,
57    ) -> BoxStream<'_, Result<Directory, directoryservice::Error>> {
58        let digest = *root_directory_digest;
59        combinators::race::race_stream(&self.services, move |svc| async move {
60            let mut stream = svc.get_recursive(&digest).peekable();
61
62            // Skip over backends that reported they don't have it.
63            if std::pin::Pin::new(&mut stream).peek().await.is_none() {
64                None
65            } else {
66                Some(stream)
67            }
68        })
69        .map_err(Error::Racing)
70        .err_into()
71        .boxed()
72    }
73
74    #[instrument(skip_all, fields(instance_name = %self.instance_name))]
75    async fn put(&self, _directory: Directory) -> Result<B3Digest, directoryservice::Error> {
76        Err(Error::Unimplemented.into())
77    }
78
79    #[instrument(skip_all)]
80    fn put_multiple_start(&self) -> Box<dyn DirectoryPutter + '_> {
81        Box::new(FailingPutter)
82    }
83}
84
85#[derive(thiserror::Error, Debug)]
86pub enum Error {
87    #[error("wrong arguments: {0}")]
88    WrongConfig(&'static str),
89
90    #[error("error from racing")]
91    Racing(#[source] combinators::race::Error<directoryservice::Error>),
92
93    #[error("puts are unimplemented")]
94    Unimplemented,
95}
96
97impl From<Error> for directoryservice::Error {
98    fn from(value: Error) -> Self {
99        Self(Box::new(value))
100    }
101}
102
103#[derive(serde::Deserialize, Debug)]
104#[serde(deny_unknown_fields)]
105pub struct RaceConfig {
106    services: Vec<String>,
107}
108
109impl TryFrom<url::Url> for RaceConfig {
110    type Error = Box<dyn std::error::Error + Send + Sync>;
111    fn try_from(url: url::Url) -> Result<Self, Self::Error> {
112        if url.has_authority() || !url.path().is_empty() {
113            return Err(Error::WrongConfig("no authority or path allowed").into());
114        }
115        Ok(serde_qs::from_str(url.query().unwrap_or_default())?)
116    }
117}
118
119#[async_trait]
120impl ServiceBuilder for RaceConfig {
121    type Output = dyn DirectoryService;
122    async fn build<'a>(
123        &'a self,
124        instance_name: &str,
125        context: &CompositionContext,
126    ) -> Result<Arc<Self::Output>, Box<dyn std::error::Error + Send + Sync>> {
127        let services =
128            futures::future::try_join_all(self.services.iter().map(|instance_ref| async move {
129                context.resolve::<Self::Output>(instance_ref).await
130            }))
131            .await?;
132
133        Ok(Arc::new(Race::new(instance_name.to_string(), services)))
134    }
135}
136
137#[cfg(test)]
138mod test {
139    use mockall::predicate;
140    use pretty_assertions::assert_matches;
141
142    use crate::{
143        combinators,
144        directoryservice::{self, DirectoryService, MockDirectoryService},
145        fixtures::DIRECTORY_WITH_KEEP,
146    };
147
148    use super::{Error, Race};
149
150    /// backends are tried exhaustively if all report None.
151    #[tokio::test]
152    async fn get_tries_exhaustively_on_none() {
153        let first = {
154            let mut svc = MockDirectoryService::new();
155            svc.expect_get()
156                .with(predicate::eq(DIRECTORY_WITH_KEEP.digest()))
157                .once()
158                .returning(|_| Ok(None));
159            svc
160        };
161        let second = {
162            let mut svc = MockDirectoryService::new();
163            svc.expect_get()
164                .with(predicate::eq(DIRECTORY_WITH_KEEP.digest()))
165                .once()
166                .returning(|_| Ok(None));
167            svc
168        };
169
170        let uut = Race::new("uut".to_string(), [first, second]);
171
172        assert!(
173            uut.get(&DIRECTORY_WITH_KEEP.digest())
174                .await
175                .expect("to succeed")
176                .is_none()
177        )
178    }
179
180    /// if one has it and one does not, we return the positive result.
181    #[tokio::test]
182    async fn get_returns_positive() {
183        let first = {
184            let mut svc = MockDirectoryService::new();
185            svc.expect_get()
186                .with(predicate::eq(DIRECTORY_WITH_KEEP.digest()))
187                .once()
188                .returning(|_| Ok(Some(DIRECTORY_WITH_KEEP.clone())));
189            svc
190        };
191
192        let second = {
193            let mut svc = MockDirectoryService::new();
194            svc.expect_get()
195                .with(predicate::eq(DIRECTORY_WITH_KEEP.digest()))
196                // We cannot be certain this is called at all, so no `once()` here.
197                .returning(|_| Ok(None));
198            svc
199        };
200
201        let uut = Race::new("uut".to_string(), [first, second]);
202
203        assert_eq!(
204            Some(DIRECTORY_WITH_KEEP.clone()),
205            uut.get(&DIRECTORY_WITH_KEEP.digest())
206                .await
207                .expect("to succeed")
208        )
209    }
210
211    /// Errors are bubbled up, and the error contains the correct service index.
212    #[tokio::test]
213    async fn get_return_error() {
214        let first = {
215            let mut svc = MockDirectoryService::new();
216            svc.expect_get()
217                .with(predicate::eq(DIRECTORY_WITH_KEEP.digest()))
218                .once()
219                .returning(|_| Err(directoryservice::Error("".into())));
220            svc
221        };
222
223        // Ideally this one would be just slower than `first`.
224        let second = {
225            let mut svc = MockDirectoryService::new();
226            svc.expect_get()
227                .with(predicate::eq(DIRECTORY_WITH_KEEP.digest()))
228                // We cannot be certain this is called at all, so no `once()` here.
229                .returning(|_| Ok(None));
230            svc
231        };
232
233        let uut = Race::new("uut".to_string(), [first, second]);
234
235        let err = uut
236            .get(&DIRECTORY_WITH_KEEP.digest())
237            .await
238            .expect_err("to fail")
239            .0;
240        let err = err.downcast_ref::<Error>().unwrap();
241        assert_matches!(err, Error::Racing(combinators::race::Error::Backend(0, _)))
242    }
243
244    // FUTUREWORK: ideally we'd be constructing mocks that take longer than others / never return,
245    // but that's not supported in automock: https://github.com/asomers/mockall/issues/189
246    // So it's a bit tough to create test cases reliably.
247}