1use alloc::vec::Vec;
2use core::num::{NonZeroU16, NonZeroU64};
3
4const KIND_MASK: u32 = 0b111;
5const GENERATION_BITS: u32 = 13;
6const GENERATION_MASK: u16 = (1 << GENERATION_BITS) - 1;
7const GENERATION_SHIFT: u32 = 3;
8const NAMESPACE_SHIFT: u32 = GENERATION_SHIFT + GENERATION_BITS;
9
10#[derive(Clone, Copy, Debug, PartialEq, Eq)]
11#[repr(u8)]
12pub enum ObjectKind {
13 Context = 1,
14 Buffer = 2,
15 Program = 3,
16 Queue = 4,
17 Event = 5,
18}
19
20impl ObjectKind {
21 const fn tag(self) -> u32 {
22 self as u32
23 }
24}
25
26#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
33#[repr(transparent)]
34pub struct ObjectNamespace(NonZeroU16);
35
36impl ObjectNamespace {
37 pub const fn new(value: u16) -> Option<Self> {
38 match NonZeroU16::new(value) {
39 Some(value) => Some(Self(value)),
40 None => None,
41 }
42 }
43
44 pub const fn get(self) -> u16 {
45 self.0.get()
46 }
47}
48
49#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
51#[repr(transparent)]
52pub struct ObjectId(NonZeroU64);
53
54impl ObjectId {
55 pub const fn from_raw(raw: u64) -> Option<Self> {
56 match NonZeroU64::new(raw) {
57 Some(raw) => Some(Self(raw)),
58 None => None,
59 }
60 }
61
62 pub const fn get(self) -> u64 {
63 self.0.get()
64 }
65
66 const fn new(index: u32, namespace: u16, generation: u16, kind: ObjectKind) -> Self {
67 let token = ((namespace as u32) << NAMESPACE_SHIFT)
68 | ((generation as u32) << GENERATION_SHIFT)
69 | kind.tag();
70 let raw = ((token as u64) << 32) | (index as u64 + 1);
71 match NonZeroU64::new(raw) {
72 Some(raw) => Self(raw),
73 None => unreachable!(),
74 }
75 }
76}
77
78#[derive(Clone, Copy, Debug, PartialEq, Eq)]
79pub enum ObjectTableError {
80 InvalidId,
81 WrongKind,
82 StaleId,
83 Full,
84 AllocationFailed,
85}
86
87struct Slot<T> {
88 generation: u16,
89 value: Option<T>,
90 retired: bool,
91}
92
93pub struct ObjectTable<T> {
98 kind: ObjectKind,
99 namespace: u16,
100 max_slots: u32,
101 live: u32,
102 slots: Vec<Slot<T>>,
103 free: Vec<u32>,
104}
105
106impl<T> ObjectTable<T> {
107 pub const fn new(kind: ObjectKind, max_slots: u32) -> Self {
108 Self {
109 kind,
110 namespace: 0,
111 max_slots,
112 live: 0,
113 slots: Vec::new(),
114 free: Vec::new(),
115 }
116 }
117
118 pub const fn with_namespace(
119 kind: ObjectKind,
120 max_slots: u32,
121 namespace: ObjectNamespace,
122 ) -> Self {
123 Self {
124 kind,
125 namespace: namespace.get(),
126 max_slots,
127 live: 0,
128 slots: Vec::new(),
129 free: Vec::new(),
130 }
131 }
132
133 pub const fn len(&self) -> u32 {
134 self.live
135 }
136
137 pub const fn is_empty(&self) -> bool {
138 self.live == 0
139 }
140
141 pub fn insert(&mut self, value: T) -> Result<ObjectId, ObjectTableError> {
142 self.try_reserve_insert()?;
143 Ok(self.insert_prepared(value))
144 }
145
146 pub fn try_reserve_insert(&mut self) -> Result<(), ObjectTableError> {
148 if !self.free.is_empty() {
149 return Ok(());
150 }
151 if self.slots.len() >= self.max_slots as usize {
152 return Err(ObjectTableError::Full);
153 }
154 self.slots
155 .try_reserve(1)
156 .map_err(|_| ObjectTableError::AllocationFailed)?;
157 let new_slot_count = self.slots.len() + 1;
158 if self.free.capacity() < new_slot_count {
159 self.free
160 .try_reserve(new_slot_count - self.free.len())
161 .map_err(|_| ObjectTableError::AllocationFailed)?;
162 }
163 Ok(())
164 }
165
166 pub(crate) fn insert_prepared(&mut self, value: T) -> ObjectId {
167 if let Some(index) = self.free.pop() {
168 let slot = &mut self.slots[index as usize];
169 debug_assert!(slot.value.is_none() && !slot.retired);
170 slot.value = Some(value);
171 self.live += 1;
172 return ObjectId::new(index, self.namespace, slot.generation, self.kind);
173 }
174
175 debug_assert!(self.slots.len() < self.max_slots as usize);
176 debug_assert!(self.slots.len() < self.slots.capacity());
177 let new_slot_count = self.slots.len() + 1;
178 debug_assert!(self.free.capacity() >= new_slot_count);
179 let index = self.slots.len() as u32;
180 let generation = 0;
181 self.slots.push(Slot {
182 generation,
183 value: Some(value),
184 retired: false,
185 });
186 self.live += 1;
187 ObjectId::new(index, self.namespace, generation, self.kind)
188 }
189
190 pub fn get(&self, id: ObjectId) -> Result<&T, ObjectTableError> {
191 let index = self.locate(id)?;
192 self.slots[index]
193 .value
194 .as_ref()
195 .ok_or(ObjectTableError::StaleId)
196 }
197
198 pub fn get_mut(&mut self, id: ObjectId) -> Result<&mut T, ObjectTableError> {
199 let index = self.locate(id)?;
200 self.slots[index]
201 .value
202 .as_mut()
203 .ok_or(ObjectTableError::StaleId)
204 }
205
206 pub fn remove(&mut self, id: ObjectId) -> Result<T, ObjectTableError> {
207 let index = self.locate(id)?;
208 let slot = &mut self.slots[index];
209 let value = slot.value.take().ok_or(ObjectTableError::StaleId)?;
210 self.live -= 1;
211
212 if slot.generation == GENERATION_MASK {
213 slot.retired = true;
214 } else {
215 slot.generation += 1;
216 self.free.push(index as u32);
217 }
218 Ok(value)
219 }
220
221 pub(crate) fn next_id_from(&self, start: usize) -> Option<(usize, ObjectId)> {
222 self.slots
223 .iter()
224 .enumerate()
225 .skip(start)
226 .find_map(|(index, slot)| {
227 if slot.retired || slot.value.is_none() {
228 return None;
229 }
230 Some((
231 index + 1,
232 ObjectId::new(index as u32, self.namespace, slot.generation, self.kind),
233 ))
234 })
235 }
236
237 fn locate(&self, id: ObjectId) -> Result<usize, ObjectTableError> {
238 let raw = id.get();
239 let slot_number = raw as u32;
240 if slot_number == 0 {
241 return Err(ObjectTableError::InvalidId);
242 }
243 let token = (raw >> 32) as u32;
244 if token & KIND_MASK != self.kind.tag() {
245 return Err(ObjectTableError::WrongKind);
246 }
247 if (token >> NAMESPACE_SHIFT) as u16 != self.namespace {
248 return Err(ObjectTableError::StaleId);
249 }
250 let generation = ((token >> GENERATION_SHIFT) as u16) & GENERATION_MASK;
251 let index = (slot_number - 1) as usize;
252 let slot = self.slots.get(index).ok_or(ObjectTableError::StaleId)?;
253 if slot.generation != generation || slot.retired || slot.value.is_none() {
254 return Err(ObjectTableError::StaleId);
255 }
256 Ok(index)
257 }
258}
259
260#[cfg(test)]
261mod tests {
262 use super::*;
263
264 #[test]
265 fn stale_ids_never_resolve_after_slot_reuse() {
266 let mut table = ObjectTable::new(ObjectKind::Buffer, 1);
267 let old = table.insert(10).unwrap();
268 assert_eq!(table.remove(old), Ok(10));
269 assert_eq!(table.get(old), Err(ObjectTableError::StaleId));
270
271 let new = table.insert(20).unwrap();
272 assert_ne!(old, new);
273 assert_eq!(table.get(new), Ok(&20));
274 assert_eq!(table.get(old), Err(ObjectTableError::StaleId));
275 }
276
277 #[test]
278 fn exhausted_generations_retire_the_slot_before_an_id_can_revive() {
279 let mut table = ObjectTable::new(ObjectKind::Buffer, 1);
280 let first = table.insert(()).unwrap();
281 table.remove(first).unwrap();
282
283 for _ in 1..=GENERATION_MASK {
284 let current = table.insert(()).unwrap();
285 assert_ne!(current, first);
286 assert_eq!(table.get(first), Err(ObjectTableError::StaleId));
287 table.remove(current).unwrap();
288 }
289
290 assert_eq!(table.insert(()), Err(ObjectTableError::Full));
291 assert_eq!(table.get(first), Err(ObjectTableError::StaleId));
292 }
293
294 #[test]
295 fn kind_tags_prevent_cross_table_aliasing() {
296 let mut contexts = ObjectTable::new(ObjectKind::Context, 1);
297 let id = contexts.insert(()).unwrap();
298 let buffers = ObjectTable::<()>::new(ObjectKind::Buffer, 1);
299 assert_eq!(buffers.get(id), Err(ObjectTableError::WrongKind));
300 }
301
302 #[test]
303 fn namespaces_prevent_cross_device_or_reset_epoch_aliasing() {
304 let first_namespace = ObjectNamespace::new(1).unwrap();
305 let second_namespace = ObjectNamespace::new(2).unwrap();
306 let mut first = ObjectTable::with_namespace(ObjectKind::Context, 1, first_namespace);
307 let id = first.insert(()).unwrap();
308 let second = ObjectTable::<()>::with_namespace(ObjectKind::Context, 1, second_namespace);
309 assert_eq!(second.get(id), Err(ObjectTableError::StaleId));
310 }
311
312 #[test]
313 fn limits_are_enforced_before_growth() {
314 let mut table = ObjectTable::new(ObjectKind::Event, 1);
315 table.insert(1).unwrap();
316 assert_eq!(table.insert(2), Err(ObjectTableError::Full));
317 assert_eq!(table.len(), 1);
318 }
319}