Skip to main content

snix_castore/directoryservice/combinators/
race.rs

1use std::sync::Arc;
2
3use futures::{StreamExt, TryFutureExt, TryStreamExt, stream::BoxStream};
4use tonic::async_trait;
5use tracing::instrument;
6
7use crate::{
8    B3Digest, Directory,
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        // prepare requests to all backends, and annotate the backend_idx in the error case.
45        let mut requests: Vec<_> = self
46            .services
47            .iter()
48            .enumerate()
49            .map(|(backend_idx, svc)| {
50                svc.get(digest)
51                    .map_err(move |err| Error::Backend(backend_idx, err))
52            })
53            .collect();
54
55        while !requests.is_empty() {
56            let (resp, _fut_idx, remaining) = futures::future::select_all(requests).await;
57
58            match resp {
59                // If this Ok(Some(_)), return, we're done
60                Ok(Some(directory)) => return Ok(Some(directory)),
61                // Skip over backends that reported they don't have it.
62                Ok(None) => {}
63                // Bubble up errors. We already mapped the backend_idx into the error.
64                Err(err) => return Err(err)?,
65            }
66
67            requests = remaining;
68        }
69
70        // if we exhausted all backends, return Ok(None).
71        return Ok(None);
72    }
73
74    #[instrument(skip_all, fields(directory.digest = %root_directory_digest, instance_name = %self.instance_name))]
75    fn get_recursive(
76        &self,
77        root_directory_digest: &B3Digest,
78    ) -> BoxStream<'_, Result<Directory, directoryservice::Error>> {
79        let digest = *root_directory_digest;
80
81        // Create a bunch of futures that return ready once they get the first element of the stream, or an EOF.
82        let mut requests: Vec<_> = self
83            .services
84            .iter()
85            .enumerate()
86            .map(|(backend_idx, svc)| {
87                Box::pin(async move {
88                    let mut stream = svc
89                        .get_recursive(&digest)
90                        .map_err(move |err| Error::Backend(backend_idx, err));
91                    if let Some(directory) = stream.try_next().await? {
92                        Ok::<_, Error>(Some((directory, stream)))
93                    } else {
94                        Ok(None)
95                    }
96                })
97            })
98            .collect();
99
100        async_stream::try_stream! {
101            while !requests.is_empty() {
102                let (resp, _fut_idx, remaining) = futures::future::select_all(requests).await;
103
104                match resp {
105                    // If this Ok(Some(_, _)), yield from that stream.
106                    Ok(Some((directory, mut stream))) => {
107                        yield directory;
108
109                        while let Some(directory) = stream.try_next().await? {
110                            yield directory
111                        }
112                    }
113                    Ok(None) => {
114                        // Skip over backends that reported they don't have it.
115                    },
116                    // Bubble up errors. We already mapped the backend_idx into the error.
117                    Err(err) => Err(directoryservice::Error::from(err))?,
118                }
119                requests = remaining;
120            }
121            // if we exhausted all backends, this returns an empty stream
122        }
123        .boxed()
124    }
125
126    #[instrument(skip_all, fields(instance_name = %self.instance_name))]
127    async fn put(&self, _directory: Directory) -> Result<B3Digest, directoryservice::Error> {
128        Err(Error::Unimplemented.into())
129    }
130
131    #[instrument(skip_all)]
132    fn put_multiple_start(&self) -> Box<dyn DirectoryPutter + '_> {
133        Box::new(FailingPutter)
134    }
135}
136
137#[derive(thiserror::Error, Debug)]
138pub enum Error {
139    #[error("wrong arguments: {0}")]
140    WrongConfig(&'static str),
141
142    #[error("error from service at index {0}")]
143    Backend(usize, #[source] directoryservice::Error),
144
145    #[error("puts are unimplemented")]
146    Unimplemented,
147}
148
149impl From<Error> for directoryservice::Error {
150    fn from(value: Error) -> Self {
151        Self(Box::new(value))
152    }
153}
154
155#[derive(serde::Deserialize, Debug)]
156#[serde(deny_unknown_fields)]
157pub struct RaceConfig {
158    services: Vec<String>,
159}
160
161impl TryFrom<url::Url> for RaceConfig {
162    type Error = Box<dyn std::error::Error + Send + Sync>;
163    fn try_from(url: url::Url) -> Result<Self, Self::Error> {
164        if url.has_authority() || !url.path().is_empty() {
165            return Err(Error::WrongConfig("no authority or path allowed").into());
166        }
167        Ok(serde_qs::from_str(url.query().unwrap_or_default())?)
168    }
169}
170
171#[async_trait]
172impl ServiceBuilder for RaceConfig {
173    type Output = dyn DirectoryService;
174    async fn build<'a>(
175        &'a self,
176        instance_name: &str,
177        context: &CompositionContext,
178    ) -> Result<Arc<Self::Output>, Box<dyn std::error::Error + Send + Sync>> {
179        let services =
180            futures::future::try_join_all(self.services.iter().map(|instance_ref| async move {
181                context.resolve::<Self::Output>(instance_ref).await
182            }))
183            .await?;
184
185        Ok(Arc::new(Race::new(instance_name.to_string(), services)))
186    }
187}
188
189#[cfg(test)]
190mod test {
191    use mockall::predicate;
192    use pretty_assertions::assert_matches;
193
194    use crate::{
195        directoryservice::{self, DirectoryService, MockDirectoryService},
196        fixtures::DIRECTORY_WITH_KEEP,
197    };
198
199    use super::{Error, Race};
200
201    /// backends are tried exhaustively if all report None.
202    #[tokio::test]
203    async fn get_tries_exhaustively_on_none() {
204        let first = {
205            let mut svc = MockDirectoryService::new();
206            svc.expect_get()
207                .with(predicate::eq(DIRECTORY_WITH_KEEP.digest()))
208                .once()
209                .returning(|_| Ok(None));
210            svc
211        };
212        let second = {
213            let mut svc = MockDirectoryService::new();
214            svc.expect_get()
215                .with(predicate::eq(DIRECTORY_WITH_KEEP.digest()))
216                .once()
217                .returning(|_| Ok(None));
218            svc
219        };
220
221        let uut = Race::new("uut".to_string(), [first, second]);
222
223        assert!(
224            uut.get(&DIRECTORY_WITH_KEEP.digest())
225                .await
226                .expect("to succeed")
227                .is_none()
228        )
229    }
230
231    /// if one has it and one does not, we return the positive result.
232    #[tokio::test]
233    async fn get_returns_positive() {
234        let first = {
235            let mut svc = MockDirectoryService::new();
236            svc.expect_get()
237                .with(predicate::eq(DIRECTORY_WITH_KEEP.digest()))
238                .once()
239                .returning(|_| Ok(Some(DIRECTORY_WITH_KEEP.clone())));
240            svc
241        };
242
243        let second = {
244            let mut svc = MockDirectoryService::new();
245            svc.expect_get()
246                .with(predicate::eq(DIRECTORY_WITH_KEEP.digest()))
247                // We cannot be certain this is called at all, so no `once()` here.
248                .returning(|_| Ok(None));
249            svc
250        };
251
252        let uut = Race::new("uut".to_string(), [first, second]);
253
254        assert_eq!(
255            Some(DIRECTORY_WITH_KEEP.clone()),
256            uut.get(&DIRECTORY_WITH_KEEP.digest())
257                .await
258                .expect("to succeed")
259        )
260    }
261
262    /// Errors are bubbled up, and the error contains the correct service index.
263    #[tokio::test]
264    async fn get_return_error() {
265        let first = {
266            let mut svc = MockDirectoryService::new();
267            svc.expect_get()
268                .with(predicate::eq(DIRECTORY_WITH_KEEP.digest()))
269                .once()
270                .returning(|_| Err(directoryservice::Error("".into())));
271            svc
272        };
273
274        // Ideally this one would be just slower than `first`.
275        let second = {
276            let mut svc = MockDirectoryService::new();
277            svc.expect_get()
278                .with(predicate::eq(DIRECTORY_WITH_KEEP.digest()))
279                // We cannot be certain this is called at all, so no `once()` here.
280                .returning(|_| Ok(None));
281            svc
282        };
283
284        let uut = Race::new("uut".to_string(), [first, second]);
285
286        let err = uut
287            .get(&DIRECTORY_WITH_KEEP.digest())
288            .await
289            .expect_err("to fail")
290            .0;
291        let err = err.downcast_ref::<Error>().unwrap();
292        assert_matches!(err, Error::Backend(0, _))
293    }
294
295    // FUTUREWORK: ideally we'd be constructing mocks that take longer than others / never return,
296    // but that's not supported in automock: https://github.com/asomers/mockall/issues/189
297    // So it's a bit tough to create test cases reliably.
298}