1use core::mem::size_of;
2use core::num::NonZeroU64;
3
4use virtio_accel_proto::{
5 KnownEventState, KnownOpcode, ObjectPayload, StatusCode, SubmitResponse, WireDeviceInfo,
6 WireEventState, read_exact,
7};
8use virtio_accel_transport::{DriverChainBuffer, QueueEpoch};
9use zerocopy::FromBytes;
10
11use crate::config::GuestConfig;
12use crate::types::{
13 Buffer, BufferDesc, Context, DeviceInfo, DeviceInfoError, Event, EventState, ExecutionQueue,
14 FailureDisposition, Handle, Program, ProgramDesc, ReadBufferOutput, SubmissionOutcome,
15};
16
17const RESPONSE_HEADER_BYTES: u64 = 16;
18const MAX_FIXED_PAYLOAD_BYTES: usize = size_of::<WireDeviceInfo>();
19
20#[derive(Debug, PartialEq, Eq)]
22pub enum ResponseError<E> {
23 UsedLength {
25 used: u32,
27 },
28 ResponseLimit,
30 HeaderAccess(E),
32 PayloadAccess(E),
34 RequestId {
36 expected: u64,
38 actual: u64,
40 },
41 Flags(u16),
43 PayloadLength {
45 expected: u64,
47 actual: u32,
49 },
50 PayloadEncoding,
52 ObjectId,
54 DeviceInfo(DeviceInfoError),
56 EventState(u16),
58 EventStatus {
60 state: u16,
62 error: StatusCode,
64 },
65}
66
67#[doc(hidden)]
68pub enum OperationResult<T> {
69 Success(T),
70 DeviceError(StatusCode),
71}
72
73mod sealed {
74 pub trait Sealed {}
75}
76
77pub trait Operation: sealed::Sealed + Sized {
79 type Output;
81
82 fn opcode(&self) -> KnownOpcode;
84
85 fn requires_discovery(&self) -> bool {
87 true
88 }
89
90 #[doc(hidden)]
91 fn decode<C: DriverChainBuffer>(
92 &self,
93 chain: &C,
94 status: StatusCode,
95 payload_bytes: u32,
96 epoch: QueueEpoch,
97 config: GuestConfig,
98 ) -> Result<OperationResult<Self::Output>, ResponseError<C::Error>>;
99
100 #[doc(hidden)]
101 fn discovered_info(_output: &Self::Output) -> Option<DeviceInfo> {
102 None
103 }
104
105 #[doc(hidden)]
106 fn failure_disposition(status: StatusCode) -> FailureDisposition {
107 if !status.is_known() {
108 FailureDisposition::Unknown
109 } else if status == StatusCode::STALE_OBJECT {
110 FailureDisposition::Invalidated
111 } else if status == StatusCode::DEVICE_LOST {
112 FailureDisposition::Indeterminate
113 } else {
114 FailureDisposition::Retryable
115 }
116 }
117
118 #[doc(hidden)]
119 fn output_requires_reset(_output: &Self::Output) -> bool {
120 false
121 }
122}
123
124macro_rules! empty_operation {
125 ($name:ident, $opcode:ident) => {
126 #[doc = concat!("Pending `", stringify!($opcode), "` operation.")]
127 #[derive(Debug)]
128 pub struct $name;
129
130 impl sealed::Sealed for $name {}
131
132 impl Operation for $name {
133 type Output = ();
134
135 fn opcode(&self) -> KnownOpcode {
136 KnownOpcode::$opcode
137 }
138
139 fn decode<C: DriverChainBuffer>(
140 &self,
141 _chain: &C,
142 status: StatusCode,
143 payload_bytes: u32,
144 _epoch: QueueEpoch,
145 _config: GuestConfig,
146 ) -> Result<OperationResult<Self::Output>, ResponseError<C::Error>> {
147 ordinary(status, payload_bytes, 0, || Ok(()))
148 }
149 }
150 };
151}
152
153#[derive(Debug)]
155pub struct GetDeviceInfo;
156
157impl sealed::Sealed for GetDeviceInfo {}
158
159impl Operation for GetDeviceInfo {
160 type Output = DeviceInfo;
161
162 fn opcode(&self) -> KnownOpcode {
163 KnownOpcode::GetDeviceInfo
164 }
165
166 fn requires_discovery(&self) -> bool {
167 false
168 }
169
170 fn decode<C: DriverChainBuffer>(
171 &self,
172 chain: &C,
173 status: StatusCode,
174 payload_bytes: u32,
175 _epoch: QueueEpoch,
176 config: GuestConfig,
177 ) -> Result<OperationResult<Self::Output>, ResponseError<C::Error>> {
178 ordinary(
179 status,
180 payload_bytes,
181 size_of::<WireDeviceInfo>() as u64,
182 || {
183 let wire = read_payload::<WireDeviceInfo, C>(chain)?;
184 DeviceInfo::from_wire(wire, config.wire().max_request_bytes.get())
185 .map_err(ResponseError::DeviceInfo)
186 },
187 )
188 }
189
190 fn discovered_info(output: &Self::Output) -> Option<DeviceInfo> {
191 Some(*output)
192 }
193}
194
195#[derive(Debug)]
197pub struct CreateContext;
198
199impl sealed::Sealed for CreateContext {}
200
201impl Operation for CreateContext {
202 type Output = Context;
203
204 fn opcode(&self) -> KnownOpcode {
205 KnownOpcode::CreateContext
206 }
207
208 fn decode<C: DriverChainBuffer>(
209 &self,
210 chain: &C,
211 status: StatusCode,
212 payload_bytes: u32,
213 epoch: QueueEpoch,
214 _config: GuestConfig,
215 ) -> Result<OperationResult<Self::Output>, ResponseError<C::Error>> {
216 ordinary_object(status, payload_bytes, chain, |id| Context {
217 handle: Handle::new(id, epoch),
218 context: None,
219 })
220 }
221}
222
223#[derive(Debug)]
225pub struct DestroyContext {
226 pub(crate) context: Context,
227}
228
229impl DestroyContext {
230 pub fn into_context(self) -> Context {
234 self.context
235 }
236}
237
238impl sealed::Sealed for DestroyContext {}
239
240impl Operation for DestroyContext {
241 type Output = ();
242
243 fn opcode(&self) -> KnownOpcode {
244 KnownOpcode::DestroyContext
245 }
246
247 fn decode<C: DriverChainBuffer>(
248 &self,
249 _chain: &C,
250 status: StatusCode,
251 payload_bytes: u32,
252 _epoch: QueueEpoch,
253 _config: GuestConfig,
254 ) -> Result<OperationResult<Self::Output>, ResponseError<C::Error>> {
255 ordinary(status, payload_bytes, 0, || Ok(()))
256 }
257}
258
259#[derive(Debug)]
261pub struct AllocateBuffer {
262 pub(crate) context: NonZeroU64,
263 pub(crate) desc: BufferDesc,
264}
265
266impl sealed::Sealed for AllocateBuffer {}
267
268impl Operation for AllocateBuffer {
269 type Output = Buffer;
270
271 fn opcode(&self) -> KnownOpcode {
272 KnownOpcode::AllocateBuffer
273 }
274
275 fn decode<C: DriverChainBuffer>(
276 &self,
277 chain: &C,
278 status: StatusCode,
279 payload_bytes: u32,
280 epoch: QueueEpoch,
281 _config: GuestConfig,
282 ) -> Result<OperationResult<Self::Output>, ResponseError<C::Error>> {
283 ordinary_object(status, payload_bytes, chain, |id| Buffer {
284 handle: Handle::new(id, epoch),
285 context: self.context,
286 desc: self.desc,
287 })
288 }
289}
290
291#[derive(Debug)]
293pub struct FreeBuffer {
294 pub(crate) buffer: Buffer,
295}
296
297impl FreeBuffer {
298 pub fn into_buffer(self) -> Buffer {
302 self.buffer
303 }
304}
305
306impl sealed::Sealed for FreeBuffer {}
307
308impl Operation for FreeBuffer {
309 type Output = ();
310
311 fn opcode(&self) -> KnownOpcode {
312 KnownOpcode::FreeBuffer
313 }
314
315 fn decode<C: DriverChainBuffer>(
316 &self,
317 _chain: &C,
318 status: StatusCode,
319 payload_bytes: u32,
320 _epoch: QueueEpoch,
321 _config: GuestConfig,
322 ) -> Result<OperationResult<Self::Output>, ResponseError<C::Error>> {
323 ordinary(status, payload_bytes, 0, || Ok(()))
324 }
325}
326
327empty_operation!(WriteBuffer, WriteBuffer);
328
329#[derive(Debug)]
331pub struct ReadBuffer {
332 pub(crate) bytes: u64,
333}
334
335impl Operation for ReadBuffer {
336 type Output = ReadBufferOutput;
337
338 fn opcode(&self) -> KnownOpcode {
339 KnownOpcode::ReadBuffer
340 }
341
342 fn decode<C: DriverChainBuffer>(
343 &self,
344 _chain: &C,
345 status: StatusCode,
346 payload_bytes: u32,
347 _epoch: QueueEpoch,
348 _config: GuestConfig,
349 ) -> Result<OperationResult<Self::Output>, ResponseError<C::Error>> {
350 ordinary(status, payload_bytes, self.bytes, || {
351 Ok(ReadBufferOutput { bytes: self.bytes })
352 })
353 }
354}
355
356impl sealed::Sealed for ReadBuffer {}
357
358#[derive(Debug)]
360pub struct LoadProgram {
361 pub(crate) context: NonZeroU64,
362 pub(crate) desc: ProgramDesc,
363}
364
365impl sealed::Sealed for LoadProgram {}
366
367impl Operation for LoadProgram {
368 type Output = Program;
369
370 fn opcode(&self) -> KnownOpcode {
371 KnownOpcode::LoadProgram
372 }
373
374 fn decode<C: DriverChainBuffer>(
375 &self,
376 chain: &C,
377 status: StatusCode,
378 payload_bytes: u32,
379 epoch: QueueEpoch,
380 _config: GuestConfig,
381 ) -> Result<OperationResult<Self::Output>, ResponseError<C::Error>> {
382 let _ = self.desc;
383 ordinary_object(status, payload_bytes, chain, |id| Program {
384 handle: Handle::new(id, epoch),
385 context: Some(self.context),
386 })
387 }
388}
389
390#[derive(Debug)]
392pub struct UnloadProgram {
393 pub(crate) program: Program,
394}
395
396impl UnloadProgram {
397 pub fn into_program(self) -> Program {
401 self.program
402 }
403}
404
405impl sealed::Sealed for UnloadProgram {}
406
407impl Operation for UnloadProgram {
408 type Output = ();
409
410 fn opcode(&self) -> KnownOpcode {
411 KnownOpcode::UnloadProgram
412 }
413
414 fn decode<C: DriverChainBuffer>(
415 &self,
416 _chain: &C,
417 status: StatusCode,
418 payload_bytes: u32,
419 _epoch: QueueEpoch,
420 _config: GuestConfig,
421 ) -> Result<OperationResult<Self::Output>, ResponseError<C::Error>> {
422 ordinary(status, payload_bytes, 0, || Ok(()))
423 }
424}
425
426#[derive(Debug)]
428pub struct CreateExecutionQueue {
429 pub(crate) context: NonZeroU64,
430}
431
432impl sealed::Sealed for CreateExecutionQueue {}
433
434impl Operation for CreateExecutionQueue {
435 type Output = ExecutionQueue;
436
437 fn opcode(&self) -> KnownOpcode {
438 KnownOpcode::CreateQueue
439 }
440
441 fn decode<C: DriverChainBuffer>(
442 &self,
443 chain: &C,
444 status: StatusCode,
445 payload_bytes: u32,
446 epoch: QueueEpoch,
447 _config: GuestConfig,
448 ) -> Result<OperationResult<Self::Output>, ResponseError<C::Error>> {
449 ordinary_object(status, payload_bytes, chain, |id| ExecutionQueue {
450 handle: Handle::new(id, epoch),
451 context: Some(self.context),
452 })
453 }
454}
455
456#[derive(Debug)]
458pub struct DestroyExecutionQueue {
459 pub(crate) queue: ExecutionQueue,
460}
461
462impl DestroyExecutionQueue {
463 pub fn into_queue(self) -> ExecutionQueue {
467 self.queue
468 }
469}
470
471impl sealed::Sealed for DestroyExecutionQueue {}
472
473impl Operation for DestroyExecutionQueue {
474 type Output = ();
475
476 fn opcode(&self) -> KnownOpcode {
477 KnownOpcode::DestroyQueue
478 }
479
480 fn decode<C: DriverChainBuffer>(
481 &self,
482 _chain: &C,
483 status: StatusCode,
484 payload_bytes: u32,
485 _epoch: QueueEpoch,
486 _config: GuestConfig,
487 ) -> Result<OperationResult<Self::Output>, ResponseError<C::Error>> {
488 ordinary(status, payload_bytes, 0, || Ok(()))
489 }
490}
491
492#[derive(Debug)]
494pub struct Submit {
495 pub(crate) context: NonZeroU64,
496}
497
498impl sealed::Sealed for Submit {}
499
500impl Operation for Submit {
501 type Output = SubmissionOutcome;
502
503 fn opcode(&self) -> KnownOpcode {
504 KnownOpcode::Submit
505 }
506
507 fn decode<C: DriverChainBuffer>(
508 &self,
509 chain: &C,
510 status: StatusCode,
511 payload_bytes: u32,
512 epoch: QueueEpoch,
513 _config: GuestConfig,
514 ) -> Result<OperationResult<Self::Output>, ResponseError<C::Error>> {
515 if status.is_success() {
516 expect_payload::<C::Error>(payload_bytes, size_of::<SubmitResponse>() as u64)?;
517 let event = decode_event(chain, epoch, self.context)?;
518 return Ok(OperationResult::Success(SubmissionOutcome::Accepted(event)));
519 }
520 if payload_bytes == 0 {
521 return Ok(OperationResult::DeviceError(status));
522 }
523 if !status.is_known() {
524 return Err(ResponseError::PayloadLength {
525 expected: 0,
526 actual: payload_bytes,
527 });
528 }
529 expect_payload::<C::Error>(payload_bytes, size_of::<SubmitResponse>() as u64)?;
530 let event = decode_event(chain, epoch, self.context)?;
531 Ok(OperationResult::Success(SubmissionOutcome::Indeterminate {
532 status,
533 event,
534 }))
535 }
536
537 fn output_requires_reset(output: &Self::Output) -> bool {
538 matches!(
539 output,
540 SubmissionOutcome::Indeterminate {
541 status: StatusCode::DEVICE_LOST,
542 ..
543 }
544 )
545 }
546}
547
548#[derive(Debug)]
550pub struct PollEvent;
551
552impl sealed::Sealed for PollEvent {}
553
554impl Operation for PollEvent {
555 type Output = EventState;
556
557 fn opcode(&self) -> KnownOpcode {
558 KnownOpcode::PollEvent
559 }
560
561 fn decode<C: DriverChainBuffer>(
562 &self,
563 chain: &C,
564 status: StatusCode,
565 payload_bytes: u32,
566 _epoch: QueueEpoch,
567 _config: GuestConfig,
568 ) -> Result<OperationResult<Self::Output>, ResponseError<C::Error>> {
569 ordinary(
570 status,
571 payload_bytes,
572 size_of::<WireEventState>() as u64,
573 || decode_event_state(chain),
574 )
575 }
576
577 fn output_requires_reset(output: &Self::Output) -> bool {
578 matches!(output, EventState::Failed(StatusCode::DEVICE_LOST))
579 }
580}
581
582empty_operation!(CancelEvent, CancelEvent);
583
584#[derive(Debug)]
586pub struct DestroyEvent {
587 pub(crate) event: Event,
588}
589
590impl DestroyEvent {
591 pub fn into_event(self) -> Event {
595 self.event
596 }
597}
598
599impl sealed::Sealed for DestroyEvent {}
600
601impl Operation for DestroyEvent {
602 type Output = ();
603
604 fn opcode(&self) -> KnownOpcode {
605 KnownOpcode::DestroyEvent
606 }
607
608 fn decode<C: DriverChainBuffer>(
609 &self,
610 _chain: &C,
611 status: StatusCode,
612 payload_bytes: u32,
613 _epoch: QueueEpoch,
614 _config: GuestConfig,
615 ) -> Result<OperationResult<Self::Output>, ResponseError<C::Error>> {
616 ordinary(status, payload_bytes, 0, || Ok(()))
617 }
618}
619
620fn ordinary<T, E>(
621 status: StatusCode,
622 payload_bytes: u32,
623 success_bytes: u64,
624 success: impl FnOnce() -> Result<T, ResponseError<E>>,
625) -> Result<OperationResult<T>, ResponseError<E>> {
626 if !status.is_success() {
627 expect_payload::<E>(payload_bytes, 0)?;
628 return Ok(OperationResult::DeviceError(status));
629 }
630 expect_payload::<E>(payload_bytes, success_bytes)?;
631 success().map(OperationResult::Success)
632}
633
634fn ordinary_object<T, C: DriverChainBuffer>(
635 status: StatusCode,
636 payload_bytes: u32,
637 chain: &C,
638 construct: impl FnOnce(NonZeroU64) -> T,
639) -> Result<OperationResult<T>, ResponseError<C::Error>> {
640 ordinary(
641 status,
642 payload_bytes,
643 size_of::<ObjectPayload>() as u64,
644 || {
645 let object = read_payload::<ObjectPayload, C>(chain)?;
646 let id = NonZeroU64::new(object.object_id.get()).ok_or(ResponseError::ObjectId)?;
647 Ok(construct(id))
648 },
649 )
650}
651
652fn expect_payload<E>(actual: u32, expected: u64) -> Result<(), ResponseError<E>> {
653 if u64::from(actual) == expected {
654 Ok(())
655 } else {
656 Err(ResponseError::PayloadLength { expected, actual })
657 }
658}
659
660fn read_payload<T: FromBytes, C: DriverChainBuffer>(
661 chain: &C,
662) -> Result<T, ResponseError<C::Error>> {
663 let bytes = size_of::<T>();
664 if bytes > MAX_FIXED_PAYLOAD_BYTES {
665 return Err(ResponseError::PayloadEncoding);
666 }
667 let mut scratch = [0_u8; MAX_FIXED_PAYLOAD_BYTES];
668 chain
669 .read_device_writable(RESPONSE_HEADER_BYTES, &mut scratch[..bytes])
670 .map_err(ResponseError::PayloadAccess)?;
671 read_exact(&scratch[..bytes]).map_err(|_| ResponseError::PayloadEncoding)
672}
673
674fn decode_event<C: DriverChainBuffer>(
675 chain: &C,
676 epoch: QueueEpoch,
677 context: NonZeroU64,
678) -> Result<Event, ResponseError<C::Error>> {
679 let response = read_payload::<SubmitResponse, C>(chain)?;
680 let id = NonZeroU64::new(response.event_id.get()).ok_or(ResponseError::ObjectId)?;
681 Ok(Event {
682 handle: Handle::new(id, epoch),
683 context: Some(context),
684 })
685}
686
687fn decode_event_state<C: DriverChainBuffer>(
688 chain: &C,
689) -> Result<EventState, ResponseError<C::Error>> {
690 let value = read_payload::<WireEventState, C>(chain)?;
691 if value.reserved.get() != 0 {
692 return Err(ResponseError::PayloadEncoding);
693 }
694 let raw_state = value.state.get();
695 let error = StatusCode(value.error.get());
696 match value
697 .known_state()
698 .map_err(|_| ResponseError::EventState(raw_state))?
699 {
700 KnownEventState::Pending if error.is_success() => Ok(EventState::Pending),
701 KnownEventState::Complete if error.is_success() => Ok(EventState::Complete),
702 KnownEventState::Failed if !error.is_success() => Ok(EventState::Failed(error)),
703 KnownEventState::Cancelled if error.is_success() => Ok(EventState::Cancelled),
704 _ => Err(ResponseError::EventStatus {
705 state: raw_state,
706 error,
707 }),
708 }
709}