use crate::hash::Hash;
use crate::object::{Object, object_id_from_parts};
use crate::serialize;
use super::{MAX_RAW_OBJECT_SIZE, ObjectSink, ObjectStore, StoreError, StoreResult};
pub trait ObjectSource {
fn read(&self, h: &Hash) -> StoreResult<Vec<u8>>;
fn read_object(&self, h: &Hash) -> StoreResult<Object> {
Ok(serialize::deserialize(&self.read(h)?)?)
}
fn read_unverified(&self, h: &Hash) -> StoreResult<Vec<u8>> {
self.read(h)
}
}
impl ObjectSource for ObjectStore {
fn read(&self, h: &Hash) -> StoreResult<Vec<u8>> {
ObjectStore::read(self, h)
}
fn read_unverified(&self, h: &Hash) -> StoreResult<Vec<u8>> {
ObjectStore::read_unverified(self, h)
}
}
#[derive(Debug)]
pub struct EphemeralSink<'s> {
store: &'s ObjectStore,
objects: std::sync::Mutex<std::collections::HashMap<Hash, Vec<u8>>>,
}
impl<'s> EphemeralSink<'s> {
#[must_use]
pub fn new(store: &'s ObjectStore) -> Self {
Self {
store,
objects: std::sync::Mutex::new(std::collections::HashMap::new()),
}
}
}
impl ObjectSink for EphemeralSink<'_> {
fn put(&self, bytes: &[u8]) -> StoreResult<Hash> {
self.put_parts(&[bytes])
}
fn put_parts(&self, parts: &[&[u8]]) -> StoreResult<Hash> {
let mut total: usize = 0;
for p in parts {
total = total
.checked_add(p.len())
.ok_or(StoreError::ObjectTooLarge)?;
}
if total > MAX_RAW_OBJECT_SIZE {
return Err(StoreError::ObjectTooLarge);
}
let h = object_id_from_parts(parts);
if self.store.contains(&h) {
return Ok(h);
}
self.objects
.lock()
.expect("ephemeral sink mutex")
.entry(h)
.or_insert_with(|| {
let mut buf = Vec::with_capacity(total);
for p in parts {
buf.extend_from_slice(p);
}
buf
});
Ok(h)
}
fn has(&self, h: &Hash) -> bool {
self.objects
.lock()
.expect("ephemeral sink mutex")
.contains_key(h)
|| self.store.contains(h)
}
}
impl EphemeralSink<'_> {
fn overlay_get(&self, h: &Hash) -> Option<Vec<u8>> {
self.objects
.lock()
.expect("ephemeral sink mutex")
.get(h)
.cloned()
}
}
impl ObjectSource for EphemeralSink<'_> {
fn read(&self, h: &Hash) -> StoreResult<Vec<u8>> {
match self.overlay_get(h) {
Some(bytes) => Ok(bytes),
None => self.store.read(h),
}
}
fn read_unverified(&self, h: &Hash) -> StoreResult<Vec<u8>> {
match self.overlay_get(h) {
Some(bytes) => Ok(bytes),
None => self.store.read_unverified(h),
}
}
}
pub struct DisplaySource<'a, S: ObjectSource + ?Sized> {
inner: &'a S,
}
impl<'a, S: ObjectSource + ?Sized> DisplaySource<'a, S> {
#[must_use]
pub fn new(inner: &'a S) -> Self {
Self { inner }
}
}
impl<S: ObjectSource + ?Sized> std::fmt::Debug for DisplaySource<'_, S> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("DisplaySource").finish_non_exhaustive()
}
}
impl<S: ObjectSource + ?Sized> ObjectSource for DisplaySource<'_, S> {
fn read(&self, h: &Hash) -> StoreResult<Vec<u8>> {
self.inner.read_unverified(h)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::layout::RepoLayout;
use std::fs::OpenOptions;
use std::io::{self, Seek, Write};
use tempfile::TempDir;
fn fresh_store() -> (TempDir, ObjectStore) {
let dir = TempDir::new().expect("tempdir");
let store = ObjectStore::init(&RepoLayout::single(dir.path())).expect("init");
(dir, store)
}
fn corrupt_first_byte(store: &ObjectStore, h: &Hash, first_byte: u8) {
let path = store.path_for(h);
let mut f = OpenOptions::new()
.read(true)
.write(true)
.open(&path)
.unwrap();
f.seek(io::SeekFrom::Start(0)).unwrap();
f.write_all(&[first_byte ^ 0xFF]).unwrap();
f.sync_all().unwrap();
}
#[test]
fn display_source_delegates_to_unverified() {
let (_dir, store) = fresh_store();
let bytes = b"trustworthy".to_vec();
let h = store.write(&bytes).unwrap();
corrupt_first_byte(&store, &h, bytes[0]);
assert!(matches!(
store.read(&h).unwrap_err(),
StoreError::HashMismatch { .. }
));
let display = DisplaySource::new(&store);
let mut corrupted = bytes.clone();
corrupted[0] ^= 0xFF;
assert_eq!(
display.read(&h).unwrap(),
corrupted,
"DisplaySource::read must delegate to the store's unverified read"
);
}
#[test]
fn ephemeral_sink_unverified_falls_through() {
let (_dir, store) = fresh_store();
let durable_bytes = b"trustworthy".to_vec();
let durable_h = store.write(&durable_bytes).unwrap();
corrupt_first_byte(&store, &durable_h, durable_bytes[0]);
let sink = EphemeralSink::new(&store);
let private_h = sink.put(b"snapshot-only").unwrap();
let display = DisplaySource::new(&sink);
assert_eq!(display.read(&private_h).unwrap(), b"snapshot-only");
let mut corrupted = durable_bytes.clone();
corrupted[0] ^= 0xFF;
assert_eq!(display.read(&durable_h).unwrap(), corrupted);
}
#[test]
fn read_unverified_default_is_verifying_read() {
struct OnlyReadImpl<'s>(&'s ObjectStore);
impl ObjectSource for OnlyReadImpl<'_> {
fn read(&self, h: &Hash) -> StoreResult<Vec<u8>> {
self.0.read(h)
}
}
let (_dir, store) = fresh_store();
let bytes = b"trustworthy".to_vec();
let h = store.write(&bytes).unwrap();
corrupt_first_byte(&store, &h, bytes[0]);
let only_read = OnlyReadImpl(&store);
assert!(matches!(
only_read.read_unverified(&h).unwrap_err(),
StoreError::HashMismatch { .. }
));
}
}