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#[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#[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#[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
131pub fn parse(bytes: &[u8]) -> Result<Model<'_>, Error> {
133 parse_with_limits(bytes, Limits::default())
134}
135
136pub 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 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}