1use alloc::vec::Vec;
10use core::mem::size_of;
11
12use virtio_accel_core::{
13 AccessMode, ArtifactFormat, BufferDesc, BufferRange, BufferUsage, ByteSource, Capabilities,
14 ContextDesc, DeviceInfo, MemoryDomain, QueueDesc, TargetIdentity, Timeout,
15};
16use virtio_accel_proto::{
17 AllocateBufferRequest, ConfigError, CreateContextRequest, CreateQueueRequest,
18 HARD_MAX_BINDINGS, KnownOpcode, LoadProgramRequest, ObjectPayload, RequestHeader, StatusCode,
19 SubmitRequest, SubmitResponse, TransferBufferRequest, WireBinding, WireConfig, WireDeviceInfo,
20 WireEventState, read_exact,
21};
22use zerocopy::FromBytes;
23
24use crate::{ObjectId, ReadableRegion};
25
26const REQUEST_HEADER_BYTES: u64 = size_of::<RequestHeader>() as u64;
27const RESPONSE_HEADER_BYTES: u64 = size_of::<virtio_accel_proto::ResponseHeader>() as u64;
28const MAX_FIXED_PREFIX_BYTES: usize = size_of::<LoadProgramRequest>();
29
30#[derive(Clone, Copy, Debug, PartialEq, Eq)]
31pub enum DecoderLimitsError {
32 Config(ConfigError),
33 BindingLimit,
34 BufferLimit,
35 ArtifactLimit,
36}
37
38#[derive(Clone, Copy, Debug, PartialEq, Eq)]
39pub struct DecoderLimits {
40 max_chain_descriptors: u16,
41 max_request_bytes: u32,
42 max_response_bytes: u32,
43 max_bindings: u32,
44 max_buffer_bytes: u64,
45 max_artifact_bytes: u64,
46 capabilities: Capabilities,
47}
48
49impl DecoderLimits {
50 pub fn new(config: &WireConfig, info: DeviceInfo) -> Result<Self, DecoderLimitsError> {
51 config.validate().map_err(DecoderLimitsError::Config)?;
52 if !(1..=HARD_MAX_BINDINGS).contains(&info.limits.max_bindings_per_submission) {
53 return Err(DecoderLimitsError::BindingLimit);
54 }
55 if info.limits.max_buffer_bytes == 0 {
56 return Err(DecoderLimitsError::BufferLimit);
57 }
58 if info.limits.max_artifact_bytes == 0 {
59 return Err(DecoderLimitsError::ArtifactLimit);
60 }
61 Ok(Self {
62 max_chain_descriptors: config.max_chain_descriptors.get(),
63 max_request_bytes: config.max_request_bytes.get(),
64 max_response_bytes: config.max_response_bytes.get(),
65 max_bindings: info.limits.max_bindings_per_submission,
66 max_buffer_bytes: info.limits.max_buffer_bytes,
67 max_artifact_bytes: info.limits.max_artifact_bytes,
68 capabilities: info.capabilities,
69 })
70 }
71
72 pub const fn max_chain_descriptors(self) -> u16 {
73 self.max_chain_descriptors
74 }
75
76 pub const fn max_request_bytes(self) -> u32 {
77 self.max_request_bytes
78 }
79
80 pub const fn max_response_bytes(self) -> u32 {
81 self.max_response_bytes
82 }
83}
84
85#[derive(Clone, Copy, Debug, PartialEq, Eq)]
86pub enum UnrecoverableDecodeError {
87 RequestHeader,
88 RequestAccess,
89 ResponseHeader,
90}
91
92#[derive(Clone, Copy, Debug, PartialEq, Eq)]
93pub enum FrameDecodeError {
94 Unrecoverable(UnrecoverableDecodeError),
95 Protocol {
96 request_id: u64,
97 status: StatusCode,
98 },
99 InsufficientResponse {
100 request_id: u64,
101 required: u64,
102 available: u64,
103 },
104}
105
106#[derive(Clone, Copy, Debug, PartialEq, Eq)]
107pub struct DecodedBinding {
108 pub buffer_id: ObjectId,
109 pub range: BufferRange,
110 pub slot: u32,
111 pub access: AccessMode,
112}
113
114#[derive(Debug)]
115pub enum DecodedRequestBody<'a> {
116 GetDeviceInfo,
117 CreateContext(ContextDesc),
118 DestroyContext {
119 context_id: ObjectId,
120 },
121 AllocateBuffer {
122 context_id: ObjectId,
123 desc: BufferDesc,
124 },
125 FreeBuffer {
126 buffer_id: ObjectId,
127 },
128 WriteBuffer {
129 buffer_id: ObjectId,
130 range: BufferRange,
131 data: ReadableRegion<'a>,
132 },
133 ReadBuffer {
134 buffer_id: ObjectId,
135 range: BufferRange,
136 },
137 LoadProgram {
138 context_id: ObjectId,
139 format: ArtifactFormat,
140 target: TargetIdentity,
141 payload: ReadableRegion<'a>,
142 resident_bytes: u64,
143 },
144 UnloadProgram {
145 program_id: ObjectId,
146 },
147 CreateQueue {
148 context_id: ObjectId,
149 desc: QueueDesc,
150 },
151 DestroyQueue {
152 queue_id: ObjectId,
153 },
154 Submit {
155 queue_id: ObjectId,
156 program_id: ObjectId,
157 bindings: Vec<DecodedBinding>,
158 timeout: Timeout,
159 },
160 PollEvent {
161 event_id: ObjectId,
162 },
163 CancelEvent {
164 event_id: ObjectId,
165 },
166 DestroyEvent {
167 event_id: ObjectId,
168 },
169}
170
171impl DecodedRequestBody<'_> {
172 pub const fn opcode(&self) -> KnownOpcode {
173 match self {
174 Self::GetDeviceInfo => KnownOpcode::GetDeviceInfo,
175 Self::CreateContext(_) => KnownOpcode::CreateContext,
176 Self::DestroyContext { .. } => KnownOpcode::DestroyContext,
177 Self::AllocateBuffer { .. } => KnownOpcode::AllocateBuffer,
178 Self::FreeBuffer { .. } => KnownOpcode::FreeBuffer,
179 Self::WriteBuffer { .. } => KnownOpcode::WriteBuffer,
180 Self::ReadBuffer { .. } => KnownOpcode::ReadBuffer,
181 Self::LoadProgram { .. } => KnownOpcode::LoadProgram,
182 Self::UnloadProgram { .. } => KnownOpcode::UnloadProgram,
183 Self::CreateQueue { .. } => KnownOpcode::CreateQueue,
184 Self::DestroyQueue { .. } => KnownOpcode::DestroyQueue,
185 Self::Submit { .. } => KnownOpcode::Submit,
186 Self::PollEvent { .. } => KnownOpcode::PollEvent,
187 Self::CancelEvent { .. } => KnownOpcode::CancelEvent,
188 Self::DestroyEvent { .. } => KnownOpcode::DestroyEvent,
189 }
190 }
191}
192
193#[derive(Debug)]
194pub struct DecodedRequest<'a> {
195 request_id: u64,
196 required_response_bytes: u32,
197 body: DecodedRequestBody<'a>,
198}
199
200impl<'a> DecodedRequest<'a> {
201 pub const fn request_id(&self) -> u64 {
202 self.request_id
203 }
204
205 pub const fn required_response_bytes(&self) -> u32 {
206 self.required_response_bytes
207 }
208
209 pub const fn body(&self) -> &DecodedRequestBody<'a> {
210 &self.body
211 }
212
213 pub fn into_body(self) -> DecodedRequestBody<'a> {
214 self.body
215 }
216}
217
218#[derive(Clone, Copy, Debug)]
224pub struct FrameDecoder {
225 limits: DecoderLimits,
226}
227
228impl FrameDecoder {
229 pub const fn new(limits: DecoderLimits) -> Self {
230 Self { limits }
231 }
232
233 pub const fn limits(&self) -> DecoderLimits {
234 self.limits
235 }
236
237 pub fn decode<'a>(
238 &self,
239 request: &'a dyn ByteSource,
240 response_capacity: u64,
241 ) -> Result<DecodedRequest<'a>, FrameDecodeError> {
242 if response_capacity < RESPONSE_HEADER_BYTES {
243 return Err(FrameDecodeError::Unrecoverable(
244 UnrecoverableDecodeError::ResponseHeader,
245 ));
246 }
247 if request.len() < REQUEST_HEADER_BYTES {
248 return Err(FrameDecodeError::Unrecoverable(
249 UnrecoverableDecodeError::RequestHeader,
250 ));
251 }
252
253 let header = read_wire(
254 request,
255 0,
256 size_of::<RequestHeader>(),
257 read_exact::<RequestHeader>,
258 )
259 .map_err(|_| FrameDecodeError::Unrecoverable(UnrecoverableDecodeError::RequestAccess))?;
260 let request_id = header.request_id.get();
261 let expected_frame_bytes = REQUEST_HEADER_BYTES + u64::from(header.payload_bytes.get());
262
263 if request.len() > u64::from(self.limits.max_request_bytes) {
264 return Err(protocol(request_id, StatusCode::RESOURCE_LIMIT));
265 }
266 if request.len() != expected_frame_bytes {
267 return Err(protocol(request_id, StatusCode::INVALID_ARGUMENT));
268 }
269 if request_id == 0 {
270 return Err(protocol(request_id, StatusCode::INVALID_ARGUMENT));
271 }
272 if header.flags.get() != 0 {
273 return Err(protocol(request_id, StatusCode::UNSUPPORTED));
274 }
275 let opcode = header
276 .known_opcode()
277 .map_err(|_| protocol(request_id, StatusCode::UNSUPPORTED))?;
278
279 let (body, required_response_bytes) = self
280 .decode_body(request, header.payload_bytes.get(), opcode)
281 .map_err(|error| match error {
282 BodyDecodeError::Protocol(status) => protocol(request_id, status),
283 BodyDecodeError::Access => {
284 FrameDecodeError::Unrecoverable(UnrecoverableDecodeError::RequestAccess)
285 }
286 })?;
287
288 if required_response_bytes > u64::from(self.limits.max_response_bytes) {
289 return Err(protocol(request_id, StatusCode::RESOURCE_LIMIT));
290 }
291 if response_capacity < required_response_bytes {
292 return Err(FrameDecodeError::InsufficientResponse {
293 request_id,
294 required: required_response_bytes,
295 available: response_capacity,
296 });
297 }
298
299 Ok(DecodedRequest {
300 request_id,
301 required_response_bytes: required_response_bytes as u32,
302 body,
303 })
304 }
305
306 fn decode_body<'a>(
307 &self,
308 request: &'a dyn ByteSource,
309 payload_bytes: u32,
310 opcode: KnownOpcode,
311 ) -> Result<(DecodedRequestBody<'a>, u64), BodyDecodeError> {
312 match opcode {
313 KnownOpcode::GetDeviceInfo => {
314 expect_payload_bytes(payload_bytes, 0)?;
315 Ok((
316 DecodedRequestBody::GetDeviceInfo,
317 RESPONSE_HEADER_BYTES + size_of::<WireDeviceInfo>() as u64,
318 ))
319 }
320 KnownOpcode::CreateContext => {
321 expect_payload_bytes(payload_bytes, size_of::<CreateContextRequest>())?;
322 let value = read_payload::<CreateContextRequest>(request)?;
323 if value.flags.get() != 0 {
324 return Err(BodyDecodeError::Protocol(StatusCode::UNSUPPORTED));
325 }
326 if value.reserved.get() != 0 {
327 return Err(invalid());
328 }
329 Ok((
330 DecodedRequestBody::CreateContext(ContextDesc::default()),
331 RESPONSE_HEADER_BYTES + size_of::<ObjectPayload>() as u64,
332 ))
333 }
334 KnownOpcode::DestroyContext => Ok((
335 DecodedRequestBody::DestroyContext {
336 context_id: decode_object(request, payload_bytes)?,
337 },
338 RESPONSE_HEADER_BYTES,
339 )),
340 KnownOpcode::AllocateBuffer => {
341 expect_payload_bytes(payload_bytes, size_of::<AllocateBufferRequest>())?;
342 let value = read_payload::<AllocateBufferRequest>(request)?;
343 let context_id = object_id(value.context_id.get())?;
344 if value.reserved0 != [0; 7] || value.reserved1.get() != 0 {
345 return Err(invalid());
346 }
347
348 let domain = MemoryDomain::try_from(value.memory_domain).map_err(|_| invalid())?;
349 let usage = BufferUsage::from_bits(value.usage.get())
350 .ok_or(BodyDecodeError::Protocol(StatusCode::UNSUPPORTED))?;
351 if usage.is_empty() {
352 return Err(invalid());
353 }
354 let desc = BufferDesc::new(value.bytes.get(), value.alignment.get(), domain, usage)
355 .map_err(|_| invalid())?;
356 if desc.bytes() > self.limits.max_buffer_bytes {
357 return Err(BodyDecodeError::Protocol(StatusCode::RESOURCE_LIMIT));
358 }
359 if !self.limits.capabilities.supports_memory_domain(domain) {
360 return Err(BodyDecodeError::Protocol(StatusCode::UNSUPPORTED));
361 }
362
363 Ok((
364 DecodedRequestBody::AllocateBuffer { context_id, desc },
365 RESPONSE_HEADER_BYTES + size_of::<ObjectPayload>() as u64,
366 ))
367 }
368 KnownOpcode::FreeBuffer => Ok((
369 DecodedRequestBody::FreeBuffer {
370 buffer_id: decode_object(request, payload_bytes)?,
371 },
372 RESPONSE_HEADER_BYTES,
373 )),
374 KnownOpcode::WriteBuffer => {
375 let (buffer_id, range) = decode_transfer(request, payload_bytes, true)?;
376 let data = ReadableRegion::new(
377 request,
378 REQUEST_HEADER_BYTES + size_of::<TransferBufferRequest>() as u64,
379 range.bytes(),
380 )
381 .map_err(|_| BodyDecodeError::Access)?;
382 Ok((
383 DecodedRequestBody::WriteBuffer {
384 buffer_id,
385 range,
386 data,
387 },
388 RESPONSE_HEADER_BYTES,
389 ))
390 }
391 KnownOpcode::ReadBuffer => {
392 let (buffer_id, range) = decode_transfer(request, payload_bytes, false)?;
393 let required_response_bytes = RESPONSE_HEADER_BYTES
394 .checked_add(range.bytes())
395 .ok_or_else(resource_limit)?;
396 Ok((
397 DecodedRequestBody::ReadBuffer { buffer_id, range },
398 required_response_bytes,
399 ))
400 }
401 KnownOpcode::LoadProgram => {
402 if u64::from(payload_bytes) < size_of::<LoadProgramRequest>() as u64 {
403 return Err(invalid());
404 }
405 let value = read_payload::<LoadProgramRequest>(request)?;
406 if value.flags.get() != 0 {
407 return Err(BodyDecodeError::Protocol(StatusCode::UNSUPPORTED));
408 }
409 let context_id = object_id(value.context_id.get())?;
410 let format = ArtifactFormat::new(value.format.get()).ok_or_else(invalid)?;
411 let payload_len = value.payload_bytes.get();
412 if payload_len == 0 || value.resident_bytes.get() == 0 {
413 return Err(invalid());
414 }
415 let expected = (size_of::<LoadProgramRequest>() as u64)
416 .checked_add(payload_len)
417 .ok_or_else(resource_limit)?;
418 if expected != u64::from(payload_bytes) {
419 return Err(invalid());
420 }
421 if payload_len > self.limits.max_artifact_bytes {
422 return Err(BodyDecodeError::Protocol(StatusCode::RESOURCE_LIMIT));
423 }
424 let payload = ReadableRegion::new(
425 request,
426 REQUEST_HEADER_BYTES + size_of::<LoadProgramRequest>() as u64,
427 payload_len,
428 )
429 .map_err(|_| BodyDecodeError::Access)?;
430 let target =
431 TargetIdentity(core::array::from_fn(|index| value.target[index].get()));
432 Ok((
433 DecodedRequestBody::LoadProgram {
434 context_id,
435 format,
436 target,
437 payload,
438 resident_bytes: value.resident_bytes.get(),
439 },
440 RESPONSE_HEADER_BYTES + size_of::<ObjectPayload>() as u64,
441 ))
442 }
443 KnownOpcode::UnloadProgram => Ok((
444 DecodedRequestBody::UnloadProgram {
445 program_id: decode_object(request, payload_bytes)?,
446 },
447 RESPONSE_HEADER_BYTES,
448 )),
449 KnownOpcode::CreateQueue => {
450 expect_payload_bytes(payload_bytes, size_of::<CreateQueueRequest>())?;
451 let value = read_payload::<CreateQueueRequest>(request)?;
452 if value.flags.get() != 0 {
453 return Err(BodyDecodeError::Protocol(StatusCode::UNSUPPORTED));
454 }
455 if value.reserved.get() != 0 {
456 return Err(invalid());
457 }
458 Ok((
459 DecodedRequestBody::CreateQueue {
460 context_id: object_id(value.context_id.get())?,
461 desc: QueueDesc::default(),
462 },
463 RESPONSE_HEADER_BYTES + size_of::<ObjectPayload>() as u64,
464 ))
465 }
466 KnownOpcode::DestroyQueue => Ok((
467 DecodedRequestBody::DestroyQueue {
468 queue_id: decode_object(request, payload_bytes)?,
469 },
470 RESPONSE_HEADER_BYTES,
471 )),
472 KnownOpcode::Submit => {
473 let (queue_id, program_id, bindings, timeout) =
474 self.decode_submit(request, payload_bytes)?;
475 Ok((
476 DecodedRequestBody::Submit {
477 queue_id,
478 program_id,
479 bindings,
480 timeout,
481 },
482 RESPONSE_HEADER_BYTES + size_of::<SubmitResponse>() as u64,
483 ))
484 }
485 KnownOpcode::PollEvent => Ok((
486 DecodedRequestBody::PollEvent {
487 event_id: decode_object(request, payload_bytes)?,
488 },
489 RESPONSE_HEADER_BYTES + size_of::<WireEventState>() as u64,
490 )),
491 KnownOpcode::CancelEvent => Ok((
492 DecodedRequestBody::CancelEvent {
493 event_id: decode_object(request, payload_bytes)?,
494 },
495 RESPONSE_HEADER_BYTES,
496 )),
497 KnownOpcode::DestroyEvent => Ok((
498 DecodedRequestBody::DestroyEvent {
499 event_id: decode_object(request, payload_bytes)?,
500 },
501 RESPONSE_HEADER_BYTES,
502 )),
503 }
504 }
505
506 fn decode_submit(
507 &self,
508 request: &dyn ByteSource,
509 payload_bytes: u32,
510 ) -> Result<(ObjectId, ObjectId, Vec<DecodedBinding>, Timeout), BodyDecodeError> {
511 if u64::from(payload_bytes) < size_of::<SubmitRequest>() as u64 {
512 return Err(invalid());
513 }
514 let value = read_payload::<SubmitRequest>(request)?;
515 if value.flags.get() != 0 {
516 return Err(BodyDecodeError::Protocol(StatusCode::UNSUPPORTED));
517 }
518
519 let binding_count = value.binding_count.get();
520 if binding_count == 0 {
521 return Err(invalid());
522 }
523 if binding_count > self.limits.max_bindings || binding_count > HARD_MAX_BINDINGS {
524 return Err(BodyDecodeError::Protocol(StatusCode::RESOURCE_LIMIT));
525 }
526 let binding_bytes = u64::from(binding_count)
527 .checked_mul(size_of::<WireBinding>() as u64)
528 .ok_or_else(resource_limit)?;
529 let expected = (size_of::<SubmitRequest>() as u64)
530 .checked_add(binding_bytes)
531 .ok_or_else(resource_limit)?;
532 if expected != u64::from(payload_bytes) {
533 return Err(invalid());
534 }
535
536 let count = binding_count as usize;
537 let mut bindings = Vec::new();
538 bindings
539 .try_reserve_exact(count)
540 .map_err(|_| BodyDecodeError::Protocol(StatusCode::OUT_OF_MEMORY))?;
541
542 let binding_base = REQUEST_HEADER_BYTES + size_of::<SubmitRequest>() as u64;
543 for index in 0..count {
544 let offset = binding_base + (index * size_of::<WireBinding>()) as u64;
545 let binding = read_wire(
546 request,
547 offset,
548 size_of::<WireBinding>(),
549 read_exact::<WireBinding>,
550 )?;
551 if binding.reserved != [0; 3] {
552 return Err(invalid());
553 }
554 let bytes = binding.bytes.get();
555 if bytes == 0 || binding.offset.get().checked_add(bytes).is_none() {
556 return Err(invalid());
557 }
558 let range = BufferRange::new(binding.offset.get(), bytes).map_err(|_| invalid())?;
559 let access = AccessMode::try_from(binding.access).map_err(|_| invalid())?;
560 bindings.push(DecodedBinding {
561 buffer_id: object_id(binding.buffer_id.get())?,
562 range,
563 slot: binding.slot.get(),
564 access,
565 });
566 }
567
568 bindings.sort_unstable_by_key(|binding| binding.slot);
569 if bindings
570 .windows(2)
571 .any(|window| window[0].slot == window[1].slot)
572 {
573 return Err(invalid());
574 }
575
576 Ok((
577 object_id(value.queue_id.get())?,
578 object_id(value.program_id.get())?,
579 bindings,
580 Timeout::from_wire_ns(value.timeout_ns.get()),
581 ))
582 }
583}
584
585#[derive(Clone, Copy, Debug, PartialEq, Eq)]
586enum BodyDecodeError {
587 Protocol(StatusCode),
588 Access,
589}
590
591fn protocol(request_id: u64, status: StatusCode) -> FrameDecodeError {
592 FrameDecodeError::Protocol { request_id, status }
593}
594
595fn invalid() -> BodyDecodeError {
596 BodyDecodeError::Protocol(StatusCode::INVALID_ARGUMENT)
597}
598
599fn resource_limit() -> BodyDecodeError {
600 BodyDecodeError::Protocol(StatusCode::RESOURCE_LIMIT)
601}
602
603fn expect_payload_bytes(payload_bytes: u32, expected: usize) -> Result<(), BodyDecodeError> {
604 if u64::from(payload_bytes) != expected as u64 {
605 return Err(invalid());
606 }
607 Ok(())
608}
609
610fn read_payload<T>(request: &dyn ByteSource) -> Result<T, BodyDecodeError>
611where
612 T: FromBytes,
613{
614 read_wire(
615 request,
616 REQUEST_HEADER_BYTES,
617 size_of::<T>(),
618 read_exact::<T>,
619 )
620}
621
622fn read_wire<T>(
623 source: &dyn ByteSource,
624 offset: u64,
625 bytes: usize,
626 decode: impl FnOnce(&[u8]) -> Result<T, virtio_accel_proto::DecodeError>,
627) -> Result<T, BodyDecodeError> {
628 if bytes > MAX_FIXED_PREFIX_BYTES {
629 return Err(BodyDecodeError::Protocol(StatusCode::INTERNAL_ERROR));
630 }
631 let mut scratch = [0_u8; MAX_FIXED_PREFIX_BYTES];
632 source
633 .read_at(offset, &mut scratch[..bytes])
634 .map_err(|_| BodyDecodeError::Access)?;
635 decode(&scratch[..bytes]).map_err(|_| invalid())
636}
637
638fn decode_object(
639 request: &dyn ByteSource,
640 payload_bytes: u32,
641) -> Result<ObjectId, BodyDecodeError> {
642 expect_payload_bytes(payload_bytes, size_of::<ObjectPayload>())?;
643 object_id(read_payload::<ObjectPayload>(request)?.object_id.get())
644}
645
646fn object_id(raw: u64) -> Result<ObjectId, BodyDecodeError> {
647 ObjectId::from_raw(raw).ok_or_else(invalid)
648}
649
650fn decode_transfer(
651 request: &dyn ByteSource,
652 payload_bytes: u32,
653 has_data: bool,
654) -> Result<(ObjectId, BufferRange), BodyDecodeError> {
655 if u64::from(payload_bytes) < size_of::<TransferBufferRequest>() as u64 {
656 return Err(invalid());
657 }
658 let value = read_payload::<TransferBufferRequest>(request)?;
659 let bytes = value.bytes.get();
660 if bytes == 0 || value.offset.get().checked_add(bytes).is_none() {
661 return Err(invalid());
662 }
663 let expected = if has_data {
664 (size_of::<TransferBufferRequest>() as u64)
665 .checked_add(bytes)
666 .ok_or_else(resource_limit)?
667 } else {
668 size_of::<TransferBufferRequest>() as u64
669 };
670 if expected != u64::from(payload_bytes) {
671 return Err(invalid());
672 }
673 Ok((
674 object_id(value.buffer_id.get())?,
675 BufferRange::new(value.offset.get(), bytes).map_err(|_| invalid())?,
676 ))
677}