1use 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#[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#[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#[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 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 pub const fn program_id(&self) -> ObjectId {
367 self.program_id
368 }
369
370 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}
385pub 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 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
438pub 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
493pub 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 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 pub fn context_ids(&self) -> impl Iterator<Item = ObjectId> + '_ {
610 self.contexts.ids()
611 }
612
613 pub fn buffer_ids(&self) -> impl Iterator<Item = ObjectId> + '_ {
616 self.buffers.ids()
617 }
618
619 pub fn program_ids(&self) -> impl Iterator<Item = ObjectId> + '_ {
621 self.programs.ids()
622 }
623
624 pub fn queue_ids(&self) -> impl Iterator<Item = ObjectId> + '_ {
626 self.queues.ids()
627 }
628
629 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 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 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}