Skip to main content

virtio_accel_guest/
client.rs

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/// Whether the client may publish new work.
32#[derive(Clone, Copy, Debug, PartialEq, Eq)]
33pub enum ClientHealth {
34    /// Queue tracking and response framing remain trustworthy.
35    Running,
36    /// An unexpected or malformed completion requires a queue reset.
37    NeedsReset,
38}
39
40/// Failure to construct a bounded guest client.
41#[derive(Clone, Copy, Debug, PartialEq, Eq)]
42pub enum ClientInitError {
43    /// Configuration or queue state is incompatible with the client.
44    Config(GuestConfigError),
45    /// Bounded in-flight tracking storage could not be allocated.
46    AllocationFailed,
47}
48
49/// Failure while restoring queue configuration after reset.
50#[derive(Clone, Copy, Debug, PartialEq, Eq)]
51pub enum QueueSetupError<E> {
52    /// Requested queue size contradicts the validated guest configuration.
53    Config(GuestConfigError),
54    /// Concrete queue configuration failed.
55    Queue(QueueError<E>),
56}
57
58/// Failure before request ownership transfers to the queue.
59#[derive(Debug, PartialEq, Eq)]
60pub enum StartErrorKind<QE, CE> {
61    /// A prior malformed completion requires reset.
62    NeedsReset,
63    /// Device information must be discovered first.
64    DiscoveryRequired,
65    /// Every bounded tracking slot is occupied.
66    InflightLimit,
67    /// No request identifier could be selected without aliasing live work.
68    RequestIdExhausted,
69    /// A typed handle belongs to a prior queue epoch.
70    StaleHandle,
71    /// Typed arguments contradict retained object bounds or usage.
72    InvalidArgument,
73    /// A request exceeds an advertised semantic device limit.
74    DeviceLimit,
75    /// The device-readable chain length is not the exact request frame length.
76    RequestCapacity {
77        /// Required exact frame length.
78        required: u64,
79        /// Supplied readable bytes.
80        available: u64,
81    },
82    /// The device-writable chain cannot hold the largest valid response.
83    ResponseCapacity {
84        /// Required response capacity.
85        required: u64,
86        /// Supplied writable bytes.
87        available: u64,
88    },
89    /// The caller-provided request source could not be read.
90    SourceAccess(ByteAccessError),
91    /// The chain request region could not be written.
92    ChainAccess(CE),
93    /// Queue publication failed before ownership transfer.
94    Publish(PublishErrorKind<QE>),
95}
96
97/// Failed request start with both caller-owned inputs returned.
98#[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    /// Borrow the failure classification.
107    pub const fn kind(&self) -> &StartErrorKind<QE, CE> {
108        &self.kind
109    }
110
111    /// Recover the unpublished chain, typed operation, and failure classification.
112    pub fn into_parts(self) -> (C, O, StartErrorKind<QE, CE>) {
113        (self.chain, self.operation, self.kind)
114    }
115}
116
117/// Result of starting one typed operation on queue `Q`.
118pub 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/// Typed operation token retained until its matching completion is observed.
129#[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    /// Nonzero wire request identifier.
139    pub const fn request_id(&self) -> u64 {
140        self.request_id
141    }
142
143    /// Queue epoch in which this request was published.
144    pub const fn epoch(&self) -> QueueEpoch {
145        self.epoch
146    }
147
148    /// Borrow operation metadata retained for response validation or recovery.
149    pub const fn operation(&self) -> &O {
150        &self.operation
151    }
152}
153
154/// Operation metadata recovered only after its queue epoch is stale.
155#[derive(Debug)]
156pub struct StaleOperation<O> {
157    request_id: u64,
158    epoch: QueueEpoch,
159    operation: O,
160}
161
162impl<O> StaleOperation<O> {
163    /// Request identifier from the invalidated queue epoch.
164    pub const fn request_id(&self) -> u64 {
165        self.request_id
166    }
167
168    /// Invalidated queue epoch.
169    pub const fn epoch(&self) -> QueueEpoch {
170        self.epoch
171    }
172
173    /// Borrow the invalidated operation metadata.
174    pub const fn operation(&self) -> &O {
175        &self.operation
176    }
177
178    /// Recover operation metadata for cleanup or inspection.
179    ///
180    /// Typed handles from this operation remain stale and cannot be republished in the new epoch.
181    pub fn into_operation(self) -> O {
182        self.operation
183    }
184}
185
186/// Terminal result for one matched request completion.
187#[derive(Debug)]
188pub enum Completion<O: Operation, C, CE> {
189    /// Protocol success with the reclaimed caller chain.
190    Success {
191        /// Typed response value. Bulk read bytes remain in `chain`.
192        output: O::Output,
193        /// Reclaimed caller-owned chain.
194        chain: C,
195    },
196    /// Well-formed protocol error; operation metadata is returned for retry decisions.
197    DeviceError {
198        /// Raw status, preserving unknown provider values.
199        status: StatusCode,
200        /// Whether typed operation ownership is retryable, invalid, or uncertain.
201        disposition: FailureDisposition,
202        /// Original typed operation.
203        operation: O,
204        /// Reclaimed caller-owned chain.
205        chain: C,
206    },
207    /// Malformed or inaccessible response; reset is required.
208    InvalidResponse {
209        /// Validation failure.
210        error: ResponseError<CE>,
211        /// Original typed operation.
212        operation: O,
213        /// Reclaimed caller-owned chain.
214        chain: C,
215    },
216}
217
218/// Result of synchronously polling one typed request.
219pub enum RequestPoll<O: Operation, C, QE, CE> {
220    /// No matching completion has arrived yet.
221    Pending(Pending<O>),
222    /// Matching request reached a terminal result.
223    Ready(Completion<O, C, CE>),
224    /// Used-ring access failed; the pending token remains valid for retry.
225    QueueError {
226        /// Original pending token.
227        pending: Pending<O>,
228        /// Queue failure.
229        error: QueueError<QE>,
230    },
231    /// Queue reset invalidated this operation.
232    Stale(StaleOperation<O>),
233    /// Queue health requires reset before more completion processing.
234    NeedsReset(Pending<O>),
235}
236
237/// Result of consuming at most one used-ring entry.
238#[derive(Clone, Copy, Debug, PartialEq, Eq)]
239pub enum PumpResult {
240    /// No completion was available.
241    Idle,
242    /// One tracked completion was retained for its pending token.
243    Completion {
244        /// Request identifier associated with the completed chain.
245        request_id: u64,
246    },
247    /// A used chain had no live matching publication.
248    UnexpectedCompletion,
249    /// Client already requires reset.
250    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
265/// Bounded, single-owner reference driver over a portable command queue.
266///
267/// The type performs no internal locking. Callers that share one client across execution contexts
268/// choose their own synchronization policy; ordinary single-owner polling needs none.
269pub 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    /// Construct a client after transport feature and queue setup.
288    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    /// Current queue lifecycle snapshot.
312    pub fn queue_state(&self) -> QueueState {
313        self.queue.state()
314    }
315
316    /// Validated guest configuration.
317    pub const fn config(&self) -> GuestConfig {
318        self.config
319    }
320
321    /// Latest discovery result in the current epoch.
322    pub const fn device_info(&self) -> Option<DeviceInfo> {
323        self.device_info
324    }
325
326    /// Current response-tracking health.
327    pub const fn health(&self) -> ClientHealth {
328        self.health
329    }
330
331    /// Disable used notifications before draining completions.
332    pub fn disable_used_notifications(&mut self) -> Result<(), QueueError<Q::Error>> {
333        self.queue.disable_used_notifications()
334    }
335
336    /// Atomically enable used notifications and check for missed completions.
337    pub fn enable_used_notifications(
338        &mut self,
339    ) -> Result<NotificationRecheck, QueueError<Q::Error>> {
340        self.queue.enable_used_notifications()
341    }
342
343    /// Consume and classify at most one used-ring entry.
344    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    /// Poll one typed operation, consuming at most one queue completion.
370    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    /// Reset the queue and invalidate all pending operations and typed handles.
409    ///
410    /// Published chains are returned by the queue's reclaimed collection. Chains already popped
411    /// from the used ring are retained for [`Self::pop_recovered_completion`].
412    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    /// Configure and ready the command queue after transport reset.
428    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    /// Recover one completion that had been popped before reset invalidated its pending token.
443    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    /// Recover the unexpected completion that forced reset.
455    pub fn take_unexpected_completion(&mut self) -> Option<UsedChain<Q::Chain>> {
456        self.unexpected.take()
457    }
458
459    /// Start protocol discovery.
460    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    /// Start context creation.
471    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    /// Start context destruction, consuming the handle until completion or failure recovery.
485    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    /// Start a bounded buffer allocation.
502    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    /// Start buffer release, consuming the handle until completion or failure recovery.
539    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    /// Start a copy from caller source bytes into a provider-owned buffer.
552    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    /// Start a zero-copy buffer write from an already prepared request-chain tail.
594    ///
595    /// Bytes beginning after the transfer prefix are left untouched and become the write payload.
596    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    /// Start a provider-buffer read whose bytes remain in the reclaimed chain.
629    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    /// Start loading an opaque program artifact directly from caller bytes.
656    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    /// Start a zero-copy program load from an already prepared request-chain tail.
706    ///
707    /// Artifact bytes beginning after the fixed load prefix are left untouched.
708    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    /// Start program release, consuming the handle until completion or failure recovery.
751    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    /// Start creation of an accelerator execution queue.
768    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    /// Start execution-queue release, consuming the handle until completion or failure recovery.
793    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    /// Start submission with bindings encoded directly into the caller chain.
810    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        // Canonical slot order proves uniqueness in one linear pass. Keep the
839        // prefix-scan fallback below so arbitrary binding order and its existing
840        // validation precedence remain unchanged without allocating scratch space.
841        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    /// Start a nonblocking event-state poll.
893    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    /// Start event cancellation.
909    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    /// Start event release, consuming the handle until completion or failure recovery.
920    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}