1use alloc::vec::Vec;
2use core::mem::{self, size_of};
3
4use virtio_accel_proto::{
5 AllocateBufferRequest, CreateContextRequest, CreateQueueRequest, Le32, Le64,
6 LoadProgramRequest, ObjectPayload, RequestFlags, RequestHeader, ResponseHeader, StatusCode,
7 SubmitRequest, TransferBufferRequest, WireBinding, WireDeviceInfo, WireEventState, read_exact,
8};
9use virtio_accel_transport::{
10 ByteAccessError, ChainId, DriverChainBuffer, DriverQueue, NotificationRecheck,
11 PublishErrorKind, QueueControl, QueueEpoch, QueueError, QueueSize, QueueState, ReadableBytes,
12 UsedChain, UsedLength,
13};
14use zerocopy::{Immutable, IntoBytes};
15
16use crate::config::{GuestConfig, GuestConfigError};
17use crate::operation::{
18 AllocateBuffer, CancelEvent, CreateContext, CreateExecutionQueue, DestroyContext, DestroyEvent,
19 DestroyExecutionQueue, FreeBuffer, GetDeviceInfo, LoadProgram, Operation, OperationResult,
20 PollEvent, ReadBuffer, ResponseError, Submit, UnloadProgram, WriteBuffer,
21};
22use crate::types::{
23 AccessMode, Binding, Buffer, BufferDesc, BufferRange, BufferUsage, Context, DeviceInfo, Event,
24 ExecutionQueue, FailureDisposition, Program, ProgramDesc,
25};
26
27const REQUEST_HEADER_BYTES: u64 = size_of::<RequestHeader>() as u64;
28const RESPONSE_HEADER_BYTES: u64 = size_of::<ResponseHeader>() as u64;
29const COPY_SCRATCH_BYTES: usize = 256;
30
31#[derive(Clone, Copy, Debug, PartialEq, Eq)]
33pub enum ClientHealth {
34 Running,
36 NeedsReset,
38}
39
40#[derive(Clone, Copy, Debug, PartialEq, Eq)]
42pub enum ClientInitError {
43 Config(GuestConfigError),
45 AllocationFailed,
47}
48
49#[derive(Clone, Copy, Debug, PartialEq, Eq)]
51pub enum QueueSetupError<E> {
52 Config(GuestConfigError),
54 Queue(QueueError<E>),
56}
57
58#[derive(Debug, PartialEq, Eq)]
60pub enum StartErrorKind<QE, CE> {
61 NeedsReset,
63 DiscoveryRequired,
65 InflightLimit,
67 RequestIdExhausted,
69 StaleHandle,
71 InvalidArgument,
73 DeviceLimit,
75 RequestCapacity {
77 required: u64,
79 available: u64,
81 },
82 ResponseCapacity {
84 required: u64,
86 available: u64,
88 },
89 SourceAccess(ByteAccessError),
91 ChainAccess(CE),
93 Publish(PublishErrorKind<QE>),
95}
96
97#[derive(Debug)]
99pub struct StartError<C, O, QE, CE> {
100 chain: C,
101 operation: O,
102 kind: StartErrorKind<QE, CE>,
103}
104
105impl<C, O, QE, CE> StartError<C, O, QE, CE> {
106 pub const fn kind(&self) -> &StartErrorKind<QE, CE> {
108 &self.kind
109 }
110
111 pub fn into_parts(self) -> (C, O, StartErrorKind<QE, CE>) {
113 (self.chain, self.operation, self.kind)
114 }
115}
116
117pub type StartResult<Q, O> = Result<
119 Pending<O>,
120 StartError<
121 <Q as DriverQueue>::Chain,
122 O,
123 <Q as DriverQueue>::Error,
124 <<Q as DriverQueue>::Chain as DriverChainBuffer>::Error,
125 >,
126>;
127
128#[derive(Debug)]
130pub struct Pending<O> {
131 slot: u16,
132 request_id: u64,
133 epoch: QueueEpoch,
134 operation: O,
135}
136
137impl<O> Pending<O> {
138 pub const fn request_id(&self) -> u64 {
140 self.request_id
141 }
142
143 pub const fn epoch(&self) -> QueueEpoch {
145 self.epoch
146 }
147
148 pub const fn operation(&self) -> &O {
150 &self.operation
151 }
152}
153
154#[derive(Debug)]
156pub struct StaleOperation<O> {
157 request_id: u64,
158 epoch: QueueEpoch,
159 operation: O,
160}
161
162impl<O> StaleOperation<O> {
163 pub const fn request_id(&self) -> u64 {
165 self.request_id
166 }
167
168 pub const fn epoch(&self) -> QueueEpoch {
170 self.epoch
171 }
172
173 pub const fn operation(&self) -> &O {
175 &self.operation
176 }
177
178 pub fn into_operation(self) -> O {
182 self.operation
183 }
184}
185
186#[derive(Debug)]
188pub enum Completion<O: Operation, C, CE> {
189 Success {
191 output: O::Output,
193 chain: C,
195 },
196 DeviceError {
198 status: StatusCode,
200 disposition: FailureDisposition,
202 operation: O,
204 chain: C,
206 },
207 InvalidResponse {
209 error: ResponseError<CE>,
211 operation: O,
213 chain: C,
215 },
216}
217
218pub enum RequestPoll<O: Operation, C, QE, CE> {
220 Pending(Pending<O>),
222 Ready(Completion<O, C, CE>),
224 QueueError {
226 pending: Pending<O>,
228 error: QueueError<QE>,
230 },
231 Stale(StaleOperation<O>),
233 NeedsReset(Pending<O>),
235}
236
237#[derive(Clone, Copy, Debug, PartialEq, Eq)]
239pub enum PumpResult {
240 Idle,
242 Completion {
244 request_id: u64,
246 },
247 UnexpectedCompletion,
249 NeedsReset,
251}
252
253enum Slot<C> {
254 Vacant,
255 InFlight { request_id: u64, chain_id: ChainId },
256 Completed { request_id: u64, used: UsedChain<C> },
257 Recovered(UsedChain<C>),
258}
259
260enum FrameWriteError<E> {
261 Chain(E),
262 Source(ByteAccessError),
263}
264
265pub struct GuestClient<Q: DriverQueue>
270where
271 Q::Chain: DriverChainBuffer,
272{
273 queue: Q,
274 config: GuestConfig,
275 device_info: Option<DeviceInfo>,
276 slots: Vec<Slot<Q::Chain>>,
277 next_request_id: u64,
278 health: ClientHealth,
279 unexpected: Option<UsedChain<Q::Chain>>,
280}
281
282impl<Q> GuestClient<Q>
283where
284 Q: DriverQueue,
285 Q::Chain: DriverChainBuffer,
286{
287 pub fn new(queue: Q, config: GuestConfig) -> Result<Self, ClientInitError> {
289 config
290 .validate_queue(queue.state())
291 .map_err(ClientInitError::Config)?;
292 let count = usize::from(config.max_inflight());
293 let mut slots = Vec::new();
294 slots
295 .try_reserve_exact(count)
296 .map_err(|_| ClientInitError::AllocationFailed)?;
297 for _ in 0..count {
298 slots.push(Slot::Vacant);
299 }
300 Ok(Self {
301 queue,
302 config,
303 device_info: None,
304 slots,
305 next_request_id: 1,
306 health: ClientHealth::Running,
307 unexpected: None,
308 })
309 }
310
311 pub fn queue_state(&self) -> QueueState {
313 self.queue.state()
314 }
315
316 pub const fn config(&self) -> GuestConfig {
318 self.config
319 }
320
321 pub const fn device_info(&self) -> Option<DeviceInfo> {
323 self.device_info
324 }
325
326 pub const fn health(&self) -> ClientHealth {
328 self.health
329 }
330
331 pub fn disable_used_notifications(&mut self) -> Result<(), QueueError<Q::Error>> {
333 self.queue.disable_used_notifications()
334 }
335
336 pub fn enable_used_notifications(
338 &mut self,
339 ) -> Result<NotificationRecheck, QueueError<Q::Error>> {
340 self.queue.enable_used_notifications()
341 }
342
343 pub fn pump(&mut self) -> Result<PumpResult, QueueError<Q::Error>> {
345 if self.health == ClientHealth::NeedsReset {
346 return Ok(PumpResult::NeedsReset);
347 }
348 let Some(used) = self.queue.pop_used()? else {
349 return Ok(PumpResult::Idle);
350 };
351 let id = used.id();
352 let Some(index) = self
353 .slots
354 .iter()
355 .position(|slot| matches!(slot, Slot::InFlight { chain_id, .. } if *chain_id == id))
356 else {
357 self.unexpected = Some(used);
358 self.health = ClientHealth::NeedsReset;
359 return Ok(PumpResult::UnexpectedCompletion);
360 };
361 let Slot::InFlight { request_id, .. } = mem::replace(&mut self.slots[index], Slot::Vacant)
362 else {
363 unreachable!("matched slot must be in flight")
364 };
365 self.slots[index] = Slot::Completed { request_id, used };
366 Ok(PumpResult::Completion { request_id })
367 }
368
369 pub fn poll<O: Operation>(
371 &mut self,
372 pending: Pending<O>,
373 ) -> RequestPoll<O, Q::Chain, Q::Error, <Q::Chain as DriverChainBuffer>::Error> {
374 if pending.epoch != self.queue.state().epoch() {
375 return RequestPoll::Stale(stale_operation(pending));
376 }
377 if self.health == ClientHealth::NeedsReset {
378 return RequestPoll::NeedsReset(pending);
379 }
380 if !self.pending_matches(&pending) {
381 return RequestPoll::Stale(stale_operation(pending));
382 }
383 if matches!(self.slots[usize::from(pending.slot)], Slot::InFlight { .. }) {
384 match self.pump() {
385 Err(error) => return RequestPoll::QueueError { pending, error },
386 Ok(PumpResult::UnexpectedCompletion | PumpResult::NeedsReset) => {
387 return RequestPoll::NeedsReset(pending);
388 }
389 Ok(PumpResult::Idle | PumpResult::Completion { .. }) => {}
390 }
391 }
392 let index = usize::from(pending.slot);
393 if !matches!(self.slots[index], Slot::Completed { .. }) {
394 return RequestPoll::Pending(pending);
395 }
396 let Slot::Completed { request_id, used } =
397 mem::replace(&mut self.slots[index], Slot::Vacant)
398 else {
399 unreachable!("checked slot must contain a completion")
400 };
401 if request_id != pending.request_id {
402 self.slots[index] = Slot::Completed { request_id, used };
403 return RequestPoll::Stale(stale_operation(pending));
404 }
405 RequestPoll::Ready(self.decode_completion(pending, used))
406 }
407
408 pub fn reset(&mut self, next_epoch: QueueEpoch) -> Result<Q::Reclaimed, QueueError<Q::Error>> {
413 let reclaimed = self.queue.reset(next_epoch)?;
414 for slot in &mut self.slots {
415 let old = mem::replace(slot, Slot::Vacant);
416 *slot = match old {
417 Slot::Completed { used, .. } | Slot::Recovered(used) => Slot::Recovered(used),
418 Slot::Vacant | Slot::InFlight { .. } => Slot::Vacant,
419 };
420 }
421 self.device_info = None;
422 self.next_request_id = 1;
423 self.health = ClientHealth::Running;
424 Ok(reclaimed)
425 }
426
427 pub fn reconfigure_queue(
429 &mut self,
430 size: QueueSize,
431 ) -> Result<(), QueueSetupError<<Q as DriverQueue>::Error>>
432 where
433 Q: QueueControl<Error = <Q as DriverQueue>::Error>,
434 {
435 self.config
436 .validate_size(size)
437 .map_err(QueueSetupError::Config)?;
438 self.queue.configure(size).map_err(QueueSetupError::Queue)?;
439 self.queue.set_ready(true).map_err(QueueSetupError::Queue)
440 }
441
442 pub fn pop_recovered_completion(&mut self) -> Option<UsedChain<Q::Chain>> {
444 let index = self
445 .slots
446 .iter()
447 .position(|slot| matches!(slot, Slot::Recovered(_)))?;
448 let Slot::Recovered(used) = mem::replace(&mut self.slots[index], Slot::Vacant) else {
449 unreachable!("matched slot must contain a recovered completion")
450 };
451 Some(used)
452 }
453
454 pub fn take_unexpected_completion(&mut self) -> Option<UsedChain<Q::Chain>> {
456 self.unexpected.take()
457 }
458
459 pub fn get_device_info(&mut self, chain: Q::Chain) -> StartResult<Q, GetDeviceInfo> {
461 self.start_frame(
462 chain,
463 GetDeviceInfo,
464 0,
465 RESPONSE_HEADER_BYTES + size_of::<WireDeviceInfo>() as u64,
466 |_| Ok(()),
467 )
468 }
469
470 pub fn create_context(&mut self, chain: Q::Chain) -> StartResult<Q, CreateContext> {
472 let request = CreateContextRequest {
473 flags: Le32::new(0),
474 reserved: Le32::new(0),
475 };
476 self.start_wire(
477 chain,
478 CreateContext,
479 &request,
480 RESPONSE_HEADER_BYTES + size_of::<ObjectPayload>() as u64,
481 )
482 }
483
484 pub fn destroy_context(
486 &mut self,
487 chain: Q::Chain,
488 context: Context,
489 ) -> StartResult<Q, DestroyContext> {
490 let stale = context.epoch() != self.queue.state().epoch();
491 let object = ObjectPayload {
492 object_id: Le64::new(context.raw()),
493 };
494 let operation = DestroyContext { context };
495 if stale {
496 return self.start_failure(chain, operation, StartErrorKind::StaleHandle);
497 }
498 self.start_wire(chain, operation, &object, RESPONSE_HEADER_BYTES)
499 }
500
501 pub fn allocate_buffer(
503 &mut self,
504 chain: Q::Chain,
505 context: &Context,
506 desc: BufferDesc,
507 ) -> StartResult<Q, AllocateBuffer> {
508 let operation = AllocateBuffer {
509 context: context.handle.id(),
510 desc,
511 };
512 if context.epoch() != self.queue.state().epoch() {
513 return self.start_failure(chain, operation, StartErrorKind::StaleHandle);
514 }
515 let Some(info) = self.device_info else {
516 return self.start_failure(chain, operation, StartErrorKind::DiscoveryRequired);
517 };
518 if desc.bytes > info.max_buffer_bytes || !info.supports_domain(desc.memory_domain) {
519 return self.start_failure(chain, operation, StartErrorKind::DeviceLimit);
520 }
521 let request = AllocateBufferRequest {
522 context_id: Le64::new(context.raw()),
523 bytes: Le64::new(desc.bytes),
524 alignment: Le64::new(desc.alignment),
525 memory_domain: desc.memory_domain as u8,
526 reserved0: [0; 7],
527 usage: Le32::new(desc.usage.bits()),
528 reserved1: Le32::new(0),
529 };
530 self.start_wire(
531 chain,
532 operation,
533 &request,
534 RESPONSE_HEADER_BYTES + size_of::<ObjectPayload>() as u64,
535 )
536 }
537
538 pub fn free_buffer(&mut self, chain: Q::Chain, buffer: Buffer) -> StartResult<Q, FreeBuffer> {
540 let stale = buffer.epoch() != self.queue.state().epoch();
541 let object = ObjectPayload {
542 object_id: Le64::new(buffer.raw()),
543 };
544 let operation = FreeBuffer { buffer };
545 if stale {
546 return self.start_failure(chain, operation, StartErrorKind::StaleHandle);
547 }
548 self.start_wire(chain, operation, &object, RESPONSE_HEADER_BYTES)
549 }
550
551 pub fn write_buffer<R: ReadableBytes + ?Sized>(
553 &mut self,
554 chain: Q::Chain,
555 buffer: &Buffer,
556 offset: u64,
557 source: &R,
558 ) -> StartResult<Q, WriteBuffer> {
559 let operation = WriteBuffer;
560 if buffer.epoch() != self.queue.state().epoch() {
561 return self.start_failure(chain, operation, StartErrorKind::StaleHandle);
562 }
563 let Ok(range) = BufferRange::new(offset, source.len()) else {
564 return self.start_failure(chain, operation, StartErrorKind::InvalidArgument);
565 };
566 if !range.fits(buffer.desc.bytes)
567 || !buffer
568 .desc
569 .usage
570 .contains(BufferUsage::TRANSFER_DESTINATION)
571 {
572 return self.start_failure(chain, operation, StartErrorKind::InvalidArgument);
573 }
574 let prefix = TransferBufferRequest {
575 buffer_id: Le64::new(buffer.raw()),
576 offset: Le64::new(offset),
577 bytes: Le64::new(source.len()),
578 };
579 let Some(payload) = (size_of::<TransferBufferRequest>() as u64).checked_add(source.len())
580 else {
581 return self.start_failure(chain, operation, StartErrorKind::DeviceLimit);
582 };
583 self.start_frame(chain, operation, payload, RESPONSE_HEADER_BYTES, |chain| {
584 write_wire(chain, REQUEST_HEADER_BYTES, &prefix).map_err(FrameWriteError::Chain)?;
585 copy_source(
586 chain,
587 REQUEST_HEADER_BYTES + size_of::<TransferBufferRequest>() as u64,
588 source,
589 )
590 })
591 }
592
593 pub fn write_buffer_prepared(
597 &mut self,
598 chain: Q::Chain,
599 buffer: &Buffer,
600 range: BufferRange,
601 ) -> StartResult<Q, WriteBuffer> {
602 let operation = WriteBuffer;
603 if buffer.epoch() != self.queue.state().epoch() {
604 return self.start_failure(chain, operation, StartErrorKind::StaleHandle);
605 }
606 if !range.fits(buffer.desc.bytes)
607 || !buffer
608 .desc
609 .usage
610 .contains(BufferUsage::TRANSFER_DESTINATION)
611 {
612 return self.start_failure(chain, operation, StartErrorKind::InvalidArgument);
613 }
614 let prefix = TransferBufferRequest {
615 buffer_id: Le64::new(buffer.raw()),
616 offset: Le64::new(range.offset),
617 bytes: Le64::new(range.bytes),
618 };
619 let Some(payload) = (size_of::<TransferBufferRequest>() as u64).checked_add(range.bytes)
620 else {
621 return self.start_failure(chain, operation, StartErrorKind::DeviceLimit);
622 };
623 self.start_frame(chain, operation, payload, RESPONSE_HEADER_BYTES, |chain| {
624 write_wire(chain, REQUEST_HEADER_BYTES, &prefix).map_err(FrameWriteError::Chain)
625 })
626 }
627
628 pub fn read_buffer(
630 &mut self,
631 chain: Q::Chain,
632 buffer: &Buffer,
633 range: BufferRange,
634 ) -> StartResult<Q, ReadBuffer> {
635 let operation = ReadBuffer { bytes: range.bytes };
636 if buffer.epoch() != self.queue.state().epoch() {
637 return self.start_failure(chain, operation, StartErrorKind::StaleHandle);
638 }
639 if !range.fits(buffer.desc.bytes)
640 || !buffer.desc.usage.contains(BufferUsage::TRANSFER_SOURCE)
641 {
642 return self.start_failure(chain, operation, StartErrorKind::InvalidArgument);
643 }
644 let request = TransferBufferRequest {
645 buffer_id: Le64::new(buffer.raw()),
646 offset: Le64::new(range.offset),
647 bytes: Le64::new(range.bytes),
648 };
649 let Some(response_bytes) = RESPONSE_HEADER_BYTES.checked_add(range.bytes) else {
650 return self.start_failure(chain, operation, StartErrorKind::DeviceLimit);
651 };
652 self.start_wire(chain, operation, &request, response_bytes)
653 }
654
655 pub fn load_program<R: ReadableBytes + ?Sized>(
657 &mut self,
658 chain: Q::Chain,
659 context: &Context,
660 desc: ProgramDesc,
661 artifact: &R,
662 ) -> StartResult<Q, LoadProgram> {
663 let operation = LoadProgram {
664 context: context.handle.id(),
665 desc,
666 };
667 if context.epoch() != self.queue.state().epoch() {
668 return self.start_failure(chain, operation, StartErrorKind::StaleHandle);
669 }
670 let Some(info) = self.device_info else {
671 return self.start_failure(chain, operation, StartErrorKind::DiscoveryRequired);
672 };
673 if artifact.is_empty() {
674 return self.start_failure(chain, operation, StartErrorKind::InvalidArgument);
675 }
676 if artifact.len() > info.max_artifact_bytes {
677 return self.start_failure(chain, operation, StartErrorKind::DeviceLimit);
678 }
679 let request = LoadProgramRequest {
680 context_id: Le64::new(context.raw()),
681 format: Le32::new(desc.format.get()),
682 flags: Le32::new(0),
683 target: core::array::from_fn(|index| Le32::new(desc.target[index])),
684 payload_bytes: Le64::new(artifact.len()),
685 resident_bytes: Le64::new(desc.resident_bytes.get()),
686 };
687 let payload = size_of::<LoadProgramRequest>() as u64 + artifact.len();
688 self.start_frame(
689 chain,
690 operation,
691 payload,
692 RESPONSE_HEADER_BYTES + size_of::<ObjectPayload>() as u64,
693 |chain| {
694 write_wire(chain, REQUEST_HEADER_BYTES, &request)
695 .map_err(FrameWriteError::Chain)?;
696 copy_source(
697 chain,
698 REQUEST_HEADER_BYTES + size_of::<LoadProgramRequest>() as u64,
699 artifact,
700 )
701 },
702 )
703 }
704
705 pub fn load_program_prepared(
709 &mut self,
710 chain: Q::Chain,
711 context: &Context,
712 desc: ProgramDesc,
713 artifact_bytes: u64,
714 ) -> StartResult<Q, LoadProgram> {
715 let operation = LoadProgram {
716 context: context.handle.id(),
717 desc,
718 };
719 if context.epoch() != self.queue.state().epoch() {
720 return self.start_failure(chain, operation, StartErrorKind::StaleHandle);
721 }
722 let Some(info) = self.device_info else {
723 return self.start_failure(chain, operation, StartErrorKind::DiscoveryRequired);
724 };
725 if artifact_bytes == 0 {
726 return self.start_failure(chain, operation, StartErrorKind::InvalidArgument);
727 }
728 if artifact_bytes > info.max_artifact_bytes {
729 return self.start_failure(chain, operation, StartErrorKind::DeviceLimit);
730 }
731 let request = LoadProgramRequest {
732 context_id: Le64::new(context.raw()),
733 format: Le32::new(desc.format.get()),
734 flags: Le32::new(0),
735 target: core::array::from_fn(|index| Le32::new(desc.target[index])),
736 payload_bytes: Le64::new(artifact_bytes),
737 resident_bytes: Le64::new(desc.resident_bytes.get()),
738 };
739 self.start_frame(
740 chain,
741 operation,
742 size_of::<LoadProgramRequest>() as u64 + artifact_bytes,
743 RESPONSE_HEADER_BYTES + size_of::<ObjectPayload>() as u64,
744 |chain| {
745 write_wire(chain, REQUEST_HEADER_BYTES, &request).map_err(FrameWriteError::Chain)
746 },
747 )
748 }
749
750 pub fn unload_program(
752 &mut self,
753 chain: Q::Chain,
754 program: Program,
755 ) -> StartResult<Q, UnloadProgram> {
756 let stale = program.epoch() != self.queue.state().epoch();
757 let object = ObjectPayload {
758 object_id: Le64::new(program.raw()),
759 };
760 let operation = UnloadProgram { program };
761 if stale {
762 return self.start_failure(chain, operation, StartErrorKind::StaleHandle);
763 }
764 self.start_wire(chain, operation, &object, RESPONSE_HEADER_BYTES)
765 }
766
767 pub fn create_execution_queue(
769 &mut self,
770 chain: Q::Chain,
771 context: &Context,
772 ) -> StartResult<Q, CreateExecutionQueue> {
773 let operation = CreateExecutionQueue {
774 context: context.handle.id(),
775 };
776 if context.epoch() != self.queue.state().epoch() {
777 return self.start_failure(chain, operation, StartErrorKind::StaleHandle);
778 }
779 let request = CreateQueueRequest {
780 context_id: Le64::new(context.raw()),
781 flags: Le32::new(0),
782 reserved: Le32::new(0),
783 };
784 self.start_wire(
785 chain,
786 operation,
787 &request,
788 RESPONSE_HEADER_BYTES + size_of::<ObjectPayload>() as u64,
789 )
790 }
791
792 pub fn destroy_execution_queue(
794 &mut self,
795 chain: Q::Chain,
796 queue: ExecutionQueue,
797 ) -> StartResult<Q, DestroyExecutionQueue> {
798 let stale = queue.epoch() != self.queue.state().epoch();
799 let object = ObjectPayload {
800 object_id: Le64::new(queue.raw()),
801 };
802 let operation = DestroyExecutionQueue { queue };
803 if stale {
804 return self.start_failure(chain, operation, StartErrorKind::StaleHandle);
805 }
806 self.start_wire(chain, operation, &object, RESPONSE_HEADER_BYTES)
807 }
808
809 pub fn submit(
811 &mut self,
812 chain: Q::Chain,
813 queue: &ExecutionQueue,
814 program: &Program,
815 bindings: &[Binding<'_>],
816 timeout_ns: u64,
817 ) -> StartResult<Q, Submit> {
818 let context = queue
819 .context
820 .expect("execution queue always retains context");
821 let operation = Submit { context };
822 let epoch = self.queue.state().epoch();
823 if queue.epoch() != epoch || program.epoch() != epoch {
824 return self.start_failure(chain, operation, StartErrorKind::StaleHandle);
825 }
826 if program.context != Some(context) || bindings.is_empty() {
827 return self.start_failure(chain, operation, StartErrorKind::InvalidArgument);
828 }
829 let Some(info) = self.device_info else {
830 return self.start_failure(chain, operation, StartErrorKind::DiscoveryRequired);
831 };
832 let Ok(binding_count) = u32::try_from(bindings.len()) else {
833 return self.start_failure(chain, operation, StartErrorKind::DeviceLimit);
834 };
835 if binding_count > info.max_bindings_per_submission {
836 return self.start_failure(chain, operation, StartErrorKind::DeviceLimit);
837 }
838 let bindings_are_canonical = bindings.windows(2).all(|pair| pair[0].slot < pair[1].slot);
842 for (index, binding) in bindings.iter().enumerate() {
843 if binding.buffer.epoch() != epoch {
844 return self.start_failure(chain, operation, StartErrorKind::StaleHandle);
845 }
846 if binding.buffer.context != context
847 || !binding.range.fits(binding.buffer.desc.bytes)
848 || !usage_allows(binding.buffer.desc.usage, binding.access)
849 || (!bindings_are_canonical
850 && bindings[..index]
851 .iter()
852 .any(|prior| prior.slot == binding.slot))
853 {
854 return self.start_failure(chain, operation, StartErrorKind::InvalidArgument);
855 }
856 }
857 let request = SubmitRequest {
858 queue_id: Le64::new(queue.raw()),
859 program_id: Le64::new(program.raw()),
860 binding_count: Le32::new(binding_count),
861 flags: Le32::new(0),
862 timeout_ns: Le64::new(timeout_ns),
863 };
864 let payload = size_of::<SubmitRequest>() as u64
865 + size_of::<WireBinding>() as u64 * u64::from(binding_count);
866 self.start_frame(
867 chain,
868 operation,
869 payload,
870 RESPONSE_HEADER_BYTES + size_of::<virtio_accel_proto::SubmitResponse>() as u64,
871 |chain| {
872 write_wire(chain, REQUEST_HEADER_BYTES, &request)
873 .map_err(FrameWriteError::Chain)?;
874 let mut offset = REQUEST_HEADER_BYTES + size_of::<SubmitRequest>() as u64;
875 for binding in bindings {
876 let wire = WireBinding {
877 buffer_id: Le64::new(binding.buffer.raw()),
878 offset: Le64::new(binding.range.offset),
879 bytes: Le64::new(binding.range.bytes),
880 slot: Le32::new(binding.slot),
881 access: binding.access as u8,
882 reserved: [0; 3],
883 };
884 write_wire(chain, offset, &wire).map_err(FrameWriteError::Chain)?;
885 offset += size_of::<WireBinding>() as u64;
886 }
887 Ok(())
888 },
889 )
890 }
891
892 pub fn poll_event(&mut self, chain: Q::Chain, event: &Event) -> StartResult<Q, PollEvent> {
894 if event.epoch() != self.queue.state().epoch() {
895 return self.start_failure(chain, PollEvent, StartErrorKind::StaleHandle);
896 }
897 let request = ObjectPayload {
898 object_id: Le64::new(event.raw()),
899 };
900 self.start_wire(
901 chain,
902 PollEvent,
903 &request,
904 RESPONSE_HEADER_BYTES + size_of::<WireEventState>() as u64,
905 )
906 }
907
908 pub fn cancel_event(&mut self, chain: Q::Chain, event: &Event) -> StartResult<Q, CancelEvent> {
910 if event.epoch() != self.queue.state().epoch() {
911 return self.start_failure(chain, CancelEvent, StartErrorKind::StaleHandle);
912 }
913 let request = ObjectPayload {
914 object_id: Le64::new(event.raw()),
915 };
916 self.start_wire(chain, CancelEvent, &request, RESPONSE_HEADER_BYTES)
917 }
918
919 pub fn destroy_event(&mut self, chain: Q::Chain, event: Event) -> StartResult<Q, DestroyEvent> {
921 let stale = event.epoch() != self.queue.state().epoch();
922 let object = ObjectPayload {
923 object_id: Le64::new(event.raw()),
924 };
925 let operation = DestroyEvent { event };
926 if stale {
927 return self.start_failure(chain, operation, StartErrorKind::StaleHandle);
928 }
929 self.start_wire(chain, operation, &object, RESPONSE_HEADER_BYTES)
930 }
931
932 fn start_wire<O: Operation, T: IntoBytes + Immutable>(
933 &mut self,
934 chain: Q::Chain,
935 operation: O,
936 value: &T,
937 response_bytes: u64,
938 ) -> StartResult<Q, O> {
939 self.start_frame(
940 chain,
941 operation,
942 size_of::<T>() as u64,
943 response_bytes,
944 |chain| write_wire(chain, REQUEST_HEADER_BYTES, value).map_err(FrameWriteError::Chain),
945 )
946 }
947
948 fn start_frame<O: Operation>(
949 &mut self,
950 mut chain: Q::Chain,
951 operation: O,
952 payload_bytes: u64,
953 response_bytes: u64,
954 write_payload: impl FnOnce(
955 &mut Q::Chain,
956 ) -> Result<
957 (),
958 FrameWriteError<<Q::Chain as DriverChainBuffer>::Error>,
959 >,
960 ) -> StartResult<Q, O> {
961 if self.health == ClientHealth::NeedsReset {
962 return self.start_failure(chain, operation, StartErrorKind::NeedsReset);
963 }
964 if operation.requires_discovery() && self.device_info.is_none() {
965 return self.start_failure(chain, operation, StartErrorKind::DiscoveryRequired);
966 }
967 let Some(slot) = self
968 .slots
969 .iter()
970 .position(|slot| matches!(slot, Slot::Vacant))
971 else {
972 return self.start_failure(chain, operation, StartErrorKind::InflightLimit);
973 };
974 let readable_bytes = chain.device_readable_len();
975 let writable_bytes = chain.device_writable_len();
976 let Some(frame_bytes) = REQUEST_HEADER_BYTES.checked_add(payload_bytes) else {
977 return self.start_failure(
978 chain,
979 operation,
980 StartErrorKind::RequestCapacity {
981 required: u64::MAX,
982 available: readable_bytes,
983 },
984 );
985 };
986 let max_request = u64::from(self.config.wire().max_request_bytes.get());
987 let max_response = u64::from(self.config.wire().max_response_bytes.get());
988 if frame_bytes > max_request || frame_bytes != readable_bytes {
989 return self.start_failure(
990 chain,
991 operation,
992 StartErrorKind::RequestCapacity {
993 required: frame_bytes,
994 available: readable_bytes,
995 },
996 );
997 }
998 if response_bytes > max_response || writable_bytes < response_bytes {
999 return self.start_failure(
1000 chain,
1001 operation,
1002 StartErrorKind::ResponseCapacity {
1003 required: response_bytes,
1004 available: writable_bytes,
1005 },
1006 );
1007 }
1008 let Ok(payload_bytes) = u32::try_from(payload_bytes) else {
1009 return self.start_failure(
1010 chain,
1011 operation,
1012 StartErrorKind::RequestCapacity {
1013 required: frame_bytes,
1014 available: readable_bytes,
1015 },
1016 );
1017 };
1018 let Some(request_id) = self.allocate_request_id() else {
1019 return self.start_failure(chain, operation, StartErrorKind::RequestIdExhausted);
1020 };
1021 let header = RequestHeader::new(
1022 operation.opcode(),
1023 RequestFlags::empty(),
1024 payload_bytes,
1025 request_id,
1026 );
1027 if let Err(error) = write_wire(&mut chain, 0, &header) {
1028 return self.start_failure(chain, operation, StartErrorKind::ChainAccess(error));
1029 }
1030 if let Err(error) = write_payload(&mut chain) {
1031 let kind = match error {
1032 FrameWriteError::Chain(error) => StartErrorKind::ChainAccess(error),
1033 FrameWriteError::Source(error) => StartErrorKind::SourceAccess(error),
1034 };
1035 return self.start_failure(chain, operation, kind);
1036 }
1037 let published = match self.queue.publish(chain) {
1038 Ok(published) => published,
1039 Err(error) => {
1040 let (chain, kind) = error.into_parts();
1041 return self.start_failure(chain, operation, StartErrorKind::Publish(kind));
1042 }
1043 };
1044 let epoch = published.id().epoch();
1045 self.slots[slot] = Slot::InFlight {
1046 request_id,
1047 chain_id: published.id(),
1048 };
1049 Ok(Pending {
1050 slot: slot as u16,
1051 request_id,
1052 epoch,
1053 operation,
1054 })
1055 }
1056
1057 fn start_failure<O>(
1058 &self,
1059 chain: Q::Chain,
1060 operation: O,
1061 kind: StartErrorKind<Q::Error, <Q::Chain as DriverChainBuffer>::Error>,
1062 ) -> StartResult<Q, O> {
1063 Err(StartError {
1064 chain,
1065 operation,
1066 kind,
1067 })
1068 }
1069
1070 fn allocate_request_id(&mut self) -> Option<u64> {
1071 for _ in 0..=self.slots.len() {
1072 let candidate = self.next_request_id.max(1);
1073 self.next_request_id = candidate.wrapping_add(1).max(1);
1074 if !self.slots.iter().any(|slot| {
1075 matches!(slot, Slot::InFlight { request_id, .. } | Slot::Completed { request_id, .. } if *request_id == candidate)
1076 }) {
1077 return Some(candidate);
1078 }
1079 }
1080 None
1081 }
1082
1083 fn pending_matches<O>(&self, pending: &Pending<O>) -> bool {
1084 self.slots
1085 .get(usize::from(pending.slot))
1086 .is_some_and(|slot| {
1087 matches!(slot, Slot::InFlight { request_id, .. } | Slot::Completed { request_id, .. } if *request_id == pending.request_id)
1088 })
1089 }
1090
1091 fn decode_completion<O: Operation>(
1092 &mut self,
1093 pending: Pending<O>,
1094 used: UsedChain<Q::Chain>,
1095 ) -> Completion<O, Q::Chain, <Q::Chain as DriverChainBuffer>::Error> {
1096 let (_, used_length, chain) = used.into_parts();
1097 let result = self.validate_response(&pending, &chain, used_length);
1098 match result {
1099 Ok(OperationResult::Success(output)) => {
1100 if O::output_requires_reset(&output) {
1101 self.health = ClientHealth::NeedsReset;
1102 }
1103 if let Some(info) = O::discovered_info(&output) {
1104 self.device_info = Some(info);
1105 }
1106 Completion::Success { output, chain }
1107 }
1108 Ok(OperationResult::DeviceError(status)) => {
1109 let disposition = O::failure_disposition(status);
1110 if disposition == FailureDisposition::Indeterminate {
1111 self.health = ClientHealth::NeedsReset;
1112 }
1113 Completion::DeviceError {
1114 status,
1115 disposition,
1116 operation: pending.operation,
1117 chain,
1118 }
1119 }
1120 Err(error) => {
1121 self.health = ClientHealth::NeedsReset;
1122 Completion::InvalidResponse {
1123 error,
1124 operation: pending.operation,
1125 chain,
1126 }
1127 }
1128 }
1129 }
1130
1131 fn validate_response<O: Operation>(
1132 &self,
1133 pending: &Pending<O>,
1134 chain: &Q::Chain,
1135 used: UsedLength,
1136 ) -> Result<OperationResult<O::Output>, ResponseError<<Q::Chain as DriverChainBuffer>::Error>>
1137 {
1138 if u64::from(used.get()) < RESPONSE_HEADER_BYTES {
1139 return Err(ResponseError::UsedLength { used: used.get() });
1140 }
1141 if used.get() > self.config.wire().max_response_bytes.get() {
1142 return Err(ResponseError::ResponseLimit);
1143 }
1144 let mut bytes = [0_u8; size_of::<ResponseHeader>()];
1145 chain
1146 .read_device_writable(0, &mut bytes)
1147 .map_err(ResponseError::HeaderAccess)?;
1148 let header =
1149 read_exact::<ResponseHeader>(&bytes).map_err(|_| ResponseError::PayloadEncoding)?;
1150 let actual_id = header.request_id.get();
1151 if actual_id != pending.request_id {
1152 return Err(ResponseError::RequestId {
1153 expected: pending.request_id,
1154 actual: actual_id,
1155 });
1156 }
1157 if header.flags.get() != 0 {
1158 return Err(ResponseError::Flags(header.flags.get()));
1159 }
1160 let expected_used = RESPONSE_HEADER_BYTES
1161 .checked_add(u64::from(header.payload_bytes.get()))
1162 .ok_or(ResponseError::UsedLength { used: used.get() })?;
1163 if expected_used != u64::from(used.get()) {
1164 return Err(ResponseError::UsedLength { used: used.get() });
1165 }
1166 pending.operation.decode(
1167 chain,
1168 StatusCode(header.status.get()),
1169 header.payload_bytes.get(),
1170 pending.epoch,
1171 self.config,
1172 )
1173 }
1174
1175 #[cfg(test)]
1176 fn set_next_request_id(&mut self, value: u64) {
1177 self.next_request_id = value;
1178 }
1179}
1180
1181fn write_wire<C: DriverChainBuffer, T: IntoBytes + Immutable>(
1182 chain: &mut C,
1183 offset: u64,
1184 value: &T,
1185) -> Result<(), C::Error> {
1186 chain.write_device_readable(offset, value.as_bytes())
1187}
1188
1189fn copy_source<C: DriverChainBuffer, R: ReadableBytes + ?Sized>(
1190 chain: &mut C,
1191 destination_offset: u64,
1192 source: &R,
1193) -> Result<(), FrameWriteError<C::Error>> {
1194 if let Some(bytes) = source.as_contiguous() {
1195 return chain
1196 .write_device_readable(destination_offset, bytes)
1197 .map_err(FrameWriteError::Chain);
1198 }
1199 let mut scratch = [0_u8; COPY_SCRATCH_BYTES];
1200 let mut copied = 0_u64;
1201 while copied < source.len() {
1202 let count = usize::try_from(core::cmp::min(
1203 source.len() - copied,
1204 COPY_SCRATCH_BYTES as u64,
1205 ))
1206 .expect("bounded copy chunk fits usize");
1207 source
1208 .read_at(copied, &mut scratch[..count])
1209 .map_err(FrameWriteError::Source)?;
1210 chain
1211 .write_device_readable(destination_offset + copied, &scratch[..count])
1212 .map_err(FrameWriteError::Chain)?;
1213 copied += count as u64;
1214 }
1215 Ok(())
1216}
1217
1218fn usage_allows(usage: BufferUsage, access: AccessMode) -> bool {
1219 match access {
1220 AccessMode::Read => {
1221 usage.intersects(BufferUsage::PROGRAM_INPUT | BufferUsage::MUTABLE_STATE)
1222 }
1223 AccessMode::Write => {
1224 usage.intersects(BufferUsage::PROGRAM_OUTPUT | BufferUsage::MUTABLE_STATE)
1225 }
1226 AccessMode::ReadWrite => usage.contains(BufferUsage::MUTABLE_STATE),
1227 }
1228}
1229
1230fn stale_operation<O>(pending: Pending<O>) -> StaleOperation<O> {
1231 StaleOperation {
1232 request_id: pending.request_id,
1233 epoch: pending.epoch,
1234 operation: pending.operation,
1235 }
1236}
1237
1238#[cfg(test)]
1239mod tests {
1240 extern crate std;
1241
1242 use core::num::{NonZeroU32, NonZeroU64};
1243 use std::vec;
1244 use std::vec::Vec;
1245
1246 use virtio_accel_proto::{Le16, PROTOCOL_MAJOR, PROTOCOL_MINOR, SubmitResponse, WireConfig};
1247 use virtio_accel_split_queue::{Descriptor, DriverChain, SplitDeviceChain, SplitQueue};
1248 use virtio_accel_transport::{
1249 DeviceChain, DeviceQueue, DriverChainBuffer, QueueControl, QueueSize, WritableBytes,
1250 };
1251
1252 use super::*;
1253 use crate::types::{DeviceInfoError, MemoryDomain, SubmissionOutcome};
1254
1255 type TestClient = GuestClient<SplitQueue>;
1256
1257 fn test_client(size: u16, max_inflight: u16, max_chain_descriptors: u16) -> TestClient {
1258 let size = QueueSize::new(size).unwrap();
1259 let mut queue = SplitQueue::new(size, max_chain_descriptors).unwrap();
1260 QueueControl::configure(&mut queue, size).unwrap();
1261 QueueControl::set_ready(&mut queue, true).unwrap();
1262 let config = GuestConfig::new(
1263 WireConfig {
1264 protocol_major: Le16::new(PROTOCOL_MAJOR),
1265 protocol_minor: Le16::new(PROTOCOL_MINOR),
1266 command_queue_count: Le16::new(1),
1267 max_chain_descriptors: Le16::new(max_chain_descriptors),
1268 max_request_bytes: Le32::new(4_096),
1269 max_response_bytes: Le32::new(4_096),
1270 },
1271 0,
1272 max_inflight,
1273 )
1274 .unwrap();
1275 GuestClient::new(queue, config).unwrap()
1276 }
1277
1278 fn chain(request_bytes: usize, response_bytes: usize) -> DriverChain {
1279 DriverChain::direct(vec![
1280 Descriptor::readable(vec![0; request_bytes]),
1281 Descriptor::writable(vec![0; response_bytes]),
1282 ])
1283 .unwrap()
1284 }
1285
1286 fn prepared_chain(prefix_bytes: usize, payload: &[u8], response_bytes: usize) -> DriverChain {
1287 DriverChain::direct(vec![
1288 Descriptor::readable(vec![0; prefix_bytes]),
1289 Descriptor::readable(payload.to_vec()),
1290 Descriptor::writable(vec![0; response_bytes]),
1291 ])
1292 .unwrap()
1293 }
1294
1295 fn device_info() -> WireDeviceInfo {
1296 WireDeviceInfo {
1297 uuid: [0x5a; 16],
1298 class: Le16::new(1),
1299 reserved: Le16::new(0),
1300 vendor_id: Le32::new(0x1234),
1301 device_id: Le32::new(0x5678),
1302 capabilities: Le64::new((1 << 0) | (1 << 1) | (1 << 2) | (1 << 5)),
1303 max_contexts: Le32::new(8),
1304 max_buffers_per_context: Le32::new(8),
1305 max_programs_per_context: Le32::new(8),
1306 max_queues_per_context: Le32::new(8),
1307 max_events_per_context: Le32::new(8),
1308 max_bindings_per_submission: Le32::new(8),
1309 max_buffer_bytes: Le64::new(4_096),
1310 max_artifact_bytes: Le64::new(4_000),
1311 }
1312 }
1313
1314 fn pop_device_chain(client: &mut TestClient) -> SplitDeviceChain {
1315 DeviceQueue::pop_available(&mut client.queue)
1316 .unwrap()
1317 .expect("published request")
1318 }
1319
1320 fn complete_chain(
1321 client: &mut TestClient,
1322 mut chain: SplitDeviceChain,
1323 status: StatusCode,
1324 payload: &[u8],
1325 response_id: Option<u64>,
1326 ) -> Vec<u8> {
1327 let request = {
1328 let (_, request, response) = chain.io().unwrap().into_parts();
1329 let mut request_bytes = vec![0; request.len() as usize];
1330 request.read_at(0, &mut request_bytes).unwrap();
1331 let header =
1332 read_exact::<RequestHeader>(&request_bytes[..size_of::<RequestHeader>()]).unwrap();
1333 let response_header = ResponseHeader::new(
1334 status,
1335 payload.len() as u32,
1336 response_id.unwrap_or_else(|| header.request_id.get()),
1337 );
1338 response.write_at(0, response_header.as_bytes()).unwrap();
1339 response.write_at(RESPONSE_HEADER_BYTES, payload).unwrap();
1340 request_bytes
1341 };
1342 DeviceQueue::complete(
1343 &mut client.queue,
1344 chain,
1345 UsedLength::new((RESPONSE_HEADER_BYTES as usize + payload.len()) as u32),
1346 )
1347 .unwrap();
1348 request
1349 }
1350
1351 fn complete_next(client: &mut TestClient, status: StatusCode, payload: &[u8]) -> Vec<u8> {
1352 let chain = pop_device_chain(client);
1353 complete_chain(client, chain, status, payload, None)
1354 }
1355
1356 fn success<O: Operation>(
1357 client: &mut TestClient,
1358 pending: Pending<O>,
1359 ) -> (O::Output, DriverChain) {
1360 match client.poll(pending) {
1361 RequestPoll::Ready(Completion::Success { output, chain }) => (output, chain),
1362 RequestPoll::Pending(_) => panic!("request remained pending"),
1363 RequestPoll::Ready(Completion::DeviceError { .. }) => panic!("unexpected device error"),
1364 RequestPoll::Ready(Completion::InvalidResponse { .. }) => {
1365 panic!("unexpected invalid response")
1366 }
1367 RequestPoll::QueueError { .. } => panic!("unexpected queue error"),
1368 RequestPoll::Stale(_) => panic!("unexpected stale request"),
1369 RequestPoll::NeedsReset(_) => panic!("unexpected reset requirement"),
1370 }
1371 }
1372
1373 fn discover(client: &mut TestClient) {
1374 let pending = client
1375 .get_device_info(chain(
1376 size_of::<RequestHeader>(),
1377 size_of::<ResponseHeader>() + size_of::<WireDeviceInfo>(),
1378 ))
1379 .unwrap();
1380 let info = device_info();
1381 complete_next(client, StatusCode::OK, info.as_bytes());
1382 let (decoded, _) = success(client, pending);
1383 assert_eq!(decoded.uuid, [0x5a; 16]);
1384 }
1385
1386 fn object_payload(id: u64) -> ObjectPayload {
1387 ObjectPayload {
1388 object_id: Le64::new(id),
1389 }
1390 }
1391
1392 #[test]
1393 fn out_of_order_completions_match_request_ids() {
1394 let mut client = test_client(16, 8, 4);
1395 let first = client
1396 .get_device_info(chain(16, 16 + size_of::<WireDeviceInfo>()))
1397 .unwrap();
1398 let second = client
1399 .get_device_info(chain(16, 16 + size_of::<WireDeviceInfo>()))
1400 .unwrap();
1401 let first_chain = pop_device_chain(&mut client);
1402 let second_chain = pop_device_chain(&mut client);
1403 let info = device_info();
1404 complete_chain(
1405 &mut client,
1406 second_chain,
1407 StatusCode::OK,
1408 info.as_bytes(),
1409 None,
1410 );
1411
1412 let first = match client.poll(first) {
1413 RequestPoll::Pending(pending) => pending,
1414 _ => panic!("first request completed out of order"),
1415 };
1416 success(&mut client, second);
1417 complete_chain(
1418 &mut client,
1419 first_chain,
1420 StatusCode::OK,
1421 info.as_bytes(),
1422 None,
1423 );
1424 success(&mut client, first);
1425 assert_eq!(client.health(), ClientHealth::Running);
1426 }
1427
1428 #[test]
1429 fn request_id_wrap_skips_live_ids() {
1430 let mut client = test_client(16, 8, 4);
1431 client.set_next_request_id(u64::MAX);
1432 let first = client.get_device_info(chain(16, 92)).unwrap();
1433 let second = client.get_device_info(chain(16, 92)).unwrap();
1434 client.set_next_request_id(u64::MAX);
1435 let third = client.get_device_info(chain(16, 92)).unwrap();
1436 assert_eq!(first.request_id(), u64::MAX);
1437 assert_eq!(second.request_id(), 1);
1438 assert_eq!(third.request_id(), 2);
1439
1440 let next = client.queue_state().epoch().checked_next().unwrap();
1441 assert_eq!(client.reset(next).unwrap().count(), 3);
1442 }
1443
1444 #[test]
1445 fn malformed_and_unknown_responses_never_create_values() {
1446 let mut client = test_client(16, 8, 4);
1447 let malformed = client.get_device_info(chain(16, 92)).unwrap();
1448 let device_chain = pop_device_chain(&mut client);
1449 let info = device_info();
1450 complete_chain(
1451 &mut client,
1452 device_chain,
1453 StatusCode::OK,
1454 info.as_bytes(),
1455 Some(malformed.request_id() + 1),
1456 );
1457 match client.poll(malformed) {
1458 RequestPoll::Ready(Completion::InvalidResponse {
1459 error: ResponseError::RequestId { .. },
1460 ..
1461 }) => {}
1462 _ => panic!("wrong request ID was accepted"),
1463 }
1464 assert_eq!(client.health(), ClientHealth::NeedsReset);
1465 let error = client.get_device_info(chain(16, 92)).unwrap_err();
1466 assert!(matches!(error.kind(), StartErrorKind::NeedsReset));
1467
1468 let next = client.queue_state().epoch().checked_next().unwrap();
1469 client.reset(next).unwrap();
1470 client
1471 .reconfigure_queue(QueueSize::new(16).unwrap())
1472 .unwrap();
1473 let unknown = client.get_device_info(chain(16, 92)).unwrap();
1474 complete_next(&mut client, StatusCode(0x1234), &[]);
1475 match client.poll(unknown) {
1476 RequestPoll::Ready(Completion::DeviceError {
1477 status,
1478 disposition,
1479 ..
1480 }) => {
1481 assert_eq!(status, StatusCode(0x1234));
1482 assert_eq!(disposition, FailureDisposition::Unknown);
1483 }
1484 _ => panic!("unknown status was not preserved"),
1485 }
1486 assert!(client.device_info().is_none());
1487 }
1488
1489 #[test]
1490 fn discovery_rejects_a_device_without_a_baseline_memory_domain() {
1491 let mut client = test_client(16, 8, 4);
1492 let pending = client.get_device_info(chain(16, 92)).unwrap();
1493 let mut info = device_info();
1494 info.capabilities = Le64::new(1 << 2);
1495 complete_next(&mut client, StatusCode::OK, info.as_bytes());
1496
1497 assert!(matches!(
1498 client.poll(pending),
1499 RequestPoll::Ready(Completion::InvalidResponse {
1500 error: ResponseError::DeviceInfo(DeviceInfoError::MissingMemoryDomain),
1501 ..
1502 })
1503 ));
1504 assert!(client.device_info().is_none());
1505 }
1506
1507 #[test]
1508 fn reset_returns_queue_and_operation_ownership() {
1509 let mut client = test_client(16, 8, 4);
1510 discover(&mut client);
1511 let create = client
1512 .create_context(chain(16 + size_of::<CreateContextRequest>(), 24))
1513 .unwrap();
1514 complete_next(&mut client, StatusCode::OK, object_payload(7).as_bytes());
1515 let (context, _) = success(&mut client, create);
1516
1517 let rejected = client.destroy_context(chain(24, 16), context).unwrap();
1518 complete_next(&mut client, StatusCode::BUSY, &[]);
1519 let context = match client.poll(rejected) {
1520 RequestPoll::Ready(Completion::DeviceError {
1521 status: StatusCode::BUSY,
1522 disposition: FailureDisposition::Retryable,
1523 operation,
1524 ..
1525 }) => operation.into_context(),
1526 _ => panic!("rejected release was not returned as retryable"),
1527 };
1528
1529 let destroy = client.destroy_context(chain(24, 16), context).unwrap();
1530 let next = client.queue_state().epoch().checked_next().unwrap();
1531 assert_eq!(client.reset(next).unwrap().count(), 1);
1532 let stale = match client.poll(destroy) {
1533 RequestPoll::Stale(pending) => pending,
1534 _ => panic!("reset did not stale pending release"),
1535 };
1536 assert_eq!(stale.into_operation().into_context().raw(), 7);
1537 }
1538
1539 #[test]
1540 fn reset_recovers_a_completion_already_popped_from_the_used_ring() {
1541 let mut client = test_client(16, 8, 4);
1542 let pending = client.get_device_info(chain(16, 92)).unwrap();
1543 let info = device_info();
1544 complete_next(&mut client, StatusCode::OK, info.as_bytes());
1545 assert_eq!(
1546 client.pump().unwrap(),
1547 PumpResult::Completion {
1548 request_id: pending.request_id()
1549 }
1550 );
1551
1552 let next = client.queue_state().epoch().checked_next().unwrap();
1553 assert_eq!(client.reset(next).unwrap().count(), 0);
1554 assert!(matches!(client.poll(pending), RequestPoll::Stale(_)));
1555 let recovered = client.pop_recovered_completion().unwrap();
1556 assert_eq!(recovered.used().get(), 92);
1557 assert!(client.pop_recovered_completion().is_none());
1558 }
1559
1560 #[test]
1561 fn device_loss_makes_release_ownership_indeterminate() {
1562 let mut client = test_client(16, 8, 4);
1563 discover(&mut client);
1564 let create = client.create_context(chain(24, 24)).unwrap();
1565 complete_next(&mut client, StatusCode::OK, object_payload(9).as_bytes());
1566 let (context, _) = success(&mut client, create);
1567
1568 let destroy = client.destroy_context(chain(24, 16), context).unwrap();
1569 complete_next(&mut client, StatusCode::DEVICE_LOST, &[]);
1570 match client.poll(destroy) {
1571 RequestPoll::Ready(Completion::DeviceError {
1572 status: StatusCode::DEVICE_LOST,
1573 disposition: FailureDisposition::Indeterminate,
1574 operation,
1575 ..
1576 }) => assert_eq!(operation.into_context().raw(), 9),
1577 _ => panic!("device loss did not preserve indeterminate ownership"),
1578 }
1579 assert_eq!(client.health(), ClientHealth::NeedsReset);
1580 }
1581
1582 #[test]
1583 fn prepublication_backpressure_returns_chain_and_operation() {
1584 let mut client = test_client(4, 4, 2);
1585 let _first = client.get_device_info(chain(16, 92)).unwrap();
1586 let _second = client.get_device_info(chain(16, 92)).unwrap();
1587 let error = client.get_device_info(chain(16, 92)).unwrap_err();
1588 let (returned, _operation, kind) = error.into_parts();
1589 assert_eq!(returned.device_readable_len(), 16);
1590 assert!(matches!(
1591 kind,
1592 StartErrorKind::Publish(PublishErrorKind::InsufficientDescriptors)
1593 ));
1594 }
1595
1596 #[test]
1597 fn complete_typed_lifecycle_uses_caller_owned_chain_storage() {
1598 let mut client = test_client(16, 8, 4);
1599 discover(&mut client);
1600
1601 let create_context = client.create_context(chain(24, 24)).unwrap();
1602 complete_next(&mut client, StatusCode::OK, object_payload(1).as_bytes());
1603 let (context, _) = success(&mut client, create_context);
1604
1605 let desc = BufferDesc::new(
1606 64,
1607 16,
1608 MemoryDomain::Host,
1609 BufferUsage::TRANSFER_SOURCE
1610 | BufferUsage::TRANSFER_DESTINATION
1611 | BufferUsage::PROGRAM_INPUT
1612 | BufferUsage::PROGRAM_OUTPUT,
1613 )
1614 .unwrap();
1615 let allocate = client
1616 .allocate_buffer(
1617 chain(16 + size_of::<AllocateBufferRequest>(), 24),
1618 &context,
1619 desc,
1620 )
1621 .unwrap();
1622 complete_next(&mut client, StatusCode::OK, object_payload(2).as_bytes());
1623 let (buffer, _) = success(&mut client, allocate);
1624
1625 let copied = b"copy";
1626 let write = client
1627 .write_buffer(
1628 chain(16 + size_of::<TransferBufferRequest>() + copied.len(), 16),
1629 &buffer,
1630 0,
1631 &copied[..],
1632 )
1633 .unwrap();
1634 let request = complete_next(&mut client, StatusCode::OK, &[]);
1635 assert_eq!(&request[40..], copied);
1636 success(&mut client, write);
1637
1638 let prepared = b"direct";
1639 let write = client
1640 .write_buffer_prepared(
1641 prepared_chain(40, prepared, 16),
1642 &buffer,
1643 BufferRange::new(4, prepared.len() as u64).unwrap(),
1644 )
1645 .unwrap();
1646 let request = complete_next(&mut client, StatusCode::OK, &[]);
1647 assert_eq!(&request[40..], prepared);
1648 success(&mut client, write);
1649
1650 let read = client
1651 .read_buffer(chain(40, 20), &buffer, BufferRange::new(0, 4).unwrap())
1652 .unwrap();
1653 complete_next(&mut client, StatusCode::OK, b"data");
1654 let (read_output, read_chain) = success(&mut client, read);
1655 assert_eq!(read_output.bytes, 4);
1656 let mut read_bytes = [0; 4];
1657 read_chain
1658 .read_device_writable(RESPONSE_HEADER_BYTES, &mut read_bytes)
1659 .unwrap();
1660 assert_eq!(&read_bytes, b"data");
1661
1662 let program_desc = ProgramDesc::new(
1663 NonZeroU32::new(1).unwrap(),
1664 [0; 12],
1665 NonZeroU64::new(32).unwrap(),
1666 );
1667 let artifact = b"program";
1668 let load = client
1669 .load_program_prepared(
1670 prepared_chain(16 + size_of::<LoadProgramRequest>(), artifact, 24),
1671 &context,
1672 program_desc,
1673 artifact.len() as u64,
1674 )
1675 .unwrap();
1676 let request = complete_next(&mut client, StatusCode::OK, object_payload(3).as_bytes());
1677 assert_eq!(&request[96..], artifact);
1678 let (program, _) = success(&mut client, load);
1679
1680 let create_queue = client
1681 .create_execution_queue(chain(16 + size_of::<CreateQueueRequest>(), 24), &context)
1682 .unwrap();
1683 complete_next(&mut client, StatusCode::OK, object_payload(4).as_bytes());
1684 let (queue, _) = success(&mut client, create_queue);
1685
1686 let bindings = [Binding {
1687 buffer: &buffer,
1688 range: BufferRange::new(0, 16).unwrap(),
1689 slot: 0,
1690 access: AccessMode::Read,
1691 }];
1692 let submit = client
1693 .submit(
1694 chain(
1695 16 + size_of::<SubmitRequest>() + size_of::<WireBinding>(),
1696 16 + size_of::<SubmitResponse>(),
1697 ),
1698 &queue,
1699 &program,
1700 &bindings,
1701 1_000_000,
1702 )
1703 .unwrap();
1704 let submit_response = SubmitResponse {
1705 event_id: Le64::new(5),
1706 };
1707 complete_next(&mut client, StatusCode::OK, submit_response.as_bytes());
1708 let (outcome, _) = success(&mut client, submit);
1709 let event = match outcome {
1710 SubmissionOutcome::Accepted(event) => event,
1711 SubmissionOutcome::Indeterminate { .. } => panic!("unexpected indeterminate submit"),
1712 };
1713
1714 let poll = client.poll_event(chain(24, 24), &event).unwrap();
1715 let state = WireEventState {
1716 state: Le16::new(1),
1717 error: Le16::new(0),
1718 reserved: Le32::new(0),
1719 };
1720 complete_next(&mut client, StatusCode::OK, state.as_bytes());
1721 assert_eq!(success(&mut client, poll).0, crate::EventState::Complete);
1722
1723 let cancel = client.cancel_event(chain(24, 16), &event).unwrap();
1724 complete_next(&mut client, StatusCode::OK, &[]);
1725 success(&mut client, cancel);
1726
1727 let destroy_event = client.destroy_event(chain(24, 16), event).unwrap();
1728 complete_next(&mut client, StatusCode::OK, &[]);
1729 success(&mut client, destroy_event);
1730
1731 let destroy_queue = client
1732 .destroy_execution_queue(chain(24, 16), queue)
1733 .unwrap();
1734 complete_next(&mut client, StatusCode::OK, &[]);
1735 success(&mut client, destroy_queue);
1736
1737 let unload = client.unload_program(chain(24, 16), program).unwrap();
1738 complete_next(&mut client, StatusCode::OK, &[]);
1739 success(&mut client, unload);
1740
1741 let free = client.free_buffer(chain(24, 16), buffer).unwrap();
1742 complete_next(&mut client, StatusCode::OK, &[]);
1743 success(&mut client, free);
1744
1745 let destroy_context = client.destroy_context(chain(24, 16), context).unwrap();
1746 complete_next(&mut client, StatusCode::OK, &[]);
1747 success(&mut client, destroy_context);
1748 }
1749
1750 #[test]
1751 fn mutable_state_allows_every_program_access_mode() {
1752 assert!(usage_allows(BufferUsage::MUTABLE_STATE, AccessMode::Read));
1753 assert!(usage_allows(BufferUsage::MUTABLE_STATE, AccessMode::Write));
1754 assert!(usage_allows(
1755 BufferUsage::MUTABLE_STATE,
1756 AccessMode::ReadWrite
1757 ));
1758 assert!(!usage_allows(
1759 BufferUsage::PROGRAM_INPUT | BufferUsage::PROGRAM_OUTPUT,
1760 AccessMode::ReadWrite
1761 ));
1762 }
1763}