Skip to main content

virtio_accel_device/
response.rs

1//! Bounded response framing over a potentially segmented destination.
2//!
3//! Callers can expose an exact payload subregion directly to a backend, then commit the 16-byte
4//! response header only after the payload is initialized. The convenience streaming path uses one
5//! fixed 4 KiB stack scratch buffer and never allocates or writes beyond the preflighted frame.
6
7use virtio_accel_core::{ByteSink, ByteSource};
8use virtio_accel_proto::{ResponseHeader, StatusCode};
9use zerocopy::IntoBytes;
10
11const RESPONSE_HEADER_BYTES: u64 =
12    core::mem::size_of::<virtio_accel_proto::ResponseHeader>() as u64;
13
14#[derive(Clone, Copy, Debug, PartialEq, Eq)]
15pub enum ResponseWriteError {
16    FrameTooLarge,
17    InsufficientCapacity,
18    SourceAccess,
19    SinkAccess,
20}
21
22/// Preflighted writer for one response frame.
23pub struct ResponseWriter<'a> {
24    sink: &'a mut dyn ByteSink,
25    max_response_bytes: u32,
26}
27
28impl core::fmt::Debug for ResponseWriter<'_> {
29    fn fmt(&self, formatter: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
30        formatter
31            .debug_struct("ResponseWriter")
32            .field("capacity", &self.sink.len())
33            .field("max_response_bytes", &self.max_response_bytes)
34            .finish()
35    }
36}
37
38impl<'a> ResponseWriter<'a> {
39    pub fn new(sink: &'a mut dyn ByteSink, max_response_bytes: u32) -> Self {
40        Self {
41            sink,
42            max_response_bytes,
43        }
44    }
45
46    pub const fn max_response_bytes(&self) -> u32 {
47        self.max_response_bytes
48    }
49
50    /// Borrow the exact success-payload destination.
51    ///
52    /// The returned guard commits the same payload length that was preflighted here, preventing a
53    /// caller from initializing one range and advertising another in the response header.
54    pub fn payload(
55        &mut self,
56        payload_bytes: u64,
57    ) -> Result<ResponsePayload<'_>, ResponseWriteError> {
58        let used = self.preflight(payload_bytes)?;
59        Ok(ResponsePayload {
60            sink: self.sink,
61            payload_bytes: u32::try_from(payload_bytes)
62                .map_err(|_| ResponseWriteError::FrameTooLarge)?,
63            used,
64        })
65    }
66
67    /// Write a bounded payload first, then atomically expose it by committing the header.
68    pub fn write_response(
69        &mut self,
70        status: StatusCode,
71        request_id: u64,
72        payload: &dyn ByteSource,
73    ) -> Result<u32, ResponseWriteError> {
74        let payload_bytes = payload.len();
75        let mut destination = self.payload(payload_bytes)?;
76
77        if let Some(contiguous) = payload.as_contiguous() {
78            destination
79                .write_at(0, contiguous)
80                .map_err(|_| ResponseWriteError::SinkAccess)?;
81        } else {
82            let mut scratch = [0_u8; 4096];
83            let mut offset = 0_u64;
84            while offset < payload_bytes {
85                let remaining = payload_bytes - offset;
86                let count = usize::try_from(remaining.min(scratch.as_slice().len() as u64))
87                    .map_err(|_| ResponseWriteError::FrameTooLarge)?;
88                payload
89                    .read_at(offset, &mut scratch[..count])
90                    .map_err(|_| ResponseWriteError::SourceAccess)?;
91                destination
92                    .write_at(offset, &scratch[..count])
93                    .map_err(|_| ResponseWriteError::SinkAccess)?;
94                offset += count as u64;
95            }
96        }
97
98        destination.commit(status, request_id)
99    }
100
101    pub fn write_empty(
102        &mut self,
103        status: StatusCode,
104        request_id: u64,
105    ) -> Result<u32, ResponseWriteError> {
106        let used = self.preflight(0)?;
107        write_header(self.sink, status, request_id, 0)?;
108        Ok(used)
109    }
110
111    fn preflight(&self, payload_bytes: u64) -> Result<u32, ResponseWriteError> {
112        let total = RESPONSE_HEADER_BYTES
113            .checked_add(payload_bytes)
114            .ok_or(ResponseWriteError::FrameTooLarge)?;
115        if total > u64::from(self.max_response_bytes) || total > u64::from(u32::MAX) {
116            return Err(ResponseWriteError::FrameTooLarge);
117        }
118        if total > self.sink.len() {
119            return Err(ResponseWriteError::InsufficientCapacity);
120        }
121        Ok(total as u32)
122    }
123}
124
125/// Exact payload destination for one preflighted response.
126pub struct ResponsePayload<'a> {
127    sink: &'a mut dyn ByteSink,
128    payload_bytes: u32,
129    used: u32,
130}
131
132impl core::fmt::Debug for ResponsePayload<'_> {
133    fn fmt(&self, formatter: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
134        formatter
135            .debug_struct("ResponsePayload")
136            .field("payload_bytes", &self.payload_bytes)
137            .field("used", &self.used)
138            .finish()
139    }
140}
141
142impl ResponsePayload<'_> {
143    /// Commit the response header after the complete payload has been initialized.
144    pub fn commit(self, status: StatusCode, request_id: u64) -> Result<u32, ResponseWriteError> {
145        write_header(self.sink, status, request_id, self.payload_bytes)?;
146        Ok(self.used)
147    }
148}
149
150impl ByteSink for ResponsePayload<'_> {
151    fn len(&self) -> u64 {
152        u64::from(self.payload_bytes)
153    }
154
155    fn write_at(
156        &mut self,
157        offset: u64,
158        source: &[u8],
159    ) -> Result<(), virtio_accel_core::BackendError> {
160        let bytes = u64::try_from(source.len())
161            .map_err(|_| virtio_accel_core::BackendError::OutOfBounds)?;
162        let end = offset
163            .checked_add(bytes)
164            .ok_or(virtio_accel_core::BackendError::OutOfBounds)?;
165        if end > self.len() {
166            return Err(virtio_accel_core::BackendError::OutOfBounds);
167        }
168        self.sink.write_at(RESPONSE_HEADER_BYTES + offset, source)
169    }
170
171    fn as_contiguous_mut(&mut self) -> Option<&mut [u8]> {
172        let sink = self.sink.as_contiguous_mut()?;
173        let end = RESPONSE_HEADER_BYTES
174            .checked_add(u64::from(self.payload_bytes))
175            .and_then(|end| usize::try_from(end).ok())?;
176        sink.get_mut(RESPONSE_HEADER_BYTES as usize..end)
177    }
178}
179
180fn write_header(
181    sink: &mut dyn ByteSink,
182    status: StatusCode,
183    request_id: u64,
184    payload_bytes: u32,
185) -> Result<(), ResponseWriteError> {
186    let header = ResponseHeader::new(status, payload_bytes, request_id);
187    sink.write_at(0, header.as_bytes())
188        .map_err(|_| ResponseWriteError::SinkAccess)
189}
190
191#[cfg(test)]
192mod tests {
193    use super::*;
194    use crate::{SegmentedSink, SegmentedSource};
195
196    #[test]
197    fn response_header_and_payload_cross_every_writable_split() {
198        let payload = *b"response";
199        let payload_segments = [&payload[..3], &payload[3..]];
200        let payload_source = SegmentedSource::new(&payload_segments).unwrap();
201        let expected_len = 16 + payload.as_slice().len();
202        for split in 1..expected_len {
203            let mut output = [0_u8; 24];
204            let (left, right) = output.split_at_mut(split);
205            let mut segments: [&mut [u8]; 2] = [left, right];
206            let mut sink = SegmentedSink::new(&mut segments).unwrap();
207            let mut writer = ResponseWriter::new(&mut sink, 1024);
208            assert_eq!(
209                writer
210                    .write_response(StatusCode::OK, 0x0102_0304_0506_0708, &payload_source)
211                    .unwrap(),
212                expected_len as u32
213            );
214
215            assert_eq!(&output[0..2], &StatusCode::OK.0.to_le_bytes());
216            assert_eq!(
217                &output[4..8],
218                &(payload.as_slice().len() as u32).to_le_bytes()
219            );
220            assert_eq!(&output[16..], &payload);
221        }
222    }
223
224    #[test]
225    fn direct_payload_region_does_not_touch_excess_capacity() {
226        let mut output = [0xaa_u8; 32];
227        let mut writer = ResponseWriter::new(&mut output, 32);
228        let mut payload = writer.payload(4).unwrap();
229        payload.write_at(0, b"data").unwrap();
230        assert_eq!(payload.commit(StatusCode::OK, 7), Ok(20));
231        assert_eq!(&output[16..20], b"data");
232        assert_eq!(&output[20..], &[0xaa; 12]);
233    }
234
235    #[test]
236    fn response_length_overflow_is_rejected_before_writing() {
237        let mut output = [0xaa_u8; 16];
238        let mut writer = ResponseWriter::new(&mut output, u32::MAX);
239        assert_eq!(
240            writer.payload(u64::MAX).unwrap_err(),
241            ResponseWriteError::FrameTooLarge
242        );
243        assert_eq!(output, [0xaa; 16]);
244    }
245
246    #[test]
247    fn payload_guard_commits_the_preflighted_length() {
248        let mut output = [0xaa_u8; 24];
249        let mut writer = ResponseWriter::new(&mut output, 24);
250        let mut payload = writer.payload(8).unwrap();
251        assert_eq!(payload.len(), 8);
252        assert!(payload.write_at(7, &[1, 2]).is_err());
253        payload.write_at(0, &[0; 8]).unwrap();
254        assert_eq!(payload.commit(StatusCode::OK, 9), Ok(24));
255        assert_eq!(&output[4..8], &8_u32.to_le_bytes());
256    }
257}