snix_castore/directoryservice/combinators/
race.rs1use 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
13pub struct Race<DS> {
18 instance_name: String,
19 services: Vec<DS>,
20}
21
22impl<DS> Race<DS> {
23 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 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 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 Ok(Some(directory)) => return Ok(Some(directory)),
61 Ok(None) => {}
63 Err(err) => return Err(err)?,
65 }
66
67 requests = remaining;
68 }
69
70 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 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 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 },
116 Err(err) => Err(directoryservice::Error::from(err))?,
118 }
119 requests = remaining;
120 }
121 }
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 #[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 #[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 .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 #[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 let second = {
276 let mut svc = MockDirectoryService::new();
277 svc.expect_get()
278 .with(predicate::eq(DIRECTORY_WITH_KEEP.digest()))
279 .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 }