snix_castore/blobservice/
memory.rs1use 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 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 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 writers: Option<(Vec<u8>, blake3::Hasher)>,
88
89 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 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 let mut db = self.db.upgradable_read();
157 if !db.contains_key(&digest) {
158 db.with_upgraded(|db| {
160 db.insert(digest, buf);
162 });
163 }
164
165 self.digest = Some(digest);
166
167 Ok(digest)
168 }
169 }
170}