Skip to main content

virtio_accel_guest/
operation.rs

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/// Malformed or inaccessible device response.
21#[derive(Debug, PartialEq, Eq)]
22pub enum ResponseError<E> {
23    /// Used length cannot contain the required response header or exact payload.
24    UsedLength {
25        /// Published used length.
26        used: u32,
27    },
28    /// The response exceeds the negotiated frame limit.
29    ResponseLimit,
30    /// The response header cannot be read.
31    HeaderAccess(E),
32    /// The response payload cannot be read.
33    PayloadAccess(E),
34    /// Response request ID does not match the chain's request.
35    RequestId {
36        /// Expected request ID.
37        expected: u64,
38        /// Device-provided request ID.
39        actual: u64,
40    },
41    /// Protocol 1.0 response flags are nonzero.
42    Flags(u16),
43    /// Payload shape does not match the operation and status.
44    PayloadLength {
45        /// Required payload bytes.
46        expected: u64,
47        /// Device-provided payload bytes.
48        actual: u32,
49    },
50    /// A fixed payload cannot be decoded from its exact bytes.
51    PayloadEncoding,
52    /// A successful object-creation response returned object ID zero.
53    ObjectId,
54    /// Device discovery returned invalid limits or reserved values.
55    DeviceInfo(DeviceInfoError),
56    /// Event state uses an unknown numeric value.
57    EventState(u16),
58    /// Event state and error status contradict each other.
59    EventStatus {
60        /// Raw event state.
61        state: u16,
62        /// Raw event error status.
63        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
77/// Sealed response contract for one typed pending operation.
78pub trait Operation: sealed::Sealed + Sized {
79    /// Validated successful response value.
80    type Output;
81
82    /// Command opcode represented by this operation.
83    fn opcode(&self) -> KnownOpcode;
84
85    /// Whether device discovery must complete before this operation can be published.
86    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/// Pending device-discovery operation.
154#[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/// Pending context-creation operation.
196#[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/// Pending context destruction; retained on failure for explicit retry.
224#[derive(Debug)]
225pub struct DestroyContext {
226    pub(crate) context: Context,
227}
228
229impl DestroyContext {
230    /// Inspect or recover the consumed context after failure.
231    ///
232    /// Retry it only when the completion disposition is `Retryable`.
233    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/// Pending buffer allocation.
260#[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/// Pending buffer release; retained on failure for explicit retry.
292#[derive(Debug)]
293pub struct FreeBuffer {
294    pub(crate) buffer: Buffer,
295}
296
297impl FreeBuffer {
298    /// Inspect or recover the consumed buffer after failure.
299    ///
300    /// Retry it only when the completion disposition is `Retryable`.
301    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/// Pending buffer-read operation.
330#[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/// Pending program load.
359#[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/// Pending program unload; retained on failure for explicit retry.
391#[derive(Debug)]
392pub struct UnloadProgram {
393    pub(crate) program: Program,
394}
395
396impl UnloadProgram {
397    /// Inspect or recover the consumed program after failure.
398    ///
399    /// Retry it only when the completion disposition is `Retryable`.
400    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/// Pending accelerator execution-queue creation.
427#[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/// Pending execution-queue destruction; retained on failure for explicit retry.
457#[derive(Debug)]
458pub struct DestroyExecutionQueue {
459    pub(crate) queue: ExecutionQueue,
460}
461
462impl DestroyExecutionQueue {
463    /// Inspect or recover the consumed queue after failure.
464    ///
465    /// Retry it only when the completion disposition is `Retryable`.
466    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/// Pending submission admission.
493#[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/// Pending nonblocking event poll.
549#[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/// Pending event destruction; retained on failure for explicit retry.
585#[derive(Debug)]
586pub struct DestroyEvent {
587    pub(crate) event: Event,
588}
589
590impl DestroyEvent {
591    /// Inspect or recover the consumed event after failure.
592    ///
593    /// Retry it only when the completion disposition is `Retryable`.
594    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}