Skip to main content

virtio_accel_cleanroom/
lib.rs

1//! Independent, dependency-free codec for the virtio-accel protocol 1.0 candidate.
2//!
3//! This crate intentionally does not depend on `virtio-accel-proto`, `zerocopy`, or any shared
4//! wire type. It implements the normative byte contract with explicit little-endian reads and
5//! writes so conformance tests can exercise interoperability through bytes alone.
6
7#![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}