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    pub fn insert(&mut self, value: T) -> Result<ObjectId, ObjectTableError> {
142        self.try_reserve_insert()?;
143        Ok(self.insert_prepared(value))
144    }
145
146    /// Reserve all capacity required by the next insertion without changing table state.
147    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}