Skip to main content

hyperlight_common/virtq/
producer.rs

1// SPDX-License-Identifier: Apache-2.0
2// Copyright 2026 The Hyperlight Authors.
3
4use alloc::collections::VecDeque;
5use alloc::vec::Vec;
6
7use bytes::Bytes;
8use smallvec::SmallVec;
9
10use super::*;
11
12/// A used chain observed by the driver (producer) side.
13///
14/// Read-only chains are returned as [`Ack`](Self::Ack). Chains with a writable
15/// buffer complete as [`Data`](Self::Data), even when the device wrote zero
16/// bytes. Non-empty segments in [`Data`](Self::Data) are backed by
17/// shared-memory pool allocations that are returned when the last clone is
18/// dropped.
19#[derive(Debug)]
20pub enum UsedChain {
21    /// Acknowledgement for a read-only/fire-and-forget chain.
22    Ack(Token),
23    /// Data written by the consumer into the chain's writable buffers.
24    ///
25    /// The payload may contain zero bytes when the consumer writes zero bytes;
26    /// that is still a data used chain because the submitted chain had
27    /// writable capacity.
28    Data(Token, Segments),
29}
30
31impl UsedChain {
32    /// Token identifying which submitted chain this used chain corresponds to.
33    pub fn token(&self) -> Token {
34        match self {
35            Self::Ack(token) | Self::Data(token, _) => *token,
36        }
37    }
38
39    /// Data written by the consumer as contiguous bytes, if this chain has data.
40    pub fn to_bytes(&self) -> Option<Bytes> {
41        match self {
42            Self::Ack(_) => None,
43            Self::Data(_, segments) => Some(segments.to_bytes()),
44        }
45    }
46
47    /// Segments written by the consumer, if this chain has data.
48    pub fn segments(&self) -> Option<&Segments> {
49        match self {
50            Self::Ack(_) => None,
51            Self::Data(_, segments) => Some(segments),
52        }
53    }
54
55    /// Consume the used chain and return written data as contiguous bytes, if present.
56    pub fn into_bytes(self) -> Option<Bytes> {
57        match self {
58            Self::Ack(_) => None,
59            Self::Data(_, segments) => Some(segments.into_bytes()),
60        }
61    }
62
63    /// Consume the used chain and return written segments, if present.
64    pub fn into_segments(self) -> Option<Segments> {
65        match self {
66            Self::Ack(_) => None,
67            Self::Data(_, segments) => Some(segments),
68        }
69    }
70}
71
72/// Allocation tracking for an in-flight descriptor chain.
73///
74/// Descriptor lengths have already been published to the ring, so in-flight
75/// state only needs the completion token and allocation ownership for later
76/// reclaim.
77#[derive(Debug)]
78pub(crate) struct Inflight {
79    token: Token,
80    chain: BufferChain,
81}
82
83/// A high-level virtqueue producer (driver side).
84///
85/// The producer sends chains to the consumer (device), and receives used chains.
86/// This is used on the driver/guest side.
87///
88/// # Threading
89///
90/// The producer is intended for single-threaded, guest-side use. Reply payloads
91/// are exposed as zero-copy [`Bytes`] via [`Bytes::from_owner`](bytes::Bytes::from_owner),
92/// which requires the owning pool to be `Send`. Do not move the producer
93/// or its replies across threads, and do not instantiate it on the multi-threaded
94/// host with those pools.
95///
96/// # Example
97///
98/// ```ignore
99/// let mut producer = VirtqProducer::new(layout, mem, notifier, pool);
100///
101/// // Build and submit a chain
102/// let mut chain = producer.chain().readable(64).writable(64).build()?;
103/// chain.write_all(b"hello")?;
104/// let token = producer.submit(chain)?;
105///
106/// // Later, poll for the used chain
107/// if let Some(used) = producer.poll()? {
108///     assert_eq!(used.token(), token);
109///     match used {
110///         UsedChain::Data(_, segments) => println!("Got used chain: {:?}", segments),
111///         UsedChain::Ack(_) => println!("Got ack"),
112///     }
113/// }
114/// ```
115pub struct VirtqProducer<M, N, P> {
116    inner: RingProducer<M>,
117    notifier: N,
118    pool: P,
119    next_token: u32,
120    inflight: Vec<Option<Inflight>>,
121    pending: VecDeque<UsedChain>,
122}
123
124impl<M, N, P> VirtqProducer<M, N, P>
125where
126    M: MemOps + Clone,
127    N: Notifier,
128    P: BufferProvider + Clone,
129{
130    /// Create a new virtqueue producer.
131    ///
132    /// # Arguments
133    ///
134    /// * `layout` - Ring memory layout (descriptor table and event suppression addresses)
135    /// * `mem` - Memory operations implementation for reading/writing to shared memory
136    /// * `notifier` - Callback for notifying the device (consumer) about new chains
137    /// * `pool` - Buffer allocator for chain payload and reply data
138    pub fn new(layout: Layout, mem: M, notifier: N, pool: P) -> Self {
139        let inner = RingProducer::new(layout, mem);
140        let ring_len = inner.len();
141
142        Self {
143            inner,
144            pool,
145            notifier,
146            next_token: 0,
147            inflight: (0..ring_len).map(|_| None).collect(),
148            pending: VecDeque::with_capacity(ring_len),
149        }
150    }
151
152    fn dealloc_elems(
153        &self,
154        elems: impl IntoIterator<Item = BufferElement>,
155    ) -> Result<(), VirtqError> {
156        let mut first_err = None;
157        for elem in elems {
158            if let Err(err) = self.pool.dealloc(elem.addr)
159                && first_err.is_none()
160            {
161                first_err = Some(VirtqError::Alloc(err));
162            }
163        }
164
165        if let Some(err) = first_err {
166            return Err(err);
167        }
168
169        Ok(())
170    }
171
172    /// Begin building a descriptor chain for submission.
173    ///
174    /// Returns a [`ChainBuilder`] that allocates buffers from the pool.
175    pub fn chain(&self) -> ChainBuilder<M, P> {
176        ChainBuilder::new(self.inner.mem().clone(), self.pool.clone())
177    }
178
179    /// Begin a batch of submissions.
180    ///
181    /// Chains submitted through the returned [`SubmitBatch`] are published to
182    /// the ring immediately, but the consumer is notified at most once when
183    /// [`SubmitBatch::finish`] is called. This mirrors the virtio pattern of
184    /// adding multiple buffers and then kicking the queue once.
185    pub fn batch(&mut self) -> SubmitBatch<'_, M, N, P> {
186        SubmitBatch::new(self)
187    }
188
189    /// Submit a [`SendChain`] to the ring.
190    ///
191    /// Publishes the descriptor chain, stores the in-flight tracking state,
192    /// and notifies the consumer if event suppression allows. Notifications
193    /// are layout-neutral; use [`batch`](Self::batch) when a higher-level
194    /// protocol wants to publish multiple chains and kick once.
195    ///
196    /// # Errors
197    ///
198    /// - [`VirtqError::PayloadTooLarge`] - written exceeds readable buffer capacity
199    /// - [`VirtqError::RingError`] - ring is full
200    /// - [`VirtqError::InvalidState`] - descriptor ID collision
201    pub fn submit(&mut self, chain: SendChain<M, P>) -> Result<Token, VirtqError> {
202        let cursor_before = self.inner.avail_cursor();
203        let token = self.publish(chain)?;
204        self.notify_since(cursor_before)?;
205        Ok(token)
206    }
207
208    fn publish(&mut self, send: SendChain<M, P>) -> Result<Token, VirtqError> {
209        let token_id = self.next_token;
210        let id = self.inner.submit_available(send.chain())?;
211        let token = Token { seq: token_id, id };
212
213        // A free descriptor id must never already be tracked as inflight.
214        if self.inflight[id as usize].is_some() {
215            return Err(VirtqError::InvalidState);
216        }
217
218        let inf = send.into_inflight(token);
219        self.inflight[id as usize] = Some(inf);
220        self.next_token = self.next_token.wrapping_add(1);
221
222        Ok(token)
223    }
224
225    fn notify_since(&mut self, cursor: RingCursor) -> Result<bool, VirtqError> {
226        let should_notify = self.inner.should_notify_since(cursor)?;
227        if should_notify {
228            self.notify_now();
229        }
230        Ok(should_notify)
231    }
232
233    fn notify_now(&self) {
234        self.notifier.notify(QueueStats {
235            num_free: self.inner.num_free(),
236            num_inflight: self.inner.num_inflight(),
237        });
238    }
239
240    /// Signal backpressure to the consumer.
241    ///
242    /// Bypasses event suppression. Call this when submit fails with a
243    /// backpressure error and the consumer needs to drain.
244    #[inline]
245    pub fn notify_backpressure(&self) {
246        self.notify_now();
247    }
248
249    /// Get the current used cursor position.
250    ///
251    /// Useful for setting up descriptor-based event suppression.
252    #[inline]
253    pub fn used_cursor(&self) -> RingCursor {
254        self.inner.used_cursor()
255    }
256
257    /// Number of free (unsubmitted) descriptors in the ring.
258    #[inline]
259    pub fn num_free(&self) -> usize {
260        self.inner.num_free()
261    }
262
263    /// Configure event suppression for used buffer notifications.
264    ///
265    /// This controls when the device (consumer) signals us about completed buffers:
266    ///
267    /// - [`SuppressionKind::Enable`]: Always signal (default) - good for latency
268    /// - [`SuppressionKind::Disable`]: Never signal - caller must poll
269    /// - [`SuppressionKind::Descriptor`]: Signal only at specific cursor position
270    ///
271    /// # Example: Used-chain batching
272    ///
273    /// ```ignore
274    /// // Submit chains, then suppress notifications until all are used
275    /// let mut se = producer.chain().readable(64).writable(128).build()?;
276    /// se.write_all(b"entry1")?;
277    /// producer.submit(se)?;
278    /// let cursor = producer.used_cursor();
279    /// producer.set_used_suppression(SuppressionKind::Descriptor(cursor))?;
280    /// // Device will notify only after reaching that cursor position
281    /// ```
282    pub fn set_used_suppression(&mut self, kind: SuppressionKind) -> Result<(), VirtqError> {
283        match kind {
284            SuppressionKind::Enable => self.inner.enable_used_notifications()?,
285            SuppressionKind::Disable => self.inner.disable_used_notifications()?,
286            SuppressionKind::Descriptor(cursor) => self
287                .inner
288                .enable_used_notifications_desc(cursor.head(), cursor.wrap())?,
289        }
290        Ok(())
291    }
292
293    /// Reset ring, inflight, and pool state to initial values.
294    ///
295    /// # Safety
296    ///
297    /// No outstanding [`UsedChain::Data`] buffers, borrowed segment views, or
298    /// peer accesses to previously submitted descriptors may exist. Resetting
299    /// recycles the same backing addresses, so outstanding zero-copy buffers or
300    /// stale descriptor users could alias memory that is handed out again.
301    ///
302    /// TODO(virtq): find a way to allow guest to keep used chains across resets.
303    pub unsafe fn reset(&mut self) {
304        self.inflight.iter_mut().for_each(|slot| *slot = None);
305        self.pending.clear();
306        self.inner.reset();
307        self.pool.reset();
308    }
309
310    /// Replace the pool and reset ring, inflight, and pending state.
311    ///
312    /// # Safety
313    ///
314    /// No outstanding [`UsedChain::Data`] buffers, borrowed segment views, or
315    /// peer accesses to previously submitted descriptors may exist. The new pool
316    /// may manage the same shared-memory addresses as the old pool, so old
317    /// zero-copy buffers must not outlive this transition.
318    pub unsafe fn reset_with_pool(&mut self, pool: P) {
319        self.pending.clear();
320        self.inflight.iter_mut().for_each(|slot| *slot = None);
321        self.inner.reset();
322        self.pool = pool;
323        self.pool.reset();
324    }
325}
326
327impl<M, N, P> VirtqProducer<M, N, P>
328where
329    M: MemOps + Clone + Send + 'static,
330    N: Notifier,
331    P: BufferProvider + Clone + Send + 'static,
332{
333    /// Poll for a single used chain from the device.
334    ///
335    /// Returns buffered used chains from prior [`reclaim`](Self::reclaim)
336    /// calls first, then checks the ring for newly used chains.
337    ///
338    /// Returns `Ok(Some(used))` if a used chain is available, `Ok(None)` if no
339    /// used chains are ready (would block), or an error if the device misbehaved.
340    ///
341    /// Data used chains contain zero-copy [`Bytes`] backed by the shared-memory
342    /// allocation via [`BufferOwner`]. The pool allocation is held alive as long
343    /// as any `Bytes` clone exists, and is returned to the pool when the last
344    /// clone is dropped.
345    ///
346    /// # Errors
347    ///
348    /// - [`VirtqError::InvalidState`] - Device returned invalid descriptor ID or
349    ///   wrote more data than the writable buffer capacity
350    pub fn poll(&mut self) -> Result<Option<UsedChain>, VirtqError> {
351        if let Some(chain) = self.pending.pop_front() {
352            return Ok(Some(chain));
353        }
354        self.poll_ring()
355    }
356
357    /// Reclaim ring slots and pool allocations from used descriptors.
358    ///
359    /// Processes all available used chains from the ring: frees readable
360    /// buffer allocations immediately, and buffers writable data for
361    /// later retrieval via [`poll`](Self::poll).
362    ///
363    /// Read-only ack used chains are discarded immediately.
364    ///
365    /// Use this to free resources under backpressure without losing
366    /// writable data. Returns the number of chains reclaimed.
367    pub fn reclaim(&mut self) -> Result<usize, VirtqError> {
368        let mut count = 0;
369        while let Some(chain) = self.poll_ring()? {
370            if matches!(chain, UsedChain::Data(_, _)) {
371                debug_assert!(self.pending.len() < self.inner.len());
372                self.pending.push_back(chain);
373            }
374            count += 1;
375        }
376        Ok(count)
377    }
378
379    /// Poll one used chain directly from the ring (bypassing pending buffer).
380    fn poll_ring(&mut self) -> Result<Option<UsedChain>, VirtqError> {
381        let used = match self.inner.poll_used() {
382            Ok(u) => u,
383            Err(RingError::WouldBlock) => return Ok(None),
384            Err(e) => return Err(e.into()),
385        };
386
387        let inf = self
388            .inflight
389            .get_mut(used.id as usize)
390            .and_then(Option::take)
391            .ok_or(VirtqError::InvalidState)?;
392
393        let written = used.len as usize;
394        let Inflight { token, chain } = inf;
395
396        self.dealloc_elems(chain.readables().iter().copied())?;
397
398        let used = if chain.writables().is_empty() {
399            UsedChain::Ack(token)
400        } else {
401            UsedChain::Data(token, self.recv_segments(chain.writables(), written)?)
402        };
403
404        Ok(Some(used))
405    }
406
407    fn recv_segments(
408        &self,
409        writables: &[BufferElement],
410        written: usize,
411    ) -> Result<Segments, VirtqError> {
412        let mut owned = SmallVec::<[(BufferElement, usize); 4]>::new();
413        let mut free = SmallVec::<[BufferElement; 4]>::new();
414        let mut remaining = written;
415
416        for &alloc in writables {
417            if remaining == 0 {
418                free.push(alloc);
419                continue;
420            }
421
422            let len = remaining.min(alloc.len as usize);
423            owned.push((alloc, len));
424            remaining -= len;
425        }
426
427        if remaining != 0 {
428            let elems = owned.iter().map(|(elem, _)| *elem).chain(free);
429            self.dealloc_elems(elems)?;
430            return Err(VirtqError::InvalidState);
431        }
432
433        for (elem, len) in &owned {
434            if unsafe { self.inner.mem().as_slice(elem.addr, *len) }.is_err() {
435                let elems = owned.iter().map(|(elem, _)| *elem).chain(free);
436                let _ = self.dealloc_elems(elems);
437                return Err(VirtqError::MemoryReadError);
438            }
439        }
440
441        let mut sgs = SmallVec::<[Bytes; 4]>::new();
442        for (elem, written) in owned {
443            let alloc = OwnedAlloc::new(
444                self.pool.clone(),
445                Allocation {
446                    addr: elem.addr,
447                    len: elem.len as usize,
448                },
449            );
450            let mem = self.inner.mem().clone();
451            let owner = BufferOwner {
452                alloc,
453                mem,
454                written,
455            };
456            sgs.push(Bytes::from_owner(owner));
457        }
458
459        self.dealloc_elems(free)?;
460
461        Ok(Segments::from_smallvec(sgs))
462    }
463
464    /// Drain all available used chains, calling the provided closure for each.
465    ///
466    /// This is a convenience method that repeatedly calls [`poll`](Self::poll)
467    /// until no more used chains are available.
468    ///
469    /// # Arguments
470    ///
471    /// * `f` - Closure called for each used chain
472    ///
473    /// # Example
474    ///
475    /// ```ignore
476    /// producer.drain(|used| {
477    ///     println!("Got used chain for {:?}", used.token());
478    /// })?;
479    /// ```
480    pub fn drain(&mut self, mut f: impl FnMut(UsedChain)) -> Result<(), VirtqError> {
481        while let Some(chain) = self.poll()? {
482            f(chain);
483        }
484
485        Ok(())
486    }
487}
488
489/// A scoped batch of producer submissions.
490///
491/// Submissions are published immediately, while notification is delayed until
492/// [`finish`](Self::finish). `finish` is explicit because the event-suppression
493/// check can fail; dropping a batch does not notify.
494#[must_use = "call finish to notify the consumer about batched submissions"]
495pub struct SubmitBatch<'a, M, N, P> {
496    producer: &'a mut VirtqProducer<M, N, P>,
497    notify_from: Option<RingCursor>,
498}
499
500impl<'a, M, N, P> SubmitBatch<'a, M, N, P>
501where
502    M: MemOps + Clone,
503    N: Notifier,
504    P: BufferProvider + Clone,
505{
506    fn new(producer: &'a mut VirtqProducer<M, N, P>) -> Self {
507        Self {
508            producer,
509            notify_from: None,
510        }
511    }
512
513    /// Begin building a descriptor chain for this batch.
514    pub fn chain(&self) -> ChainBuilder<M, P> {
515        self.producer.chain()
516    }
517
518    /// Publish a chain as part of this batch without notifying yet.
519    pub fn submit(&mut self, chain: SendChain<M, P>) -> Result<Token, VirtqError> {
520        let cursor_before = self.producer.inner.avail_cursor();
521        let token = self.producer.publish(chain)?;
522
523        if self.notify_from.is_none() {
524            self.notify_from = Some(cursor_before);
525        }
526        Ok(token)
527    }
528
529    /// Finish the batch and notify the consumer once if event suppression
530    /// requires it for the whole published range.
531    ///
532    /// Returns `true` if a notification was sent.
533    pub fn finish(mut self) -> Result<bool, VirtqError> {
534        let Some(notify_from) = self.notify_from.take() else {
535            return Ok(false);
536        };
537
538        self.producer.notify_since(notify_from)
539    }
540}
541
542/// Builder for configuring a descriptor chain's buffer layout.
543///
544/// If dropped without building, no resources are leaked (allocations are
545/// deferred to [`build`](Self::build)).
546#[must_use = "call .build() to create a SendChain"]
547pub struct ChainBuilder<M: MemOps, P: BufferProvider + Clone> {
548    mem: M,
549    pool: P,
550    rd_caps: SmallVec<[usize; 4]>,
551    wr_caps: SmallVec<[usize; 4]>,
552}
553
554impl<M: MemOps, P: BufferProvider + Clone> ChainBuilder<M, P> {
555    fn new(mem: M, pool: P) -> Self {
556        Self {
557            mem,
558            pool,
559            rd_caps: SmallVec::new(),
560            wr_caps: SmallVec::new(),
561        }
562    }
563
564    /// Request a device-readable buffer of `cap` bytes.
565    ///
566    /// The producer writes data into readable buffers before submission; the
567    /// consumer reads that data after polling the chain.
568    /// The actual allocation is deferred to [`build`](Self::build).
569    pub fn readable(mut self, cap: usize) -> Self {
570        self.rd_caps.push(cap);
571        self
572    }
573
574    /// Request a device-writable buffer of `cap` bytes.
575    ///
576    /// The writable buffer is filled by the consumer and returned via
577    /// [`VirtqProducer::poll`] as [`UsedChain`].
578    ///
579    /// Multiple writable buffers are completed as ordered [`Segments`]. The
580    /// consumer writes them sequentially, because the virtio used ring reports
581    /// one aggregate written length rather than per-descriptor lengths.
582    pub fn writable(mut self, cap: usize) -> Self {
583        self.wr_caps.push(cap);
584        self
585    }
586
587    /// Allocate buffers and return a [`SendChain`] for writing.
588    ///
589    /// # Errors
590    ///
591    /// - [`VirtqError::InvalidState`] - No buffers requested
592    /// - [`VirtqError::Alloc`] - Pool exhausted
593    pub fn build(self) -> Result<SendChain<M, P>, VirtqError> {
594        if self.rd_caps.is_empty() && self.wr_caps.is_empty() {
595            return Err(VirtqError::InvalidState);
596        }
597
598        let mut rollback = Rollback::new(&self.pool);
599        let mut rd_caps = SmallVec::<[usize; 4]>::new();
600        let mut rd_elems = SmallVec::<[BufferElement; 4]>::new();
601        let mut wr_elems = SmallVec::<[BufferElement; 4]>::new();
602
603        // Allocate readable buffers, splitting into multiple descriptors if needed.
604        // The buffer element lengths are initialized to zero and updated as the
605        // `SendChain` writes.
606        for &cap in &self.rd_caps {
607            let sgs = self.pool.alloc_sg(cap)?;
608            let mut remaining = cap;
609
610            for alloc in sgs {
611                let _ = checked_descriptor_len(alloc.len)?;
612                let seg_cap = remaining.min(alloc.len);
613
614                rd_caps.push(seg_cap);
615                rd_elems.push(BufferElement {
616                    addr: alloc.addr,
617                    len: 0,
618                    writable: false,
619                });
620                remaining -= seg_cap;
621                rollback.allocs.push(alloc);
622            }
623
624            if remaining != 0 {
625                return Err(VirtqError::InvalidState);
626            }
627        }
628
629        // Allocate writable buffers, with the same caveat about splitting as readable buffers.
630        // Writable buffer elements are initialized with their full capacity for the device to
631        // write into.
632        for &cap in &self.wr_caps {
633            let sgs = self.pool.alloc_sg(cap)?;
634            for alloc in sgs {
635                let len = checked_descriptor_len(alloc.len)?;
636                wr_elems.push(BufferElement {
637                    addr: alloc.addr,
638                    len,
639                    writable: true,
640                });
641                rollback.allocs.push(alloc);
642            }
643        }
644
645        let chain = BufferChainBuilder::new()
646            .readables(rd_elems)
647            .writables(wr_elems)
648            .build()?;
649
650        rollback.release();
651
652        Ok(SendChain {
653            mem: self.mem,
654            pool: self.pool,
655            chain: Some(chain),
656            rd_caps,
657            rd_capacity: self.rd_caps.iter().sum(),
658            rd_written: 0,
659            write_mode: WriteMode::Unset,
660        })
661    }
662}
663
664struct Rollback<'a, P: BufferProvider> {
665    pool: &'a P,
666    allocs: SmallVec<[Allocation; 8]>,
667}
668
669impl<'a, P: BufferProvider> Rollback<'a, P> {
670    fn new(pool: &'a P) -> Self {
671        Self {
672            pool,
673            allocs: SmallVec::new(),
674        }
675    }
676
677    fn release(mut self) {
678        self.allocs.clear();
679    }
680}
681
682impl<P: BufferProvider> Drop for Rollback<'_, P> {
683    fn drop(&mut self) {
684        for alloc in self.allocs.drain(..) {
685            let result = self.pool.dealloc(alloc.addr);
686            debug_assert!(result.is_ok(), "rollback dealloc failed: {result:?}");
687        }
688    }
689}
690
691/// Tracks which write API a [`SendChain`] payload uses, so the two paths are
692/// not mixed.
693///
694/// Copy writes ([`SendChain::write`]/[`SendChain::write_all`]) append at an
695/// aggregate cursor, while direct writes
696/// ([`SendChain::write_seg`]/[`SendChain::with_seg`]) set per-segment lengths
697/// absolutely. Mixing them would corrupt the written-length accounting.
698#[derive(Debug, Clone, Copy, PartialEq, Eq)]
699enum WriteMode {
700    Unset,
701    Append,
702    Direct,
703}
704
705/// A configured send chain ready for writing and submission.
706///
707/// Created by [`ChainBuilder::build`]. Write readable payload bytes directly on
708/// the chain, then submit via [`VirtqProducer::submit`].
709///
710/// # Examples
711///
712/// ```ignore
713/// let mut sc = producer.chain().readable(64).writable(128).build()?;
714/// sc.write_all(b"header")?.write_all(b" body")?;
715/// let tok = producer.submit(sc)?;
716///
717/// let mut sc = producer.chain().readable(128).build()?;
718/// sc.with_seg(0, |buf| serialize_into(buf))?;
719/// let tok = producer.submit(sc)?;
720/// ```
721///
722/// Copy writes (`write`/`write_all`) and direct writes (`write_seg`/`with_seg`)
723/// must not be mixed on the same chain; doing so panics in debug builds.
724///
725/// If dropped without submitting, allocated buffers are returned to the pool.
726#[must_use = "dropping without submitting deallocates the buffers"]
727pub struct SendChain<M: MemOps, P: BufferProvider> {
728    mem: M,
729    pool: P,
730    chain: Option<BufferChain>,
731    rd_caps: SmallVec<[usize; 4]>,
732    rd_capacity: usize,
733    rd_written: usize,
734    write_mode: WriteMode,
735}
736
737// `chain` is wrapped in `Option` only so `into_inflight` can `take()` it
738// without moving out of this `Drop` type; it stays `Some` for a chain's whole
739// public lifetime, so these `expect`s cannot fail.
740#[allow(clippy::expect_used)]
741impl<M: MemOps, P: BufferProvider> SendChain<M, P> {
742    fn chain(&self) -> &BufferChain {
743        self.chain.as_ref().expect("SendChain missing BufferChain")
744    }
745
746    fn chain_mut(&mut self) -> &mut BufferChain {
747        self.chain.as_mut().expect("SendChain missing BufferChain")
748    }
749
750    /// Record that this chain uses `mode`, asserting it is not mixed with the
751    /// other write path.
752    fn note_write_mode(&mut self, mode: WriteMode) {
753        debug_assert!(
754            self.write_mode == WriteMode::Unset || self.write_mode == mode,
755            "SendChain mixes copy writes (write/write_all) with direct writes (write_seg/with_seg)"
756        );
757        self.write_mode = mode;
758    }
759
760    fn into_inflight(mut self, token: Token) -> Inflight {
761        let chain = self.chain.take().expect("SendChain missing BufferChain");
762        Inflight { token, chain }
763    }
764
765    /// Number of producer-written readable segments in this chain.
766    pub fn segment_count(&self) -> usize {
767        self.chain().readables().len()
768    }
769
770    /// Total producer-written readable capacity in bytes.
771    pub fn capacity(&self) -> usize {
772        self.rd_capacity
773    }
774
775    /// Number of producer-written readable bytes written so far.
776    pub fn written(&self) -> usize {
777        self.rd_written
778    }
779
780    /// Remaining producer-written readable capacity.
781    pub fn remaining(&self) -> usize {
782        self.capacity() - self.written()
783    }
784
785    /// Write bytes into payload segments, returning how many bytes were written.
786    ///
787    /// Appends at the current aggregate write position and scatters across
788    /// readable segments in chain order. Uses [`MemOps::write`] (volatile on
789    /// host side). If `buf` is larger than the remaining capacity, writes as
790    /// many bytes as will fit.
791    ///
792    /// # Errors
793    ///
794    /// - [`VirtqError::NoPayloadSegment`] - no readable buffer allocated
795    /// - [`VirtqError::MemoryWriteError`] - underlying write failed
796    pub fn write(&mut self, buf: &[u8]) -> Result<usize, VirtqError> {
797        if self.segment_count() == 0 {
798            return Err(VirtqError::NoPayloadSegment);
799        }
800
801        self.note_write_mode(WriteMode::Append);
802
803        let mut remaining = &buf[..buf.len().min(self.remaining())];
804        let mut written = 0;
805
806        let SendChain {
807            mem,
808            chain,
809            rd_caps,
810            ..
811        } = self;
812
813        let readables = chain
814            .as_mut()
815            .expect("SendChain missing BufferChain")
816            .readables_mut();
817
818        for (readable, &cap) in readables.iter_mut().zip(rd_caps.iter()) {
819            if remaining.is_empty() {
820                break;
821            }
822
823            let written_len = readable.len as usize;
824            let free = cap - written_len;
825            if free == 0 {
826                continue;
827            }
828
829            let n = free.min(remaining.len());
830            let addr = readable.addr + written_len as u64;
831            mem.write(addr, &remaining[..n])
832                .map_err(|_| VirtqError::MemoryWriteError)?;
833
834            readable.len += n as u32;
835            written += n;
836            remaining = &remaining[n..];
837        }
838
839        self.rd_written += written;
840        Ok(written)
841    }
842
843    /// Write the entire buffer into payload segments.
844    ///
845    /// Appends at the current aggregate write position and scatters across
846    /// readable segments in chain order. Uses [`MemOps::write`] (volatile on
847    /// host side).
848    ///
849    /// # Errors
850    ///
851    /// - [`VirtqError::PayloadTooLarge`] - buf exceeds remaining capacity
852    /// - [`VirtqError::NoPayloadSegment`] - no readable buffer allocated
853    /// - [`VirtqError::MemoryWriteError`] - underlying write failed
854    pub fn write_all(&mut self, buf: &[u8]) -> Result<&mut Self, VirtqError> {
855        if self.segment_count() == 0 {
856            return Err(VirtqError::NoPayloadSegment);
857        }
858
859        if buf.len() > self.remaining() {
860            return Err(VirtqError::PayloadTooLarge {
861                recv: buf.len(),
862                limit: self.remaining(),
863            });
864        }
865
866        let written = self.write(buf)?;
867        debug_assert_eq!(written, buf.len());
868        Ok(self)
869    }
870
871    /// Write bytes into one readable segment by index.
872    ///
873    /// Writes from the start of the selected segment and records `buf.len()` as
874    /// that segment's descriptor length.
875    ///
876    /// # Errors
877    ///
878    /// - [`VirtqError::NoPayloadSegment`] - `index` does not name a payload segment
879    /// - [`VirtqError::PayloadTooLarge`] - `buf` exceeds the segment capacity
880    /// - [`VirtqError::MemoryWriteError`] - underlying write failed
881    pub fn write_seg(&mut self, index: usize, buf: &[u8]) -> Result<&mut Self, VirtqError> {
882        self.note_write_mode(WriteMode::Direct);
883
884        let cap = *self
885            .rd_caps
886            .get(index)
887            .ok_or(VirtqError::NoPayloadSegment)?;
888
889        if buf.len() > cap {
890            return Err(VirtqError::PayloadTooLarge {
891                recv: buf.len(),
892                limit: cap,
893            });
894        }
895
896        let addr = self
897            .chain()
898            .readables()
899            .get(index)
900            .ok_or(VirtqError::NoPayloadSegment)?
901            .addr;
902
903        self.mem
904            .write(addr, buf)
905            .map_err(|_| VirtqError::MemoryWriteError)?;
906
907        let previous = self.chain().readables()[index].len as usize;
908        self.chain_mut().readables_mut()[index].len = checked_descriptor_len(buf.len())?;
909        self.rd_written = self.rd_written - previous + buf.len();
910        Ok(self)
911    }
912
913    /// Serialize directly into one readable segment by index.
914    ///
915    /// The closure returns the number of valid bytes it wrote. The written
916    /// length for that segment is recorded on success.
917    ///
918    /// # Errors
919    ///
920    /// - [`VirtqError::NoPayloadSegment`] - `index` does not name a payload segment
921    /// - [`VirtqError::PayloadTooLarge`] - closure reports more bytes than segment capacity
922    /// - [`VirtqError::MemoryWriteError`] - the memory backend cannot expose a mutable slice
923    pub fn with_seg<E>(
924        &mut self,
925        index: usize,
926        f: impl FnOnce(&mut [u8]) -> Result<usize, E>,
927    ) -> Result<&mut Self, E>
928    where
929        E: From<VirtqError>,
930    {
931        self.note_write_mode(WriteMode::Direct);
932
933        let cap = *self
934            .rd_caps
935            .get(index)
936            .ok_or_else(|| E::from(VirtqError::NoPayloadSegment))?;
937
938        let addr = self
939            .chain()
940            .readables()
941            .get(index)
942            .ok_or_else(|| E::from(VirtqError::NoPayloadSegment))?
943            .addr;
944
945        let buf = unsafe {
946            self.mem
947                .as_mut_slice(addr, cap)
948                .map_err(|_| E::from(VirtqError::MemoryWriteError))?
949        };
950
951        let written = f(buf)?;
952        if written > buf.len() {
953            return Err(E::from(VirtqError::PayloadTooLarge {
954                recv: written,
955                limit: buf.len(),
956            }));
957        }
958
959        let previous = self.chain().readables()[index].len as usize;
960        // SAFETY: index was validated by the earlier get() call, so the readable element exists.
961        self.chain_mut().readables_mut()[index].len =
962            checked_descriptor_len(written).map_err(E::from)?;
963        self.rd_written = self.rd_written - previous + written;
964
965        Ok(self)
966    }
967}
968
969impl<M: MemOps, P: BufferProvider> Drop for SendChain<M, P> {
970    fn drop(&mut self) {
971        if let Some(chain) = self.chain.take() {
972            for elem in chain.elems() {
973                let result = self.pool.dealloc(elem.addr);
974                debug_assert!(result.is_ok(), "SendChain drop dealloc failed: {result:?}");
975            }
976        }
977    }
978}
979
980fn checked_descriptor_len(len: usize) -> Result<u32, VirtqError> {
981    if len > u32::MAX as usize {
982        return Err(VirtqError::PayloadTooLarge {
983            recv: len,
984            limit: u32::MAX as usize,
985        });
986    }
987    Ok(len as u32)
988}
989
990#[cfg(test)]
991mod tests {
992    use super::*;
993    use crate::virtq::ring::tests::{TestMem, make_consumer, make_ring};
994    use crate::virtq::test_utils::*;
995
996    fn poll_received<M: MemOps + Clone, N: Notifier>(
997        consumer: &mut VirtqConsumer<M, N>,
998    ) -> (RecvChain, ReplyChain<M>) {
999        consumer.poll(1024).unwrap().unwrap()
1000    }
1001
1002    #[derive(Clone)]
1003    struct NoDirectSliceMem(TestMem);
1004
1005    // SAFETY: Delegates all non-slice memory operations to TestMem. Direct
1006    // slices are intentionally unsupported to exercise producer error handling.
1007    unsafe impl MemOps for NoDirectSliceMem {
1008        type Error = ();
1009
1010        fn read(&self, addr: u64, dst: &mut [u8]) -> Result<(), Self::Error> {
1011            self.0.read(addr, dst).map_err(|e| match e {})
1012        }
1013
1014        fn write(&self, addr: u64, src: &[u8]) -> Result<(), Self::Error> {
1015            self.0.write(addr, src).map_err(|e| match e {})
1016        }
1017
1018        fn load_acquire(&self, addr: u64) -> Result<u16, Self::Error> {
1019            self.0.load_acquire(addr).map_err(|e| match e {})
1020        }
1021
1022        fn store_release(&self, addr: u64, val: u16) -> Result<(), Self::Error> {
1023            self.0.store_release(addr, val).map_err(|e| match e {})
1024        }
1025
1026        unsafe fn as_slice(&self, _addr: u64, _len: usize) -> Result<&[u8], Self::Error> {
1027            Err(())
1028        }
1029
1030        unsafe fn as_mut_slice(&self, _addr: u64, _len: usize) -> Result<&mut [u8], Self::Error> {
1031            Err(())
1032        }
1033    }
1034
1035    #[test]
1036    fn test_chain_readwrite_build() {
1037        let ring = make_ring(16);
1038        let (producer, _consumer, _notifier) = make_test_producer(&ring);
1039
1040        let se = producer.chain().readable(64).writable(128).build().unwrap();
1041        assert_eq!(se.capacity(), 64);
1042        assert_eq!(se.written(), 0);
1043        assert_eq!(se.remaining(), 64);
1044    }
1045
1046    #[test]
1047    fn test_chain_readable_writable_names_build() {
1048        let ring = make_ring(16);
1049        let (producer, _consumer, _notifier) = make_test_producer(&ring);
1050
1051        let se = producer.chain().readable(16).writable(32).build().unwrap();
1052        assert_eq!(se.segment_count(), 1);
1053        assert_eq!(se.capacity(), 16);
1054    }
1055
1056    #[test]
1057    fn test_chain_multi_readable_write_all_scatters() {
1058        let ring = make_ring(16);
1059        let (mut producer, mut consumer, _notifier) = make_test_producer(&ring);
1060
1061        let mut se = producer
1062            .chain()
1063            .readable(5)
1064            .readable(6)
1065            .writable(32)
1066            .build()
1067            .unwrap();
1068        se.write_all(b"hello world").unwrap();
1069
1070        assert_eq!(se.written(), 11);
1071
1072        let token = producer.submit(se).unwrap();
1073        let (recv, reply) = poll_received(&mut consumer);
1074        assert_eq!(recv.token(), token);
1075        assert_eq!(recv.to_bytes().as_ref(), b"hello world");
1076        assert_eq!(recv.segments().segment_count(), 2);
1077        assert_eq!(recv.segments().as_slice()[0].as_ref(), b"hello");
1078        assert_eq!(recv.segments().as_slice()[1].as_ref(), b" world");
1079        consumer.complete(reply).unwrap();
1080    }
1081
1082    #[test]
1083    fn test_chain_readable_splits_logical_capacity() {
1084        let ring = make_ring(16);
1085        let layout = ring.layout();
1086        let mem = ring.mem();
1087        let pool_base = mem.base_addr() + Layout::query_size(ring.len()) as u64 + 0x100;
1088        let pool = TestPool::new_with_max_alloc_len(pool_base, 0x8000, 4);
1089        let notifier = TestNotifier::new();
1090        let mut producer = VirtqProducer::new(layout, mem.clone(), notifier.clone(), pool);
1091        let mut consumer = VirtqConsumer::new(layout, mem, notifier);
1092
1093        let mut se = producer.chain().readable(10).writable(32).build().unwrap();
1094
1095        assert_eq!(se.segment_count(), 3);
1096        assert_eq!(se.capacity(), 10);
1097
1098        se.write_all(b"abcdefghij").unwrap();
1099        assert_eq!(se.written(), 10);
1100
1101        let token = producer.submit(se).unwrap();
1102        let (recv, reply) = poll_received(&mut consumer);
1103        assert_eq!(recv.token(), token);
1104        assert_eq!(recv.to_bytes().as_ref(), b"abcdefghij");
1105        assert_eq!(recv.segments().segment_count(), 3);
1106        assert_eq!(recv.segments().as_slice()[0].as_ref(), b"abcd");
1107        assert_eq!(recv.segments().as_slice()[1].as_ref(), b"efgh");
1108        assert_eq!(recv.segments().as_slice()[2].as_ref(), b"ij");
1109        consumer.complete(reply).unwrap();
1110    }
1111
1112    #[test]
1113    fn test_chain_readable_rejects_zero_capacity_on_build() {
1114        let ring = make_ring(16);
1115        let (producer, _consumer, _notifier) = make_test_producer(&ring);
1116
1117        assert!(matches!(
1118            producer.chain().readable(0).build(),
1119            Err(VirtqError::Alloc(AllocError::InvalidArg))
1120        ));
1121    }
1122
1123    #[test]
1124    fn test_chain_writable_splits_logical_capacity() {
1125        let ring = make_ring(16);
1126        let layout = ring.layout();
1127        let mem = ring.mem();
1128        let pool_base = mem.base_addr() + Layout::query_size(ring.len()) as u64 + 0x100;
1129        let pool = TestPool::new_with_max_alloc_len(pool_base, 0x8000, 4);
1130        let notifier = TestNotifier::new();
1131        let mut producer = VirtqProducer::new(layout, mem.clone(), notifier.clone(), pool);
1132        let mut consumer = VirtqConsumer::new(layout, mem, notifier);
1133
1134        let se = producer.chain().writable(10).build().unwrap();
1135        let token = producer.submit(se).unwrap();
1136
1137        let (_recv, reply) = poll_received(&mut consumer);
1138        let ReplyChain::Writable(mut wc) = reply else {
1139            panic!("expected writable reply");
1140        };
1141        assert_eq!(wc.capacity(), 10);
1142
1143        wc.write_all(b"abcdefghij").unwrap();
1144        consumer.complete(wc).unwrap();
1145
1146        let used = producer.poll().unwrap().unwrap();
1147        assert_eq!(used.token(), token);
1148        let segments = used.segments().unwrap();
1149        assert_eq!(segments.segment_count(), 3);
1150        assert_eq!(segments.as_slice()[0].as_ref(), b"abcd");
1151        assert_eq!(segments.as_slice()[1].as_ref(), b"efgh");
1152        assert_eq!(segments.as_slice()[2].as_ref(), b"ij");
1153    }
1154
1155    #[test]
1156    fn test_chain_writable_rejects_zero_capacity_on_build() {
1157        let ring = make_ring(16);
1158        let (producer, _consumer, _notifier) = make_test_producer(&ring);
1159
1160        assert!(matches!(
1161            producer.chain().writable(0).build(),
1162            Err(VirtqError::Alloc(AllocError::InvalidArg))
1163        ));
1164    }
1165
1166    #[test]
1167    fn test_chain_multi_readable_write_all_preserves_segments() {
1168        let ring = make_ring(16);
1169        let (mut producer, mut consumer, _notifier) = make_test_producer(&ring);
1170
1171        let mut se = producer.chain().readable(4).readable(4).build().unwrap();
1172        se.write_all(b"headbody").unwrap();
1173
1174        producer.submit(se).unwrap();
1175        let (recv, reply) = poll_received(&mut consumer);
1176        assert_eq!(recv.to_bytes().as_ref(), b"headbody");
1177        consumer.complete(reply).unwrap();
1178    }
1179
1180    #[test]
1181    fn test_chain_payload_segment_writer_serializes_directly() {
1182        let ring = make_ring(16);
1183        let (mut producer, mut consumer, _notifier) = make_test_producer(&ring);
1184
1185        let mut se = producer.chain().readable(4).readable(4).build().unwrap();
1186        se.with_seg(0, |segment| {
1187            segment.copy_from_slice(b"head");
1188            Ok::<usize, VirtqError>(4)
1189        })
1190        .unwrap();
1191        se.with_seg(1, |segment| {
1192            segment.copy_from_slice(b"body");
1193            Ok::<usize, VirtqError>(4)
1194        })
1195        .unwrap();
1196
1197        producer.submit(se).unwrap();
1198        let (recv, reply) = poll_received(&mut consumer);
1199        assert_eq!(recv.to_bytes().as_ref(), b"headbody");
1200        assert_eq!(recv.segments().segment_count(), 2);
1201        consumer.complete(reply).unwrap();
1202    }
1203
1204    #[test]
1205    fn test_chain_payload_segment_write_copies_directly() {
1206        let ring = make_ring(16);
1207        let (mut producer, mut consumer, _notifier) = make_test_producer(&ring);
1208
1209        let mut se = producer.chain().readable(4).readable(4).build().unwrap();
1210        se.write_seg(0, b"head").unwrap();
1211        se.write_seg(1, b"body").unwrap();
1212
1213        producer.submit(se).unwrap();
1214        let (recv, reply) = poll_received(&mut consumer);
1215        assert_eq!(recv.to_bytes().as_ref(), b"headbody");
1216        assert_eq!(recv.segments().segment_count(), 2);
1217        consumer.complete(reply).unwrap();
1218    }
1219
1220    #[test]
1221    fn test_chain_multi_writable_used_returns_segments() {
1222        let ring = make_ring(16);
1223        let (mut producer, mut consumer, _notifier) = make_test_producer(&ring);
1224
1225        let se = producer.chain().writable(5).writable(6).build().unwrap();
1226        let token = producer.submit(se).unwrap();
1227
1228        let (_recv, reply) = poll_received(&mut consumer);
1229        let ReplyChain::Writable(mut wc) = reply else {
1230            panic!("expected writable reply");
1231        };
1232        assert_eq!(wc.capacity(), 11);
1233
1234        wc.write_all(b"hello world").unwrap();
1235        consumer.complete(wc).unwrap();
1236
1237        let used = producer.poll().unwrap().unwrap();
1238        assert_eq!(used.token(), token);
1239        let segments = used.segments().unwrap();
1240        assert_eq!(segments.segment_count(), 2);
1241        assert_eq!(segments.as_slice()[0].as_ref(), b"hello");
1242        assert_eq!(segments.as_slice()[1].as_ref(), b" world");
1243        assert_eq!(segments.to_bytes().as_ref(), b"hello world");
1244    }
1245
1246    #[test]
1247    fn test_chain_multi_writable_short_used_truncates_last_segment() {
1248        let ring = make_ring(16);
1249        let (mut producer, mut consumer, _notifier) = make_test_producer(&ring);
1250
1251        let se = producer.chain().writable(5).writable(6).build().unwrap();
1252        producer.submit(se).unwrap();
1253
1254        let (_recv, reply) = poll_received(&mut consumer);
1255        let ReplyChain::Writable(mut wc) = reply else {
1256            panic!("expected writable reply");
1257        };
1258
1259        wc.write_all(b"hello wo").unwrap();
1260        consumer.complete(wc).unwrap();
1261
1262        let used = producer.poll().unwrap().unwrap();
1263        let segments = used.segments().unwrap();
1264        assert_eq!(segments.segment_count(), 2);
1265        assert_eq!(segments.as_slice()[0].as_ref(), b"hello");
1266        assert_eq!(segments.as_slice()[1].as_ref(), b" wo");
1267        assert_eq!(segments.to_bytes().as_ref(), b"hello wo");
1268    }
1269
1270    #[test]
1271    fn test_chain_multi_writable_zero_used_returns_empty_segments() {
1272        let ring = make_ring(16);
1273        let (mut producer, mut consumer, _notifier) = make_test_producer(&ring);
1274
1275        let se = producer.chain().writable(5).writable(6).build().unwrap();
1276        producer.submit(se).unwrap();
1277
1278        let (_recv, reply) = poll_received(&mut consumer);
1279        consumer.complete(reply).unwrap();
1280
1281        let used = producer.poll().unwrap().unwrap();
1282        let segments = used.segments().unwrap();
1283        assert_eq!(segments.segment_count(), 0);
1284        assert!(segments.is_empty());
1285        assert!(segments.to_bytes().is_empty());
1286    }
1287
1288    #[test]
1289    fn test_chain_readable_only_build() {
1290        let ring = make_ring(16);
1291        let (producer, _consumer, _notifier) = make_test_producer(&ring);
1292
1293        let se = producer.chain().readable(32).build().unwrap();
1294        assert_eq!(se.capacity(), 32);
1295    }
1296
1297    #[test]
1298    fn test_chain_writable_only_build() {
1299        let ring = make_ring(16);
1300        let (producer, _consumer, _notifier) = make_test_producer(&ring);
1301
1302        let se = producer.chain().writable(64).build().unwrap();
1303        assert_eq!(se.capacity(), 0);
1304    }
1305
1306    #[test]
1307    fn test_chain_empty_build_fails() {
1308        let ring = make_ring(16);
1309        let (producer, _consumer, _notifier) = make_test_producer(&ring);
1310
1311        let result = producer.chain().build();
1312        assert!(matches!(result, Err(VirtqError::InvalidState)));
1313    }
1314
1315    #[test]
1316    fn test_send_chain_write_all_and_submit() {
1317        let ring = make_ring(16);
1318        let (mut producer, mut consumer, _notifier) = make_test_producer(&ring);
1319
1320        let mut se = producer.chain().readable(64).writable(128).build().unwrap();
1321
1322        se.write_all(b"hello")
1323            .unwrap()
1324            .write_all(b" world")
1325            .unwrap();
1326        assert_eq!(se.written(), 11);
1327        assert_eq!(se.remaining(), 53);
1328        let tok = producer.submit(se).unwrap();
1329
1330        let (recv, reply) = poll_received(&mut consumer);
1331        assert_eq!(recv.token(), tok);
1332        assert_eq!(recv.to_bytes().as_ref(), b"hello world");
1333        consumer.complete(reply).unwrap();
1334    }
1335
1336    #[test]
1337    fn test_send_payload_write_all_fluent() {
1338        let ring = make_ring(16);
1339        let (mut producer, mut consumer, _notifier) = make_test_producer(&ring);
1340
1341        let mut se = producer.chain().readable(64).writable(128).build().unwrap();
1342        se.write_all(b"hello")
1343            .unwrap()
1344            .write_all(b" world")
1345            .unwrap();
1346        assert_eq!(se.written(), 11);
1347        assert_eq!(se.remaining(), 53);
1348        let tok = producer.submit(se).unwrap();
1349
1350        let (recv, reply) = poll_received(&mut consumer);
1351        assert_eq!(recv.token(), tok);
1352        assert_eq!(recv.to_bytes().as_ref(), b"hello world");
1353        consumer.complete(reply).unwrap();
1354    }
1355
1356    #[test]
1357    fn test_send_payload_partial_write() {
1358        let ring = make_ring(16);
1359        let (mut producer, mut consumer, _notifier) = make_test_producer(&ring);
1360
1361        let mut se = producer.chain().readable(8).build().unwrap();
1362        let written = se.write(b"hello world").unwrap();
1363        assert_eq!(written, 8);
1364        assert_eq!(se.remaining(), 0);
1365
1366        producer.submit(se).unwrap();
1367        let (recv, reply) = poll_received(&mut consumer);
1368        assert_eq!(recv.to_bytes().as_ref(), b"hello wo");
1369        consumer.complete(reply).unwrap();
1370    }
1371
1372    #[test]
1373    fn test_send_payload_write_with_serializes_directly() {
1374        let ring = make_ring(16);
1375        let (mut producer, mut consumer, _notifier) = make_test_producer(&ring);
1376
1377        let mut se = producer.chain().readable(64).writable(128).build().unwrap();
1378        se.with_seg(0, |buf| {
1379            buf[..5].copy_from_slice(b"hello");
1380            Ok::<usize, VirtqError>(5)
1381        })
1382        .unwrap();
1383
1384        let _tok = producer.submit(se).unwrap();
1385
1386        let (recv, reply) = poll_received(&mut consumer);
1387        assert_eq!(recv.to_bytes().as_ref(), b"hello");
1388        consumer.complete(reply).unwrap();
1389    }
1390
1391    #[test]
1392    fn test_send_chain_single_segment_writer_serializes_directly() {
1393        let ring = make_ring(16);
1394        let (mut producer, mut consumer, _notifier) = make_test_producer(&ring);
1395
1396        let mut se = producer.chain().readable(64).writable(128).build().unwrap();
1397        se.with_seg(0, |segment| {
1398            assert_eq!(segment.len(), 64);
1399            segment[..5].copy_from_slice(b"hello");
1400            Ok::<usize, VirtqError>(5)
1401        })
1402        .unwrap();
1403
1404        let _tok = producer.submit(se).unwrap();
1405
1406        let (recv, reply) = poll_received(&mut consumer);
1407        assert_eq!(recv.to_bytes().as_ref(), b"hello");
1408        consumer.complete(reply).unwrap();
1409    }
1410
1411    #[test]
1412    fn test_send_chain_single_segment_writer_rejects_multi_segment() {
1413        let ring = make_ring(16);
1414        let (producer, _consumer, _notifier) = make_test_producer(&ring);
1415
1416        let mut se = producer.chain().readable(4).readable(4).build().unwrap();
1417        assert!(matches!(
1418            se.with_seg(2, |_| Ok::<usize, VirtqError>(0)),
1419            Err(VirtqError::NoPayloadSegment)
1420        ));
1421    }
1422
1423    #[test]
1424    fn test_send_chain_single_segment_writer_rejects_auto_split_chain() {
1425        let ring = make_ring(16);
1426        let layout = ring.layout();
1427        let mem = ring.mem();
1428        let pool_base = mem.base_addr() + Layout::query_size(ring.len()) as u64 + 0x100;
1429        let pool = TestPool::new_with_max_alloc_len(pool_base, 0x8000, 4);
1430        let notifier = TestNotifier::new();
1431        let producer = VirtqProducer::new(layout, mem, notifier, pool);
1432
1433        let mut se = producer.chain().readable(8).build().unwrap();
1434        assert_eq!(se.segment_count(), 2);
1435        assert!(matches!(
1436            se.with_seg(2, |_| Ok::<usize, VirtqError>(0)),
1437            Err(VirtqError::NoPayloadSegment)
1438        ));
1439    }
1440
1441    #[test]
1442    fn test_send_payload_segment_set_written_too_large() {
1443        let ring = make_ring(16);
1444        let (producer, _consumer, _notifier) = make_test_producer(&ring);
1445
1446        let mut se = producer.chain().readable(32).writable(64).build().unwrap();
1447        let err = se
1448            .with_seg(0, |_| Ok::<usize, VirtqError>(64))
1449            .err()
1450            .unwrap();
1451        assert!(matches!(
1452            err,
1453            VirtqError::PayloadTooLarge {
1454                recv: 64,
1455                limit: 32
1456            }
1457        ));
1458    }
1459
1460    #[test]
1461    fn test_send_chain_write_too_large() {
1462        let ring = make_ring(16);
1463        let (producer, _consumer, _notifier) = make_test_producer(&ring);
1464
1465        let mut se = producer.chain().readable(4).build().unwrap();
1466        let err = se.write_all(b"too long").err().unwrap();
1467        assert!(matches!(
1468            err,
1469            VirtqError::PayloadTooLarge { recv: 8, limit: 4 }
1470        ));
1471    }
1472
1473    #[test]
1474    fn test_writeonly_has_no_readable_buffer() {
1475        let ring = make_ring(16);
1476        let (producer, _consumer, _notifier) = make_test_producer(&ring);
1477
1478        let mut se = producer.chain().writable(32).build().unwrap();
1479        let err = se.write_all(b"data").err().unwrap();
1480        assert!(matches!(err, VirtqError::NoPayloadSegment));
1481    }
1482
1483    #[test]
1484    fn test_drop_chain_builder_deallocs() {
1485        let ring = make_ring(16);
1486        let (mut producer, _consumer, _notifier) = make_test_producer(&ring);
1487
1488        {
1489            let _builder = producer.chain().readable(64).writable(128);
1490            // dropped without build
1491        }
1492
1493        // Ring should still be fully usable
1494        let se = producer.chain().readable(64).writable(128).build().unwrap();
1495        let tok = producer.submit(se).unwrap();
1496        assert!(tok.id < 16);
1497    }
1498
1499    #[test]
1500    fn test_drop_send_chain_deallocs() {
1501        let ring = make_ring(16);
1502        let (mut producer, _consumer, _notifier) = make_test_producer(&ring);
1503
1504        {
1505            let _se = producer.chain().readable(64).writable(128).build().unwrap();
1506            // dropped without submit
1507        }
1508
1509        // Ring should still be fully usable
1510        let se = producer.chain().readable(64).writable(128).build().unwrap();
1511        let tok = producer.submit(se).unwrap();
1512        assert!(tok.id < 16);
1513    }
1514
1515    #[test]
1516    fn test_submit_notifies() {
1517        let ring = make_ring(16);
1518        let (mut producer, _consumer, notifier) = make_test_producer(&ring);
1519
1520        let initial_count = notifier.notification_count();
1521
1522        let mut se = producer.chain().readable(64).writable(128).build().unwrap();
1523        se.write_all(b"hello").unwrap();
1524        producer.submit(se).unwrap();
1525
1526        assert!(notifier.notification_count() > initial_count);
1527    }
1528
1529    #[test]
1530    fn test_submit_read_only_notifies_by_default() {
1531        let ring = make_ring(16);
1532        let (mut producer, _consumer, notifier) = make_test_producer(&ring);
1533
1534        let initial_count = notifier.notification_count();
1535
1536        let mut se = producer.chain().readable(64).build().unwrap();
1537        se.write_all(b"fire-and-forget").unwrap();
1538        producer.submit(se).unwrap();
1539
1540        assert!(notifier.notification_count() > initial_count);
1541    }
1542
1543    #[test]
1544    fn test_submit_write_only_notifies_by_default() {
1545        let ring = make_ring(16);
1546        let (mut producer, _consumer, notifier) = make_test_producer(&ring);
1547
1548        let initial_count = notifier.notification_count();
1549
1550        let se = producer.chain().writable(128).build().unwrap();
1551        producer.submit(se).unwrap();
1552
1553        assert!(notifier.notification_count() > initial_count);
1554    }
1555
1556    #[test]
1557    fn test_batch_notifies_once_on_finish() {
1558        let ring = make_ring(16);
1559        let (mut producer, mut consumer, notifier) = make_test_producer(&ring);
1560
1561        let initial_count = notifier.notification_count();
1562
1563        let mut batch = producer.batch();
1564
1565        let mut first = batch.chain().readable(64).build().unwrap();
1566        first.write_all(b"first").unwrap();
1567        batch.submit(first).unwrap();
1568
1569        let mut second = batch.chain().readable(64).build().unwrap();
1570        second.write_all(b"second").unwrap();
1571        batch.submit(second).unwrap();
1572
1573        assert_eq!(notifier.notification_count(), initial_count);
1574        assert!(batch.finish().unwrap());
1575        assert_eq!(notifier.notification_count(), initial_count + 1);
1576
1577        let (recv, reply) = poll_received(&mut consumer);
1578        assert_eq!(recv.to_bytes().as_ref(), b"first");
1579        consumer.complete(reply).unwrap();
1580
1581        let (recv, reply) = poll_received(&mut consumer);
1582        assert_eq!(recv.to_bytes().as_ref(), b"second");
1583        consumer.complete(reply).unwrap();
1584    }
1585
1586    #[test]
1587    fn test_batch_finish_notifies_from_batch_start_cursor() {
1588        let ring = make_ring(16);
1589        let (mut producer, mut consumer, notifier) = make_test_producer(&ring);
1590
1591        let cursor = consumer.avail_cursor();
1592        consumer
1593            .set_avail_suppression(SuppressionKind::Descriptor(cursor))
1594            .unwrap();
1595
1596        let mut batch = producer.batch();
1597
1598        let mut first = batch.chain().readable(64).build().unwrap();
1599        first.write_all(b"first").unwrap();
1600        batch.submit(first).unwrap();
1601        assert_eq!(notifier.notification_count(), 0);
1602
1603        let mut second = batch.chain().readable(64).writable(64).build().unwrap();
1604        second.write_all(b"second").unwrap();
1605        batch.submit(second).unwrap();
1606
1607        assert!(batch.finish().unwrap());
1608
1609        assert_eq!(notifier.notification_count(), 1);
1610    }
1611
1612    #[test]
1613    fn test_empty_batch_finish_does_not_notify() {
1614        let ring = make_ring(16);
1615        let (mut producer, _consumer, notifier) = make_test_producer(&ring);
1616
1617        let batch = producer.batch();
1618        assert!(!batch.finish().unwrap());
1619        assert_eq!(notifier.notification_count(), 0);
1620    }
1621
1622    #[test]
1623    fn test_write_only_round_trip() {
1624        let ring = make_ring(16);
1625        let (mut producer, mut consumer, _notifier) = make_test_producer(&ring);
1626
1627        let se = producer.chain().writable(32).build().unwrap();
1628        let token = producer.submit(se).unwrap();
1629
1630        let (recv, reply) = poll_received(&mut consumer);
1631        assert_eq!(recv.token(), token);
1632        assert!(recv.to_bytes().is_empty());
1633
1634        if let ReplyChain::Writable(mut wc) = reply {
1635            wc.write_all(b"filled-by-consumer").unwrap();
1636            consumer.complete(wc).unwrap();
1637        } else {
1638            panic!("expected Writable");
1639        }
1640
1641        let used = producer.poll().unwrap().unwrap();
1642        assert_eq!(used.token(), token);
1643        assert_eq!(used.to_bytes().unwrap().len(), b"filled-by-consumer".len());
1644        assert_eq!(used.to_bytes().unwrap().as_ref(), b"filled-by-consumer");
1645    }
1646
1647    #[test]
1648    fn test_read_only_round_trip() {
1649        let ring = make_ring(16);
1650        let (mut producer, mut consumer, _notifier) = make_test_producer(&ring);
1651
1652        let mut se = producer.chain().readable(32).build().unwrap();
1653        se.write_all(b"fire-and-forget").unwrap();
1654        let token = producer.submit(se).unwrap();
1655
1656        let (recv, reply) = poll_received(&mut consumer);
1657        assert_eq!(recv.token(), token);
1658        assert_eq!(recv.to_bytes().as_ref(), b"fire-and-forget");
1659        assert!(matches!(reply, ReplyChain::Ack(_)));
1660        consumer.complete(reply).unwrap();
1661
1662        let used = producer.poll().unwrap().unwrap();
1663        assert!(matches!(used, UsedChain::Ack(t) if t == token));
1664    }
1665
1666    #[test]
1667    fn test_readwrite_round_trip() {
1668        let ring = make_ring(16);
1669        let (mut producer, mut consumer, _notifier) = make_test_producer(&ring);
1670
1671        let mut se = producer.chain().readable(64).writable(128).build().unwrap();
1672        se.write_all(b"request data").unwrap();
1673        let token = producer.submit(se).unwrap();
1674
1675        let (recv, reply) = poll_received(&mut consumer);
1676        assert_eq!(recv.to_bytes().as_ref(), b"request data");
1677        if let ReplyChain::Writable(mut wc) = reply {
1678            wc.write_all(b"response data").unwrap();
1679            consumer.complete(wc).unwrap();
1680        } else {
1681            panic!("expected Writable");
1682        }
1683
1684        let used = producer.poll().unwrap().unwrap();
1685        assert_eq!(used.token(), token);
1686        assert_eq!(used.to_bytes().unwrap().as_ref(), b"response data");
1687    }
1688
1689    #[test]
1690    fn test_poll_used_requires_direct_slice() {
1691        let ring = make_ring(16);
1692        let layout = ring.layout();
1693        let test_mem = ring.mem();
1694        let pool_base = test_mem.base_addr() + Layout::query_size(ring.len()) as u64 + 0x100;
1695        let pool = TestPool::new(pool_base, 0x8000);
1696        let notifier = TestNotifier::new();
1697        let mem = NoDirectSliceMem(test_mem);
1698        let mut producer = VirtqProducer::new(layout, mem.clone(), notifier.clone(), pool);
1699        let mut consumer = VirtqConsumer::new(layout, mem, notifier);
1700
1701        let mut se = producer.chain().readable(64).writable(128).build().unwrap();
1702        se.write_all(b"request data").unwrap();
1703        producer.submit(se).unwrap();
1704
1705        let (_recv, reply) = poll_received(&mut consumer);
1706        if let ReplyChain::Writable(mut wc) = reply {
1707            wc.write_all(b"response data").unwrap();
1708            consumer.complete(wc).unwrap();
1709        } else {
1710            panic!("expected Writable");
1711        }
1712
1713        assert!(matches!(producer.poll(), Err(VirtqError::MemoryReadError)));
1714    }
1715
1716    #[test]
1717    fn test_villain_used_len_exceeding_writable_capacity_is_rejected_and_released() {
1718        let ring = make_ring(16);
1719        let (mut producer, _consumer, _notifier) = make_test_producer(&ring);
1720        let mut ring_consumer = make_consumer(&ring);
1721
1722        let se = producer.chain().writable(4).build().unwrap();
1723        producer.submit(se).unwrap();
1724
1725        let (id, _) = ring_consumer.poll_available().unwrap();
1726        ring_consumer.submit_used(id, 8).unwrap();
1727
1728        assert!(matches!(producer.poll(), Err(VirtqError::InvalidState)));
1729        assert_eq!(producer.inner.num_inflight(), 0);
1730
1731        let se = producer.chain().writable(4).build().unwrap();
1732        producer.submit(se).unwrap();
1733        assert_eq!(producer.inner.num_inflight(), 1);
1734    }
1735
1736    #[test]
1737    fn test_villain_used_descriptor_with_invalid_id_is_rejected() {
1738        let ring = make_ring(16);
1739        let (mut producer, _consumer, _notifier) = make_test_producer(&ring);
1740
1741        let se = producer.chain().writable(4).build().unwrap();
1742        producer.submit(se).unwrap();
1743
1744        let mut desc = Descriptor::new(0, 0, ring.len() as u16, DescFlags::empty());
1745        desc.mark_used(true);
1746        ring.write_desc(0, desc);
1747
1748        assert!(matches!(
1749            producer.poll(),
1750            Err(VirtqError::RingError(RingError::InvalidState))
1751        ));
1752        assert_eq!(producer.inner.num_inflight(), 1);
1753    }
1754
1755    #[test]
1756    fn test_virtq_producer_reset() {
1757        let ring = make_ring(16);
1758        let (mut producer, mut consumer, _notifier) = make_test_producer(&ring);
1759
1760        // Submit and complete a round trip
1761        let mut se = producer.chain().readable(32).writable(64).build().unwrap();
1762        se.write_all(b"hello").unwrap();
1763        producer.submit(se).unwrap();
1764
1765        let (recv, reply) = poll_received(&mut consumer);
1766        assert_eq!(recv.to_bytes().as_ref(), b"hello");
1767        consumer.complete(reply).unwrap();
1768        let _ = producer.poll().unwrap().unwrap();
1769
1770        // Now reset
1771        // SAFETY: the used chain was dropped before reset and no peer can
1772        // access the reset test ring concurrently.
1773        unsafe {
1774            producer.reset();
1775        }
1776
1777        // All inflight slots should be cleared
1778        assert_eq!(producer.inner.num_inflight(), 0);
1779        // Ring state should be back to initial
1780        assert_eq!(producer.inner.num_free(), producer.inner.len());
1781    }
1782
1783    #[test]
1784    fn test_virtq_producer_reset_clears_inflight() {
1785        let ring = make_ring(16);
1786        let (mut producer, _consumer, _notifier) = make_test_producer(&ring);
1787
1788        // Submit without completing
1789        let se = producer.chain().writable(64).build().unwrap();
1790        producer.submit(se).unwrap();
1791
1792        assert_eq!(producer.inner.num_inflight(), 1);
1793
1794        // SAFETY: no peer can access the reset test ring concurrently.
1795        unsafe {
1796            producer.reset();
1797        }
1798
1799        assert_eq!(producer.inner.num_inflight(), 0);
1800        assert_eq!(producer.inner.num_free(), producer.inner.len());
1801    }
1802}