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    buffer_ids: Vec<ObjectId>,
349}
350
351impl<E> EventRecord<E> {
352    pub fn resource(&self) -> Result<&E, DeviceStateError> {
353        self.resource.get()
354    }
355
356    pub const fn context_id(&self) -> ObjectId {
357        self.context_id
358    }
359
360    pub const fn queue_id(&self) -> ObjectId {
361        self.queue_id
362    }
363
364    pub const fn program_id(&self) -> ObjectId {
365        self.program_id
366    }
367
368    pub fn buffer_ids(&self) -> &[ObjectId] {
369        &self.buffer_ids
370    }
371
372    pub const fn release_state(&self) -> ReleaseState {
373        self.resource.release
374    }
375}
376
377/// Validated provider resources for one event-producing submission.
378pub struct SubmissionResources<'a, B, P, Q> {
379    context_id: ObjectId,
380    queue: &'a Q,
381    program: &'a P,
382    buffers: &'a ObjectTable<BufferRecord<B>>,
383    buffer_ids: &'a [ObjectId],
384}
385
386impl<B, P, Q> core::fmt::Debug for SubmissionResources<'_, B, P, Q> {
387    fn fmt(&self, formatter: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
388        formatter
389            .debug_struct("SubmissionResources")
390            .field("context_id", &self.context_id)
391            .field("buffer_ids", &self.buffer_ids)
392            .finish_non_exhaustive()
393    }
394}
395
396impl<'a, B, P, Q> SubmissionResources<'a, B, P, Q> {
397    pub const fn context_id(&self) -> ObjectId {
398        self.context_id
399    }
400
401    pub const fn queue(&self) -> &'a Q {
402        self.queue
403    }
404
405    pub const fn program(&self) -> &'a P {
406        self.program
407    }
408
409    /// Retained buffer IDs sorted by raw object ID, including duplicate bindings.
410    pub fn buffer_ids(&self) -> &[ObjectId] {
411        self.buffer_ids
412    }
413
414    pub fn buffer_by_id(&self, id: ObjectId) -> Result<&'a B, DeviceStateError> {
415        self.buffer_with_info_by_id(id).map(|(buffer, _)| buffer)
416    }
417
418    pub(crate) fn buffer_with_info_by_id(
419        &self,
420        id: ObjectId,
421    ) -> Result<(&'a B, BufferInfo), DeviceStateError> {
422        self.buffer_ids
423            .binary_search_by_key(&id.get(), |candidate| candidate.get())
424            .map_err(|_| DeviceStateError::InvalidArgument)?;
425        let record = self.buffers.get(id).map_err(map_table_error)?;
426        Ok((record.resource()?, record.info()))
427    }
428}
429
430/// Complete typed object graph for one device instance.
431pub struct DeviceState<C, B, P, Q, E> {
432    namespace: ObjectNamespace,
433    limits: DeviceLimits,
434    policy: ResourcePolicy,
435    retained: RetainedBytes,
436    contexts: ObjectTable<ContextRecord<C>>,
437    buffers: ObjectTable<BufferRecord<B>>,
438    programs: ObjectTable<ProgramRecord<P>>,
439    queues: ObjectTable<QueueRecord<Q>>,
440    events: ObjectTable<EventRecord<E>>,
441}
442
443impl<C, B, P, Q, E> DeviceState<C, B, P, Q, E> {
444    pub fn new(
445        namespace: ObjectNamespace,
446        limits: DeviceLimits,
447        policy: ResourcePolicy,
448    ) -> Result<Self, DeviceStateConfigError> {
449        if limits.max_contexts == 0
450            || limits.max_buffers_per_context == 0
451            || limits.max_programs_per_context == 0
452            || limits.max_queues_per_context == 0
453            || limits.max_events_per_context == 0
454            || limits.max_buffer_bytes == 0
455            || limits.max_artifact_bytes == 0
456        {
457            return Err(DeviceStateConfigError::ZeroLimit);
458        }
459        if !(1..=HARD_MAX_BINDINGS).contains(&limits.max_bindings_per_submission) {
460            return Err(DeviceStateConfigError::BindingLimit);
461        }
462        limits
463            .max_events_per_context
464            .checked_mul(limits.max_bindings_per_submission)
465            .ok_or(DeviceStateConfigError::ReferenceCountOverflow)?;
466
467        let buffers = aggregate_slots(limits.max_contexts, limits.max_buffers_per_context)?;
468        let programs = aggregate_slots(limits.max_contexts, limits.max_programs_per_context)?;
469        let queues = aggregate_slots(limits.max_contexts, limits.max_queues_per_context)?;
470        let events = aggregate_slots(limits.max_contexts, limits.max_events_per_context)?;
471
472        Ok(Self {
473            namespace,
474            limits,
475            policy,
476            retained: RetainedBytes::default(),
477            contexts: ObjectTable::with_namespace(
478                ObjectKind::Context,
479                limits.max_contexts,
480                namespace,
481            ),
482            buffers: ObjectTable::with_namespace(ObjectKind::Buffer, buffers, namespace),
483            programs: ObjectTable::with_namespace(ObjectKind::Program, programs, namespace),
484            queues: ObjectTable::with_namespace(ObjectKind::Queue, queues, namespace),
485            events: ObjectTable::with_namespace(ObjectKind::Event, events, namespace),
486        })
487    }
488
489    pub const fn limits(&self) -> DeviceLimits {
490        self.limits
491    }
492
493    pub const fn resource_policy(&self) -> ResourcePolicy {
494        self.policy
495    }
496
497    /// Bulk bytes still represented by live or releasing provider handles.
498    pub const fn retained_bytes(&self) -> RetainedBytes {
499        self.retained
500    }
501
502    pub const fn namespace(&self) -> ObjectNamespace {
503        self.namespace
504    }
505
506    pub const fn resource_counts(&self) -> ResourceCounts {
507        ResourceCounts {
508            contexts: self.context_count() as u64,
509            buffers: self.buffer_count() as u64,
510            programs: self.program_count() as u64,
511            queues: self.queue_count() as u64,
512            events: self.event_count() as u64,
513        }
514    }
515
516    pub const fn is_empty(&self) -> bool {
517        self.contexts.is_empty()
518            && self.buffers.is_empty()
519            && self.programs.is_empty()
520            && self.queues.is_empty()
521            && self.events.is_empty()
522    }
523
524    pub const fn context_count(&self) -> u32 {
525        self.contexts.len()
526    }
527
528    pub const fn buffer_count(&self) -> u32 {
529        self.buffers.len()
530    }
531
532    pub const fn program_count(&self) -> u32 {
533        self.programs.len()
534    }
535
536    pub const fn queue_count(&self) -> u32 {
537        self.queues.len()
538    }
539
540    pub const fn event_count(&self) -> u32 {
541        self.events.len()
542    }
543
544    pub fn context_record(&self, id: ObjectId) -> Result<&ContextRecord<C>, DeviceStateError> {
545        self.contexts.get(id).map_err(map_table_error)
546    }
547
548    pub fn buffer_record(&self, id: ObjectId) -> Result<&BufferRecord<B>, DeviceStateError> {
549        self.buffers.get(id).map_err(map_table_error)
550    }
551
552    pub fn buffer_record_mut(
553        &mut self,
554        id: ObjectId,
555    ) -> Result<&mut BufferRecord<B>, DeviceStateError> {
556        self.buffers.get_mut(id).map_err(map_table_error)
557    }
558
559    pub fn program_record(&self, id: ObjectId) -> Result<&ProgramRecord<P>, DeviceStateError> {
560        self.programs.get(id).map_err(map_table_error)
561    }
562
563    pub fn queue_record(&self, id: ObjectId) -> Result<&QueueRecord<Q>, DeviceStateError> {
564        self.queues.get(id).map_err(map_table_error)
565    }
566
567    pub fn event_record(&self, id: ObjectId) -> Result<&EventRecord<E>, DeviceStateError> {
568        self.events.get(id).map_err(map_table_error)
569    }
570
571    pub(crate) fn next_context_id(&self, start: usize) -> Option<(usize, ObjectId)> {
572        self.contexts.next_id_from(start)
573    }
574
575    pub(crate) fn next_buffer_id(&self, start: usize) -> Option<(usize, ObjectId)> {
576        self.buffers.next_id_from(start)
577    }
578
579    pub(crate) fn next_program_id(&self, start: usize) -> Option<(usize, ObjectId)> {
580        self.programs.next_id_from(start)
581    }
582
583    pub(crate) fn next_queue_id(&self, start: usize) -> Option<(usize, ObjectId)> {
584        self.queues.next_id_from(start)
585    }
586
587    pub(crate) fn next_event_id(&self, start: usize) -> Option<(usize, ObjectId)> {
588        self.events.next_id_from(start)
589    }
590
591    pub fn create_context_with<ProviderError>(
592        &mut self,
593        create: impl FnOnce() -> Result<C, ProviderError>,
594    ) -> Result<ObjectId, CreateError<ProviderError>> {
595        if self.contexts.len() >= self.limits.max_contexts {
596            return Err(CreateError::State(DeviceStateError::ResourceLimit));
597        }
598        self.contexts
599            .try_reserve_insert()
600            .map_err(|error| CreateError::State(map_table_error(error)))?;
601        let resource = create().map_err(CreateError::Provider)?;
602        Ok(self.contexts.insert_prepared(ContextRecord {
603            resource: ResourceSlot::new(resource),
604            children: ChildCounts::default(),
605        }))
606    }
607
608    pub fn create_buffer_with<ProviderError>(
609        &mut self,
610        context_id: ObjectId,
611        desc: BufferDesc,
612        create: impl FnOnce(&C, BufferDesc) -> Result<(B, BufferInfo), ProviderError>,
613    ) -> Result<BufferCreateOutcome, CreateError<ProviderError>> {
614        if desc.bytes() > self.limits.max_buffer_bytes {
615            return Err(CreateError::State(DeviceStateError::ResourceLimit));
616        }
617        let minimum_retained = self
618            .retained
619            .buffer_backing
620            .saturating_add(u128::from(desc.bytes()));
621        if minimum_retained > u128::from(self.policy.max_buffer_backing_bytes()) {
622            return Err(CreateError::State(DeviceStateError::ResourceLimit));
623        }
624        self.check_child_admission(
625            context_id,
626            ChildKind::Buffer,
627            self.limits.max_buffers_per_context,
628        )
629        .map_err(CreateError::State)?;
630        self.buffers
631            .try_reserve_insert()
632            .map_err(|error| CreateError::State(map_table_error(error)))?;
633
634        let context = self
635            .contexts
636            .get_mut(context_id)
637            .map_err(|error| CreateError::State(map_table_error(error)))?;
638        let (resource, info) = create(context.resource()?, desc).map_err(CreateError::Provider)?;
639        let retained = self
640            .retained
641            .buffer_backing
642            .saturating_add(u128::from(info.allocation_bytes()));
643        let cleanup_error = if info.desc() != desc {
644            Some(BackendError::Incompatible)
645        } else if retained > u128::from(self.policy.max_buffer_backing_bytes()) {
646            Some(BackendError::ResourceLimit)
647        } else {
648            None
649        };
650        let id = self.buffers.insert_prepared(BufferRecord {
651            resource: ResourceSlot::new(resource),
652            context_id,
653            info,
654            in_flight: 0,
655        });
656        self.retained.buffer_backing = retained;
657        context.children.buffers += 1;
658        Ok(match cleanup_error {
659            Some(error) => BufferCreateOutcome::CleanupRequired { id, error },
660            None => BufferCreateOutcome::Admitted(id),
661        })
662    }
663
664    pub fn create_program_with<ProviderError>(
665        &mut self,
666        context_id: ObjectId,
667        artifact_bytes: u64,
668        resident_bytes: u64,
669        create: impl FnOnce(&C) -> Result<P, ProviderError>,
670    ) -> Result<ObjectId, CreateError<ProviderError>> {
671        if artifact_bytes == 0 || resident_bytes == 0 {
672            return Err(CreateError::State(DeviceStateError::InvalidArgument));
673        }
674        if artifact_bytes > self.limits.max_artifact_bytes {
675            return Err(CreateError::State(DeviceStateError::ResourceLimit));
676        }
677        let retained = self
678            .retained
679            .program_resident
680            .saturating_add(u128::from(resident_bytes));
681        if retained > u128::from(self.policy.max_program_resident_bytes()) {
682            return Err(CreateError::State(DeviceStateError::ResourceLimit));
683        }
684        self.check_child_admission(
685            context_id,
686            ChildKind::Program,
687            self.limits.max_programs_per_context,
688        )
689        .map_err(CreateError::State)?;
690        self.programs
691            .try_reserve_insert()
692            .map_err(|error| CreateError::State(map_table_error(error)))?;
693
694        let context = self
695            .contexts
696            .get_mut(context_id)
697            .map_err(|error| CreateError::State(map_table_error(error)))?;
698        let resource = create(context.resource()?).map_err(CreateError::Provider)?;
699        let id = self.programs.insert_prepared(ProgramRecord {
700            resource: ResourceSlot::new(resource),
701            context_id,
702            resident_bytes,
703            in_flight: 0,
704        });
705        self.retained.program_resident = retained;
706        context.children.programs += 1;
707        Ok(id)
708    }
709
710    pub fn create_queue_with<ProviderError>(
711        &mut self,
712        context_id: ObjectId,
713        create: impl FnOnce(&C) -> Result<Q, ProviderError>,
714    ) -> Result<ObjectId, CreateError<ProviderError>> {
715        self.check_child_admission(
716            context_id,
717            ChildKind::Queue,
718            self.limits.max_queues_per_context,
719        )
720        .map_err(CreateError::State)?;
721        self.queues
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 = create(context.resource()?).map_err(CreateError::Provider)?;
730        let id = self.queues.insert_prepared(QueueRecord {
731            resource: ResourceSlot::new(resource),
732            context_id,
733            in_flight: 0,
734        });
735        context.children.queues += 1;
736        Ok(id)
737    }
738
739    pub fn create_event_with<ProviderError>(
740        &mut self,
741        queue_id: ObjectId,
742        program_id: ObjectId,
743        mut buffer_ids: Vec<ObjectId>,
744        create: impl FnOnce(SubmissionResources<'_, B, P, Q>) -> Result<E, ProviderError>,
745    ) -> Result<ObjectId, CreateError<ProviderError>> {
746        if buffer_ids.is_empty() {
747            return Err(CreateError::State(DeviceStateError::InvalidArgument));
748        }
749        if buffer_ids.len() > self.limits.max_bindings_per_submission as usize {
750            return Err(CreateError::State(DeviceStateError::ResourceLimit));
751        }
752
753        let context_id = self
754            .validate_submission(queue_id, program_id, &buffer_ids)
755            .map_err(CreateError::State)?;
756        let context = self
757            .contexts
758            .get(context_id)
759            .map_err(|error| CreateError::State(map_table_error(error)))?;
760        if context.children.events >= self.limits.max_events_per_context {
761            return Err(CreateError::State(DeviceStateError::ResourceLimit));
762        }
763        self.events
764            .try_reserve_insert()
765            .map_err(|error| CreateError::State(map_table_error(error)))?;
766
767        buffer_ids.sort_unstable_by_key(|id| id.get());
768        self.check_reference_increments(queue_id, program_id, &buffer_ids)
769            .map_err(CreateError::State)?;
770        self.increment_event_references(context_id, queue_id, program_id, &buffer_ids)
771            .map_err(CreateError::State)?;
772
773        let event_result = {
774            let queue = self
775                .queues
776                .get(queue_id)
777                .map_err(|error| CreateError::State(map_table_error(error)))?
778                .resource()?;
779            let program = self
780                .programs
781                .get(program_id)
782                .map_err(|error| CreateError::State(map_table_error(error)))?
783                .resource()?;
784            create(SubmissionResources {
785                context_id,
786                queue,
787                program,
788                buffers: &self.buffers,
789                buffer_ids: &buffer_ids,
790            })
791        };
792        let event = match event_result {
793            Ok(event) => event,
794            Err(error) => {
795                self.decrement_event_references(context_id, queue_id, program_id, &buffer_ids)
796                    .map_err(CreateError::State)?;
797                return Err(CreateError::Provider(error));
798            }
799        };
800
801        let id = self.events.insert_prepared(EventRecord {
802            resource: ResourceSlot::new(event),
803            context_id,
804            queue_id,
805            program_id,
806            buffer_ids,
807        });
808        Ok(id)
809    }
810
811    pub fn begin_context_release(&mut self, id: ObjectId) -> Result<C, DeviceStateError> {
812        let record = self.contexts.get_mut(id).map_err(map_table_error)?;
813        if !record.children.is_empty() {
814            return Err(DeviceStateError::Busy);
815        }
816        record.resource.begin_release()
817    }
818
819    pub fn restore_context_release(
820        &mut self,
821        id: ObjectId,
822        resource: C,
823    ) -> Result<(), RestoreError<C>> {
824        restore_resource(&mut self.contexts, id, resource)
825    }
826
827    pub fn commit_context_release(&mut self, id: ObjectId) -> Result<(), DeviceStateError> {
828        ensure_releasing(self.contexts.get(id).map_err(map_table_error)?)?;
829        self.contexts.remove(id).map_err(map_table_error)?;
830        Ok(())
831    }
832
833    pub fn begin_buffer_release(&mut self, id: ObjectId) -> Result<B, DeviceStateError> {
834        let record = self.buffers.get_mut(id).map_err(map_table_error)?;
835        if record.in_flight != 0 {
836            return Err(DeviceStateError::Busy);
837        }
838        record.resource.begin_release()
839    }
840
841    pub fn restore_buffer_release(
842        &mut self,
843        id: ObjectId,
844        resource: B,
845    ) -> Result<(), RestoreError<B>> {
846        restore_resource(&mut self.buffers, id, resource)
847    }
848
849    pub fn commit_buffer_release(&mut self, id: ObjectId) -> Result<(), DeviceStateError> {
850        let (context_id, allocation_bytes) = {
851            let record = self.buffers.get(id).map_err(map_table_error)?;
852            ensure_releasing(record)?;
853            (record.context_id, record.info.allocation_bytes())
854        };
855        self.buffers.remove(id).map_err(map_table_error)?;
856        self.retained.buffer_backing -= u128::from(allocation_bytes);
857        self.contexts
858            .get_mut(context_id)
859            .map_err(map_table_error)?
860            .children
861            .buffers -= 1;
862        Ok(())
863    }
864
865    pub fn begin_program_release(&mut self, id: ObjectId) -> Result<P, DeviceStateError> {
866        let record = self.programs.get_mut(id).map_err(map_table_error)?;
867        if record.in_flight != 0 {
868            return Err(DeviceStateError::Busy);
869        }
870        record.resource.begin_release()
871    }
872
873    pub fn restore_program_release(
874        &mut self,
875        id: ObjectId,
876        resource: P,
877    ) -> Result<(), RestoreError<P>> {
878        restore_resource(&mut self.programs, id, resource)
879    }
880
881    pub fn commit_program_release(&mut self, id: ObjectId) -> Result<(), DeviceStateError> {
882        let (context_id, resident_bytes) = {
883            let record = self.programs.get(id).map_err(map_table_error)?;
884            ensure_releasing(record)?;
885            (record.context_id, record.resident_bytes)
886        };
887        self.programs.remove(id).map_err(map_table_error)?;
888        self.retained.program_resident -= u128::from(resident_bytes);
889        self.contexts
890            .get_mut(context_id)
891            .map_err(map_table_error)?
892            .children
893            .programs -= 1;
894        Ok(())
895    }
896
897    pub fn begin_queue_release(&mut self, id: ObjectId) -> Result<Q, DeviceStateError> {
898        let record = self.queues.get_mut(id).map_err(map_table_error)?;
899        if record.in_flight != 0 {
900            return Err(DeviceStateError::Busy);
901        }
902        record.resource.begin_release()
903    }
904
905    pub fn restore_queue_release(
906        &mut self,
907        id: ObjectId,
908        resource: Q,
909    ) -> Result<(), RestoreError<Q>> {
910        restore_resource(&mut self.queues, id, resource)
911    }
912
913    pub fn commit_queue_release(&mut self, id: ObjectId) -> Result<(), DeviceStateError> {
914        let context_id = {
915            let record = self.queues.get(id).map_err(map_table_error)?;
916            ensure_releasing(record)?;
917            record.context_id
918        };
919        self.queues.remove(id).map_err(map_table_error)?;
920        self.contexts
921            .get_mut(context_id)
922            .map_err(map_table_error)?
923            .children
924            .queues -= 1;
925        Ok(())
926    }
927
928    pub fn begin_event_release(&mut self, id: ObjectId) -> Result<E, DeviceStateError> {
929        self.events
930            .get_mut(id)
931            .map_err(map_table_error)?
932            .resource
933            .begin_release()
934    }
935
936    pub fn restore_event_release(
937        &mut self,
938        id: ObjectId,
939        resource: E,
940    ) -> Result<(), RestoreError<E>> {
941        restore_resource(&mut self.events, id, resource)
942    }
943
944    pub fn commit_event_release(&mut self, id: ObjectId) -> Result<(), DeviceStateError> {
945        {
946            let record = self.events.get(id).map_err(map_table_error)?;
947            ensure_releasing(record)?;
948            self.validate_submission(record.queue_id, record.program_id, &record.buffer_ids)?;
949        }
950        let record = self.events.remove(id).map_err(map_table_error)?;
951        self.decrement_event_references(
952            record.context_id,
953            record.queue_id,
954            record.program_id,
955            &record.buffer_ids,
956        )
957    }
958
959    fn check_child_admission(
960        &self,
961        context_id: ObjectId,
962        child: ChildKind,
963        limit: u32,
964    ) -> Result<(), DeviceStateError> {
965        let context = self.contexts.get(context_id).map_err(map_table_error)?;
966        context.resource()?;
967        let count = match child {
968            ChildKind::Buffer => context.children.buffers,
969            ChildKind::Program => context.children.programs,
970            ChildKind::Queue => context.children.queues,
971        };
972        if count >= limit {
973            return Err(DeviceStateError::ResourceLimit);
974        }
975        Ok(())
976    }
977
978    fn validate_submission(
979        &self,
980        queue_id: ObjectId,
981        program_id: ObjectId,
982        buffer_ids: &[ObjectId],
983    ) -> Result<ObjectId, DeviceStateError> {
984        let queue = self.queues.get(queue_id).map_err(map_table_error)?;
985        queue.resource()?;
986        let program = self.programs.get(program_id).map_err(map_table_error)?;
987        program.resource()?;
988        if program.context_id != queue.context_id {
989            return Err(DeviceStateError::ContextMismatch);
990        }
991        for buffer_id in buffer_ids {
992            let buffer = self.buffers.get(*buffer_id).map_err(map_table_error)?;
993            buffer.resource()?;
994            if buffer.context_id != queue.context_id {
995                return Err(DeviceStateError::ContextMismatch);
996            }
997        }
998        Ok(queue.context_id)
999    }
1000
1001    fn check_reference_increments(
1002        &self,
1003        queue_id: ObjectId,
1004        program_id: ObjectId,
1005        sorted_buffer_ids: &[ObjectId],
1006    ) -> Result<(), DeviceStateError> {
1007        self.queues
1008            .get(queue_id)
1009            .map_err(map_table_error)?
1010            .in_flight
1011            .checked_add(1)
1012            .ok_or(DeviceStateError::ReferenceCountOverflow)?;
1013        self.programs
1014            .get(program_id)
1015            .map_err(map_table_error)?
1016            .in_flight
1017            .checked_add(1)
1018            .ok_or(DeviceStateError::ReferenceCountOverflow)?;
1019
1020        let mut index = 0;
1021        while index < sorted_buffer_ids.len() {
1022            let id = sorted_buffer_ids[index];
1023            let mut end = index + 1;
1024            while end < sorted_buffer_ids.len() && sorted_buffer_ids[end] == id {
1025                end += 1;
1026            }
1027            let count =
1028                u32::try_from(end - index).map_err(|_| DeviceStateError::ReferenceCountOverflow)?;
1029            self.buffers
1030                .get(id)
1031                .map_err(map_table_error)?
1032                .in_flight
1033                .checked_add(count)
1034                .ok_or(DeviceStateError::ReferenceCountOverflow)?;
1035            index = end;
1036        }
1037        Ok(())
1038    }
1039
1040    fn increment_event_references(
1041        &mut self,
1042        context_id: ObjectId,
1043        queue_id: ObjectId,
1044        program_id: ObjectId,
1045        buffer_ids: &[ObjectId],
1046    ) -> Result<(), DeviceStateError> {
1047        self.queues
1048            .get_mut(queue_id)
1049            .map_err(map_table_error)?
1050            .in_flight += 1;
1051        self.programs
1052            .get_mut(program_id)
1053            .map_err(map_table_error)?
1054            .in_flight += 1;
1055        for buffer_id in buffer_ids {
1056            self.buffers
1057                .get_mut(*buffer_id)
1058                .map_err(map_table_error)?
1059                .in_flight += 1;
1060        }
1061        self.contexts
1062            .get_mut(context_id)
1063            .map_err(map_table_error)?
1064            .children
1065            .events += 1;
1066        Ok(())
1067    }
1068
1069    fn decrement_event_references(
1070        &mut self,
1071        context_id: ObjectId,
1072        queue_id: ObjectId,
1073        program_id: ObjectId,
1074        buffer_ids: &[ObjectId],
1075    ) -> Result<(), DeviceStateError> {
1076        self.queues
1077            .get_mut(queue_id)
1078            .map_err(map_table_error)?
1079            .in_flight -= 1;
1080        self.programs
1081            .get_mut(program_id)
1082            .map_err(map_table_error)?
1083            .in_flight -= 1;
1084        for buffer_id in buffer_ids {
1085            self.buffers
1086                .get_mut(*buffer_id)
1087                .map_err(map_table_error)?
1088                .in_flight -= 1;
1089        }
1090        self.contexts
1091            .get_mut(context_id)
1092            .map_err(map_table_error)?
1093            .children
1094            .events -= 1;
1095        Ok(())
1096    }
1097}
1098
1099#[derive(Clone, Copy)]
1100enum ChildKind {
1101    Buffer,
1102    Program,
1103    Queue,
1104}
1105
1106trait ReleasableRecord {
1107    fn release_state(&self) -> ReleaseState;
1108}
1109
1110impl<C> ReleasableRecord for ContextRecord<C> {
1111    fn release_state(&self) -> ReleaseState {
1112        self.release_state()
1113    }
1114}
1115
1116impl<B> ReleasableRecord for BufferRecord<B> {
1117    fn release_state(&self) -> ReleaseState {
1118        self.release_state()
1119    }
1120}
1121
1122impl<P> ReleasableRecord for ProgramRecord<P> {
1123    fn release_state(&self) -> ReleaseState {
1124        self.release_state()
1125    }
1126}
1127
1128impl<Q> ReleasableRecord for QueueRecord<Q> {
1129    fn release_state(&self) -> ReleaseState {
1130        self.release_state()
1131    }
1132}
1133
1134impl<E> ReleasableRecord for EventRecord<E> {
1135    fn release_state(&self) -> ReleaseState {
1136        self.release_state()
1137    }
1138}
1139
1140fn ensure_releasing(record: &impl ReleasableRecord) -> Result<(), DeviceStateError> {
1141    if record.release_state() != ReleaseState::Releasing {
1142        return Err(DeviceStateError::InvalidTransition);
1143    }
1144    Ok(())
1145}
1146
1147trait ResourceRecord<R> {
1148    fn resource_mut(&mut self) -> &mut ResourceSlot<R>;
1149}
1150
1151impl<C> ResourceRecord<C> for ContextRecord<C> {
1152    fn resource_mut(&mut self) -> &mut ResourceSlot<C> {
1153        &mut self.resource
1154    }
1155}
1156
1157impl<B> ResourceRecord<B> for BufferRecord<B> {
1158    fn resource_mut(&mut self) -> &mut ResourceSlot<B> {
1159        &mut self.resource
1160    }
1161}
1162
1163impl<P> ResourceRecord<P> for ProgramRecord<P> {
1164    fn resource_mut(&mut self) -> &mut ResourceSlot<P> {
1165        &mut self.resource
1166    }
1167}
1168
1169impl<Q> ResourceRecord<Q> for QueueRecord<Q> {
1170    fn resource_mut(&mut self) -> &mut ResourceSlot<Q> {
1171        &mut self.resource
1172    }
1173}
1174
1175impl<E> ResourceRecord<E> for EventRecord<E> {
1176    fn resource_mut(&mut self) -> &mut ResourceSlot<E> {
1177        &mut self.resource
1178    }
1179}
1180
1181fn restore_resource<R, Record: ResourceRecord<R>>(
1182    table: &mut ObjectTable<Record>,
1183    id: ObjectId,
1184    resource: R,
1185) -> Result<(), RestoreError<R>> {
1186    let record = match table.get_mut(id) {
1187        Ok(record) => record,
1188        Err(error) => {
1189            return Err(RestoreError {
1190                error: map_table_error(error),
1191                resource,
1192            });
1193        }
1194    };
1195    record.resource_mut().restore(resource)
1196}
1197
1198fn aggregate_slots(contexts: u32, per_context: u32) -> Result<u32, DeviceStateConfigError> {
1199    contexts
1200        .checked_mul(per_context)
1201        .ok_or(DeviceStateConfigError::CountOverflow)
1202}
1203
1204fn map_table_error(error: ObjectTableError) -> DeviceStateError {
1205    match error {
1206        ObjectTableError::InvalidId => DeviceStateError::InvalidObject,
1207        ObjectTableError::WrongKind | ObjectTableError::StaleId => DeviceStateError::StaleObject,
1208        ObjectTableError::Full => DeviceStateError::ResourceLimit,
1209        ObjectTableError::AllocationFailed => DeviceStateError::OutOfMemory,
1210    }
1211}
1212
1213#[cfg(test)]
1214mod tests {
1215    use super::*;
1216    use core::cell::Cell;
1217    use virtio_accel_core::{BufferDesc, BufferProperties, BufferUsage, MemoryDomain};
1218
1219    type TestState = DeviceState<u32, u32, u32, u32, u32>;
1220
1221    fn limits(max_contexts: u32) -> DeviceLimits {
1222        DeviceLimits {
1223            max_contexts,
1224            max_buffers_per_context: 1,
1225            max_programs_per_context: 1,
1226            max_queues_per_context: 1,
1227            max_events_per_context: 1,
1228            max_bindings_per_submission: 4,
1229            max_buffer_bytes: 1 << 20,
1230            max_artifact_bytes: 1 << 20,
1231        }
1232    }
1233
1234    fn state(namespace: u16, max_contexts: u32) -> TestState {
1235        DeviceState::new(
1236            ObjectNamespace::new(namespace).unwrap(),
1237            limits(max_contexts),
1238            resource_policy(),
1239        )
1240        .unwrap()
1241    }
1242
1243    fn resource_policy() -> ResourcePolicy {
1244        ResourcePolicy::new(1 << 30, 1 << 30).unwrap()
1245    }
1246
1247    fn admitted(outcome: BufferCreateOutcome) -> ObjectId {
1248        match outcome {
1249            BufferCreateOutcome::Admitted(id) => id,
1250            BufferCreateOutcome::CleanupRequired { .. } => {
1251                panic!("test allocation unexpectedly required cleanup")
1252            }
1253        }
1254    }
1255
1256    fn buffer_desc() -> BufferDesc {
1257        BufferDesc::new(
1258            4096,
1259            64,
1260            MemoryDomain::Shared,
1261            BufferUsage::TRANSFER_SOURCE
1262                | BufferUsage::TRANSFER_DESTINATION
1263                | BufferUsage::PROGRAM_INPUT,
1264        )
1265        .unwrap()
1266    }
1267
1268    fn buffer_info(desc: BufferDesc) -> BufferInfo {
1269        BufferInfo::new(
1270            desc,
1271            4096,
1272            64,
1273            BufferProperties::HOST_VISIBLE | BufferProperties::DIRECT_BINDING,
1274        )
1275        .unwrap()
1276    }
1277
1278    fn create_context(state: &mut TestState, resource: u32) -> ObjectId {
1279        state
1280            .create_context_with(|| Ok::<_, &'static str>(resource))
1281            .unwrap()
1282    }
1283
1284    #[test]
1285    fn complete_lifecycle_tracks_children_references_and_release_rollback() {
1286        let mut state = state(1, 1);
1287        let context = create_context(&mut state, 10);
1288        let buffer = admitted(
1289            state
1290                .create_buffer_with(context, buffer_desc(), |context, desc| {
1291                    assert_eq!(*context, 10);
1292                    Ok::<_, &'static str>((20, buffer_info(desc)))
1293                })
1294                .unwrap(),
1295        );
1296        *state
1297            .buffer_record_mut(buffer)
1298            .unwrap()
1299            .resource_mut()
1300            .unwrap() = 21;
1301        let program = state
1302            .create_program_with(context, 4096, 8192, |context| {
1303                assert_eq!(*context, 10);
1304                Ok::<_, &'static str>(30)
1305            })
1306            .unwrap();
1307        let queue = state
1308            .create_queue_with(context, |context| {
1309                assert_eq!(*context, 10);
1310                Ok::<_, &'static str>(40)
1311            })
1312            .unwrap();
1313        let event = state
1314            .create_event_with(queue, program, alloc::vec![buffer, buffer], |resources| {
1315                assert_eq!(resources.context_id(), context);
1316                assert_eq!(*resources.queue(), 40);
1317                assert_eq!(*resources.program(), 30);
1318                assert_eq!(*resources.buffer_by_id(buffer).unwrap(), 21);
1319                Ok::<_, &'static str>(50)
1320            })
1321            .unwrap();
1322
1323        assert_eq!(
1324            state.context_record(context).unwrap().children(),
1325            ChildCounts {
1326                buffers: 1,
1327                programs: 1,
1328                queues: 1,
1329                events: 1,
1330            }
1331        );
1332        assert_eq!(state.buffer_record(buffer).unwrap().in_flight(), 2);
1333        assert_eq!(state.program_record(program).unwrap().in_flight(), 1);
1334        assert_eq!(state.queue_record(queue).unwrap().in_flight(), 1);
1335        assert_eq!(
1336            state.begin_context_release(context),
1337            Err(DeviceStateError::Busy)
1338        );
1339        assert_eq!(
1340            state.begin_buffer_release(buffer),
1341            Err(DeviceStateError::Busy)
1342        );
1343        assert_eq!(
1344            state.begin_program_release(program),
1345            Err(DeviceStateError::Busy)
1346        );
1347        assert_eq!(
1348            state.begin_queue_release(queue),
1349            Err(DeviceStateError::Busy)
1350        );
1351
1352        let event_resource = state.begin_event_release(event).unwrap();
1353        assert_eq!(event_resource, 50);
1354        assert_eq!(
1355            state.event_record(event).unwrap().release_state(),
1356            ReleaseState::Releasing
1357        );
1358        state.restore_event_release(event, event_resource).unwrap();
1359        assert_eq!(*state.event_record(event).unwrap().resource().unwrap(), 50);
1360        let event_resource = state.begin_event_release(event).unwrap();
1361        assert_eq!(event_resource, 50);
1362        state.commit_event_release(event).unwrap();
1363        assert!(matches!(
1364            state.event_record(event),
1365            Err(DeviceStateError::StaleObject)
1366        ));
1367        assert_eq!(state.buffer_record(buffer).unwrap().in_flight(), 0);
1368        assert_eq!(state.program_record(program).unwrap().in_flight(), 0);
1369        assert_eq!(state.queue_record(queue).unwrap().in_flight(), 0);
1370
1371        let buffer_resource = state.begin_buffer_release(buffer).unwrap();
1372        state
1373            .restore_buffer_release(buffer, buffer_resource)
1374            .unwrap();
1375        let buffer_resource = state.begin_buffer_release(buffer).unwrap();
1376        assert_eq!(buffer_resource, 21);
1377        state.commit_buffer_release(buffer).unwrap();
1378
1379        let program_resource = state.begin_program_release(program).unwrap();
1380        state
1381            .restore_program_release(program, program_resource)
1382            .unwrap();
1383        let program_resource = state.begin_program_release(program).unwrap();
1384        assert_eq!(program_resource, 30);
1385        state.commit_program_release(program).unwrap();
1386
1387        let queue_resource = state.begin_queue_release(queue).unwrap();
1388        state.restore_queue_release(queue, queue_resource).unwrap();
1389        let queue_resource = state.begin_queue_release(queue).unwrap();
1390        assert_eq!(queue_resource, 40);
1391        state.commit_queue_release(queue).unwrap();
1392
1393        assert!(state.context_record(context).unwrap().children().is_empty());
1394        let context_resource = state.begin_context_release(context).unwrap();
1395        state
1396            .restore_context_release(context, context_resource)
1397            .unwrap();
1398        let context_resource = state.begin_context_release(context).unwrap();
1399        assert_eq!(context_resource, 10);
1400        state.commit_context_release(context).unwrap();
1401        assert!(matches!(
1402            state.context_record(context),
1403            Err(DeviceStateError::StaleObject)
1404        ));
1405        assert_eq!(state.context_count(), 0);
1406        assert_eq!(state.buffer_count(), 0);
1407        assert_eq!(state.program_count(), 0);
1408        assert_eq!(state.queue_count(), 0);
1409        assert_eq!(state.event_count(), 0);
1410    }
1411
1412    #[test]
1413    fn quota_exhaustion_never_invokes_provider_callbacks() {
1414        let mut state = state(1, 1);
1415        let calls = Cell::new(0_u32);
1416        let context = state
1417            .create_context_with(|| {
1418                calls.set(calls.get() + 1);
1419                Ok::<_, &'static str>(10)
1420            })
1421            .unwrap();
1422        assert!(matches!(
1423            state.create_context_with(|| {
1424                calls.set(calls.get() + 1);
1425                Ok::<_, &'static str>(11)
1426            }),
1427            Err(CreateError::State(DeviceStateError::ResourceLimit))
1428        ));
1429
1430        let buffer = admitted(
1431            state
1432                .create_buffer_with(context, buffer_desc(), |_, desc| {
1433                    calls.set(calls.get() + 1);
1434                    Ok::<_, &'static str>((20, buffer_info(desc)))
1435                })
1436                .unwrap(),
1437        );
1438        assert!(matches!(
1439            state.create_buffer_with(context, buffer_desc(), |_, desc| {
1440                calls.set(calls.get() + 1);
1441                Ok::<_, &'static str>((21, buffer_info(desc)))
1442            }),
1443            Err(CreateError::State(DeviceStateError::ResourceLimit))
1444        ));
1445
1446        let program = state
1447            .create_program_with(context, 1, 1, |_| {
1448                calls.set(calls.get() + 1);
1449                Ok::<_, &'static str>(30)
1450            })
1451            .unwrap();
1452        assert!(matches!(
1453            state.create_program_with(context, 1, 1, |_| {
1454                calls.set(calls.get() + 1);
1455                Ok::<_, &'static str>(31)
1456            }),
1457            Err(CreateError::State(DeviceStateError::ResourceLimit))
1458        ));
1459
1460        let queue = state
1461            .create_queue_with(context, |_| {
1462                calls.set(calls.get() + 1);
1463                Ok::<_, &'static str>(40)
1464            })
1465            .unwrap();
1466        assert!(matches!(
1467            state.create_queue_with(context, |_| {
1468                calls.set(calls.get() + 1);
1469                Ok::<_, &'static str>(41)
1470            }),
1471            Err(CreateError::State(DeviceStateError::ResourceLimit))
1472        ));
1473
1474        state
1475            .create_event_with(queue, program, alloc::vec![buffer], |_| {
1476                calls.set(calls.get() + 1);
1477                Ok::<_, &'static str>(50)
1478            })
1479            .unwrap();
1480        assert!(matches!(
1481            state.create_event_with(queue, program, alloc::vec![buffer], |_| {
1482                calls.set(calls.get() + 1);
1483                Ok::<_, &'static str>(51)
1484            }),
1485            Err(CreateError::State(DeviceStateError::ResourceLimit))
1486        ));
1487        assert_eq!(calls.get(), 5);
1488        assert_eq!(state.context_count(), 1);
1489        assert_eq!(state.buffer_count(), 1);
1490        assert_eq!(state.program_count(), 1);
1491        assert_eq!(state.queue_count(), 1);
1492        assert_eq!(state.event_count(), 1);
1493    }
1494
1495    #[test]
1496    fn provider_rejection_rolls_back_every_creation_path() {
1497        let mut state = state(1, 1);
1498        assert!(matches!(
1499            state.create_context_with(|| Err::<u32, _>("context")),
1500            Err(CreateError::Provider("context"))
1501        ));
1502        assert_eq!(state.context_count(), 0);
1503
1504        let context = create_context(&mut state, 10);
1505        assert!(matches!(
1506            state.create_buffer_with(context, buffer_desc(), |_, _| {
1507                Err::<(u32, BufferInfo), _>("buffer")
1508            }),
1509            Err(CreateError::Provider("buffer"))
1510        ));
1511        assert!(matches!(
1512            state.create_program_with(context, 1, 1, |_| Err::<u32, _>("program")),
1513            Err(CreateError::Provider("program"))
1514        ));
1515        assert!(matches!(
1516            state.create_queue_with(context, |_| Err::<u32, _>("queue")),
1517            Err(CreateError::Provider("queue"))
1518        ));
1519        assert_eq!(
1520            state.context_record(context).unwrap().children(),
1521            ChildCounts::default()
1522        );
1523
1524        let buffer = admitted(
1525            state
1526                .create_buffer_with(context, buffer_desc(), |_, desc| {
1527                    Ok::<_, &'static str>((20, buffer_info(desc)))
1528                })
1529                .unwrap(),
1530        );
1531        let program = state
1532            .create_program_with(context, 1, 1, |_| Ok::<_, &'static str>(30))
1533            .unwrap();
1534        let queue = state
1535            .create_queue_with(context, |_| Ok::<_, &'static str>(40))
1536            .unwrap();
1537        assert!(matches!(
1538            state.create_event_with(queue, program, alloc::vec![buffer], |_| {
1539                Err::<u32, _>("event")
1540            }),
1541            Err(CreateError::Provider("event"))
1542        ));
1543        assert_eq!(state.event_count(), 0);
1544        assert_eq!(state.buffer_record(buffer).unwrap().in_flight(), 0);
1545        assert_eq!(state.program_record(program).unwrap().in_flight(), 0);
1546        assert_eq!(state.queue_record(queue).unwrap().in_flight(), 0);
1547        assert_eq!(state.context_record(context).unwrap().children().events, 0);
1548    }
1549
1550    #[test]
1551    fn wrong_kind_cross_context_and_cross_device_ids_fail_before_provider_use() {
1552        let mut first = state(1, 2);
1553        let first_context = create_context(&mut first, 10);
1554        let second_context = create_context(&mut first, 11);
1555        let buffer = admitted(
1556            first
1557                .create_buffer_with(first_context, buffer_desc(), |_, desc| {
1558                    Ok::<_, &'static str>((20, buffer_info(desc)))
1559                })
1560                .unwrap(),
1561        );
1562        let program = first
1563            .create_program_with(second_context, 1, 1, |_| Ok::<_, &'static str>(30))
1564            .unwrap();
1565        let queue = first
1566            .create_queue_with(first_context, |_| Ok::<_, &'static str>(40))
1567            .unwrap();
1568        let called = Cell::new(false);
1569        assert!(matches!(
1570            first.create_event_with(queue, program, alloc::vec![buffer], |_| {
1571                called.set(true);
1572                Ok::<_, &'static str>(50)
1573            }),
1574            Err(CreateError::State(DeviceStateError::ContextMismatch))
1575        ));
1576        assert!(!called.get());
1577        assert!(matches!(
1578            first.buffer_record(queue),
1579            Err(DeviceStateError::StaleObject)
1580        ));
1581
1582        let second = state(2, 2);
1583        assert!(matches!(
1584            second.context_record(first_context),
1585            Err(DeviceStateError::StaleObject)
1586        ));
1587    }
1588
1589    #[test]
1590    fn invalid_limits_are_rejected_before_tables_exist() {
1591        let namespace = ObjectNamespace::new(1).unwrap();
1592        let mut invalid = limits(1);
1593        invalid.max_bindings_per_submission = 0;
1594        assert!(matches!(
1595            TestState::new(namespace, invalid, resource_policy()),
1596            Err(DeviceStateConfigError::BindingLimit)
1597        ));
1598
1599        let mut invalid = limits(u32::MAX);
1600        invalid.max_buffers_per_context = 2;
1601        assert!(matches!(
1602            TestState::new(namespace, invalid, resource_policy()),
1603            Err(DeviceStateConfigError::CountOverflow)
1604        ));
1605
1606        let mut invalid = limits(1);
1607        invalid.max_events_per_context = u32::MAX;
1608        invalid.max_bindings_per_submission = HARD_MAX_BINDINGS;
1609        assert!(matches!(
1610            TestState::new(namespace, invalid, resource_policy()),
1611            Err(DeviceStateConfigError::ReferenceCountOverflow)
1612        ));
1613    }
1614
1615    #[test]
1616    fn byte_limits_and_resident_charges_have_distinct_semantics() {
1617        let mut state = state(1, 1);
1618        let context = create_context(&mut state, 10);
1619        let calls = Cell::new(0_u32);
1620        let oversized_buffer = BufferDesc::new(
1621            state.limits().max_buffer_bytes + 1,
1622            1,
1623            MemoryDomain::Host,
1624            BufferUsage::TRANSFER_SOURCE,
1625        )
1626        .unwrap();
1627
1628        assert!(matches!(
1629            state.create_buffer_with(context, oversized_buffer, |_, desc| {
1630                calls.set(calls.get() + 1);
1631                Ok::<_, &'static str>((20, buffer_info(desc)))
1632            }),
1633            Err(CreateError::State(DeviceStateError::ResourceLimit))
1634        ));
1635        assert!(matches!(
1636            state.create_program_with(context, state.limits().max_artifact_bytes + 1, 1, |_| {
1637                calls.set(calls.get() + 1);
1638                Ok::<_, &'static str>(30)
1639            }),
1640            Err(CreateError::State(DeviceStateError::ResourceLimit))
1641        ));
1642        let resident_bytes = state.limits().max_artifact_bytes + 1;
1643        let program = state
1644            .create_program_with(context, 1, resident_bytes, |_| {
1645                calls.set(calls.get() + 1);
1646                Ok::<_, &'static str>(30)
1647            })
1648            .unwrap();
1649        assert_eq!(calls.get(), 1);
1650        assert_eq!(state.buffer_count(), 0);
1651        assert_eq!(state.program_count(), 1);
1652        assert_eq!(
1653            state.program_record(program).unwrap().resident_bytes(),
1654            resident_bytes
1655        );
1656        assert_eq!(
1657            state.context_record(context).unwrap().children(),
1658            ChildCounts {
1659                programs: 1,
1660                ..ChildCounts::default()
1661            }
1662        );
1663        assert_eq!(
1664            state.retained_bytes(),
1665            RetainedBytes {
1666                buffer_backing: 0,
1667                program_resident: u128::from(resident_bytes),
1668            }
1669        );
1670    }
1671
1672    #[test]
1673    fn aggregate_policy_uses_actual_backing_and_charges_until_release_commits() {
1674        assert!(ResourcePolicy::new(0, 1).is_none());
1675        assert!(ResourcePolicy::new(1, 0).is_none());
1676
1677        let mut state = TestState::new(
1678            ObjectNamespace::new(1).unwrap(),
1679            limits(1),
1680            ResourcePolicy::new(4096, 4096).unwrap(),
1681        )
1682        .unwrap();
1683        let context = create_context(&mut state, 10);
1684        let desc = buffer_desc();
1685        let padded = BufferInfo::new(
1686            desc,
1687            8192,
1688            64,
1689            BufferProperties::HOST_VISIBLE | BufferProperties::DIRECT_BINDING,
1690        )
1691        .unwrap();
1692        let outcome = state
1693            .create_buffer_with(context, desc, |_, _| Ok::<_, &'static str>((20, padded)))
1694            .unwrap();
1695        let BufferCreateOutcome::CleanupRequired {
1696            id,
1697            error: BackendError::ResourceLimit,
1698        } = outcome
1699        else {
1700            panic!("padded allocation did not require cleanup");
1701        };
1702        assert_eq!(
1703            state.retained_bytes(),
1704            RetainedBytes {
1705                buffer_backing: 8192,
1706                program_resident: 0,
1707            }
1708        );
1709
1710        let resource = state.begin_buffer_release(id).unwrap();
1711        assert_eq!(state.retained_bytes().buffer_backing, 8192);
1712        state.restore_buffer_release(id, resource).unwrap();
1713        let resource = state.begin_buffer_release(id).unwrap();
1714        assert_eq!(resource, 20);
1715        state.commit_buffer_release(id).unwrap();
1716        assert!(state.retained_bytes().is_empty());
1717    }
1718
1719    #[test]
1720    fn program_policy_rejects_before_provider_invocation() {
1721        let mut state = TestState::new(
1722            ObjectNamespace::new(1).unwrap(),
1723            limits(1),
1724            ResourcePolicy::new(4096, 8).unwrap(),
1725        )
1726        .unwrap();
1727        let context = create_context(&mut state, 10);
1728        let called = Cell::new(false);
1729        assert!(matches!(
1730            state.create_program_with(context, 1, 9, |_| {
1731                called.set(true);
1732                Ok::<_, &'static str>(30)
1733            }),
1734            Err(CreateError::State(DeviceStateError::ResourceLimit))
1735        ));
1736        assert!(!called.get());
1737        assert!(state.retained_bytes().is_empty());
1738    }
1739}