use std::{
collections::HashMap,
sync::{Arc, RwLock},
};
use crate::{Object, ObjectId, Result, error::invalid};
pub trait ObjectBackend: Send + Sync {
fn contains(&self, id: ObjectId) -> Result<bool>;
fn read(&self, id: ObjectId, max_object_bytes: usize) -> Result<Option<Object>>;
}
#[derive(Default)]
pub struct MemoryObjectBackend {
objects: RwLock<HashMap<ObjectId, Object>>,
}
impl MemoryObjectBackend {
#[must_use]
pub fn new(objects: impl IntoIterator<Item = Object>) -> Self {
Self {
objects: RwLock::new(
objects
.into_iter()
.map(|object| (object.id, object))
.collect(),
),
}
}
pub fn insert(&self, object: Object) {
self.objects
.write()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.insert(object.id, object);
}
}
impl ObjectBackend for MemoryObjectBackend {
fn contains(&self, id: ObjectId) -> Result<bool> {
Ok(self
.objects
.read()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.contains_key(&id))
}
fn read(&self, id: ObjectId, max_object_bytes: usize) -> Result<Option<Object>> {
let object = self
.objects
.read()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.get(&id)
.cloned();
if object
.as_ref()
.is_some_and(|object| object.data.len() > max_object_bytes)
{
return Err(crate::GitError::LimitExceeded {
resource: "backend object bytes",
limit: max_object_bytes,
});
}
Ok(object)
}
}
pub(crate) fn read(
backends: &[Arc<dyn ObjectBackend>],
id: ObjectId,
max_object_bytes: usize,
) -> Result<Option<Object>> {
for backend in backends {
let Some(object) = backend.read(id, max_object_bytes)? else {
continue;
};
if object.id != id {
return Err(invalid("object backend returned a different identifier"));
}
if object.data.len() > max_object_bytes {
return Err(crate::GitError::LimitExceeded {
resource: "backend object bytes",
limit: max_object_bytes,
});
}
return Ok(Some(object));
}
Ok(None)
}
pub(crate) fn contains(backends: &[Arc<dyn ObjectBackend>], id: ObjectId) -> Result<bool> {
for backend in backends {
if backend.contains(id)? {
return Ok(true);
}
}
Ok(false)
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use super::{MemoryObjectBackend, ObjectBackend};
use crate::{Object, ObjectId, ObjectKind};
#[test]
fn memory_backend_enforces_object_limit() {
let id: ObjectId = "1111111111111111111111111111111111111111".parse().unwrap();
let backend = MemoryObjectBackend::new([Object {
id,
kind: ObjectKind::Blob,
data: b"data".to_vec(),
}]);
assert!(backend.contains(id).unwrap());
assert_eq!(backend.read(id, 4).unwrap().unwrap().data, b"data");
assert!(backend.read(id, 3).is_err());
}
struct WrongBackend(Object);
impl ObjectBackend for WrongBackend {
fn contains(&self, _id: ObjectId) -> crate::Result<bool> {
Ok(false)
}
fn read(&self, _id: ObjectId, _limit: usize) -> crate::Result<Option<Object>> {
Ok(Some(self.0.clone()))
}
}
#[test]
fn rejects_backend_identifier_mismatch() {
let requested: ObjectId = "1111111111111111111111111111111111111111".parse().unwrap();
let returned: ObjectId = "2222222222222222222222222222222222222222".parse().unwrap();
let backends: Vec<Arc<dyn ObjectBackend>> = vec![Arc::new(WrongBackend(Object {
id: returned,
kind: ObjectKind::Blob,
data: Vec::new(),
}))];
assert!(super::read(&backends, requested, 1).is_err());
assert!(!super::contains(&backends, requested).unwrap());
}
}