Skip to main content

snix_castore/blobservice/
memory.rs

1use parking_lot::RwLock;
2use std::io::{self, Cursor, Write};
3use std::task::Poll;
4use std::{collections::HashMap, sync::Arc};
5use tonic::async_trait;
6use tracing::{Level, instrument};
7
8use super::{BlobReader, BlobService, BlobWriter};
9use crate::B3Digest;
10use crate::composition::{CompositionContext, ServiceBuilder};
11
12#[derive(Clone)]
13pub struct MemoryBlobService {
14    instance_name: String,
15    db: Arc<RwLock<HashMap<B3Digest, Vec<u8>>>>,
16}
17
18#[cfg(any(test, feature = "mocks"))]
19impl MemoryBlobService {
20    /// Only used by `gen_test_blob_service`.
21    pub(crate) fn new_mock() -> Self {
22        Self {
23            instance_name: "mock".to_string(),
24            db: Default::default(),
25        }
26    }
27}
28
29#[async_trait]
30impl BlobService for MemoryBlobService {
31    #[instrument(skip_all, ret(level = Level::TRACE), err, fields(blob.digest=%digest, instance_name=%self.instance_name))]
32    async fn has(&self, digest: &B3Digest) -> io::Result<bool> {
33        let db = self.db.read();
34        Ok(db.contains_key(digest))
35    }
36
37    #[instrument(skip_all, err, fields(blob.digest=%digest, instance_name=%self.instance_name))]
38    async fn open_read(&self, digest: &B3Digest) -> io::Result<Option<Box<dyn BlobReader>>> {
39        let db = self.db.read();
40
41        match db.get(digest).map(|x| Cursor::new(x.clone())) {
42            Some(result) => Ok(Some(Box::new(result))),
43            None => Ok(None),
44        }
45    }
46
47    #[instrument(skip_all, fields(instance_name=%self.instance_name))]
48    async fn open_write(&self) -> Box<dyn BlobWriter> {
49        Box::new(MemoryBlobWriter::new(self.db.clone()))
50    }
51}
52
53#[derive(serde::Deserialize, Debug)]
54#[serde(deny_unknown_fields)]
55pub struct MemoryBlobServiceConfig {}
56
57impl TryFrom<url::Url> for MemoryBlobServiceConfig {
58    type Error = Box<dyn std::error::Error + Send + Sync>;
59    fn try_from(url: url::Url) -> Result<Self, Self::Error> {
60        // memory doesn't support authority or path in the URL.
61        if url.has_authority() || !url.path().is_empty() {
62            return Err("invalid url".into());
63        }
64        Ok(MemoryBlobServiceConfig {})
65    }
66}
67
68#[async_trait]
69impl ServiceBuilder for MemoryBlobServiceConfig {
70    type Output = dyn BlobService;
71    async fn build<'a>(
72        &'a self,
73        instance_name: &str,
74        _context: &CompositionContext,
75    ) -> Result<Arc<Self::Output>, Box<dyn std::error::Error + Send + Sync>> {
76        Ok(Arc::new(MemoryBlobService {
77            instance_name: instance_name.to_string(),
78            db: Default::default(),
79        }))
80    }
81}
82
83pub struct MemoryBlobWriter {
84    db: Arc<RwLock<HashMap<B3Digest, Vec<u8>>>>,
85
86    /// Contains the buffer Vec and hasher, or None if already closed
87    writers: Option<(Vec<u8>, blake3::Hasher)>,
88
89    /// The digest that has been returned, if we successfully closed.
90    digest: Option<B3Digest>,
91}
92
93impl MemoryBlobWriter {
94    fn new(db: Arc<RwLock<HashMap<B3Digest, Vec<u8>>>>) -> Self {
95        Self {
96            db,
97            writers: Some((Vec::new(), blake3::Hasher::new())),
98            digest: None,
99        }
100    }
101}
102impl tokio::io::AsyncWrite for MemoryBlobWriter {
103    fn poll_write(
104        mut self: std::pin::Pin<&mut Self>,
105        _cx: &mut std::task::Context<'_>,
106        b: &[u8],
107    ) -> std::task::Poll<Result<usize, io::Error>> {
108        Poll::Ready(match &mut self.writers {
109            None => Err(io::Error::new(
110                io::ErrorKind::NotConnected,
111                "already closed",
112            )),
113            Some((buf, hasher)) => {
114                let bytes_written = buf.write(b)?;
115                hasher.write(&b[..bytes_written])
116            }
117        })
118    }
119
120    fn poll_flush(
121        self: std::pin::Pin<&mut Self>,
122        _cx: &mut std::task::Context<'_>,
123    ) -> std::task::Poll<Result<(), io::Error>> {
124        Poll::Ready(match self.writers {
125            None => Err(io::Error::new(
126                io::ErrorKind::NotConnected,
127                "already closed",
128            )),
129            Some(_) => Ok(()),
130        })
131    }
132
133    fn poll_shutdown(
134        self: std::pin::Pin<&mut Self>,
135        _cx: &mut std::task::Context<'_>,
136    ) -> std::task::Poll<Result<(), io::Error>> {
137        // shutdown is "instantaneous", we only write to memory.
138        Poll::Ready(Ok(()))
139    }
140}
141
142#[async_trait]
143impl BlobWriter for MemoryBlobWriter {
144    async fn close(&mut self) -> io::Result<B3Digest> {
145        if self.writers.is_none() {
146            match &self.digest {
147                Some(digest) => Ok(*digest),
148                None => Err(io::Error::new(io::ErrorKind::BrokenPipe, "already closed")),
149            }
150        } else {
151            let (buf, hasher) = self.writers.take().unwrap();
152
153            let digest: B3Digest = hasher.finalize().as_bytes().into();
154
155            // Only insert if the blob doesn't already exist.
156            let mut db = self.db.upgradable_read();
157            if !db.contains_key(&digest) {
158                // open the database for writing.
159                db.with_upgraded(|db| {
160                    // and put buf in there. This will move buf out.
161                    db.insert(digest, buf);
162                });
163            }
164
165            self.digest = Some(digest);
166
167            Ok(digest)
168        }
169    }
170}