1#![no_std]
8#![forbid(unsafe_code)]
9
10use core::convert::TryFrom;
11
12pub const PROTOCOL_MAJOR: u16 = 1;
13pub const PROTOCOL_MINOR: u16 = 0;
14pub const REQUEST_HEADER_BYTES: usize = 16;
15pub const RESPONSE_HEADER_BYTES: usize = 16;
16pub const HARD_MAX_CHAIN_DESCRIPTORS: u16 = 256;
17pub const HARD_MAX_REQUEST_BYTES: u32 = 16 * 1024 * 1024;
18pub const HARD_MAX_RESPONSE_BYTES: u32 = 16 * 1024 * 1024;
19pub const HARD_MAX_BINDINGS: u32 = 4_096;
20pub const MIN_MAX_REQUEST_BYTES: u32 = 97;
21pub const MIN_MAX_RESPONSE_BYTES: u32 = 92;
22pub const KNOWN_BUFFER_USAGE_BITS: u32 = 0x1f;
23
24#[derive(Clone, Copy, Debug, PartialEq, Eq)]
25pub enum Error {
26 Size,
27 Version,
28 CommandQueueCount,
29 ChainDescriptorLimit,
30 RequestByteLimit,
31 ResponseByteLimit,
32 Unsupported,
33 InvalidArgument,
34 ResourceLimit,
35 RecoveryRequired,
36 OutputTooSmall,
37}
38
39#[derive(Clone, Copy, Debug, PartialEq, Eq)]
40pub struct Config {
41 pub protocol_major: u16,
42 pub protocol_minor: u16,
43 pub command_queue_count: u16,
44 pub max_chain_descriptors: u16,
45 pub max_request_bytes: u32,
46 pub max_response_bytes: u32,
47}
48
49impl Config {
50 pub const ENCODED_BYTES: usize = 16;
51
52 pub fn validate(self, queue_size: u16) -> Result<(), Error> {
53 if self.protocol_major != PROTOCOL_MAJOR {
54 return Err(Error::Version);
55 }
56 if self.command_queue_count != 1 {
57 return Err(Error::CommandQueueCount);
58 }
59 if !(2..=HARD_MAX_CHAIN_DESCRIPTORS).contains(&self.max_chain_descriptors)
60 || self.max_chain_descriptors > queue_size
61 {
62 return Err(Error::ChainDescriptorLimit);
63 }
64 if !(MIN_MAX_REQUEST_BYTES..=HARD_MAX_REQUEST_BYTES).contains(&self.max_request_bytes) {
65 return Err(Error::RequestByteLimit);
66 }
67 if !(MIN_MAX_RESPONSE_BYTES..=HARD_MAX_RESPONSE_BYTES).contains(&self.max_response_bytes) {
68 return Err(Error::ResponseByteLimit);
69 }
70 Ok(())
71 }
72
73 pub fn encode(self, queue_size: u16, output: &mut [u8]) -> Result<usize, Error> {
74 self.validate(queue_size)?;
75 require_output(output, Self::ENCODED_BYTES)?;
76 write_u16(output, 0, self.protocol_major)?;
77 write_u16(output, 2, self.protocol_minor)?;
78 write_u16(output, 4, self.command_queue_count)?;
79 write_u16(output, 6, self.max_chain_descriptors)?;
80 write_u32(output, 8, self.max_request_bytes)?;
81 write_u32(output, 12, self.max_response_bytes)?;
82 Ok(Self::ENCODED_BYTES)
83 }
84}
85
86pub fn decode_config(bytes: &[u8], queue_size: u16) -> Result<Config, Error> {
87 if bytes.len() != Config::ENCODED_BYTES {
88 return Err(Error::Size);
89 }
90 let config = Config {
91 protocol_major: read_u16(bytes, 0)?,
92 protocol_minor: read_u16(bytes, 2)?,
93 command_queue_count: read_u16(bytes, 4)?,
94 max_chain_descriptors: read_u16(bytes, 6)?,
95 max_request_bytes: read_u32(bytes, 8)?,
96 max_response_bytes: read_u32(bytes, 12)?,
97 };
98 config.validate(queue_size)?;
99 Ok(config)
100}
101
102pub fn validate_features(features: u64) -> Result<(), Error> {
103 if features == 0 {
104 Ok(())
105 } else {
106 Err(Error::Unsupported)
107 }
108}
109
110#[derive(Clone, Copy, Debug, PartialEq, Eq)]
111#[repr(u16)]
112pub enum Opcode {
113 GetDeviceInfo = 0x0001,
114 CreateContext = 0x0100,
115 DestroyContext = 0x0101,
116 AllocateBuffer = 0x0200,
117 FreeBuffer = 0x0201,
118 WriteBuffer = 0x0202,
119 ReadBuffer = 0x0203,
120 LoadProgram = 0x0300,
121 UnloadProgram = 0x0301,
122 CreateQueue = 0x0400,
123 DestroyQueue = 0x0401,
124 Submit = 0x0500,
125 PollEvent = 0x0501,
126 CancelEvent = 0x0502,
127 DestroyEvent = 0x0503,
128}
129
130impl TryFrom<u16> for Opcode {
131 type Error = Error;
132
133 fn try_from(value: u16) -> Result<Self, Self::Error> {
134 match value {
135 0x0001 => Ok(Self::GetDeviceInfo),
136 0x0100 => Ok(Self::CreateContext),
137 0x0101 => Ok(Self::DestroyContext),
138 0x0200 => Ok(Self::AllocateBuffer),
139 0x0201 => Ok(Self::FreeBuffer),
140 0x0202 => Ok(Self::WriteBuffer),
141 0x0203 => Ok(Self::ReadBuffer),
142 0x0300 => Ok(Self::LoadProgram),
143 0x0301 => Ok(Self::UnloadProgram),
144 0x0400 => Ok(Self::CreateQueue),
145 0x0401 => Ok(Self::DestroyQueue),
146 0x0500 => Ok(Self::Submit),
147 0x0501 => Ok(Self::PollEvent),
148 0x0502 => Ok(Self::CancelEvent),
149 0x0503 => Ok(Self::DestroyEvent),
150 _ => Err(Error::Unsupported),
151 }
152 }
153}
154
155#[derive(Clone, Copy, Debug, PartialEq, Eq)]
156pub struct AllocateBuffer {
157 pub context_id: u64,
158 pub bytes: u64,
159 pub alignment: u64,
160 pub memory_domain: u8,
161 pub usage: u32,
162}
163
164#[derive(Clone, Copy, Debug, PartialEq, Eq)]
165pub struct TransferBuffer {
166 pub buffer_id: u64,
167 pub offset: u64,
168 pub bytes: u64,
169}
170
171#[derive(Clone, Copy, Debug, PartialEq, Eq)]
172pub struct LoadProgram<'a> {
173 pub context_id: u64,
174 pub format: u32,
175 pub target: [u32; 12],
176 pub resident_bytes: u64,
177 pub artifact: &'a [u8],
178}
179
180#[derive(Clone, Copy, Debug, PartialEq, Eq)]
181pub struct Binding {
182 pub buffer_id: u64,
183 pub offset: u64,
184 pub bytes: u64,
185 pub slot: u32,
186 pub access: u8,
187}
188
189#[derive(Clone, Copy, Debug, PartialEq, Eq)]
190pub enum Bindings<'a> {
191 Encoded { bytes: &'a [u8], count: u32 },
192 Values(&'a [Binding]),
193}
194
195impl<'a> Bindings<'a> {
196 pub fn count(self) -> Result<u32, Error> {
197 match self {
198 Self::Encoded { count, .. } => Ok(count),
199 Self::Values(values) => u32::try_from(values.len()).map_err(|_| Error::ResourceLimit),
200 }
201 }
202
203 pub fn get(self, index: u32) -> Result<Binding, Error> {
204 match self {
205 Self::Encoded { bytes, count } => {
206 if index >= count {
207 return Err(Error::InvalidArgument);
208 }
209 let offset = usize::try_from(index)
210 .map_err(|_| Error::ResourceLimit)?
211 .checked_mul(32)
212 .ok_or(Error::ResourceLimit)?;
213 decode_binding(bytes, offset)
214 }
215 Self::Values(values) => values
216 .get(usize::try_from(index).map_err(|_| Error::ResourceLimit)?)
217 .copied()
218 .ok_or(Error::InvalidArgument),
219 }
220 }
221}
222
223#[derive(Clone, Copy, Debug, PartialEq, Eq)]
224pub struct Submit<'a> {
225 pub queue_id: u64,
226 pub program_id: u64,
227 pub timeout_ns: u64,
228 pub bindings: Bindings<'a>,
229}
230
231#[derive(Clone, Copy, Debug, PartialEq, Eq)]
232pub enum RequestBody<'a> {
233 GetDeviceInfo,
234 CreateContext,
235 DestroyContext {
236 context_id: u64,
237 },
238 AllocateBuffer(AllocateBuffer),
239 FreeBuffer {
240 buffer_id: u64,
241 },
242 WriteBuffer {
243 transfer: TransferBuffer,
244 data: &'a [u8],
245 },
246 ReadBuffer(TransferBuffer),
247 LoadProgram(LoadProgram<'a>),
248 UnloadProgram {
249 program_id: u64,
250 },
251 CreateQueue {
252 context_id: u64,
253 },
254 DestroyQueue {
255 queue_id: u64,
256 },
257 Submit(Submit<'a>),
258 PollEvent {
259 event_id: u64,
260 },
261 CancelEvent {
262 event_id: u64,
263 },
264 DestroyEvent {
265 event_id: u64,
266 },
267}
268
269impl RequestBody<'_> {
270 pub const fn opcode(&self) -> Opcode {
271 match self {
272 Self::GetDeviceInfo => Opcode::GetDeviceInfo,
273 Self::CreateContext => Opcode::CreateContext,
274 Self::DestroyContext { .. } => Opcode::DestroyContext,
275 Self::AllocateBuffer(_) => Opcode::AllocateBuffer,
276 Self::FreeBuffer { .. } => Opcode::FreeBuffer,
277 Self::WriteBuffer { .. } => Opcode::WriteBuffer,
278 Self::ReadBuffer(_) => Opcode::ReadBuffer,
279 Self::LoadProgram(_) => Opcode::LoadProgram,
280 Self::UnloadProgram { .. } => Opcode::UnloadProgram,
281 Self::CreateQueue { .. } => Opcode::CreateQueue,
282 Self::DestroyQueue { .. } => Opcode::DestroyQueue,
283 Self::Submit(_) => Opcode::Submit,
284 Self::PollEvent { .. } => Opcode::PollEvent,
285 Self::CancelEvent { .. } => Opcode::CancelEvent,
286 Self::DestroyEvent { .. } => Opcode::DestroyEvent,
287 }
288 }
289}
290
291#[derive(Clone, Copy, Debug, PartialEq, Eq)]
292pub struct Request<'a> {
293 pub request_id: u64,
294 pub body: RequestBody<'a>,
295}
296
297impl Request<'_> {
298 pub fn encoded_len(&self) -> Result<usize, Error> {
299 let payload = request_payload_len(&self.body)?;
300 let total = REQUEST_HEADER_BYTES
301 .checked_add(payload)
302 .ok_or(Error::ResourceLimit)?;
303 if total > HARD_MAX_REQUEST_BYTES as usize {
304 return Err(Error::ResourceLimit);
305 }
306 Ok(total)
307 }
308
309 pub fn encode(&self, output: &mut [u8]) -> Result<usize, Error> {
310 if self.request_id == 0 {
311 return Err(Error::InvalidArgument);
312 }
313 let payload_len = request_payload_len(&self.body)?;
314 let total = REQUEST_HEADER_BYTES
315 .checked_add(payload_len)
316 .ok_or(Error::ResourceLimit)?;
317 let payload_u32 = u32::try_from(payload_len).map_err(|_| Error::ResourceLimit)?;
318 if total > HARD_MAX_REQUEST_BYTES as usize {
319 return Err(Error::ResourceLimit);
320 }
321 require_output(output, total)?;
322 write_u16(output, 0, self.body.opcode() as u16)?;
323 write_u16(output, 2, 0)?;
324 write_u32(output, 4, payload_u32)?;
325 write_u64(output, 8, self.request_id)?;
326 encode_request_body(&self.body, &mut output[REQUEST_HEADER_BYTES..total])?;
327 Ok(total)
328 }
329}
330
331pub fn decode_request(bytes: &[u8]) -> Result<Request<'_>, Error> {
332 if bytes.len() < REQUEST_HEADER_BYTES {
333 return Err(Error::Size);
334 }
335 let opcode = Opcode::try_from(read_u16(bytes, 0)?)?;
336 if read_u16(bytes, 2)? != 0 {
337 return Err(Error::Unsupported);
338 }
339 let payload_len = usize::try_from(read_u32(bytes, 4)?).map_err(|_| Error::ResourceLimit)?;
340 let total = REQUEST_HEADER_BYTES
341 .checked_add(payload_len)
342 .ok_or(Error::ResourceLimit)?;
343 if total != bytes.len() {
344 return Err(Error::InvalidArgument);
345 }
346 if total > HARD_MAX_REQUEST_BYTES as usize {
347 return Err(Error::ResourceLimit);
348 }
349 let request_id = read_u64(bytes, 8)?;
350 if request_id == 0 {
351 return Err(Error::InvalidArgument);
352 }
353 let payload = &bytes[REQUEST_HEADER_BYTES..];
354 let body = decode_request_body(opcode, payload)?;
355 if request_payload_len(&body)? != payload_len {
356 return Err(Error::InvalidArgument);
357 }
358 Ok(Request { request_id, body })
359}
360
361fn decode_request_body(opcode: Opcode, payload: &[u8]) -> Result<RequestBody<'_>, Error> {
362 match opcode {
363 Opcode::GetDeviceInfo => {
364 require_exact(payload, 0)?;
365 Ok(RequestBody::GetDeviceInfo)
366 }
367 Opcode::CreateContext => {
368 require_exact(payload, 8)?;
369 if read_u32(payload, 0)? != 0 {
370 return Err(Error::Unsupported);
371 }
372 if read_u32(payload, 4)? != 0 {
373 return Err(Error::InvalidArgument);
374 }
375 Ok(RequestBody::CreateContext)
376 }
377 Opcode::DestroyContext => Ok(RequestBody::DestroyContext {
378 context_id: decode_object(payload)?,
379 }),
380 Opcode::AllocateBuffer => {
381 require_exact(payload, 40)?;
382 if payload[25..32].iter().any(|byte| *byte != 0) || read_u32(payload, 36)? != 0 {
383 return Err(Error::InvalidArgument);
384 }
385 Ok(RequestBody::AllocateBuffer(AllocateBuffer {
386 context_id: read_u64(payload, 0)?,
387 bytes: read_u64(payload, 8)?,
388 alignment: read_u64(payload, 16)?,
389 memory_domain: payload[24],
390 usage: read_u32(payload, 32)?,
391 }))
392 }
393 Opcode::FreeBuffer => Ok(RequestBody::FreeBuffer {
394 buffer_id: decode_object(payload)?,
395 }),
396 Opcode::WriteBuffer => {
397 if payload.len() < 24 {
398 return Err(Error::InvalidArgument);
399 }
400 Ok(RequestBody::WriteBuffer {
401 transfer: decode_transfer(payload)?,
402 data: &payload[24..],
403 })
404 }
405 Opcode::ReadBuffer => {
406 require_exact(payload, 24)?;
407 Ok(RequestBody::ReadBuffer(decode_transfer(payload)?))
408 }
409 Opcode::LoadProgram => {
410 if payload.len() < 80 {
411 return Err(Error::InvalidArgument);
412 }
413 if read_u32(payload, 12)? != 0 {
414 return Err(Error::Unsupported);
415 }
416 let mut target = [0_u32; 12];
417 for (index, word) in target.iter_mut().enumerate() {
418 *word = read_u32(payload, 16 + index * 4)?;
419 }
420 let artifact = &payload[80..];
421 let declared_bytes =
422 usize::try_from(read_u64(payload, 64)?).map_err(|_| Error::ResourceLimit)?;
423 if declared_bytes != artifact.len() {
424 return Err(Error::InvalidArgument);
425 }
426 Ok(RequestBody::LoadProgram(LoadProgram {
427 context_id: read_u64(payload, 0)?,
428 format: read_u32(payload, 8)?,
429 target,
430 resident_bytes: read_u64(payload, 72)?,
431 artifact,
432 }))
433 }
434 Opcode::UnloadProgram => Ok(RequestBody::UnloadProgram {
435 program_id: decode_object(payload)?,
436 }),
437 Opcode::CreateQueue => {
438 require_exact(payload, 16)?;
439 if read_u32(payload, 8)? != 0 {
440 return Err(Error::Unsupported);
441 }
442 if read_u32(payload, 12)? != 0 {
443 return Err(Error::InvalidArgument);
444 }
445 Ok(RequestBody::CreateQueue {
446 context_id: read_u64(payload, 0)?,
447 })
448 }
449 Opcode::DestroyQueue => Ok(RequestBody::DestroyQueue {
450 queue_id: decode_object(payload)?,
451 }),
452 Opcode::Submit => {
453 if payload.len() < 32 {
454 return Err(Error::InvalidArgument);
455 }
456 let count = read_u32(payload, 16)?;
457 if count == 0 {
458 return Err(Error::InvalidArgument);
459 }
460 if count > HARD_MAX_BINDINGS {
461 return Err(Error::ResourceLimit);
462 }
463 if read_u32(payload, 20)? != 0 {
464 return Err(Error::Unsupported);
465 }
466 Ok(RequestBody::Submit(Submit {
467 queue_id: read_u64(payload, 0)?,
468 program_id: read_u64(payload, 8)?,
469 timeout_ns: read_u64(payload, 24)?,
470 bindings: Bindings::Encoded {
471 bytes: &payload[32..],
472 count,
473 },
474 }))
475 }
476 Opcode::PollEvent => Ok(RequestBody::PollEvent {
477 event_id: decode_object(payload)?,
478 }),
479 Opcode::CancelEvent => Ok(RequestBody::CancelEvent {
480 event_id: decode_object(payload)?,
481 }),
482 Opcode::DestroyEvent => Ok(RequestBody::DestroyEvent {
483 event_id: decode_object(payload)?,
484 }),
485 }
486}
487
488fn request_payload_len(body: &RequestBody<'_>) -> Result<usize, Error> {
489 match body {
490 RequestBody::GetDeviceInfo => Ok(0),
491 RequestBody::CreateContext => Ok(8),
492 RequestBody::DestroyContext { context_id } => object_len(*context_id),
493 RequestBody::AllocateBuffer(request) => {
494 validate_allocate(*request)?;
495 Ok(40)
496 }
497 RequestBody::FreeBuffer { buffer_id } => object_len(*buffer_id),
498 RequestBody::WriteBuffer { transfer, data } => {
499 validate_transfer(*transfer)?;
500 let bytes = usize::try_from(transfer.bytes).map_err(|_| Error::ResourceLimit)?;
501 if bytes != data.len() {
502 return Err(Error::InvalidArgument);
503 }
504 24_usize.checked_add(bytes).ok_or(Error::ResourceLimit)
505 }
506 RequestBody::ReadBuffer(transfer) => {
507 validate_transfer(*transfer)?;
508 Ok(24)
509 }
510 RequestBody::LoadProgram(request) => {
511 if request.context_id == 0
512 || request.format == 0
513 || request.artifact.is_empty()
514 || request.resident_bytes == 0
515 {
516 return Err(Error::InvalidArgument);
517 }
518 80_usize
519 .checked_add(request.artifact.len())
520 .ok_or(Error::ResourceLimit)
521 }
522 RequestBody::UnloadProgram { program_id } => object_len(*program_id),
523 RequestBody::CreateQueue { context_id } => {
524 object_len(*context_id)?;
525 Ok(16)
526 }
527 RequestBody::DestroyQueue { queue_id } => object_len(*queue_id),
528 RequestBody::Submit(request) => {
529 if request.queue_id == 0 || request.program_id == 0 {
530 return Err(Error::InvalidArgument);
531 }
532 validate_bindings(request.bindings)
533 }
534 RequestBody::PollEvent { event_id }
535 | RequestBody::CancelEvent { event_id }
536 | RequestBody::DestroyEvent { event_id } => object_len(*event_id),
537 }
538}
539
540fn encode_request_body(body: &RequestBody<'_>, output: &mut [u8]) -> Result<(), Error> {
541 output.fill(0);
542 match body {
543 RequestBody::GetDeviceInfo | RequestBody::CreateContext => {}
544 RequestBody::DestroyContext { context_id } => write_u64(output, 0, *context_id)?,
545 RequestBody::AllocateBuffer(request) => {
546 write_u64(output, 0, request.context_id)?;
547 write_u64(output, 8, request.bytes)?;
548 write_u64(output, 16, request.alignment)?;
549 output[24] = request.memory_domain;
550 write_u32(output, 32, request.usage)?;
551 }
552 RequestBody::FreeBuffer { buffer_id } => write_u64(output, 0, *buffer_id)?,
553 RequestBody::WriteBuffer { transfer, data } => {
554 encode_transfer(*transfer, output)?;
555 write_slice(output, 24, data)?;
556 }
557 RequestBody::ReadBuffer(transfer) => encode_transfer(*transfer, output)?,
558 RequestBody::LoadProgram(request) => {
559 write_u64(output, 0, request.context_id)?;
560 write_u32(output, 8, request.format)?;
561 for (index, word) in request.target.iter().enumerate() {
562 write_u32(output, 16 + index * 4, *word)?;
563 }
564 write_u64(
565 output,
566 64,
567 u64::try_from(request.artifact.len()).map_err(|_| Error::ResourceLimit)?,
568 )?;
569 write_u64(output, 72, request.resident_bytes)?;
570 write_slice(output, 80, request.artifact)?;
571 }
572 RequestBody::UnloadProgram { program_id } => write_u64(output, 0, *program_id)?,
573 RequestBody::CreateQueue { context_id } => write_u64(output, 0, *context_id)?,
574 RequestBody::DestroyQueue { queue_id } => write_u64(output, 0, *queue_id)?,
575 RequestBody::Submit(request) => {
576 write_u64(output, 0, request.queue_id)?;
577 write_u64(output, 8, request.program_id)?;
578 let count = request.bindings.count()?;
579 write_u32(output, 16, count)?;
580 write_u64(output, 24, request.timeout_ns)?;
581 for index in 0..count {
582 encode_binding(
583 request.bindings.get(index)?,
584 output,
585 32 + index as usize * 32,
586 )?;
587 }
588 }
589 RequestBody::PollEvent { event_id }
590 | RequestBody::CancelEvent { event_id }
591 | RequestBody::DestroyEvent { event_id } => write_u64(output, 0, *event_id)?,
592 }
593 Ok(())
594}
595
596fn validate_allocate(request: AllocateBuffer) -> Result<(), Error> {
597 if request.context_id == 0
598 || request.bytes == 0
599 || request.alignment == 0
600 || !request.alignment.is_power_of_two()
601 || !(1..=3).contains(&request.memory_domain)
602 || request.usage == 0
603 {
604 return Err(Error::InvalidArgument);
605 }
606 if request.usage & !KNOWN_BUFFER_USAGE_BITS != 0 {
607 return Err(Error::Unsupported);
608 }
609 Ok(())
610}
611
612fn decode_transfer(bytes: &[u8]) -> Result<TransferBuffer, Error> {
613 Ok(TransferBuffer {
614 buffer_id: read_u64(bytes, 0)?,
615 offset: read_u64(bytes, 8)?,
616 bytes: read_u64(bytes, 16)?,
617 })
618}
619
620fn validate_transfer(transfer: TransferBuffer) -> Result<(), Error> {
621 if transfer.buffer_id == 0 || transfer.bytes == 0 {
622 return Err(Error::InvalidArgument);
623 }
624 transfer
625 .offset
626 .checked_add(transfer.bytes)
627 .ok_or(Error::InvalidArgument)?;
628 Ok(())
629}
630
631fn encode_transfer(transfer: TransferBuffer, output: &mut [u8]) -> Result<(), Error> {
632 write_u64(output, 0, transfer.buffer_id)?;
633 write_u64(output, 8, transfer.offset)?;
634 write_u64(output, 16, transfer.bytes)
635}
636
637fn validate_bindings(bindings: Bindings<'_>) -> Result<usize, Error> {
638 let count = bindings.count()?;
639 if count == 0 {
640 return Err(Error::InvalidArgument);
641 }
642 if count > HARD_MAX_BINDINGS {
643 return Err(Error::ResourceLimit);
644 }
645 if let Bindings::Encoded { bytes, count } = bindings {
646 let expected = usize::try_from(count)
647 .map_err(|_| Error::ResourceLimit)?
648 .checked_mul(32)
649 .ok_or(Error::ResourceLimit)?;
650 if bytes.len() != expected {
651 return Err(Error::InvalidArgument);
652 }
653 }
654 for index in 0..count {
655 let binding = bindings.get(index)?;
656 validate_binding(binding)?;
657 for previous in 0..index {
658 if bindings.get(previous)?.slot == binding.slot {
659 return Err(Error::InvalidArgument);
660 }
661 }
662 }
663 32_usize
664 .checked_add(
665 usize::try_from(count)
666 .map_err(|_| Error::ResourceLimit)?
667 .checked_mul(32)
668 .ok_or(Error::ResourceLimit)?,
669 )
670 .ok_or(Error::ResourceLimit)
671}
672
673fn decode_binding(bytes: &[u8], offset: usize) -> Result<Binding, Error> {
674 let end = offset.checked_add(32).ok_or(Error::ResourceLimit)?;
675 let binding = bytes.get(offset..end).ok_or(Error::InvalidArgument)?;
676 if binding[29..32].iter().any(|byte| *byte != 0) {
677 return Err(Error::InvalidArgument);
678 }
679 Ok(Binding {
680 buffer_id: read_u64(binding, 0)?,
681 offset: read_u64(binding, 8)?,
682 bytes: read_u64(binding, 16)?,
683 slot: read_u32(binding, 24)?,
684 access: binding[28],
685 })
686}
687
688fn validate_binding(binding: Binding) -> Result<(), Error> {
689 if binding.buffer_id == 0 || binding.bytes == 0 || !(1..=3).contains(&binding.access) {
690 return Err(Error::InvalidArgument);
691 }
692 binding
693 .offset
694 .checked_add(binding.bytes)
695 .ok_or(Error::InvalidArgument)?;
696 Ok(())
697}
698
699fn encode_binding(binding: Binding, output: &mut [u8], offset: usize) -> Result<(), Error> {
700 write_u64(output, offset, binding.buffer_id)?;
701 write_u64(output, offset + 8, binding.offset)?;
702 write_u64(output, offset + 16, binding.bytes)?;
703 write_u32(output, offset + 24, binding.slot)?;
704 *output.get_mut(offset + 28).ok_or(Error::OutputTooSmall)? = binding.access;
705 Ok(())
706}
707
708fn decode_object(payload: &[u8]) -> Result<u64, Error> {
709 require_exact(payload, 8)?;
710 let object = read_u64(payload, 0)?;
711 object_len(object)?;
712 Ok(object)
713}
714
715fn object_len(object: u64) -> Result<usize, Error> {
716 if object == 0 {
717 Err(Error::InvalidArgument)
718 } else {
719 Ok(8)
720 }
721}
722
723#[derive(Clone, Copy, Debug, PartialEq, Eq)]
724pub struct DeviceInfo {
725 pub uuid: [u8; 16],
726 pub class: u16,
727 pub vendor_id: u32,
728 pub device_id: u32,
729 pub capabilities: u64,
730 pub max_contexts: u32,
731 pub max_buffers_per_context: u32,
732 pub max_programs_per_context: u32,
733 pub max_queues_per_context: u32,
734 pub max_events_per_context: u32,
735 pub max_bindings_per_submission: u32,
736 pub max_buffer_bytes: u64,
737 pub max_artifact_bytes: u64,
738}
739
740impl DeviceInfo {
741 pub const ENCODED_BYTES: usize = 76;
742
743 fn validate(self) -> Result<(), Error> {
744 if self.max_contexts == 0
745 || self.max_buffers_per_context == 0
746 || self.max_programs_per_context == 0
747 || self.max_queues_per_context == 0
748 || self.max_events_per_context == 0
749 || self.max_bindings_per_submission == 0
750 || self.max_buffer_bytes == 0
751 || self.max_artifact_bytes == 0
752 {
753 return Err(Error::InvalidArgument);
754 }
755 if self.max_bindings_per_submission > HARD_MAX_BINDINGS {
756 return Err(Error::ResourceLimit);
757 }
758 Ok(())
759 }
760}
761
762#[derive(Clone, Copy, Debug, PartialEq, Eq)]
763#[repr(u16)]
764pub enum EventKind {
765 Pending = 0,
766 Complete = 1,
767 Failed = 2,
768 Cancelled = 3,
769}
770
771#[derive(Clone, Copy, Debug, PartialEq, Eq)]
772pub struct EventState {
773 pub kind: EventKind,
774 pub error: u16,
775}
776
777impl EventState {
778 fn validate(self) -> Result<(), Error> {
779 match self.kind {
780 EventKind::Failed if self.error != 0 => Ok(()),
781 EventKind::Pending | EventKind::Complete | EventKind::Cancelled if self.error == 0 => {
782 Ok(())
783 }
784 _ => Err(Error::InvalidArgument),
785 }
786 }
787}
788
789#[derive(Clone, Copy, Debug, PartialEq, Eq)]
790pub enum ResponseBody<'a> {
791 Empty,
792 DeviceInfo(DeviceInfo),
793 Object(u64),
794 Data(&'a [u8]),
795 EventId(u64),
796 EventState(EventState),
797}
798
799#[derive(Clone, Copy, Debug, PartialEq, Eq)]
800pub struct Response<'a> {
801 pub status: u16,
802 pub request_id: u64,
803 pub body: ResponseBody<'a>,
804}
805
806#[derive(Clone, Copy, Debug, PartialEq, Eq)]
807pub struct ResponseContext {
808 pub opcode: Option<Opcode>,
809 pub request_id: u64,
810 pub read_bytes: Option<usize>,
811}
812
813impl Response<'_> {
814 pub fn encoded_len(&self, context: ResponseContext) -> Result<usize, Error> {
815 let payload = response_payload_len(self, context)?;
816 let total = RESPONSE_HEADER_BYTES
817 .checked_add(payload)
818 .ok_or(Error::ResourceLimit)?;
819 if total > HARD_MAX_RESPONSE_BYTES as usize {
820 return Err(Error::ResourceLimit);
821 }
822 Ok(total)
823 }
824
825 pub fn encode(&self, context: ResponseContext, output: &mut [u8]) -> Result<usize, Error> {
826 let payload_len = response_payload_len(self, context)?;
827 let total = RESPONSE_HEADER_BYTES
828 .checked_add(payload_len)
829 .ok_or(Error::ResourceLimit)?;
830 if total > HARD_MAX_RESPONSE_BYTES as usize {
831 return Err(Error::ResourceLimit);
832 }
833 require_output(output, total)?;
834 write_u16(output, 0, self.status)?;
835 write_u16(output, 2, 0)?;
836 write_u32(
837 output,
838 4,
839 u32::try_from(payload_len).map_err(|_| Error::ResourceLimit)?,
840 )?;
841 write_u64(output, 8, self.request_id)?;
842 encode_response_body(self.body, &mut output[RESPONSE_HEADER_BYTES..total])?;
843 Ok(total)
844 }
845}
846
847pub fn decode_response(bytes: &[u8], context: ResponseContext) -> Result<Response<'_>, Error> {
848 if bytes.len() < RESPONSE_HEADER_BYTES {
849 return Err(Error::Size);
850 }
851 let status = read_u16(bytes, 0)?;
852 if read_u16(bytes, 2)? != 0 {
853 return Err(Error::InvalidArgument);
854 }
855 let payload_len = usize::try_from(read_u32(bytes, 4)?).map_err(|_| Error::ResourceLimit)?;
856 let total = RESPONSE_HEADER_BYTES
857 .checked_add(payload_len)
858 .ok_or(Error::ResourceLimit)?;
859 if total != bytes.len() {
860 return Err(Error::InvalidArgument);
861 }
862 if total > HARD_MAX_RESPONSE_BYTES as usize {
863 return Err(Error::ResourceLimit);
864 }
865 let request_id = read_u64(bytes, 8)?;
866 if request_id != context.request_id || request_id == 0 {
867 return Err(Error::InvalidArgument);
868 }
869 let payload = &bytes[RESPONSE_HEADER_BYTES..];
870 let body = decode_response_body(status, payload, context)?;
871 let response = Response {
872 status,
873 request_id,
874 body,
875 };
876 if response_payload_len(&response, context)? != payload_len {
877 return Err(Error::InvalidArgument);
878 }
879 Ok(response)
880}
881
882fn decode_response_body<'a>(
883 status: u16,
884 payload: &'a [u8],
885 context: ResponseContext,
886) -> Result<ResponseBody<'a>, Error> {
887 if status != 0 {
888 if context.opcode == Some(Opcode::Submit) && payload.len() == 8 {
889 return Ok(ResponseBody::EventId(decode_object(payload)?));
890 }
891 require_exact(payload, 0)?;
892 return Ok(ResponseBody::Empty);
893 }
894
895 match context.opcode.ok_or(Error::InvalidArgument)? {
896 Opcode::GetDeviceInfo => Ok(ResponseBody::DeviceInfo(decode_device_info(payload)?)),
897 Opcode::CreateContext
898 | Opcode::AllocateBuffer
899 | Opcode::LoadProgram
900 | Opcode::CreateQueue => Ok(ResponseBody::Object(decode_object(payload)?)),
901 Opcode::DestroyContext
902 | Opcode::FreeBuffer
903 | Opcode::WriteBuffer
904 | Opcode::UnloadProgram
905 | Opcode::DestroyQueue
906 | Opcode::CancelEvent
907 | Opcode::DestroyEvent => {
908 require_exact(payload, 0)?;
909 Ok(ResponseBody::Empty)
910 }
911 Opcode::ReadBuffer => {
912 let expected = context.read_bytes.ok_or(Error::InvalidArgument)?;
913 require_exact(payload, expected)?;
914 Ok(ResponseBody::Data(payload))
915 }
916 Opcode::Submit => Ok(ResponseBody::EventId(decode_object(payload)?)),
917 Opcode::PollEvent => Ok(ResponseBody::EventState(decode_event_state(payload)?)),
918 }
919}
920
921fn response_payload_len(response: &Response<'_>, context: ResponseContext) -> Result<usize, Error> {
922 if response.request_id == 0 || response.request_id != context.request_id {
923 return Err(Error::InvalidArgument);
924 }
925 if response.status != 0 {
926 return match (context.opcode, response.body) {
927 (Some(Opcode::Submit), ResponseBody::EventId(event_id)) => object_len(event_id),
928 (_, ResponseBody::Empty) => Ok(0),
929 _ => Err(Error::InvalidArgument),
930 };
931 }
932
933 match (context.opcode, response.body) {
934 (Some(Opcode::GetDeviceInfo), ResponseBody::DeviceInfo(info)) => {
935 info.validate()?;
936 Ok(DeviceInfo::ENCODED_BYTES)
937 }
938 (
939 Some(
940 Opcode::CreateContext
941 | Opcode::AllocateBuffer
942 | Opcode::LoadProgram
943 | Opcode::CreateQueue,
944 ),
945 ResponseBody::Object(object),
946 ) => object_len(object),
947 (
948 Some(
949 Opcode::DestroyContext
950 | Opcode::FreeBuffer
951 | Opcode::WriteBuffer
952 | Opcode::UnloadProgram
953 | Opcode::DestroyQueue
954 | Opcode::CancelEvent
955 | Opcode::DestroyEvent,
956 ),
957 ResponseBody::Empty,
958 ) => Ok(0),
959 (Some(Opcode::ReadBuffer), ResponseBody::Data(data)) => {
960 if Some(data.len()) != context.read_bytes {
961 return Err(Error::InvalidArgument);
962 }
963 Ok(data.len())
964 }
965 (Some(Opcode::Submit), ResponseBody::EventId(event_id)) => object_len(event_id),
966 (Some(Opcode::PollEvent), ResponseBody::EventState(state)) => {
967 state.validate()?;
968 Ok(8)
969 }
970 _ => Err(Error::InvalidArgument),
971 }
972}
973
974fn decode_device_info(payload: &[u8]) -> Result<DeviceInfo, Error> {
975 require_exact(payload, DeviceInfo::ENCODED_BYTES)?;
976 if read_u16(payload, 18)? != 0 {
977 return Err(Error::InvalidArgument);
978 }
979 let mut uuid = [0_u8; 16];
980 uuid.copy_from_slice(&payload[..16]);
981 let info = DeviceInfo {
982 uuid,
983 class: read_u16(payload, 16)?,
984 vendor_id: read_u32(payload, 20)?,
985 device_id: read_u32(payload, 24)?,
986 capabilities: read_u64(payload, 28)?,
987 max_contexts: read_u32(payload, 36)?,
988 max_buffers_per_context: read_u32(payload, 40)?,
989 max_programs_per_context: read_u32(payload, 44)?,
990 max_queues_per_context: read_u32(payload, 48)?,
991 max_events_per_context: read_u32(payload, 52)?,
992 max_bindings_per_submission: read_u32(payload, 56)?,
993 max_buffer_bytes: read_u64(payload, 60)?,
994 max_artifact_bytes: read_u64(payload, 68)?,
995 };
996 info.validate()?;
997 Ok(info)
998}
999
1000fn decode_event_state(payload: &[u8]) -> Result<EventState, Error> {
1001 require_exact(payload, 8)?;
1002 if read_u32(payload, 4)? != 0 {
1003 return Err(Error::InvalidArgument);
1004 }
1005 let kind = match read_u16(payload, 0)? {
1006 0 => EventKind::Pending,
1007 1 => EventKind::Complete,
1008 2 => EventKind::Failed,
1009 3 => EventKind::Cancelled,
1010 _ => return Err(Error::RecoveryRequired),
1011 };
1012 let state = EventState {
1013 kind,
1014 error: read_u16(payload, 2)?,
1015 };
1016 state.validate()?;
1017 Ok(state)
1018}
1019
1020fn encode_response_body(body: ResponseBody<'_>, output: &mut [u8]) -> Result<(), Error> {
1021 output.fill(0);
1022 match body {
1023 ResponseBody::Empty => {}
1024 ResponseBody::DeviceInfo(info) => {
1025 write_slice(output, 0, &info.uuid)?;
1026 write_u16(output, 16, info.class)?;
1027 write_u32(output, 20, info.vendor_id)?;
1028 write_u32(output, 24, info.device_id)?;
1029 write_u64(output, 28, info.capabilities)?;
1030 write_u32(output, 36, info.max_contexts)?;
1031 write_u32(output, 40, info.max_buffers_per_context)?;
1032 write_u32(output, 44, info.max_programs_per_context)?;
1033 write_u32(output, 48, info.max_queues_per_context)?;
1034 write_u32(output, 52, info.max_events_per_context)?;
1035 write_u32(output, 56, info.max_bindings_per_submission)?;
1036 write_u64(output, 60, info.max_buffer_bytes)?;
1037 write_u64(output, 68, info.max_artifact_bytes)?;
1038 }
1039 ResponseBody::Object(object) | ResponseBody::EventId(object) => {
1040 write_u64(output, 0, object)?;
1041 }
1042 ResponseBody::Data(data) => write_slice(output, 0, data)?,
1043 ResponseBody::EventState(state) => {
1044 write_u16(output, 0, state.kind as u16)?;
1045 write_u16(output, 2, state.error)?;
1046 }
1047 }
1048 Ok(())
1049}
1050
1051fn require_exact(bytes: &[u8], expected: usize) -> Result<(), Error> {
1052 if bytes.len() == expected {
1053 Ok(())
1054 } else {
1055 Err(Error::InvalidArgument)
1056 }
1057}
1058
1059fn require_output(output: &[u8], required: usize) -> Result<(), Error> {
1060 if output.len() >= required {
1061 Ok(())
1062 } else {
1063 Err(Error::OutputTooSmall)
1064 }
1065}
1066
1067fn read_u16(bytes: &[u8], offset: usize) -> Result<u16, Error> {
1068 let raw: [u8; 2] = bytes
1069 .get(offset..offset.checked_add(2).ok_or(Error::Size)?)
1070 .ok_or(Error::Size)?
1071 .try_into()
1072 .map_err(|_| Error::Size)?;
1073 Ok(u16::from_le_bytes(raw))
1074}
1075
1076fn read_u32(bytes: &[u8], offset: usize) -> Result<u32, Error> {
1077 let raw: [u8; 4] = bytes
1078 .get(offset..offset.checked_add(4).ok_or(Error::Size)?)
1079 .ok_or(Error::Size)?
1080 .try_into()
1081 .map_err(|_| Error::Size)?;
1082 Ok(u32::from_le_bytes(raw))
1083}
1084
1085fn read_u64(bytes: &[u8], offset: usize) -> Result<u64, Error> {
1086 let raw: [u8; 8] = bytes
1087 .get(offset..offset.checked_add(8).ok_or(Error::Size)?)
1088 .ok_or(Error::Size)?
1089 .try_into()
1090 .map_err(|_| Error::Size)?;
1091 Ok(u64::from_le_bytes(raw))
1092}
1093
1094fn write_u16(output: &mut [u8], offset: usize, value: u16) -> Result<(), Error> {
1095 write_slice(output, offset, &value.to_le_bytes())
1096}
1097
1098fn write_u32(output: &mut [u8], offset: usize, value: u32) -> Result<(), Error> {
1099 write_slice(output, offset, &value.to_le_bytes())
1100}
1101
1102fn write_u64(output: &mut [u8], offset: usize, value: u64) -> Result<(), Error> {
1103 write_slice(output, offset, &value.to_le_bytes())
1104}
1105
1106fn write_slice(output: &mut [u8], offset: usize, value: &[u8]) -> Result<(), Error> {
1107 let end = offset
1108 .checked_add(value.len())
1109 .ok_or(Error::OutputTooSmall)?;
1110 output
1111 .get_mut(offset..end)
1112 .ok_or(Error::OutputTooSmall)?
1113 .copy_from_slice(value);
1114 Ok(())
1115}