Skip to main content

virtio_accel_device/
state.rs

1//! Context-scoped device ownership, quotas, references, and release transitions.
2//!
3//! `DeviceState` contains no locks or interior mutability. Every transition requires exclusive
4//! access, giving a future concurrent command engine one outer synchronization boundary and no
5//! internal lock-ordering graph. Creation methods validate quotas and reserve table capacity before
6//! invoking a provider closure. Release methods move resources through an explicit `Releasing`
7//! state so rejected provider releases can restore ownership without reviving a stale ID.
8
9use alloc::vec::Vec;
10use core::num::NonZeroU64;
11
12use virtio_accel_core::{BackendError, BufferDesc, BufferInfo, DeviceLimits};
13use virtio_accel_proto::HARD_MAX_BINDINGS;
14
15use crate::{ObjectId, ObjectKind, ObjectNamespace, ObjectTable, ObjectTableError};
16
17#[derive(Clone, Copy, Debug, PartialEq, Eq)]
18pub enum DeviceStateConfigError {
19    ZeroLimit,
20    BindingLimit,
21    CountOverflow,
22    ReferenceCountOverflow,
23}
24
25#[derive(Clone, Copy, Debug, PartialEq, Eq)]
26pub enum DeviceStateError {
27    InvalidArgument,
28    InvalidObject,
29    StaleObject,
30    ContextMismatch,
31    Busy,
32    ResourceLimit,
33    OutOfMemory,
34    Releasing,
35    InvalidTransition,
36    ReferenceCountOverflow,
37}
38
39#[derive(Debug)]
40pub enum CreateError<E> {
41    State(DeviceStateError),
42    Provider(E),
43}
44
45impl<E> From<DeviceStateError> for CreateError<E> {
46    fn from(error: DeviceStateError) -> Self {
47        Self::State(error)
48    }
49}
50
51/// Result of allocating provider backing before its guest-visible ID is published.
52///
53/// `CleanupRequired` retains the object in the state graph so the command engine can release it
54/// through the provider's ownership boundary. The ID must never be exposed to the guest.
55#[derive(Clone, Copy, Debug, PartialEq, Eq)]
56pub enum BufferCreateOutcome {
57    Admitted(ObjectId),
58    CleanupRequired { id: ObjectId, error: BackendError },
59}
60
61#[derive(Debug)]
62pub struct RestoreError<R> {
63    pub error: DeviceStateError,
64    pub resource: R,
65}
66
67#[derive(Clone, Copy, Debug, PartialEq, Eq)]
68pub enum ReleaseState {
69    Live,
70    Releasing,
71}
72
73#[derive(Debug)]
74struct ResourceSlot<R> {
75    resource: Option<R>,
76    release: ReleaseState,
77}
78
79impl<R> ResourceSlot<R> {
80    fn new(resource: R) -> Self {
81        Self {
82            resource: Some(resource),
83            release: ReleaseState::Live,
84        }
85    }
86
87    fn get(&self) -> Result<&R, DeviceStateError> {
88        self.resource.as_ref().ok_or(DeviceStateError::Releasing)
89    }
90
91    fn get_mut(&mut self) -> Result<&mut R, DeviceStateError> {
92        self.resource.as_mut().ok_or(DeviceStateError::Releasing)
93    }
94
95    fn begin_release(&mut self) -> Result<R, DeviceStateError> {
96        if self.release != ReleaseState::Live {
97            return Err(DeviceStateError::Releasing);
98        }
99        let resource = self.resource.take().ok_or(DeviceStateError::Releasing)?;
100        self.release = ReleaseState::Releasing;
101        Ok(resource)
102    }
103
104    fn restore(&mut self, resource: R) -> Result<(), RestoreError<R>> {
105        if self.release != ReleaseState::Releasing || self.resource.is_some() {
106            return Err(RestoreError {
107                error: DeviceStateError::InvalidTransition,
108                resource,
109            });
110        }
111        self.resource = Some(resource);
112        self.release = ReleaseState::Live;
113        Ok(())
114    }
115}
116
117#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
118pub struct ChildCounts {
119    pub buffers: u32,
120    pub programs: u32,
121    pub queues: u32,
122    pub events: u32,
123}
124
125impl ChildCounts {
126    pub const fn is_empty(self) -> bool {
127        self.buffers == 0 && self.programs == 0 && self.queues == 0 && self.events == 0
128    }
129}
130
131#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
132pub struct ResourceCounts {
133    pub contexts: u64,
134    pub buffers: u64,
135    pub programs: u64,
136    pub queues: u64,
137    pub events: u64,
138}
139
140impl ResourceCounts {
141    pub const fn is_empty(self) -> bool {
142        self.contexts == 0
143            && self.buffers == 0
144            && self.programs == 0
145            && self.queues == 0
146            && self.events == 0
147    }
148
149    pub const fn total(self) -> u64 {
150        self.contexts
151            .saturating_add(self.buffers)
152            .saturating_add(self.programs)
153            .saturating_add(self.queues)
154            .saturating_add(self.events)
155    }
156
157    pub(crate) const fn saturating_add(self, other: Self) -> Self {
158        Self {
159            contexts: self.contexts.saturating_add(other.contexts),
160            buffers: self.buffers.saturating_add(other.buffers),
161            programs: self.programs.saturating_add(other.programs),
162            queues: self.queues.saturating_add(other.queues),
163            events: self.events.saturating_add(other.events),
164        }
165    }
166}
167
168/// Device-private aggregate limits for provider-retained bulk storage.
169///
170/// Per-object and object-count limits remain in [`DeviceLimits`]. These limits let one device
171/// integration constrain the aggregate host/provider memory exposed to an untrusted guest without
172/// adding host policy to the wire ABI.
173#[derive(Clone, Copy, Debug, PartialEq, Eq)]
174pub struct ResourcePolicy {
175    max_buffer_backing_bytes: NonZeroU64,
176    max_program_resident_bytes: NonZeroU64,
177}
178
179impl ResourcePolicy {
180    pub const fn new(
181        max_buffer_backing_bytes: u64,
182        max_program_resident_bytes: u64,
183    ) -> Option<Self> {
184        match (
185            NonZeroU64::new(max_buffer_backing_bytes),
186            NonZeroU64::new(max_program_resident_bytes),
187        ) {
188            (Some(max_buffer_backing_bytes), Some(max_program_resident_bytes)) => Some(Self {
189                max_buffer_backing_bytes,
190                max_program_resident_bytes,
191            }),
192            _ => None,
193        }
194    }
195
196    pub const fn max_buffer_backing_bytes(self) -> u64 {
197        self.max_buffer_backing_bytes.get()
198    }
199
200    pub const fn max_program_resident_bytes(self) -> u64 {
201        self.max_program_resident_bytes.get()
202    }
203}
204
205/// Exact retained bulk-storage charges represented by a device state graph.
206///
207/// The totals use `u128` because a valid graph may contain up to `u32::MAX` objects, each carrying
208/// a `u64` charge. This keeps accounting exact even when the sum is intentionally rejected by a
209/// `u64` policy limit.
210#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
211pub struct RetainedBytes {
212    pub buffer_backing: u128,
213    pub program_resident: u128,
214}
215
216impl RetainedBytes {
217    pub const fn is_empty(self) -> bool {
218        self.buffer_backing == 0 && self.program_resident == 0
219    }
220
221    pub(crate) const fn saturating_add(self, other: Self) -> Self {
222        Self {
223            buffer_backing: self.buffer_backing.saturating_add(other.buffer_backing),
224            program_resident: self.program_resident.saturating_add(other.program_resident),
225        }
226    }
227}
228
229#[derive(Debug)]
230pub struct ContextRecord<C> {
231    resource: ResourceSlot<C>,
232    children: ChildCounts,
233}
234
235impl<C> ContextRecord<C> {
236    pub fn resource(&self) -> Result<&C, DeviceStateError> {
237        self.resource.get()
238    }
239
240    pub const fn release_state(&self) -> ReleaseState {
241        self.resource.release
242    }
243
244    pub const fn children(&self) -> ChildCounts {
245        self.children
246    }
247}
248
249#[derive(Debug)]
250pub struct BufferRecord<B> {
251    resource: ResourceSlot<B>,
252    context_id: ObjectId,
253    info: BufferInfo,
254    in_flight: u32,
255}
256
257impl<B> BufferRecord<B> {
258    pub fn resource(&self) -> Result<&B, DeviceStateError> {
259        self.resource.get()
260    }
261
262    /// Borrow the live provider buffer for an explicit mutating operation.
263    ///
264    /// The command engine remains responsible for validating the operation and any in-flight
265    /// access policy before invoking the provider.
266    pub fn resource_mut(&mut self) -> Result<&mut B, DeviceStateError> {
267        self.resource.get_mut()
268    }
269
270    pub const fn context_id(&self) -> ObjectId {
271        self.context_id
272    }
273
274    pub const fn info(&self) -> BufferInfo {
275        self.info
276    }
277
278    pub const fn in_flight(&self) -> u32 {
279        self.in_flight
280    }
281
282    pub const fn release_state(&self) -> ReleaseState {
283        self.resource.release
284    }
285}
286
287#[derive(Debug)]
288pub struct ProgramRecord<P> {
289    resource: ResourceSlot<P>,
290    context_id: ObjectId,
291    resident_bytes: u64,
292    in_flight: u32,
293}
294
295impl<P> ProgramRecord<P> {
296    pub fn resource(&self) -> Result<&P, DeviceStateError> {
297        self.resource.get()
298    }
299
300    pub const fn context_id(&self) -> ObjectId {
301        self.context_id
302    }
303
304    pub const fn resident_bytes(&self) -> u64 {
305        self.resident_bytes
306    }
307
308    pub const fn in_flight(&self) -> u32 {
309        self.in_flight
310    }
311
312    pub const fn release_state(&self) -> ReleaseState {
313        self.resource.release
314    }
315}
316
317#[derive(Debug)]
318pub struct QueueRecord<Q> {
319    resource: ResourceSlot<Q>,
320    context_id: ObjectId,
321    in_flight: u32,
322}
323
324impl<Q> QueueRecord<Q> {
325    pub fn resource(&self) -> Result<&Q, DeviceStateError> {
326        self.resource.get()
327    }
328
329    pub const fn context_id(&self) -> ObjectId {
330        self.context_id
331    }
332
333    pub const fn in_flight(&self) -> u32 {
334        self.in_flight
335    }
336
337    pub const fn release_state(&self) -> ReleaseState {
338        self.resource.release
339    }
340}
341
342#[derive(Debug)]
343pub struct EventRecord<E> {
344    resource: ResourceSlot<E>,
345    context_id: ObjectId,
346    queue_id: ObjectId,
347    program_id: ObjectId,
348    sequence_program_ids: Option<Vec<ObjectId>>,
349    buffer_ids: Vec<ObjectId>,
350}
351
352impl<E> EventRecord<E> {
353    pub fn resource(&self) -> Result<&E, DeviceStateError> {
354        self.resource.get()
355    }
356
357    pub const fn context_id(&self) -> ObjectId {
358        self.context_id
359    }
360
361    pub const fn queue_id(&self) -> ObjectId {
362        self.queue_id
363    }
364
365    /// First program in submission order. Every event has at least one.
366    pub const fn program_id(&self) -> ObjectId {
367        self.program_id
368    }
369
370    /// Programs retained by this event, in provider execution order.
371    pub fn program_ids(&self) -> &[ObjectId] {
372        self.sequence_program_ids
373            .as_deref()
374            .unwrap_or(core::slice::from_ref(&self.program_id))
375    }
376
377    pub fn buffer_ids(&self) -> &[ObjectId] {
378        &self.buffer_ids
379    }
380
381    pub const fn release_state(&self) -> ReleaseState {
382        self.resource.release
383    }
384}
385/// Validated provider resources for one event-producing submission.
386pub struct SubmissionResources<'a, B, P, Q> {
387    context_id: ObjectId,
388    queue: &'a Q,
389    program: &'a P,
390    buffers: &'a ObjectTable<BufferRecord<B>>,
391    buffer_ids: &'a [ObjectId],
392}
393
394impl<B, P, Q> core::fmt::Debug for SubmissionResources<'_, B, P, Q> {
395    fn fmt(&self, formatter: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
396        formatter
397            .debug_struct("SubmissionResources")
398            .field("context_id", &self.context_id)
399            .field("buffer_ids", &self.buffer_ids)
400            .finish_non_exhaustive()
401    }
402}
403
404impl<'a, B, P, Q> SubmissionResources<'a, B, P, Q> {
405    pub const fn context_id(&self) -> ObjectId {
406        self.context_id
407    }
408
409    pub const fn queue(&self) -> &'a Q {
410        self.queue
411    }
412
413    pub const fn program(&self) -> &'a P {
414        self.program
415    }
416
417    /// Retained buffer IDs sorted by raw object ID, including duplicate bindings.
418    pub fn buffer_ids(&self) -> &[ObjectId] {
419        self.buffer_ids
420    }
421
422    pub fn buffer_by_id(&self, id: ObjectId) -> Result<&'a B, DeviceStateError> {
423        self.buffer_with_info_by_id(id).map(|(buffer, _)| buffer)
424    }
425
426    pub(crate) fn buffer_with_info_by_id(
427        &self,
428        id: ObjectId,
429    ) -> Result<(&'a B, BufferInfo), DeviceStateError> {
430        self.buffer_ids
431            .binary_search_by_key(&id.get(), |candidate| candidate.get())
432            .map_err(|_| DeviceStateError::InvalidArgument)?;
433        let record = self.buffers.get(id).map_err(map_table_error)?;
434        Ok((record.resource()?, record.info()))
435    }
436}
437
438/// Validated provider resources for one ordered, event-producing submission.
439/// Program order is preserved; retained buffer IDs are sorted by object ID.
440pub struct SubmissionSequenceResources<'a, B, P, Q> {
441    context_id: ObjectId,
442    queue: &'a Q,
443    programs: &'a ObjectTable<ProgramRecord<P>>,
444    program_ids: &'a [ObjectId],
445    buffers: &'a ObjectTable<BufferRecord<B>>,
446    buffer_ids: &'a [ObjectId],
447}
448
449impl<B, P, Q> core::fmt::Debug for SubmissionSequenceResources<'_, B, P, Q> {
450    fn fmt(&self, formatter: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
451        formatter
452            .debug_struct("SubmissionSequenceResources")
453            .field("context_id", &self.context_id)
454            .field("program_ids", &self.program_ids)
455            .field("buffer_ids", &self.buffer_ids)
456            .finish_non_exhaustive()
457    }
458}
459
460impl<'a, B, P, Q> SubmissionSequenceResources<'a, B, P, Q> {
461    pub const fn context_id(&self) -> ObjectId {
462        self.context_id
463    }
464
465    pub const fn queue(&self) -> &'a Q {
466        self.queue
467    }
468
469    pub fn program_ids(&self) -> &[ObjectId] {
470        self.program_ids
471    }
472
473    pub fn program(&self, index: usize) -> Result<&'a P, DeviceStateError> {
474        let id = *self
475            .program_ids
476            .get(index)
477            .ok_or(DeviceStateError::InvalidArgument)?;
478        self.programs.get(id).map_err(map_table_error)?.resource()
479    }
480
481    pub fn buffer_ids(&self) -> &[ObjectId] {
482        self.buffer_ids
483    }
484
485    pub fn buffer_by_id(&self, id: ObjectId) -> Result<&'a B, DeviceStateError> {
486        self.buffer_ids
487            .binary_search_by_key(&id.get(), |candidate| candidate.get())
488            .map_err(|_| DeviceStateError::InvalidArgument)?;
489        self.buffers.get(id).map_err(map_table_error)?.resource()
490    }
491}
492
493/// Complete typed object graph for one device instance.
494pub struct DeviceState<C, B, P, Q, E> {
495    namespace: ObjectNamespace,
496    limits: DeviceLimits,
497    policy: ResourcePolicy,
498    retained: RetainedBytes,
499    contexts: ObjectTable<ContextRecord<C>>,
500    buffers: ObjectTable<BufferRecord<B>>,
501    programs: ObjectTable<ProgramRecord<P>>,
502    queues: ObjectTable<QueueRecord<Q>>,
503    events: ObjectTable<EventRecord<E>>,
504}
505
506impl<C, B, P, Q, E> DeviceState<C, B, P, Q, E> {
507    pub fn new(
508        namespace: ObjectNamespace,
509        limits: DeviceLimits,
510        policy: ResourcePolicy,
511    ) -> Result<Self, DeviceStateConfigError> {
512        if limits.max_contexts == 0
513            || limits.max_buffers_per_context == 0
514            || limits.max_programs_per_context == 0
515            || limits.max_queues_per_context == 0
516            || limits.max_events_per_context == 0
517            || limits.max_buffer_bytes == 0
518            || limits.max_artifact_bytes == 0
519        {
520            return Err(DeviceStateConfigError::ZeroLimit);
521        }
522        if !(1..=HARD_MAX_BINDINGS).contains(&limits.max_bindings_per_submission) {
523            return Err(DeviceStateConfigError::BindingLimit);
524        }
525        limits
526            .max_events_per_context
527            .checked_mul(limits.max_bindings_per_submission)
528            .ok_or(DeviceStateConfigError::ReferenceCountOverflow)?;
529
530        let buffers = aggregate_slots(limits.max_contexts, limits.max_buffers_per_context)?;
531        let programs = aggregate_slots(limits.max_contexts, limits.max_programs_per_context)?;
532        let queues = aggregate_slots(limits.max_contexts, limits.max_queues_per_context)?;
533        let events = aggregate_slots(limits.max_contexts, limits.max_events_per_context)?;
534
535        Ok(Self {
536            namespace,
537            limits,
538            policy,
539            retained: RetainedBytes::default(),
540            contexts: ObjectTable::with_namespace(
541                ObjectKind::Context,
542                limits.max_contexts,
543                namespace,
544            ),
545            buffers: ObjectTable::with_namespace(ObjectKind::Buffer, buffers, namespace),
546            programs: ObjectTable::with_namespace(ObjectKind::Program, programs, namespace),
547            queues: ObjectTable::with_namespace(ObjectKind::Queue, queues, namespace),
548            events: ObjectTable::with_namespace(ObjectKind::Event, events, namespace),
549        })
550    }
551
552    pub const fn limits(&self) -> DeviceLimits {
553        self.limits
554    }
555
556    pub const fn resource_policy(&self) -> ResourcePolicy {
557        self.policy
558    }
559
560    /// Bulk bytes still represented by live or releasing provider handles.
561    pub const fn retained_bytes(&self) -> RetainedBytes {
562        self.retained
563    }
564
565    pub const fn namespace(&self) -> ObjectNamespace {
566        self.namespace
567    }
568
569    pub const fn resource_counts(&self) -> ResourceCounts {
570        ResourceCounts {
571            contexts: self.context_count() as u64,
572            buffers: self.buffer_count() as u64,
573            programs: self.program_count() as u64,
574            queues: self.queue_count() as u64,
575            events: self.event_count() as u64,
576        }
577    }
578
579    pub const fn is_empty(&self) -> bool {
580        self.contexts.is_empty()
581            && self.buffers.is_empty()
582            && self.programs.is_empty()
583            && self.queues.is_empty()
584            && self.events.is_empty()
585    }
586
587    pub const fn context_count(&self) -> u32 {
588        self.contexts.len()
589    }
590
591    pub const fn buffer_count(&self) -> u32 {
592        self.buffers.len()
593    }
594
595    pub const fn program_count(&self) -> u32 {
596        self.programs.len()
597    }
598
599    pub const fn queue_count(&self) -> u32 {
600        self.queues.len()
601    }
602
603    pub const fn event_count(&self) -> u32 {
604        self.events.len()
605    }
606
607    /// Occupied context identities, including records awaiting release commit.
608    /// Iteration allocates nothing and does not authorize release or bypass child checks.
609    pub fn context_ids(&self) -> impl Iterator<Item = ObjectId> + '_ {
610        self.contexts.ids()
611    }
612
613    /// Occupied buffer identities, including records awaiting release commit.
614    /// Iteration allocates nothing and does not bypass in-flight reference checks.
615    pub fn buffer_ids(&self) -> impl Iterator<Item = ObjectId> + '_ {
616        self.buffers.ids()
617    }
618
619    /// Occupied program identities, including records awaiting release commit.
620    pub fn program_ids(&self) -> impl Iterator<Item = ObjectId> + '_ {
621        self.programs.ids()
622    }
623
624    /// Occupied queue identities, including records awaiting release commit.
625    pub fn queue_ids(&self) -> impl Iterator<Item = ObjectId> + '_ {
626        self.queues.ids()
627    }
628
629    /// Occupied event identities, including records awaiting release commit.
630    /// An enumerated event may still be pending; enumeration is not completion evidence.
631    pub fn event_ids(&self) -> impl Iterator<Item = ObjectId> + '_ {
632        self.events.ids()
633    }
634
635    pub fn context_record(&self, id: ObjectId) -> Result<&ContextRecord<C>, DeviceStateError> {
636        self.contexts.get(id).map_err(map_table_error)
637    }
638
639    pub fn buffer_record(&self, id: ObjectId) -> Result<&BufferRecord<B>, DeviceStateError> {
640        self.buffers.get(id).map_err(map_table_error)
641    }
642
643    pub fn buffer_record_mut(
644        &mut self,
645        id: ObjectId,
646    ) -> Result<&mut BufferRecord<B>, DeviceStateError> {
647        self.buffers.get_mut(id).map_err(map_table_error)
648    }
649
650    pub fn program_record(&self, id: ObjectId) -> Result<&ProgramRecord<P>, DeviceStateError> {
651        self.programs.get(id).map_err(map_table_error)
652    }
653
654    pub fn queue_record(&self, id: ObjectId) -> Result<&QueueRecord<Q>, DeviceStateError> {
655        self.queues.get(id).map_err(map_table_error)
656    }
657
658    pub fn event_record(&self, id: ObjectId) -> Result<&EventRecord<E>, DeviceStateError> {
659        self.events.get(id).map_err(map_table_error)
660    }
661
662    pub(crate) fn next_context_id(&self, start: usize) -> Option<(usize, ObjectId)> {
663        self.contexts.next_id_from(start)
664    }
665
666    pub(crate) fn next_buffer_id(&self, start: usize) -> Option<(usize, ObjectId)> {
667        self.buffers.next_id_from(start)
668    }
669
670    pub(crate) fn next_program_id(&self, start: usize) -> Option<(usize, ObjectId)> {
671        self.programs.next_id_from(start)
672    }
673
674    pub(crate) fn next_queue_id(&self, start: usize) -> Option<(usize, ObjectId)> {
675        self.queues.next_id_from(start)
676    }
677
678    pub(crate) fn next_event_id(&self, start: usize) -> Option<(usize, ObjectId)> {
679        self.events.next_id_from(start)
680    }
681
682    pub fn create_context_with<ProviderError>(
683        &mut self,
684        create: impl FnOnce() -> Result<C, ProviderError>,
685    ) -> Result<ObjectId, CreateError<ProviderError>> {
686        if self.contexts.len() >= self.limits.max_contexts {
687            return Err(CreateError::State(DeviceStateError::ResourceLimit));
688        }
689        self.contexts
690            .try_reserve_insert()
691            .map_err(|error| CreateError::State(map_table_error(error)))?;
692        let resource = create().map_err(CreateError::Provider)?;
693        Ok(self.contexts.insert_prepared(ContextRecord {
694            resource: ResourceSlot::new(resource),
695            children: ChildCounts::default(),
696        }))
697    }
698
699    pub fn create_buffer_with<ProviderError>(
700        &mut self,
701        context_id: ObjectId,
702        desc: BufferDesc,
703        create: impl FnOnce(&C, BufferDesc) -> Result<(B, BufferInfo), ProviderError>,
704    ) -> Result<BufferCreateOutcome, CreateError<ProviderError>> {
705        if desc.bytes() > self.limits.max_buffer_bytes {
706            return Err(CreateError::State(DeviceStateError::ResourceLimit));
707        }
708        let minimum_retained = self
709            .retained
710            .buffer_backing
711            .saturating_add(u128::from(desc.bytes()));
712        if minimum_retained > u128::from(self.policy.max_buffer_backing_bytes()) {
713            return Err(CreateError::State(DeviceStateError::ResourceLimit));
714        }
715        self.check_child_admission(
716            context_id,
717            ChildKind::Buffer,
718            self.limits.max_buffers_per_context,
719        )
720        .map_err(CreateError::State)?;
721        self.buffers
722            .try_reserve_insert()
723            .map_err(|error| CreateError::State(map_table_error(error)))?;
724
725        let context = self
726            .contexts
727            .get_mut(context_id)
728            .map_err(|error| CreateError::State(map_table_error(error)))?;
729        let (resource, info) = create(context.resource()?, desc).map_err(CreateError::Provider)?;
730        let retained = self
731            .retained
732            .buffer_backing
733            .saturating_add(u128::from(info.allocation_bytes()));
734        let cleanup_error = if info.desc() != desc {
735            Some(BackendError::Incompatible)
736        } else if retained > u128::from(self.policy.max_buffer_backing_bytes()) {
737            Some(BackendError::ResourceLimit)
738        } else {
739            None
740        };
741        let id = self.buffers.insert_prepared(BufferRecord {
742            resource: ResourceSlot::new(resource),
743            context_id,
744            info,
745            in_flight: 0,
746        });
747        self.retained.buffer_backing = retained;
748        context.children.buffers += 1;
749        Ok(match cleanup_error {
750            Some(error) => BufferCreateOutcome::CleanupRequired { id, error },
751            None => BufferCreateOutcome::Admitted(id),
752        })
753    }
754
755    pub fn create_program_with<ProviderError>(
756        &mut self,
757        context_id: ObjectId,
758        artifact_bytes: u64,
759        resident_bytes: u64,
760        create: impl FnOnce(&C) -> Result<P, ProviderError>,
761    ) -> Result<ObjectId, CreateError<ProviderError>> {
762        if artifact_bytes == 0 || resident_bytes == 0 {
763            return Err(CreateError::State(DeviceStateError::InvalidArgument));
764        }
765        if artifact_bytes > self.limits.max_artifact_bytes {
766            return Err(CreateError::State(DeviceStateError::ResourceLimit));
767        }
768        let retained = self
769            .retained
770            .program_resident
771            .saturating_add(u128::from(resident_bytes));
772        if retained > u128::from(self.policy.max_program_resident_bytes()) {
773            return Err(CreateError::State(DeviceStateError::ResourceLimit));
774        }
775        self.check_child_admission(
776            context_id,
777            ChildKind::Program,
778            self.limits.max_programs_per_context,
779        )
780        .map_err(CreateError::State)?;
781        self.programs
782            .try_reserve_insert()
783            .map_err(|error| CreateError::State(map_table_error(error)))?;
784
785        let context = self
786            .contexts
787            .get_mut(context_id)
788            .map_err(|error| CreateError::State(map_table_error(error)))?;
789        let resource = create(context.resource()?).map_err(CreateError::Provider)?;
790        let id = self.programs.insert_prepared(ProgramRecord {
791            resource: ResourceSlot::new(resource),
792            context_id,
793            resident_bytes,
794            in_flight: 0,
795        });
796        self.retained.program_resident = retained;
797        context.children.programs += 1;
798        Ok(id)
799    }
800
801    pub fn create_queue_with<ProviderError>(
802        &mut self,
803        context_id: ObjectId,
804        create: impl FnOnce(&C) -> Result<Q, ProviderError>,
805    ) -> Result<ObjectId, CreateError<ProviderError>> {
806        self.check_child_admission(
807            context_id,
808            ChildKind::Queue,
809            self.limits.max_queues_per_context,
810        )
811        .map_err(CreateError::State)?;
812        self.queues
813            .try_reserve_insert()
814            .map_err(|error| CreateError::State(map_table_error(error)))?;
815
816        let context = self
817            .contexts
818            .get_mut(context_id)
819            .map_err(|error| CreateError::State(map_table_error(error)))?;
820        let resource = create(context.resource()?).map_err(CreateError::Provider)?;
821        let id = self.queues.insert_prepared(QueueRecord {
822            resource: ResourceSlot::new(resource),
823            context_id,
824            in_flight: 0,
825        });
826        context.children.queues += 1;
827        Ok(id)
828    }
829
830    pub fn create_event_with<ProviderError>(
831        &mut self,
832        queue_id: ObjectId,
833        program_id: ObjectId,
834        mut buffer_ids: Vec<ObjectId>,
835        create: impl FnOnce(SubmissionResources<'_, B, P, Q>) -> Result<E, ProviderError>,
836    ) -> Result<ObjectId, CreateError<ProviderError>> {
837        if buffer_ids.is_empty() {
838            return Err(CreateError::State(DeviceStateError::InvalidArgument));
839        }
840        if buffer_ids.len() > self.limits.max_bindings_per_submission as usize {
841            return Err(CreateError::State(DeviceStateError::ResourceLimit));
842        }
843        let context_id = self
844            .validate_submission(queue_id, program_id, &buffer_ids)
845            .map_err(CreateError::State)?;
846        let context = self
847            .contexts
848            .get(context_id)
849            .map_err(|error| CreateError::State(map_table_error(error)))?;
850        if context.children.events >= self.limits.max_events_per_context {
851            return Err(CreateError::State(DeviceStateError::ResourceLimit));
852        }
853        self.events
854            .try_reserve_insert()
855            .map_err(|error| CreateError::State(map_table_error(error)))?;
856
857        buffer_ids.sort_unstable_by_key(|id| id.get());
858        self.check_reference_increments(queue_id, core::slice::from_ref(&program_id), &buffer_ids)
859            .map_err(CreateError::State)?;
860        self.increment_event_references(
861            context_id,
862            queue_id,
863            core::slice::from_ref(&program_id),
864            &buffer_ids,
865        )
866        .map_err(CreateError::State)?;
867
868        let event_result = {
869            let queue = self
870                .queues
871                .get(queue_id)
872                .map_err(|error| CreateError::State(map_table_error(error)))?
873                .resource()?;
874            let program = self
875                .programs
876                .get(program_id)
877                .map_err(|error| CreateError::State(map_table_error(error)))?
878                .resource()?;
879            create(SubmissionResources {
880                context_id,
881                queue,
882                program,
883                buffers: &self.buffers,
884                buffer_ids: &buffer_ids,
885            })
886        };
887        let event = match event_result {
888            Ok(event) => event,
889            Err(error) => {
890                self.decrement_event_references(
891                    context_id,
892                    queue_id,
893                    core::slice::from_ref(&program_id),
894                    &buffer_ids,
895                )
896                .map_err(CreateError::State)?;
897                return Err(CreateError::Provider(error));
898            }
899        };
900
901        let id = self.events.insert_prepared(EventRecord {
902            resource: ResourceSlot::new(event),
903            context_id,
904            queue_id,
905            program_id,
906            sequence_program_ids: None,
907            buffer_ids,
908        });
909        Ok(id)
910    }
911
912    /// Admit one event that retains an ordered sequence of programs.
913    ///
914    /// This is a provider-side lifecycle primitive, not a wire operation. The
915    /// caller chooses a bounded encoding and validates each invocation's
916    /// bindings before entering this method. All programs, the queue and every
917    /// buffer must belong to one context; none can be released until the event
918    /// is committed.
919    pub fn create_event_sequence_with<ProviderError>(
920        &mut self,
921        queue_id: ObjectId,
922        program_ids: Vec<ObjectId>,
923        mut buffer_ids: Vec<ObjectId>,
924        create: impl FnOnce(SubmissionSequenceResources<'_, B, P, Q>) -> Result<E, ProviderError>,
925    ) -> Result<ObjectId, CreateError<ProviderError>> {
926        if program_ids.is_empty() || buffer_ids.is_empty() {
927            return Err(CreateError::State(DeviceStateError::InvalidArgument));
928        }
929        let context_id = self
930            .validate_submission_sequence(queue_id, &program_ids, &buffer_ids)
931            .map_err(CreateError::State)?;
932        let context = self
933            .contexts
934            .get(context_id)
935            .map_err(|error| CreateError::State(map_table_error(error)))?;
936        if context.children.events >= self.limits.max_events_per_context {
937            return Err(CreateError::State(DeviceStateError::ResourceLimit));
938        }
939        self.events
940            .try_reserve_insert()
941            .map_err(|error| CreateError::State(map_table_error(error)))?;
942
943        buffer_ids.sort_unstable_by_key(|id| id.get());
944        self.check_reference_increments(queue_id, &program_ids, &buffer_ids)
945            .map_err(CreateError::State)?;
946        self.increment_event_references(context_id, queue_id, &program_ids, &buffer_ids)
947            .map_err(CreateError::State)?;
948
949        let event_result = {
950            let queue = self
951                .queues
952                .get(queue_id)
953                .map_err(|error| CreateError::State(map_table_error(error)))?
954                .resource()?;
955            create(SubmissionSequenceResources {
956                context_id,
957                queue,
958                programs: &self.programs,
959                program_ids: &program_ids,
960                buffers: &self.buffers,
961                buffer_ids: &buffer_ids,
962            })
963        };
964        let event = match event_result {
965            Ok(event) => event,
966            Err(error) => {
967                self.decrement_event_references(context_id, queue_id, &program_ids, &buffer_ids)
968                    .map_err(CreateError::State)?;
969                return Err(CreateError::Provider(error));
970            }
971        };
972
973        let id = self.events.insert_prepared(EventRecord {
974            resource: ResourceSlot::new(event),
975            context_id,
976            queue_id,
977            program_id: program_ids[0],
978            sequence_program_ids: Some(program_ids),
979            buffer_ids,
980        });
981        Ok(id)
982    }
983
984    pub fn begin_context_release(&mut self, id: ObjectId) -> Result<C, DeviceStateError> {
985        let record = self.contexts.get_mut(id).map_err(map_table_error)?;
986        if !record.children.is_empty() {
987            return Err(DeviceStateError::Busy);
988        }
989        record.resource.begin_release()
990    }
991
992    pub fn restore_context_release(
993        &mut self,
994        id: ObjectId,
995        resource: C,
996    ) -> Result<(), RestoreError<C>> {
997        restore_resource(&mut self.contexts, id, resource)
998    }
999
1000    pub fn commit_context_release(&mut self, id: ObjectId) -> Result<(), DeviceStateError> {
1001        ensure_releasing(self.contexts.get(id).map_err(map_table_error)?)?;
1002        self.contexts.remove(id).map_err(map_table_error)?;
1003        Ok(())
1004    }
1005
1006    pub fn begin_buffer_release(&mut self, id: ObjectId) -> Result<B, DeviceStateError> {
1007        let record = self.buffers.get_mut(id).map_err(map_table_error)?;
1008        if record.in_flight != 0 {
1009            return Err(DeviceStateError::Busy);
1010        }
1011        record.resource.begin_release()
1012    }
1013
1014    pub fn restore_buffer_release(
1015        &mut self,
1016        id: ObjectId,
1017        resource: B,
1018    ) -> Result<(), RestoreError<B>> {
1019        restore_resource(&mut self.buffers, id, resource)
1020    }
1021
1022    pub fn commit_buffer_release(&mut self, id: ObjectId) -> Result<(), DeviceStateError> {
1023        let (context_id, allocation_bytes) = {
1024            let record = self.buffers.get(id).map_err(map_table_error)?;
1025            ensure_releasing(record)?;
1026            (record.context_id, record.info.allocation_bytes())
1027        };
1028        self.buffers.remove(id).map_err(map_table_error)?;
1029        self.retained.buffer_backing -= u128::from(allocation_bytes);
1030        self.contexts
1031            .get_mut(context_id)
1032            .map_err(map_table_error)?
1033            .children
1034            .buffers -= 1;
1035        Ok(())
1036    }
1037
1038    pub fn begin_program_release(&mut self, id: ObjectId) -> Result<P, DeviceStateError> {
1039        let record = self.programs.get_mut(id).map_err(map_table_error)?;
1040        if record.in_flight != 0 {
1041            return Err(DeviceStateError::Busy);
1042        }
1043        record.resource.begin_release()
1044    }
1045
1046    pub fn restore_program_release(
1047        &mut self,
1048        id: ObjectId,
1049        resource: P,
1050    ) -> Result<(), RestoreError<P>> {
1051        restore_resource(&mut self.programs, id, resource)
1052    }
1053
1054    pub fn commit_program_release(&mut self, id: ObjectId) -> Result<(), DeviceStateError> {
1055        let (context_id, resident_bytes) = {
1056            let record = self.programs.get(id).map_err(map_table_error)?;
1057            ensure_releasing(record)?;
1058            (record.context_id, record.resident_bytes)
1059        };
1060        self.programs.remove(id).map_err(map_table_error)?;
1061        self.retained.program_resident -= u128::from(resident_bytes);
1062        self.contexts
1063            .get_mut(context_id)
1064            .map_err(map_table_error)?
1065            .children
1066            .programs -= 1;
1067        Ok(())
1068    }
1069
1070    pub fn begin_queue_release(&mut self, id: ObjectId) -> Result<Q, DeviceStateError> {
1071        let record = self.queues.get_mut(id).map_err(map_table_error)?;
1072        if record.in_flight != 0 {
1073            return Err(DeviceStateError::Busy);
1074        }
1075        record.resource.begin_release()
1076    }
1077
1078    pub fn restore_queue_release(
1079        &mut self,
1080        id: ObjectId,
1081        resource: Q,
1082    ) -> Result<(), RestoreError<Q>> {
1083        restore_resource(&mut self.queues, id, resource)
1084    }
1085
1086    pub fn commit_queue_release(&mut self, id: ObjectId) -> Result<(), DeviceStateError> {
1087        let context_id = {
1088            let record = self.queues.get(id).map_err(map_table_error)?;
1089            ensure_releasing(record)?;
1090            record.context_id
1091        };
1092        self.queues.remove(id).map_err(map_table_error)?;
1093        self.contexts
1094            .get_mut(context_id)
1095            .map_err(map_table_error)?
1096            .children
1097            .queues -= 1;
1098        Ok(())
1099    }
1100
1101    pub fn begin_event_release(&mut self, id: ObjectId) -> Result<E, DeviceStateError> {
1102        self.events
1103            .get_mut(id)
1104            .map_err(map_table_error)?
1105            .resource
1106            .begin_release()
1107    }
1108
1109    pub fn restore_event_release(
1110        &mut self,
1111        id: ObjectId,
1112        resource: E,
1113    ) -> Result<(), RestoreError<E>> {
1114        restore_resource(&mut self.events, id, resource)
1115    }
1116
1117    pub fn commit_event_release(&mut self, id: ObjectId) -> Result<(), DeviceStateError> {
1118        {
1119            let record = self.events.get(id).map_err(map_table_error)?;
1120            ensure_releasing(record)?;
1121            self.validate_submission_sequence(
1122                record.queue_id,
1123                record.program_ids(),
1124                &record.buffer_ids,
1125            )?;
1126        }
1127        let record = self.events.remove(id).map_err(map_table_error)?;
1128        self.decrement_event_references(
1129            record.context_id,
1130            record.queue_id,
1131            record.program_ids(),
1132            &record.buffer_ids,
1133        )
1134    }
1135
1136    fn check_child_admission(
1137        &self,
1138        context_id: ObjectId,
1139        child: ChildKind,
1140        limit: u32,
1141    ) -> Result<(), DeviceStateError> {
1142        let context = self.contexts.get(context_id).map_err(map_table_error)?;
1143        context.resource()?;
1144        let count = match child {
1145            ChildKind::Buffer => context.children.buffers,
1146            ChildKind::Program => context.children.programs,
1147            ChildKind::Queue => context.children.queues,
1148        };
1149        if count >= limit {
1150            return Err(DeviceStateError::ResourceLimit);
1151        }
1152        Ok(())
1153    }
1154
1155    fn validate_submission(
1156        &self,
1157        queue_id: ObjectId,
1158        program_id: ObjectId,
1159        buffer_ids: &[ObjectId],
1160    ) -> Result<ObjectId, DeviceStateError> {
1161        let queue = self.queues.get(queue_id).map_err(map_table_error)?;
1162        queue.resource()?;
1163        let program = self.programs.get(program_id).map_err(map_table_error)?;
1164        program.resource()?;
1165        if program.context_id != queue.context_id {
1166            return Err(DeviceStateError::ContextMismatch);
1167        }
1168        for buffer_id in buffer_ids {
1169            let buffer = self.buffers.get(*buffer_id).map_err(map_table_error)?;
1170            buffer.resource()?;
1171            if buffer.context_id != queue.context_id {
1172                return Err(DeviceStateError::ContextMismatch);
1173            }
1174        }
1175        Ok(queue.context_id)
1176    }
1177
1178    fn validate_submission_sequence(
1179        &self,
1180        queue_id: ObjectId,
1181        program_ids: &[ObjectId],
1182        buffer_ids: &[ObjectId],
1183    ) -> Result<ObjectId, DeviceStateError> {
1184        let first = *program_ids
1185            .first()
1186            .ok_or(DeviceStateError::InvalidArgument)?;
1187        let context_id = self.validate_submission(queue_id, first, buffer_ids)?;
1188        for program_id in &program_ids[1..] {
1189            let program = self.programs.get(*program_id).map_err(map_table_error)?;
1190            program.resource()?;
1191            if program.context_id != context_id {
1192                return Err(DeviceStateError::ContextMismatch);
1193            }
1194        }
1195        Ok(context_id)
1196    }
1197
1198    fn check_reference_increments(
1199        &self,
1200        queue_id: ObjectId,
1201        program_ids: &[ObjectId],
1202        sorted_buffer_ids: &[ObjectId],
1203    ) -> Result<(), DeviceStateError> {
1204        self.queues
1205            .get(queue_id)
1206            .map_err(map_table_error)?
1207            .in_flight
1208            .checked_add(1)
1209            .ok_or(DeviceStateError::ReferenceCountOverflow)?;
1210        for (index, id) in program_ids.iter().copied().enumerate() {
1211            if program_ids[..index].contains(&id) {
1212                continue;
1213            }
1214            let count = u32::try_from(program_ids[index..].iter().filter(|&&p| p == id).count())
1215                .map_err(|_| DeviceStateError::ReferenceCountOverflow)?;
1216            self.programs
1217                .get(id)
1218                .map_err(map_table_error)?
1219                .in_flight
1220                .checked_add(count)
1221                .ok_or(DeviceStateError::ReferenceCountOverflow)?;
1222        }
1223
1224        let mut index = 0;
1225        while index < sorted_buffer_ids.len() {
1226            let id = sorted_buffer_ids[index];
1227            let mut end = index + 1;
1228            while end < sorted_buffer_ids.len() && sorted_buffer_ids[end] == id {
1229                end += 1;
1230            }
1231            let count =
1232                u32::try_from(end - index).map_err(|_| DeviceStateError::ReferenceCountOverflow)?;
1233            self.buffers
1234                .get(id)
1235                .map_err(map_table_error)?
1236                .in_flight
1237                .checked_add(count)
1238                .ok_or(DeviceStateError::ReferenceCountOverflow)?;
1239            index = end;
1240        }
1241        Ok(())
1242    }
1243
1244    fn increment_event_references(
1245        &mut self,
1246        context_id: ObjectId,
1247        queue_id: ObjectId,
1248        program_ids: &[ObjectId],
1249        buffer_ids: &[ObjectId],
1250    ) -> Result<(), DeviceStateError> {
1251        self.queues
1252            .get_mut(queue_id)
1253            .map_err(map_table_error)?
1254            .in_flight += 1;
1255        for program_id in program_ids {
1256            self.programs
1257                .get_mut(*program_id)
1258                .map_err(map_table_error)?
1259                .in_flight += 1;
1260        }
1261        for buffer_id in buffer_ids {
1262            self.buffers
1263                .get_mut(*buffer_id)
1264                .map_err(map_table_error)?
1265                .in_flight += 1;
1266        }
1267        self.contexts
1268            .get_mut(context_id)
1269            .map_err(map_table_error)?
1270            .children
1271            .events += 1;
1272        Ok(())
1273    }
1274
1275    fn decrement_event_references(
1276        &mut self,
1277        context_id: ObjectId,
1278        queue_id: ObjectId,
1279        program_ids: &[ObjectId],
1280        buffer_ids: &[ObjectId],
1281    ) -> Result<(), DeviceStateError> {
1282        self.queues
1283            .get_mut(queue_id)
1284            .map_err(map_table_error)?
1285            .in_flight -= 1;
1286        for program_id in program_ids {
1287            self.programs
1288                .get_mut(*program_id)
1289                .map_err(map_table_error)?
1290                .in_flight -= 1;
1291        }
1292        for buffer_id in buffer_ids {
1293            self.buffers
1294                .get_mut(*buffer_id)
1295                .map_err(map_table_error)?
1296                .in_flight -= 1;
1297        }
1298        self.contexts
1299            .get_mut(context_id)
1300            .map_err(map_table_error)?
1301            .children
1302            .events -= 1;
1303        Ok(())
1304    }
1305}
1306
1307#[derive(Clone, Copy)]
1308enum ChildKind {
1309    Buffer,
1310    Program,
1311    Queue,
1312}
1313
1314trait ReleasableRecord {
1315    fn release_state(&self) -> ReleaseState;
1316}
1317
1318impl<C> ReleasableRecord for ContextRecord<C> {
1319    fn release_state(&self) -> ReleaseState {
1320        self.release_state()
1321    }
1322}
1323
1324impl<B> ReleasableRecord for BufferRecord<B> {
1325    fn release_state(&self) -> ReleaseState {
1326        self.release_state()
1327    }
1328}
1329
1330impl<P> ReleasableRecord for ProgramRecord<P> {
1331    fn release_state(&self) -> ReleaseState {
1332        self.release_state()
1333    }
1334}
1335
1336impl<Q> ReleasableRecord for QueueRecord<Q> {
1337    fn release_state(&self) -> ReleaseState {
1338        self.release_state()
1339    }
1340}
1341
1342impl<E> ReleasableRecord for EventRecord<E> {
1343    fn release_state(&self) -> ReleaseState {
1344        self.release_state()
1345    }
1346}
1347
1348fn ensure_releasing(record: &impl ReleasableRecord) -> Result<(), DeviceStateError> {
1349    if record.release_state() != ReleaseState::Releasing {
1350        return Err(DeviceStateError::InvalidTransition);
1351    }
1352    Ok(())
1353}
1354
1355trait ResourceRecord<R> {
1356    fn resource_mut(&mut self) -> &mut ResourceSlot<R>;
1357}
1358
1359impl<C> ResourceRecord<C> for ContextRecord<C> {
1360    fn resource_mut(&mut self) -> &mut ResourceSlot<C> {
1361        &mut self.resource
1362    }
1363}
1364
1365impl<B> ResourceRecord<B> for BufferRecord<B> {
1366    fn resource_mut(&mut self) -> &mut ResourceSlot<B> {
1367        &mut self.resource
1368    }
1369}
1370
1371impl<P> ResourceRecord<P> for ProgramRecord<P> {
1372    fn resource_mut(&mut self) -> &mut ResourceSlot<P> {
1373        &mut self.resource
1374    }
1375}
1376
1377impl<Q> ResourceRecord<Q> for QueueRecord<Q> {
1378    fn resource_mut(&mut self) -> &mut ResourceSlot<Q> {
1379        &mut self.resource
1380    }
1381}
1382
1383impl<E> ResourceRecord<E> for EventRecord<E> {
1384    fn resource_mut(&mut self) -> &mut ResourceSlot<E> {
1385        &mut self.resource
1386    }
1387}
1388
1389fn restore_resource<R, Record: ResourceRecord<R>>(
1390    table: &mut ObjectTable<Record>,
1391    id: ObjectId,
1392    resource: R,
1393) -> Result<(), RestoreError<R>> {
1394    let record = match table.get_mut(id) {
1395        Ok(record) => record,
1396        Err(error) => {
1397            return Err(RestoreError {
1398                error: map_table_error(error),
1399                resource,
1400            });
1401        }
1402    };
1403    record.resource_mut().restore(resource)
1404}
1405
1406fn aggregate_slots(contexts: u32, per_context: u32) -> Result<u32, DeviceStateConfigError> {
1407    contexts
1408        .checked_mul(per_context)
1409        .ok_or(DeviceStateConfigError::CountOverflow)
1410}
1411
1412fn map_table_error(error: ObjectTableError) -> DeviceStateError {
1413    match error {
1414        ObjectTableError::InvalidId => DeviceStateError::InvalidObject,
1415        ObjectTableError::WrongKind | ObjectTableError::StaleId => DeviceStateError::StaleObject,
1416        ObjectTableError::Full => DeviceStateError::ResourceLimit,
1417        ObjectTableError::AllocationFailed => DeviceStateError::OutOfMemory,
1418    }
1419}
1420
1421#[cfg(test)]
1422mod tests {
1423    use super::*;
1424    use core::cell::Cell;
1425    use virtio_accel_core::{BufferDesc, BufferProperties, BufferUsage, MemoryDomain};
1426
1427    type TestState = DeviceState<u32, u32, u32, u32, u32>;
1428
1429    fn limits(max_contexts: u32) -> DeviceLimits {
1430        DeviceLimits {
1431            max_contexts,
1432            max_buffers_per_context: 1,
1433            max_programs_per_context: 1,
1434            max_queues_per_context: 1,
1435            max_events_per_context: 1,
1436            max_bindings_per_submission: 4,
1437            max_buffer_bytes: 1 << 20,
1438            max_artifact_bytes: 1 << 20,
1439        }
1440    }
1441
1442    fn state(namespace: u16, max_contexts: u32) -> TestState {
1443        DeviceState::new(
1444            ObjectNamespace::new(namespace).unwrap(),
1445            limits(max_contexts),
1446            resource_policy(),
1447        )
1448        .unwrap()
1449    }
1450
1451    fn resource_policy() -> ResourcePolicy {
1452        ResourcePolicy::new(1 << 30, 1 << 30).unwrap()
1453    }
1454
1455    fn admitted(outcome: BufferCreateOutcome) -> ObjectId {
1456        match outcome {
1457            BufferCreateOutcome::Admitted(id) => id,
1458            BufferCreateOutcome::CleanupRequired { .. } => {
1459                panic!("test allocation unexpectedly required cleanup")
1460            }
1461        }
1462    }
1463
1464    fn buffer_desc() -> BufferDesc {
1465        BufferDesc::new(
1466            4096,
1467            64,
1468            MemoryDomain::Shared,
1469            BufferUsage::TRANSFER_SOURCE
1470                | BufferUsage::TRANSFER_DESTINATION
1471                | BufferUsage::PROGRAM_INPUT,
1472        )
1473        .unwrap()
1474    }
1475
1476    fn buffer_info(desc: BufferDesc) -> BufferInfo {
1477        BufferInfo::new(
1478            desc,
1479            4096,
1480            64,
1481            BufferProperties::HOST_VISIBLE | BufferProperties::DIRECT_BINDING,
1482        )
1483        .unwrap()
1484    }
1485
1486    fn create_context(state: &mut TestState, resource: u32) -> ObjectId {
1487        state
1488            .create_context_with(|| Ok::<_, &'static str>(resource))
1489            .unwrap()
1490    }
1491
1492    #[test]
1493    fn complete_lifecycle_tracks_children_references_and_release_rollback() {
1494        let mut state = state(1, 1);
1495        let context = create_context(&mut state, 10);
1496        let buffer = admitted(
1497            state
1498                .create_buffer_with(context, buffer_desc(), |context, desc| {
1499                    assert_eq!(*context, 10);
1500                    Ok::<_, &'static str>((20, buffer_info(desc)))
1501                })
1502                .unwrap(),
1503        );
1504        *state
1505            .buffer_record_mut(buffer)
1506            .unwrap()
1507            .resource_mut()
1508            .unwrap() = 21;
1509        let program = state
1510            .create_program_with(context, 4096, 8192, |context| {
1511                assert_eq!(*context, 10);
1512                Ok::<_, &'static str>(30)
1513            })
1514            .unwrap();
1515        let queue = state
1516            .create_queue_with(context, |context| {
1517                assert_eq!(*context, 10);
1518                Ok::<_, &'static str>(40)
1519            })
1520            .unwrap();
1521        let event = state
1522            .create_event_with(queue, program, alloc::vec![buffer, buffer], |resources| {
1523                assert_eq!(resources.context_id(), context);
1524                assert_eq!(*resources.queue(), 40);
1525                assert_eq!(*resources.program(), 30);
1526                assert_eq!(*resources.buffer_by_id(buffer).unwrap(), 21);
1527                Ok::<_, &'static str>(50)
1528            })
1529            .unwrap();
1530
1531        assert_eq!(
1532            state.context_record(context).unwrap().children(),
1533            ChildCounts {
1534                buffers: 1,
1535                programs: 1,
1536                queues: 1,
1537                events: 1,
1538            }
1539        );
1540        assert_eq!(state.buffer_record(buffer).unwrap().in_flight(), 2);
1541        assert_eq!(state.program_record(program).unwrap().in_flight(), 1);
1542        assert_eq!(state.queue_record(queue).unwrap().in_flight(), 1);
1543        assert_eq!(state.context_ids().collect::<Vec<_>>(), [context]);
1544        assert_eq!(state.buffer_ids().collect::<Vec<_>>(), [buffer]);
1545        assert_eq!(state.program_ids().collect::<Vec<_>>(), [program]);
1546        assert_eq!(state.queue_ids().collect::<Vec<_>>(), [queue]);
1547        assert_eq!(state.event_ids().collect::<Vec<_>>(), [event]);
1548        assert_eq!(
1549            state.begin_context_release(context),
1550            Err(DeviceStateError::Busy)
1551        );
1552        assert_eq!(
1553            state.begin_buffer_release(buffer),
1554            Err(DeviceStateError::Busy)
1555        );
1556        assert_eq!(
1557            state.begin_program_release(program),
1558            Err(DeviceStateError::Busy)
1559        );
1560        assert_eq!(
1561            state.begin_queue_release(queue),
1562            Err(DeviceStateError::Busy)
1563        );
1564
1565        let event_resource = state.begin_event_release(event).unwrap();
1566        assert_eq!(event_resource, 50);
1567        // Releasing records still own references until the release commits.
1568        assert_eq!(state.event_ids().collect::<Vec<_>>(), [event]);
1569        assert_eq!(
1570            state.event_record(event).unwrap().release_state(),
1571            ReleaseState::Releasing
1572        );
1573        state.restore_event_release(event, event_resource).unwrap();
1574        assert_eq!(*state.event_record(event).unwrap().resource().unwrap(), 50);
1575        let event_resource = state.begin_event_release(event).unwrap();
1576        assert_eq!(event_resource, 50);
1577        state.commit_event_release(event).unwrap();
1578        assert_eq!(state.event_ids().next(), None);
1579        assert!(matches!(
1580            state.event_record(event),
1581            Err(DeviceStateError::StaleObject)
1582        ));
1583        assert_eq!(state.buffer_record(buffer).unwrap().in_flight(), 0);
1584        assert_eq!(state.program_record(program).unwrap().in_flight(), 0);
1585        assert_eq!(state.queue_record(queue).unwrap().in_flight(), 0);
1586
1587        let buffer_resource = state.begin_buffer_release(buffer).unwrap();
1588        state
1589            .restore_buffer_release(buffer, buffer_resource)
1590            .unwrap();
1591        let buffer_resource = state.begin_buffer_release(buffer).unwrap();
1592        assert_eq!(buffer_resource, 21);
1593        state.commit_buffer_release(buffer).unwrap();
1594
1595        let program_resource = state.begin_program_release(program).unwrap();
1596        state
1597            .restore_program_release(program, program_resource)
1598            .unwrap();
1599        let program_resource = state.begin_program_release(program).unwrap();
1600        assert_eq!(program_resource, 30);
1601        state.commit_program_release(program).unwrap();
1602
1603        let queue_resource = state.begin_queue_release(queue).unwrap();
1604        state.restore_queue_release(queue, queue_resource).unwrap();
1605        let queue_resource = state.begin_queue_release(queue).unwrap();
1606        assert_eq!(queue_resource, 40);
1607        state.commit_queue_release(queue).unwrap();
1608
1609        assert!(state.context_record(context).unwrap().children().is_empty());
1610        let context_resource = state.begin_context_release(context).unwrap();
1611        state
1612            .restore_context_release(context, context_resource)
1613            .unwrap();
1614        let context_resource = state.begin_context_release(context).unwrap();
1615        assert_eq!(context_resource, 10);
1616        state.commit_context_release(context).unwrap();
1617        assert!(matches!(
1618            state.context_record(context),
1619            Err(DeviceStateError::StaleObject)
1620        ));
1621        assert_eq!(state.context_count(), 0);
1622        assert_eq!(state.buffer_count(), 0);
1623        assert_eq!(state.program_count(), 0);
1624        assert_eq!(state.queue_count(), 0);
1625        assert_eq!(state.event_count(), 0);
1626        assert_eq!(state.context_ids().next(), None);
1627        assert_eq!(state.buffer_ids().next(), None);
1628        assert_eq!(state.program_ids().next(), None);
1629        assert_eq!(state.queue_ids().next(), None);
1630    }
1631
1632    #[test]
1633    fn ordered_event_retains_every_repeated_program_and_buffer_until_commit() {
1634        let mut state = state(1, 1);
1635        let context = create_context(&mut state, 10);
1636        let buffer = admitted(
1637            state
1638                .create_buffer_with(context, buffer_desc(), |_, desc| {
1639                    Ok::<_, &'static str>((20, buffer_info(desc)))
1640                })
1641                .unwrap(),
1642        );
1643        let program = state
1644            .create_program_with(context, 1, 1, |_| Ok::<_, &'static str>(30))
1645            .unwrap();
1646        let queue = state
1647            .create_queue_with(context, |_| Ok::<_, &'static str>(40))
1648            .unwrap();
1649
1650        let event = state
1651            .create_event_sequence_with(
1652                queue,
1653                alloc::vec![program, program],
1654                alloc::vec![buffer, buffer],
1655                |resources| {
1656                    assert_eq!(resources.program_ids(), &[program, program]);
1657                    assert_eq!(*resources.program(0).unwrap(), 30);
1658                    assert_eq!(*resources.program(1).unwrap(), 30);
1659                    assert_eq!(*resources.buffer_by_id(buffer).unwrap(), 20);
1660                    Ok::<_, &'static str>(50)
1661                },
1662            )
1663            .unwrap();
1664        assert_eq!(state.program_record(program).unwrap().in_flight(), 2);
1665        assert_eq!(state.buffer_record(buffer).unwrap().in_flight(), 2);
1666        assert_eq!(state.queue_record(queue).unwrap().in_flight(), 1);
1667        assert_eq!(
1668            state.event_record(event).unwrap().program_ids(),
1669            &[program, program]
1670        );
1671        assert_eq!(state.event_record(event).unwrap().program_id(), program);
1672        assert_eq!(
1673            state.begin_program_release(program),
1674            Err(DeviceStateError::Busy)
1675        );
1676        assert_eq!(
1677            state.begin_buffer_release(buffer),
1678            Err(DeviceStateError::Busy)
1679        );
1680
1681        assert_eq!(state.begin_event_release(event).unwrap(), 50);
1682        state.commit_event_release(event).unwrap();
1683        assert_eq!(state.program_record(program).unwrap().in_flight(), 0);
1684        assert_eq!(state.buffer_record(buffer).unwrap().in_flight(), 0);
1685        assert_eq!(state.queue_record(queue).unwrap().in_flight(), 0);
1686    }
1687
1688    #[test]
1689    fn rejected_ordered_event_restores_every_reference() {
1690        let mut state = state(1, 1);
1691        let context = create_context(&mut state, 10);
1692        let buffer = admitted(
1693            state
1694                .create_buffer_with(context, buffer_desc(), |_, desc| {
1695                    Ok::<_, &'static str>((20, buffer_info(desc)))
1696                })
1697                .unwrap(),
1698        );
1699        let program = state
1700            .create_program_with(context, 1, 1, |_| Ok::<_, &'static str>(30))
1701            .unwrap();
1702        let queue = state
1703            .create_queue_with(context, |_| Ok::<_, &'static str>(40))
1704            .unwrap();
1705
1706        let result = state.create_event_sequence_with(
1707            queue,
1708            alloc::vec![program, program],
1709            alloc::vec![buffer, buffer],
1710            |_| Err::<u32, _>("rejected"),
1711        );
1712        assert!(matches!(result, Err(CreateError::Provider("rejected"))));
1713        assert_eq!(state.event_count(), 0);
1714        assert_eq!(state.program_record(program).unwrap().in_flight(), 0);
1715        assert_eq!(state.buffer_record(buffer).unwrap().in_flight(), 0);
1716        assert_eq!(state.queue_record(queue).unwrap().in_flight(), 0);
1717    }
1718
1719    #[test]
1720    fn quota_exhaustion_never_invokes_provider_callbacks() {
1721        let mut state = state(1, 1);
1722        let calls = Cell::new(0_u32);
1723        let context = state
1724            .create_context_with(|| {
1725                calls.set(calls.get() + 1);
1726                Ok::<_, &'static str>(10)
1727            })
1728            .unwrap();
1729        assert!(matches!(
1730            state.create_context_with(|| {
1731                calls.set(calls.get() + 1);
1732                Ok::<_, &'static str>(11)
1733            }),
1734            Err(CreateError::State(DeviceStateError::ResourceLimit))
1735        ));
1736
1737        let buffer = admitted(
1738            state
1739                .create_buffer_with(context, buffer_desc(), |_, desc| {
1740                    calls.set(calls.get() + 1);
1741                    Ok::<_, &'static str>((20, buffer_info(desc)))
1742                })
1743                .unwrap(),
1744        );
1745        assert!(matches!(
1746            state.create_buffer_with(context, buffer_desc(), |_, desc| {
1747                calls.set(calls.get() + 1);
1748                Ok::<_, &'static str>((21, buffer_info(desc)))
1749            }),
1750            Err(CreateError::State(DeviceStateError::ResourceLimit))
1751        ));
1752
1753        let program = state
1754            .create_program_with(context, 1, 1, |_| {
1755                calls.set(calls.get() + 1);
1756                Ok::<_, &'static str>(30)
1757            })
1758            .unwrap();
1759        assert!(matches!(
1760            state.create_program_with(context, 1, 1, |_| {
1761                calls.set(calls.get() + 1);
1762                Ok::<_, &'static str>(31)
1763            }),
1764            Err(CreateError::State(DeviceStateError::ResourceLimit))
1765        ));
1766
1767        let queue = state
1768            .create_queue_with(context, |_| {
1769                calls.set(calls.get() + 1);
1770                Ok::<_, &'static str>(40)
1771            })
1772            .unwrap();
1773        assert!(matches!(
1774            state.create_queue_with(context, |_| {
1775                calls.set(calls.get() + 1);
1776                Ok::<_, &'static str>(41)
1777            }),
1778            Err(CreateError::State(DeviceStateError::ResourceLimit))
1779        ));
1780
1781        state
1782            .create_event_with(queue, program, alloc::vec![buffer], |_| {
1783                calls.set(calls.get() + 1);
1784                Ok::<_, &'static str>(50)
1785            })
1786            .unwrap();
1787        assert!(matches!(
1788            state.create_event_with(queue, program, alloc::vec![buffer], |_| {
1789                calls.set(calls.get() + 1);
1790                Ok::<_, &'static str>(51)
1791            }),
1792            Err(CreateError::State(DeviceStateError::ResourceLimit))
1793        ));
1794        assert_eq!(calls.get(), 5);
1795        assert_eq!(state.context_count(), 1);
1796        assert_eq!(state.buffer_count(), 1);
1797        assert_eq!(state.program_count(), 1);
1798        assert_eq!(state.queue_count(), 1);
1799        assert_eq!(state.event_count(), 1);
1800    }
1801
1802    #[test]
1803    fn provider_rejection_rolls_back_every_creation_path() {
1804        let mut state = state(1, 1);
1805        assert!(matches!(
1806            state.create_context_with(|| Err::<u32, _>("context")),
1807            Err(CreateError::Provider("context"))
1808        ));
1809        assert_eq!(state.context_count(), 0);
1810
1811        let context = create_context(&mut state, 10);
1812        assert!(matches!(
1813            state.create_buffer_with(context, buffer_desc(), |_, _| {
1814                Err::<(u32, BufferInfo), _>("buffer")
1815            }),
1816            Err(CreateError::Provider("buffer"))
1817        ));
1818        assert!(matches!(
1819            state.create_program_with(context, 1, 1, |_| Err::<u32, _>("program")),
1820            Err(CreateError::Provider("program"))
1821        ));
1822        assert!(matches!(
1823            state.create_queue_with(context, |_| Err::<u32, _>("queue")),
1824            Err(CreateError::Provider("queue"))
1825        ));
1826        assert_eq!(
1827            state.context_record(context).unwrap().children(),
1828            ChildCounts::default()
1829        );
1830
1831        let buffer = admitted(
1832            state
1833                .create_buffer_with(context, buffer_desc(), |_, desc| {
1834                    Ok::<_, &'static str>((20, buffer_info(desc)))
1835                })
1836                .unwrap(),
1837        );
1838        let program = state
1839            .create_program_with(context, 1, 1, |_| Ok::<_, &'static str>(30))
1840            .unwrap();
1841        let queue = state
1842            .create_queue_with(context, |_| Ok::<_, &'static str>(40))
1843            .unwrap();
1844        assert!(matches!(
1845            state.create_event_with(queue, program, alloc::vec![buffer], |_| {
1846                Err::<u32, _>("event")
1847            }),
1848            Err(CreateError::Provider("event"))
1849        ));
1850        assert_eq!(state.event_count(), 0);
1851        assert_eq!(state.buffer_record(buffer).unwrap().in_flight(), 0);
1852        assert_eq!(state.program_record(program).unwrap().in_flight(), 0);
1853        assert_eq!(state.queue_record(queue).unwrap().in_flight(), 0);
1854        assert_eq!(state.context_record(context).unwrap().children().events, 0);
1855    }
1856
1857    #[test]
1858    fn wrong_kind_cross_context_and_cross_device_ids_fail_before_provider_use() {
1859        let mut first = state(1, 2);
1860        let first_context = create_context(&mut first, 10);
1861        let second_context = create_context(&mut first, 11);
1862        let buffer = admitted(
1863            first
1864                .create_buffer_with(first_context, buffer_desc(), |_, desc| {
1865                    Ok::<_, &'static str>((20, buffer_info(desc)))
1866                })
1867                .unwrap(),
1868        );
1869        let program = first
1870            .create_program_with(second_context, 1, 1, |_| Ok::<_, &'static str>(30))
1871            .unwrap();
1872        let queue = first
1873            .create_queue_with(first_context, |_| Ok::<_, &'static str>(40))
1874            .unwrap();
1875        let called = Cell::new(false);
1876        assert!(matches!(
1877            first.create_event_with(queue, program, alloc::vec![buffer], |_| {
1878                called.set(true);
1879                Ok::<_, &'static str>(50)
1880            }),
1881            Err(CreateError::State(DeviceStateError::ContextMismatch))
1882        ));
1883        assert!(!called.get());
1884        assert!(matches!(
1885            first.buffer_record(queue),
1886            Err(DeviceStateError::StaleObject)
1887        ));
1888
1889        let second = state(2, 2);
1890        assert!(matches!(
1891            second.context_record(first_context),
1892            Err(DeviceStateError::StaleObject)
1893        ));
1894    }
1895
1896    #[test]
1897    fn invalid_limits_are_rejected_before_tables_exist() {
1898        let namespace = ObjectNamespace::new(1).unwrap();
1899        let mut invalid = limits(1);
1900        invalid.max_bindings_per_submission = 0;
1901        assert!(matches!(
1902            TestState::new(namespace, invalid, resource_policy()),
1903            Err(DeviceStateConfigError::BindingLimit)
1904        ));
1905
1906        let mut invalid = limits(u32::MAX);
1907        invalid.max_buffers_per_context = 2;
1908        assert!(matches!(
1909            TestState::new(namespace, invalid, resource_policy()),
1910            Err(DeviceStateConfigError::CountOverflow)
1911        ));
1912
1913        let mut invalid = limits(1);
1914        invalid.max_events_per_context = u32::MAX;
1915        invalid.max_bindings_per_submission = HARD_MAX_BINDINGS;
1916        assert!(matches!(
1917            TestState::new(namespace, invalid, resource_policy()),
1918            Err(DeviceStateConfigError::ReferenceCountOverflow)
1919        ));
1920    }
1921
1922    #[test]
1923    fn byte_limits_and_resident_charges_have_distinct_semantics() {
1924        let mut state = state(1, 1);
1925        let context = create_context(&mut state, 10);
1926        let calls = Cell::new(0_u32);
1927        let oversized_buffer = BufferDesc::new(
1928            state.limits().max_buffer_bytes + 1,
1929            1,
1930            MemoryDomain::Host,
1931            BufferUsage::TRANSFER_SOURCE,
1932        )
1933        .unwrap();
1934
1935        assert!(matches!(
1936            state.create_buffer_with(context, oversized_buffer, |_, desc| {
1937                calls.set(calls.get() + 1);
1938                Ok::<_, &'static str>((20, buffer_info(desc)))
1939            }),
1940            Err(CreateError::State(DeviceStateError::ResourceLimit))
1941        ));
1942        assert!(matches!(
1943            state.create_program_with(context, state.limits().max_artifact_bytes + 1, 1, |_| {
1944                calls.set(calls.get() + 1);
1945                Ok::<_, &'static str>(30)
1946            }),
1947            Err(CreateError::State(DeviceStateError::ResourceLimit))
1948        ));
1949        let resident_bytes = state.limits().max_artifact_bytes + 1;
1950        let program = state
1951            .create_program_with(context, 1, resident_bytes, |_| {
1952                calls.set(calls.get() + 1);
1953                Ok::<_, &'static str>(30)
1954            })
1955            .unwrap();
1956        assert_eq!(calls.get(), 1);
1957        assert_eq!(state.buffer_count(), 0);
1958        assert_eq!(state.program_count(), 1);
1959        assert_eq!(
1960            state.program_record(program).unwrap().resident_bytes(),
1961            resident_bytes
1962        );
1963        assert_eq!(
1964            state.context_record(context).unwrap().children(),
1965            ChildCounts {
1966                programs: 1,
1967                ..ChildCounts::default()
1968            }
1969        );
1970        assert_eq!(
1971            state.retained_bytes(),
1972            RetainedBytes {
1973                buffer_backing: 0,
1974                program_resident: u128::from(resident_bytes),
1975            }
1976        );
1977    }
1978
1979    #[test]
1980    fn aggregate_policy_uses_actual_backing_and_charges_until_release_commits() {
1981        assert!(ResourcePolicy::new(0, 1).is_none());
1982        assert!(ResourcePolicy::new(1, 0).is_none());
1983
1984        let mut state = TestState::new(
1985            ObjectNamespace::new(1).unwrap(),
1986            limits(1),
1987            ResourcePolicy::new(4096, 4096).unwrap(),
1988        )
1989        .unwrap();
1990        let context = create_context(&mut state, 10);
1991        let desc = buffer_desc();
1992        let padded = BufferInfo::new(
1993            desc,
1994            8192,
1995            64,
1996            BufferProperties::HOST_VISIBLE | BufferProperties::DIRECT_BINDING,
1997        )
1998        .unwrap();
1999        let outcome = state
2000            .create_buffer_with(context, desc, |_, _| Ok::<_, &'static str>((20, padded)))
2001            .unwrap();
2002        let BufferCreateOutcome::CleanupRequired {
2003            id,
2004            error: BackendError::ResourceLimit,
2005        } = outcome
2006        else {
2007            panic!("padded allocation did not require cleanup");
2008        };
2009        assert_eq!(
2010            state.retained_bytes(),
2011            RetainedBytes {
2012                buffer_backing: 8192,
2013                program_resident: 0,
2014            }
2015        );
2016
2017        let resource = state.begin_buffer_release(id).unwrap();
2018        assert_eq!(state.retained_bytes().buffer_backing, 8192);
2019        state.restore_buffer_release(id, resource).unwrap();
2020        let resource = state.begin_buffer_release(id).unwrap();
2021        assert_eq!(resource, 20);
2022        state.commit_buffer_release(id).unwrap();
2023        assert!(state.retained_bytes().is_empty());
2024    }
2025
2026    #[test]
2027    fn program_policy_rejects_before_provider_invocation() {
2028        let mut state = TestState::new(
2029            ObjectNamespace::new(1).unwrap(),
2030            limits(1),
2031            ResourcePolicy::new(4096, 8).unwrap(),
2032        )
2033        .unwrap();
2034        let context = create_context(&mut state, 10);
2035        let called = Cell::new(false);
2036        assert!(matches!(
2037            state.create_program_with(context, 1, 9, |_| {
2038                called.set(true);
2039                Ok::<_, &'static str>(30)
2040            }),
2041            Err(CreateError::State(DeviceStateError::ResourceLimit))
2042        ));
2043        assert!(!called.get());
2044        assert!(state.retained_bytes().is_empty());
2045    }
2046}