Skip to main content

virtio_accel_tosa/
validate.rs

1use alloc::vec::Vec;
2use core::fmt;
3
4use crate::generated::tosa as wire;
5use crate::view::SliceSource;
6use crate::{DType, Model, Op, Version};
7
8/// Resource ceilings applied before and during graph traversal.
9///
10/// Every field is public so a provider can derive stricter limits from its admission policy. The
11/// defaults are finite and intentionally much smaller than the FlatBuffers runtime defaults.
12/// `max_rank` bounds tensor ranks; serialized `shape_t` value lists may contain up to twice that
13/// many entries because `PAD` carries a before/after pair for each dimension.
14#[derive(Clone, Copy, Debug, PartialEq, Eq)]
15pub struct Limits {
16    pub max_model_bytes: usize,
17    pub max_apparent_bytes: usize,
18    pub max_flatbuffer_depth: usize,
19    pub max_flatbuffer_tables: usize,
20    pub max_regions: usize,
21    pub max_blocks: usize,
22    pub max_tensors: usize,
23    pub max_shapes: usize,
24    pub max_operators: usize,
25    pub max_edges: usize,
26    pub max_name_bytes: usize,
27    pub max_rank: usize,
28    pub max_constant_bytes: usize,
29}
30
31impl Default for Limits {
32    fn default() -> Self {
33        Self {
34            max_model_bytes: 256 * 1024 * 1024,
35            max_apparent_bytes: 512 * 1024 * 1024,
36            max_flatbuffer_depth: 64,
37            max_flatbuffer_tables: 1_000_000,
38            max_regions: 64,
39            max_blocks: 4_096,
40            max_tensors: 262_144,
41            max_shapes: 65_536,
42            max_operators: 1_000_000,
43            max_edges: 8_000_000,
44            max_name_bytes: 1_024,
45            max_rank: 32,
46            max_constant_bytes: 256 * 1024 * 1024,
47        }
48    }
49}
50
51/// Counts collected during the bounded validation pass.
52#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
53pub struct Stats {
54    pub regions: usize,
55    pub blocks: usize,
56    pub tensors: usize,
57    pub shapes: usize,
58    pub operators: usize,
59    pub edges: usize,
60    pub constant_bytes: usize,
61}
62
63#[derive(Clone, Copy, Debug, PartialEq, Eq)]
64pub enum Resource {
65    ModelBytes,
66    Regions,
67    Blocks,
68    Tensors,
69    Shapes,
70    Operators,
71    Edges,
72    NameBytes,
73    Rank,
74    ConstantBytes,
75}
76
77#[derive(Clone, Copy, Debug, PartialEq, Eq)]
78pub enum NameKind {
79    Region,
80    Block,
81    Tensor,
82    Shape,
83    Symbol,
84    Reference,
85}
86
87/// Failure returned for malformed, unsupported, or over-budget input.
88#[derive(Clone, Copy, Debug, PartialEq, Eq)]
89pub enum Error {
90    EmptyInput,
91    MissingIdentifier,
92    InvalidFlatbuffer,
93    UnsupportedVersion {
94        major: i32,
95        minor: i32,
96        patch: i32,
97        draft: bool,
98    },
99    LimitExceeded {
100        resource: Resource,
101        limit: usize,
102    },
103    AllocationFailed(Resource),
104    MissingName(NameKind),
105    EmptyName(NameKind),
106    DuplicateName(NameKind),
107    UnknownDataType(u32),
108    UnsupportedDataType(u32),
109    RankedTensorWithoutShape,
110    UnrankedTensorWithDimensions,
111    InvalidDimension(i32),
112    InvalidShapeData,
113    ExternalDataRange,
114    UnknownOperator(u32),
115    UnsupportedOperator(u32),
116    MissingAttribute(Op),
117    AttributeMismatch {
118        op: Op,
119        attribute: u8,
120    },
121    UnknownSymbol,
122    MultipleProducers,
123}
124
125impl fmt::Display for Error {
126    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
127        write!(formatter, "{self:?}")
128    }
129}
130
131/// Parse with finite production defaults.
132pub fn parse(bytes: &[u8]) -> Result<Model<'_>, Error> {
133    parse_with_limits(bytes, Limits::default())
134}
135
136/// Verify and parse a stable TOSA 1.0 graph under caller-selected resource ceilings.
137pub fn parse_with_limits(bytes: &[u8], limits: Limits) -> Result<Model<'_>, Error> {
138    if bytes.is_empty() {
139        return Err(Error::EmptyInput);
140    }
141    check_limit(bytes.len(), limits.max_model_bytes, Resource::ModelBytes)?;
142    // The FlatBuffers identifier helper asserts its eight-byte minimum instead of returning false.
143    if bytes.len() < 8 || !wire::tosa_graph_buffer_has_identifier(bytes) {
144        return Err(Error::MissingIdentifier);
145    }
146
147    let verifier = flatbuffers::VerifierOptions {
148        max_depth: limits.max_flatbuffer_depth,
149        max_tables: limits.max_flatbuffer_tables,
150        max_apparent_size: limits.max_apparent_bytes,
151        ignore_missing_null_terminator: false,
152    };
153    let graph = wire::root_as_tosa_graph_with_opts(&verifier, bytes)
154        .map_err(|_| Error::InvalidFlatbuffer)?;
155    let raw_version = graph.version();
156    let (major, minor, patch, draft) = (
157        raw_version._major(),
158        raw_version._minor(),
159        raw_version._patch(),
160        raw_version._draft(),
161    );
162    if (major, minor, patch, draft) != (1, 0, 0, false) {
163        return Err(Error::UnsupportedVersion {
164            major,
165            minor,
166            patch,
167            draft,
168        });
169    }
170
171    let mut validator = StructuralValidator {
172        bytes,
173        limits,
174        stats: Stats::default(),
175    };
176    validator.validate_graph(graph)?;
177
178    Ok(Model {
179        graph,
180        bytes,
181        version: Version::TOSA_1_0,
182        stats: validator.stats,
183        source: SliceSource(bytes),
184    })
185}
186
187struct StructuralValidator<'a> {
188    bytes: &'a [u8],
189    limits: Limits,
190    stats: Stats,
191}
192
193impl StructuralValidator<'_> {
194    fn validate_graph(&mut self, graph: wire::TosaGraph<'_>) -> Result<(), Error> {
195        let regions = graph.regions();
196        let region_count = regions.map_or(0, |items| items.len());
197        self.add(Resource::Regions, region_count)?;
198
199        let mut region_names = Vec::new();
200        reserve(&mut region_names, region_count, Resource::Regions)?;
201        if let Some(regions) = regions {
202            for region in regions {
203                let name = self.name(region.name(), NameKind::Region)?;
204                region_names.push(name);
205                self.validate_region(region)?;
206            }
207        }
208        reject_duplicates(&mut region_names, NameKind::Region)?;
209        Ok(())
210    }
211
212    fn validate_region(&mut self, region: wire::TosaRegion<'_>) -> Result<(), Error> {
213        let blocks = region.blocks();
214        self.add(Resource::Blocks, blocks.map_or(0, |items| items.len()))?;
215
216        let block_count = blocks.map_or(0, |items| items.len());
217        let mut block_names = Vec::new();
218        reserve(&mut block_names, block_count, Resource::Blocks)?;
219        if let Some(blocks) = blocks {
220            for block in blocks {
221                let name = self.name(block.name(), NameKind::Block)?;
222                block_names.push(name);
223                self.validate_block(block)?;
224            }
225        }
226        reject_duplicates(&mut block_names, NameKind::Block)?;
227        Ok(())
228    }
229
230    fn validate_block(&mut self, block: wire::TosaBasicBlock<'_>) -> Result<(), Error> {
231        let tensors = block.tensors();
232        let shapes = block.shapes();
233        let operators = block.operators();
234        let tensor_count = tensors.map_or(0, |items| items.len());
235        let shape_count = shapes.map_or(0, |items| items.len());
236        self.add(Resource::Tensors, tensor_count)?;
237        self.add(Resource::Shapes, shape_count)?;
238        self.add(
239            Resource::Operators,
240            operators.map_or(0, |items| items.len()),
241        )?;
242
243        let symbol_count = tensor_count
244            .checked_add(shape_count)
245            .ok_or(Error::LimitExceeded {
246                resource: Resource::Tensors,
247                limit: self.limits.max_tensors,
248            })?;
249        let mut tensor_names = Vec::new();
250        reserve(&mut tensor_names, tensor_count, Resource::Tensors)?;
251        let mut variable_names = Vec::new();
252        reserve(&mut variable_names, tensor_count, Resource::Tensors)?;
253        if let Some(tensors) = tensors {
254            for tensor in tensors {
255                let name = self.name(tensor.name(), NameKind::Tensor)?;
256                tensor_names.push(name);
257                if tensor.variable() {
258                    variable_names.push(name);
259                }
260                self.validate_tensor(tensor)?;
261            }
262        }
263        reject_duplicates(&mut tensor_names, NameKind::Tensor)?;
264        variable_names.sort_unstable();
265
266        let mut shape_names = Vec::new();
267        reserve(&mut shape_names, shape_count, Resource::Shapes)?;
268        if let Some(shapes) = shapes {
269            for shape in shapes {
270                let name = self.name(shape.name(), NameKind::Shape)?;
271                shape_names.push(name);
272                self.validate_shape(shape)?;
273            }
274        }
275        reject_duplicates(&mut shape_names, NameKind::Shape)?;
276
277        let mut symbols = Vec::new();
278        reserve(&mut symbols, symbol_count, Resource::Tensors)?;
279        symbols.extend_from_slice(&tensor_names);
280        symbols.extend_from_slice(&shape_names);
281        reject_duplicates(&mut symbols, NameKind::Symbol)?;
282
283        self.validate_references(block.inputs(), &symbols)?;
284        self.validate_references(block.outputs(), &symbols)?;
285
286        let producer_count = operators.map_or(0, |items| {
287            items
288                .iter()
289                .try_fold(0_usize, |count, operator| {
290                    count.checked_add(operator.outputs().map_or(0, |outputs| outputs.len()))
291                })
292                .unwrap_or(usize::MAX)
293        });
294        check_limit(producer_count, self.limits.max_edges, Resource::Edges)?;
295        let mut produced = Vec::new();
296        reserve(&mut produced, producer_count, Resource::Edges)?;
297        if let Some(operators) = operators {
298            for operator in operators {
299                self.validate_operator(operator, &symbols, &variable_names, &mut produced)?;
300            }
301        }
302        produced.sort_unstable();
303        if produced.windows(2).any(|names| names[0] == names[1]) {
304            return Err(Error::MultipleProducers);
305        }
306        Ok(())
307    }
308
309    fn validate_tensor(&mut self, tensor: wire::TosaTensor<'_>) -> Result<(), Error> {
310        let dtype = DType::new(tensor.type_().0);
311        if dtype.get() == 0 || dtype.name().is_none() {
312            return Err(Error::UnknownDataType(dtype.get()));
313        }
314        if !dtype.is_tosa_1_0() {
315            return Err(Error::UnsupportedDataType(dtype.get()));
316        }
317
318        match (tensor.is_unranked(), tensor.shape()) {
319            (true, Some(shape)) if !shape.is_empty() => {
320                return Err(Error::UnrankedTensorWithDimensions);
321            }
322            (false, None) => return Err(Error::RankedTensorWithoutShape),
323            (false, Some(shape)) => {
324                check_limit(shape.len(), self.limits.max_rank, Resource::Rank)?;
325                for dimension in shape {
326                    if dimension < 1 {
327                        return Err(Error::InvalidDimension(dimension));
328                    }
329                }
330            }
331            (true, Some(_) | None) => {}
332        }
333
334        let embedded = tensor.data().map_or(0, |data| data.len());
335        let offset = tensor.offset();
336        let external = usize::try_from(tensor.size()).map_err(|_| Error::ExternalDataRange)?;
337        if external != 0 {
338            if embedded != 0 || offset <= 1 {
339                return Err(Error::ExternalDataRange);
340            }
341            let start = usize::try_from(offset).map_err(|_| Error::ExternalDataRange)?;
342            let end = start
343                .checked_add(external)
344                .filter(|end| *end <= self.bytes.len())
345                .ok_or(Error::ExternalDataRange)?;
346            let _ = end;
347        } else if offset != 0 {
348            return Err(Error::ExternalDataRange);
349        }
350        let constant_bytes = embedded.checked_add(external).ok_or(Error::LimitExceeded {
351            resource: Resource::ConstantBytes,
352            limit: self.limits.max_constant_bytes,
353        })?;
354        self.add(Resource::ConstantBytes, constant_bytes)
355    }
356
357    fn validate_shape(&mut self, shape: wire::TosaShape<'_>) -> Result<(), Error> {
358        let rank = usize::try_from(shape.rank()).map_err(|_| Error::LimitExceeded {
359            resource: Resource::Rank,
360            limit: self.limits.max_rank,
361        })?;
362        let max_shape_values = self.limits.max_rank.saturating_mul(2);
363        check_limit(rank, max_shape_values, Resource::Rank)?;
364        let data_len = shape.data().map_or(0, |data| data.len());
365        if data_len != 0 {
366            let required = rank
367                .checked_mul(core::mem::size_of::<i64>())
368                .ok_or(Error::InvalidShapeData)?;
369            if data_len != required {
370                return Err(Error::InvalidShapeData);
371            }
372        }
373        self.add(Resource::ConstantBytes, data_len)
374    }
375
376    fn validate_operator<'a>(
377        &mut self,
378        operator: wire::TosaOperator<'a>,
379        symbols: &[&'a str],
380        variable_names: &[&'a str],
381        produced: &mut Vec<&'a str>,
382    ) -> Result<(), Error> {
383        let op = Op::new(operator.op().0);
384        if op.get() == 0 || op.name().is_none() {
385            return Err(Error::UnknownOperator(op.get()));
386        }
387        if !op.is_tosa_1_0() {
388            return Err(Error::UnsupportedOperator(op.get()));
389        }
390        let attribute = operator.attribute_type().0;
391        if u32::from(attribute) != op.get() {
392            return Err(Error::AttributeMismatch { op, attribute });
393        }
394        if operator.attribute().is_none() {
395            return Err(Error::MissingAttribute(op));
396        }
397
398        self.validate_references(operator.inputs(), symbols)?;
399        self.validate_references(operator.outputs(), symbols)?;
400        if let Some(outputs) = operator.outputs() {
401            for output in outputs {
402                if variable_names.binary_search(&output).is_err() {
403                    produced.push(output);
404                }
405            }
406        }
407        Ok(())
408    }
409
410    fn validate_references<'a>(
411        &mut self,
412        references: Option<flatbuffers::Vector<'a, flatbuffers::ForwardsUOffset<&'a str>>>,
413        symbols: &[&'a str],
414    ) -> Result<(), Error> {
415        self.add(Resource::Edges, references.map_or(0, |items| items.len()))?;
416        if let Some(references) = references {
417            for reference in references {
418                self.name(Some(reference), NameKind::Reference)?;
419                if symbols.binary_search(&reference).is_err() {
420                    return Err(Error::UnknownSymbol);
421                }
422            }
423        }
424        Ok(())
425    }
426
427    fn name<'a>(&self, name: Option<&'a str>, kind: NameKind) -> Result<&'a str, Error> {
428        let name = name.ok_or(Error::MissingName(kind))?;
429        if name.is_empty() {
430            return Err(Error::EmptyName(kind));
431        }
432        check_limit(name.len(), self.limits.max_name_bytes, Resource::NameBytes)?;
433        Ok(name)
434    }
435
436    fn add(&mut self, resource: Resource, amount: usize) -> Result<(), Error> {
437        let (value, limit) = match resource {
438            Resource::Regions => (&mut self.stats.regions, self.limits.max_regions),
439            Resource::Blocks => (&mut self.stats.blocks, self.limits.max_blocks),
440            Resource::Tensors => (&mut self.stats.tensors, self.limits.max_tensors),
441            Resource::Shapes => (&mut self.stats.shapes, self.limits.max_shapes),
442            Resource::Operators => (&mut self.stats.operators, self.limits.max_operators),
443            Resource::Edges => (&mut self.stats.edges, self.limits.max_edges),
444            Resource::ConstantBytes => (
445                &mut self.stats.constant_bytes,
446                self.limits.max_constant_bytes,
447            ),
448            Resource::ModelBytes | Resource::NameBytes | Resource::Rank => unreachable!(),
449        };
450        *value = value
451            .checked_add(amount)
452            .ok_or(Error::LimitExceeded { resource, limit })?;
453        check_limit(*value, limit, resource)
454    }
455}
456
457fn check_limit(actual: usize, limit: usize, resource: Resource) -> Result<(), Error> {
458    if actual > limit {
459        Err(Error::LimitExceeded { resource, limit })
460    } else {
461        Ok(())
462    }
463}
464
465fn reserve<T>(values: &mut Vec<T>, additional: usize, resource: Resource) -> Result<(), Error> {
466    values
467        .try_reserve_exact(additional)
468        .map_err(|_| Error::AllocationFailed(resource))
469}
470
471fn reject_duplicates(values: &mut [&str], kind: NameKind) -> Result<(), Error> {
472    values.sort_unstable();
473    if values.windows(2).any(|names| names[0] == names[1]) {
474        Err(Error::DuplicateName(kind))
475    } else {
476        Ok(())
477    }
478}