Skip to main content

virtio_accel_device/
object_table.rs

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/// Device-instance namespace encoded into every object ID.
27///
28/// A transport integration assigns a distinct nonzero namespace to each device reset epoch and
29/// does not reuse it while an ID from that epoch could still be presented. IDs from different
30/// devices or reset epochs can therefore never resolve even when their slot, kind, and generation
31/// are otherwise identical.
32#[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/// Opaque guest-visible identifier. Its encoding is device-private, not part of the wire ABI.
50#[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
93/// Bounded generational object table for one resource kind.
94///
95/// Kind tags, generations, and device namespaces occupy separate token fields. A slot is retired
96/// before generation overflow, so an old ID cannot become valid again after wraparound.
97pub 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    /// Borrow the currently occupied identities without allocating.
142    ///
143    /// IDs retain this table's namespace, resource kind and current generation.
144    /// Vacant and permanently retired slots are skipped. IDs are yielded
145    /// in ascending slot index order. The iterator borrows the table, so
146    /// mutation requires ending the iteration first; a saved ID must still pass ordinary lookup checks after subsequent mutation.
147    pub fn ids(&self) -> impl Iterator<Item = ObjectId> + '_ {
148        self.slots.iter().enumerate().filter_map(|(index, slot)| {
149            slot.value
150                .as_ref()
151                .map(|_| ObjectId::new(index as u32, self.namespace, slot.generation, self.kind))
152        })
153    }
154
155    pub fn insert(&mut self, value: T) -> Result<ObjectId, ObjectTableError> {
156        self.try_reserve_insert()?;
157        Ok(self.insert_prepared(value))
158    }
159
160    /// Reserve all capacity required by the next insertion without changing table state.
161    pub fn try_reserve_insert(&mut self) -> Result<(), ObjectTableError> {
162        if !self.free.is_empty() {
163            return Ok(());
164        }
165        if self.slots.len() >= self.max_slots as usize {
166            return Err(ObjectTableError::Full);
167        }
168        self.slots
169            .try_reserve(1)
170            .map_err(|_| ObjectTableError::AllocationFailed)?;
171        let new_slot_count = self.slots.len() + 1;
172        if self.free.capacity() < new_slot_count {
173            self.free
174                .try_reserve(new_slot_count - self.free.len())
175                .map_err(|_| ObjectTableError::AllocationFailed)?;
176        }
177        Ok(())
178    }
179
180    pub(crate) fn insert_prepared(&mut self, value: T) -> ObjectId {
181        if let Some(index) = self.free.pop() {
182            let slot = &mut self.slots[index as usize];
183            debug_assert!(slot.value.is_none() && !slot.retired);
184            slot.value = Some(value);
185            self.live += 1;
186            return ObjectId::new(index, self.namespace, slot.generation, self.kind);
187        }
188
189        debug_assert!(self.slots.len() < self.max_slots as usize);
190        debug_assert!(self.slots.len() < self.slots.capacity());
191        let new_slot_count = self.slots.len() + 1;
192        debug_assert!(self.free.capacity() >= new_slot_count);
193        let index = self.slots.len() as u32;
194        let generation = 0;
195        self.slots.push(Slot {
196            generation,
197            value: Some(value),
198            retired: false,
199        });
200        self.live += 1;
201        ObjectId::new(index, self.namespace, generation, self.kind)
202    }
203
204    pub fn get(&self, id: ObjectId) -> Result<&T, ObjectTableError> {
205        let index = self.locate(id)?;
206        self.slots[index]
207            .value
208            .as_ref()
209            .ok_or(ObjectTableError::StaleId)
210    }
211
212    pub fn get_mut(&mut self, id: ObjectId) -> Result<&mut T, ObjectTableError> {
213        let index = self.locate(id)?;
214        self.slots[index]
215            .value
216            .as_mut()
217            .ok_or(ObjectTableError::StaleId)
218    }
219
220    pub fn remove(&mut self, id: ObjectId) -> Result<T, ObjectTableError> {
221        let index = self.locate(id)?;
222        let slot = &mut self.slots[index];
223        let value = slot.value.take().ok_or(ObjectTableError::StaleId)?;
224        self.live -= 1;
225
226        if slot.generation == GENERATION_MASK {
227            slot.retired = true;
228        } else {
229            slot.generation += 1;
230            self.free.push(index as u32);
231        }
232        Ok(value)
233    }
234
235    pub(crate) fn next_id_from(&self, start: usize) -> Option<(usize, ObjectId)> {
236        self.slots
237            .iter()
238            .enumerate()
239            .skip(start)
240            .find_map(|(index, slot)| {
241                if slot.retired || slot.value.is_none() {
242                    return None;
243                }
244                Some((
245                    index + 1,
246                    ObjectId::new(index as u32, self.namespace, slot.generation, self.kind),
247                ))
248            })
249    }
250
251    fn locate(&self, id: ObjectId) -> Result<usize, ObjectTableError> {
252        let raw = id.get();
253        let slot_number = raw as u32;
254        if slot_number == 0 {
255            return Err(ObjectTableError::InvalidId);
256        }
257        let token = (raw >> 32) as u32;
258        if token & KIND_MASK != self.kind.tag() {
259            return Err(ObjectTableError::WrongKind);
260        }
261        if (token >> NAMESPACE_SHIFT) as u16 != self.namespace {
262            return Err(ObjectTableError::StaleId);
263        }
264        let generation = ((token >> GENERATION_SHIFT) as u16) & GENERATION_MASK;
265        let index = (slot_number - 1) as usize;
266        let slot = self.slots.get(index).ok_or(ObjectTableError::StaleId)?;
267        if slot.generation != generation || slot.retired || slot.value.is_none() {
268            return Err(ObjectTableError::StaleId);
269        }
270        Ok(index)
271    }
272}
273
274#[cfg(test)]
275mod tests {
276    use super::*;
277
278    #[test]
279    fn iteration_preserves_identity_through_vacancy_reuse_and_exhaustion() {
280        let namespace = ObjectNamespace::new(19).unwrap();
281        let mut table = ObjectTable::with_namespace(ObjectKind::Buffer, 2, namespace);
282        assert_eq!(table.ids().next(), None);
283        let first = table.insert(1).unwrap();
284        let second = table.insert(2).unwrap();
285        assert_eq!(table.ids().collect::<Vec<_>>(), [first, second]);
286        table.remove(first).unwrap();
287        assert_eq!(table.ids().collect::<Vec<_>>(), [second]);
288        let replacement = table.insert(3).unwrap();
289        assert_ne!(first, replacement);
290        assert_eq!(table.ids().collect::<Vec<_>>(), [replacement, second]);
291        assert_eq!(table.get(first), Err(ObjectTableError::StaleId));
292        for id in table.ids() {
293            assert!(table.get(id).is_ok());
294            let other = ObjectTable::<u32>::with_namespace(
295                ObjectKind::Buffer,
296                2,
297                ObjectNamespace::new(20).unwrap(),
298            );
299            assert!(other.get(id).is_err());
300        }
301        table.slots[0].generation = GENERATION_MASK;
302        let exhausted = table.ids().next().unwrap();
303        table.remove(exhausted).unwrap();
304        assert_eq!(table.ids().collect::<Vec<_>>(), [second]);
305        assert_eq!(table.insert(4), Err(ObjectTableError::Full));
306    }
307
308    #[test]
309    fn stale_ids_never_resolve_after_slot_reuse() {
310        let mut table = ObjectTable::new(ObjectKind::Buffer, 1);
311        let old = table.insert(10).unwrap();
312        assert_eq!(table.remove(old), Ok(10));
313        assert_eq!(table.get(old), Err(ObjectTableError::StaleId));
314
315        let new = table.insert(20).unwrap();
316        assert_ne!(old, new);
317        assert_eq!(table.get(new), Ok(&20));
318        assert_eq!(table.get(old), Err(ObjectTableError::StaleId));
319    }
320
321    #[test]
322    fn exhausted_generations_retire_the_slot_before_an_id_can_revive() {
323        let mut table = ObjectTable::new(ObjectKind::Buffer, 1);
324        let first = table.insert(()).unwrap();
325        table.remove(first).unwrap();
326
327        for _ in 1..=GENERATION_MASK {
328            let current = table.insert(()).unwrap();
329            assert_ne!(current, first);
330            assert_eq!(table.get(first), Err(ObjectTableError::StaleId));
331            table.remove(current).unwrap();
332        }
333
334        assert_eq!(table.insert(()), Err(ObjectTableError::Full));
335        assert_eq!(table.get(first), Err(ObjectTableError::StaleId));
336    }
337
338    #[test]
339    fn kind_tags_prevent_cross_table_aliasing() {
340        let mut contexts = ObjectTable::new(ObjectKind::Context, 1);
341        let id = contexts.insert(()).unwrap();
342        let buffers = ObjectTable::<()>::new(ObjectKind::Buffer, 1);
343        assert_eq!(buffers.get(id), Err(ObjectTableError::WrongKind));
344    }
345
346    #[test]
347    fn namespaces_prevent_cross_device_or_reset_epoch_aliasing() {
348        let first_namespace = ObjectNamespace::new(1).unwrap();
349        let second_namespace = ObjectNamespace::new(2).unwrap();
350        let mut first = ObjectTable::with_namespace(ObjectKind::Context, 1, first_namespace);
351        let id = first.insert(()).unwrap();
352        let second = ObjectTable::<()>::with_namespace(ObjectKind::Context, 1, second_namespace);
353        assert_eq!(second.get(id), Err(ObjectTableError::StaleId));
354    }
355
356    #[test]
357    fn limits_are_enforced_before_growth() {
358        let mut table = ObjectTable::new(ObjectKind::Event, 1);
359        table.insert(1).unwrap();
360        assert_eq!(table.insert(2), Err(ObjectTableError::Full));
361        assert_eq!(table.len(), 1);
362    }
363}