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 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
377pub 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 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
430pub 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 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}