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}