1use 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
22pub 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 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 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
125pub 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 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}