Skip to main content

weavatrix_git/
backend.rs

1use std::{
2    collections::HashMap,
3    sync::{Arc, RwLock},
4};
5
6use crate::{Object, ObjectId, Result, error::invalid};
7
8pub trait ObjectBackend: Send + Sync {
9    fn contains(&self, id: ObjectId) -> Result<bool>;
10    fn read(&self, id: ObjectId, max_object_bytes: usize) -> Result<Option<Object>>;
11}
12
13#[derive(Default)]
14pub struct MemoryObjectBackend {
15    objects: RwLock<HashMap<ObjectId, Object>>,
16}
17
18impl MemoryObjectBackend {
19    #[must_use]
20    pub fn new(objects: impl IntoIterator<Item = Object>) -> Self {
21        Self {
22            objects: RwLock::new(
23                objects
24                    .into_iter()
25                    .map(|object| (object.id, object))
26                    .collect(),
27            ),
28        }
29    }
30
31    pub fn insert(&self, object: Object) {
32        self.objects
33            .write()
34            .unwrap_or_else(std::sync::PoisonError::into_inner)
35            .insert(object.id, object);
36    }
37}
38
39impl ObjectBackend for MemoryObjectBackend {
40    fn contains(&self, id: ObjectId) -> Result<bool> {
41        Ok(self
42            .objects
43            .read()
44            .unwrap_or_else(std::sync::PoisonError::into_inner)
45            .contains_key(&id))
46    }
47
48    fn read(&self, id: ObjectId, max_object_bytes: usize) -> Result<Option<Object>> {
49        let object = self
50            .objects
51            .read()
52            .unwrap_or_else(std::sync::PoisonError::into_inner)
53            .get(&id)
54            .cloned();
55        if object
56            .as_ref()
57            .is_some_and(|object| object.data.len() > max_object_bytes)
58        {
59            return Err(crate::GitError::LimitExceeded {
60                resource: "backend object bytes",
61                limit: max_object_bytes,
62            });
63        }
64        Ok(object)
65    }
66}
67
68pub(crate) fn read(
69    backends: &[Arc<dyn ObjectBackend>],
70    id: ObjectId,
71    max_object_bytes: usize,
72) -> Result<Option<Object>> {
73    for backend in backends {
74        let Some(object) = backend.read(id, max_object_bytes)? else {
75            continue;
76        };
77        if object.id != id {
78            return Err(invalid("object backend returned a different identifier"));
79        }
80        if object.data.len() > max_object_bytes {
81            return Err(crate::GitError::LimitExceeded {
82                resource: "backend object bytes",
83                limit: max_object_bytes,
84            });
85        }
86        return Ok(Some(object));
87    }
88    Ok(None)
89}
90
91pub(crate) fn contains(backends: &[Arc<dyn ObjectBackend>], id: ObjectId) -> Result<bool> {
92    for backend in backends {
93        if backend.contains(id)? {
94            return Ok(true);
95        }
96    }
97    Ok(false)
98}
99
100#[cfg(test)]
101mod tests {
102    use std::sync::Arc;
103
104    use super::{MemoryObjectBackend, ObjectBackend};
105    use crate::{Object, ObjectId, ObjectKind};
106
107    #[test]
108    fn memory_backend_enforces_object_limit() {
109        let id: ObjectId = "1111111111111111111111111111111111111111".parse().unwrap();
110        let backend = MemoryObjectBackend::new([Object {
111            id,
112            kind: ObjectKind::Blob,
113            data: b"data".to_vec(),
114        }]);
115        assert!(backend.contains(id).unwrap());
116        assert_eq!(backend.read(id, 4).unwrap().unwrap().data, b"data");
117        assert!(backend.read(id, 3).is_err());
118    }
119
120    struct WrongBackend(Object);
121
122    impl ObjectBackend for WrongBackend {
123        fn contains(&self, _id: ObjectId) -> crate::Result<bool> {
124            Ok(false)
125        }
126
127        fn read(&self, _id: ObjectId, _limit: usize) -> crate::Result<Option<Object>> {
128            Ok(Some(self.0.clone()))
129        }
130    }
131
132    #[test]
133    fn rejects_backend_identifier_mismatch() {
134        let requested: ObjectId = "1111111111111111111111111111111111111111".parse().unwrap();
135        let returned: ObjectId = "2222222222222222222222222222222222222222".parse().unwrap();
136        let backends: Vec<Arc<dyn ObjectBackend>> = vec![Arc::new(WrongBackend(Object {
137            id: returned,
138            kind: ObjectKind::Blob,
139            data: Vec::new(),
140        }))];
141        assert!(super::read(&backends, requested, 1).is_err());
142        assert!(!super::contains(&backends, requested).unwrap());
143    }
144}