Skip to main content

snix_castore/directoryservice/
grpc.rs

1use super::{Directory, DirectoryPutter, DirectoryService};
2use crate::B3Digest;
3use crate::composition::{CompositionContext, ServiceBuilder};
4use crate::directoryservice::order_validator::{self, OrderValidator, RootToLeaves};
5use crate::proto::{self, get_directory_request::ByWhat};
6use futures::StreamExt;
7use futures::stream::BoxStream;
8use std::sync::Arc;
9use tokio::spawn;
10use tokio::sync::mpsc::UnboundedSender;
11use tokio::task::JoinHandle;
12use tokio_stream::wrappers::UnboundedReceiverStream;
13use tonic::{Code, Status, async_trait};
14use tracing::{Instrument as _, instrument, warn};
15
16/// Maximum size of proto messages endoded and decoded.
17/// Unfortunately there's directories out there larger than 4MiB, so bump this
18/// by a bit.
19pub const MAX_DECODING_MESSAGE_SIZE: usize = 10 * 1024 * 1024;
20
21/// Connects to a (remote) snix-store DirectoryService over gRPC.
22#[derive(Clone)]
23pub struct GRPCDirectoryService<T> {
24    instance_name: String,
25    /// The internal reference to a gRPC client.
26    /// Cloning it is cheap, and it internally handles concurrent requests.
27    grpc_client: proto::directory_service_client::DirectoryServiceClient<T>,
28}
29
30impl<T> GRPCDirectoryService<T> {
31    /// construct a [GRPCDirectoryService] from a [proto::directory_service_client::DirectoryServiceClient].
32    /// panics if called outside the context of a tokio runtime.
33    pub fn from_client(
34        instance_name: String,
35        grpc_client: proto::directory_service_client::DirectoryServiceClient<T>,
36    ) -> Self {
37        Self {
38            instance_name,
39            grpc_client,
40        }
41    }
42}
43
44#[async_trait]
45impl<T> DirectoryService for GRPCDirectoryService<T>
46where
47    T: tonic::client::GrpcService<tonic::body::Body> + Send + Sync + Clone + 'static,
48    T::ResponseBody: tonic::codegen::Body<Data = tonic::codegen::Bytes> + Send + 'static,
49    <T::ResponseBody as tonic::codegen::Body>::Error: Into<tonic::codegen::StdError> + Send,
50    T::Future: Send,
51{
52    #[instrument(level = "trace", skip_all, fields(directory.digest = %digest, instance_name = %self.instance_name))]
53    async fn get(&self, digest: &B3Digest) -> Result<Option<Directory>, super::Error> {
54        // clone the client, as it takes a &mut.
55        // We retrieve the first message only, then close the stream (we set recursive to false)
56        match self
57            .grpc_client
58            .clone()
59            .get(proto::GetDirectoryRequest {
60                recursive: false,
61                by_what: Some(ByWhat::Digest((*digest).into())),
62            })
63            .await
64            .map_err(Error::Tonic)?
65            .into_inner()
66            .message()
67            .await
68        {
69            Ok(Some(proto_directory)) => {
70                // Validate the retrieved Directory indeed has the
71                // digest we expect it to have, to detect corruptions.
72                let actual_digest = proto_directory.digest();
73                if &actual_digest != digest {
74                    Err(Error::WrongDigest {
75                        expected: *digest,
76                        actual: actual_digest,
77                    })?
78                } else {
79                    let directory =
80                        Directory::try_from(proto_directory).map_err(Error::DirectoryValidation)?;
81                    Ok(Some(directory))
82                }
83            }
84            Ok(None) => Ok(None),
85            Err(e) if e.code() == Code::NotFound => Ok(None),
86            Err(e) => Err(Error::Tonic(e))?,
87        }
88    }
89
90    #[instrument(level = "trace", skip_all, fields(directory.digest = %directory.digest(), instance_name = %self.instance_name))]
91    async fn put(&self, directory: Directory) -> Result<B3Digest, super::Error> {
92        let resp = self
93            .grpc_client
94            .clone()
95            .put(tokio_stream::once(proto::Directory::from(directory)))
96            .await
97            .map_err(Error::Tonic)?;
98
99        let digest = resp
100            .into_inner()
101            .root_digest
102            .try_into()
103            .map_err(|_| Error::InvalidDigestLen)?;
104
105        Ok(digest)
106    }
107
108    #[instrument(level = "trace", skip_all, fields(directory.digest = %root_directory_digest, instance_name = %self.instance_name))]
109    fn get_recursive(
110        &self,
111        root_directory_digest: &B3Digest,
112    ) -> BoxStream<'static, Result<Directory, super::Error>> {
113        let mut grpc_client = self.grpc_client.clone();
114        let root_directory_digest = *root_directory_digest;
115
116        let mut order_validator = RootToLeaves::new_with_root_digest(root_directory_digest);
117
118        async_stream::try_stream! {
119            let mut directories = grpc_client
120                .get(proto::GetDirectoryRequest {
121                    recursive: true,
122                    by_what: Some(ByWhat::Digest((root_directory_digest).into())),
123                })
124                .await
125                .map_err(Error::Tonic)?
126                .into_inner();
127
128            while let Some(proto_directory) = directories.message().await.map_err(Error::Tonic)? {
129                let directory = Directory::try_from(proto_directory).map_err(Error::DirectoryValidation)?;
130                order_validator.try_accept(&directory).map_err(Error::DirectoryOrdering)?;
131
132                yield directory;
133            }
134        }
135        .boxed()
136    }
137
138    #[instrument(skip_all)]
139    fn put_multiple_start(&self) -> Box<dyn DirectoryPutter + 'static> {
140        let (tx, rx) = tokio::sync::mpsc::unbounded_channel();
141
142        let task = spawn({
143            let mut grpc_client = self.grpc_client.clone();
144
145            async move {
146                Ok::<_, Status>(
147                    grpc_client
148                        .put(UnboundedReceiverStream::new(rx))
149                        .await?
150                        .into_inner(),
151                )
152            }
153            // instrument the task with the current span, this is not done by default
154            .in_current_span()
155        });
156
157        Box::new(GRPCPutter {
158            rq: Some((task, tx)),
159        })
160    }
161}
162
163/// Allows uploading multiple Directory messages in the same gRPC stream.
164pub struct GRPCPutter {
165    /// Data about the current request - a handle to the task, and the tx part
166    /// of the channel.
167    /// The tx part of the pipe is used to send [proto::Directory] to the ongoing request.
168    /// The task will yield a [proto::PutDirectoryResponse] once the stream is closed.
169    #[allow(clippy::type_complexity)] // lol
170    rq: Option<(
171        JoinHandle<Result<proto::PutDirectoryResponse, Status>>,
172        UnboundedSender<proto::Directory>,
173    )>,
174}
175
176#[async_trait]
177impl DirectoryPutter for GRPCPutter {
178    #[instrument(level = "trace", skip_all, fields(directory.digest=%directory.digest()), err)]
179    async fn put(&mut self, directory: Directory) -> Result<(), super::Error> {
180        let (_, directory_sender) = self
181            .rq
182            .as_ref()
183            .ok_or_else(|| Error::DirectoryPutterAlreadyClosed)?;
184        // If we're not already closed, send the directory to directory_sender.
185        if directory_sender.send(directory.into()).is_err() {
186            // If the channel has been prematurely closed, invoke close (so we can peek at the error code)
187            // That error code is much more helpful, because it
188            // contains the error message from the server.
189            self.close().await?;
190        }
191        Ok(())
192    }
193
194    /// Closes the stream for sending, and returns the value.
195    #[instrument(level = "trace", skip_all, ret, err)]
196    async fn close(&mut self) -> Result<B3Digest, super::Error> {
197        // get self.rq, and replace it with None.
198        // This ensures we can only close it once.
199        let (task, directory_sender) =
200            std::mem::take(&mut self.rq).ok_or_else(|| Error::DirectoryPutterAlreadyClosed)?;
201
202        // close directory_sender, so blocking on task will finish.
203        drop(directory_sender);
204
205        let resp = task
206            .await
207            .map_err(Error::TokioJoin)?
208            .map_err(Error::Tonic)?;
209
210        Ok(B3Digest::try_from(resp.root_digest).map_err(|_| Error::InvalidDigestLen)?)
211    }
212}
213
214#[derive(thiserror::Error, Debug)]
215pub enum Error {
216    #[error("Directory Graph ordering error: {0}")]
217    DirectoryOrdering(#[from] order_validator::OrderingError),
218
219    #[error("DirectoryPutter already closed")]
220    DirectoryPutterAlreadyClosed,
221
222    #[error("requested directory has wrong digest, expected {expected}, actual {actual}")]
223    WrongDigest {
224        expected: B3Digest,
225        actual: B3Digest,
226    },
227
228    #[error("tonic status: {0}")]
229    Tonic(#[from] tonic::Status),
230
231    #[error("invalid digest length returned from put")]
232    InvalidDigestLen,
233
234    #[error("failed to decode protobuf: {0}")]
235    ProtobufDecode(#[from] prost::DecodeError),
236    #[error("failed to validate directory: {0}")]
237    DirectoryValidation(#[from] crate::DirectoryError),
238
239    #[error("join error: {0}")]
240    TokioJoin(#[from] tokio::task::JoinError),
241    #[error("io error: {0}")]
242    IO(#[from] std::io::Error),
243}
244
245impl From<Error> for super::Error {
246    fn from(value: Error) -> Self {
247        Self(Box::new(value))
248    }
249}
250
251#[derive(serde::Deserialize, Debug)]
252#[serde(deny_unknown_fields)]
253pub struct GRPCDirectoryServiceConfig {
254    url: String,
255}
256
257impl TryFrom<url::Url> for GRPCDirectoryServiceConfig {
258    type Error = Box<dyn std::error::Error + Send + Sync>;
259    fn try_from(url: url::Url) -> Result<Self, Self::Error> {
260        //   This is normally grpc+unix for unix sockets, and grpc+http(s) for the HTTP counterparts.
261        // - In the case of unix sockets, there must be a path, but may not be a host.
262        // - In the case of non-unix sockets, there must be a host, but no path.
263        // Constructing the channel is handled by snix_castore::channel::from_url.
264        Ok(GRPCDirectoryServiceConfig {
265            url: url.to_string(),
266        })
267    }
268}
269
270#[async_trait]
271impl ServiceBuilder for GRPCDirectoryServiceConfig {
272    type Output = dyn DirectoryService;
273    async fn build<'a>(
274        &'a self,
275        instance_name: &str,
276        _context: &CompositionContext,
277    ) -> Result<Arc<Self::Output>, Box<dyn std::error::Error + Send + Sync>> {
278        let client = proto::directory_service_client::DirectoryServiceClient::with_interceptor(
279            crate::tonic::channel_from_url(&self.url.parse()?).await?,
280            snix_tracing::propagate::tonic::send_trace,
281        );
282
283        let client = client.max_decoding_message_size(MAX_DECODING_MESSAGE_SIZE);
284
285        Ok(Arc::new(GRPCDirectoryService::from_client(
286            instance_name.to_string(),
287            client,
288        )))
289    }
290}
291#[cfg(test)]
292mod tests {
293    use std::time::Duration;
294    use tempfile::TempDir;
295    use tokio::net::UnixListener;
296    use tokio_retry::{Retry, strategy::ExponentialBackoff};
297    use tokio_stream::wrappers::UnixListenerStream;
298
299    use crate::{
300        directoryservice::{DirectoryService, GRPCDirectoryService},
301        fixtures,
302        proto::{GRPCDirectoryServiceWrapper, directory_service_client::DirectoryServiceClient},
303        utils::gen_test_directory_service,
304    };
305
306    /// This ensures connecting via gRPC works as expected.
307    #[tokio::test]
308    async fn test_valid_unix_path_ping_pong() {
309        let tmpdir = TempDir::new().unwrap();
310        let socket_path = tmpdir.path().join("daemon");
311
312        let path_clone = socket_path.clone();
313
314        // Spin up a server
315        tokio::spawn(async {
316            let uds = UnixListener::bind(path_clone).unwrap();
317            let uds_stream = UnixListenerStream::new(uds);
318
319            // spin up a new server
320            let mut server = tonic::transport::Server::builder();
321            let router = server.add_service(
322                crate::proto::directory_service_server::DirectoryServiceServer::new(
323                    GRPCDirectoryServiceWrapper::new(Box::new(gen_test_directory_service())),
324                ),
325            );
326            router.serve_with_incoming(uds_stream).await
327        });
328
329        // wait for the socket to be created
330        Retry::start(
331            ExponentialBackoff::from_millis(20).max_delay(Duration::from_secs(10)),
332            || async {
333                if socket_path.exists() {
334                    Ok(())
335                } else {
336                    Err(())
337                }
338            },
339        )
340        .await
341        .expect("failed to wait for socket");
342
343        // prepare a client
344        let grpc_client = {
345            let url = url::Url::parse(&format!(
346                "grpc+unix:{}?wait-connect=1",
347                socket_path.display()
348            ))
349            .expect("must parse");
350            let client = DirectoryServiceClient::new(
351                crate::tonic::channel_from_url(&url)
352                    .await
353                    .expect("must succeed"),
354            );
355            GRPCDirectoryService::from_client("test-instance".into(), client)
356        };
357
358        assert!(
359            grpc_client
360                .get(&fixtures::DIRECTORY_A.digest())
361                .await
362                .expect("must not fail")
363                .is_none()
364        )
365    }
366}