1use bincode::{Decode, Encode};
27use thiserror::Error;
28
29use crate::attribute::{AggregateKind, Aggregation, Attribute, Bound};
30use crate::entity::{EntityDef, InverseAttribute, UniqueRule, WhereRule};
31use crate::registry::Schema;
32use crate::types::{TypeDef, TypeKind};
33
34const MAGIC: [u8; 8] = *b"NEHSCHM\0";
35const FORMAT_VERSION: u16 = 3;
36const MIN_HEADER_BYTES: usize = MAGIC.len() + 1;
37const MAX_ARTIFACT_BYTES: usize = 8 * 1024 * 1024;
38
39#[derive(Encode, Decode)]
41enum WireAggregateKind {
42 List,
43 Set,
44 Bag,
45 Array,
46}
47
48#[derive(Encode, Decode)]
50enum WireBound {
51 Integer(u64),
52 Unbounded,
53 Expression(String),
54}
55
56#[derive(Encode, Decode)]
58struct WireAggregation {
59 kind: WireAggregateKind,
60 lower: WireBound,
61 upper: WireBound,
62 unique: bool,
63 optional_elements: bool,
64}
65
66#[derive(Encode, Decode)]
68struct WireAttribute {
69 name: String,
70 type_name: String,
71 optional: bool,
72 aggregate: bool,
73 aggregation: Vec<WireAggregation>,
74}
75
76#[derive(Encode, Decode)]
78struct WireInverse {
79 name: String,
80 redeclares: Option<String>,
81 entity: String,
82 for_attribute: String,
83 aggregation: Option<WireAggregation>,
84}
85
86#[derive(Encode, Decode)]
88struct WireUnique {
89 label: Option<String>,
90 attributes: Vec<String>,
91}
92
93#[derive(Encode, Decode)]
95struct WireWhereRule {
96 label: String,
97 expression: String,
98}
99
100#[derive(Encode, Decode)]
102struct WireEntity {
103 name: String,
104 supertype: Option<String>,
105 abstract_: bool,
106 attributes: Vec<WireAttribute>,
107 derived: Vec<String>,
108 where_rules: Vec<WireWhereRule>,
109 inverses: Vec<WireInverse>,
110 unique_rules: Vec<WireUnique>,
111}
112
113#[derive(Encode, Decode)]
115enum WireTypeKind {
116 Defined(String),
117 Enumeration(Vec<String>),
118 Select(Vec<String>),
119}
120
121#[derive(Encode, Decode)]
123struct WireType {
124 name: String,
125 kind: WireTypeKind,
126}
127
128#[derive(Encode, Decode)]
129struct WireSchema {
130 name: String,
131 entities: Vec<WireEntity>,
132 types: Vec<WireType>,
133}
134
135#[derive(Encode, Decode)]
137struct WireAttributeV2 {
138 name: String,
139 type_name: String,
140 optional: bool,
141 aggregate: bool,
142}
143
144#[derive(Encode, Decode)]
146struct WireEntityV2 {
147 name: String,
148 supertype: Option<String>,
149 abstract_: bool,
150 attributes: Vec<WireAttributeV2>,
151 derived: Vec<String>,
152 where_rules: Vec<WireWhereRule>,
153}
154
155#[derive(Encode, Decode)]
157struct WireEntityV1 {
158 name: String,
159 supertype: Option<String>,
160 abstract_: bool,
161 attributes: Vec<WireAttributeV2>,
162 derived: Vec<String>,
163}
164
165#[derive(Encode, Decode)]
166struct WireSchemaV2 {
167 name: String,
168 entities: Vec<WireEntityV2>,
169 types: Vec<WireType>,
170}
171
172#[derive(Encode, Decode)]
173struct WireSchemaV1 {
174 name: String,
175 entities: Vec<WireEntityV1>,
176 types: Vec<WireType>,
177}
178
179impl From<WireEntityV1> for WireEntityV2 {
180 fn from(old: WireEntityV1) -> Self {
181 Self {
182 name: old.name,
183 supertype: old.supertype,
184 abstract_: old.abstract_,
185 attributes: old.attributes,
186 derived: old.derived,
187 where_rules: Vec::new(),
188 }
189 }
190}
191
192impl From<WireSchemaV1> for WireSchemaV2 {
193 fn from(old: WireSchemaV1) -> Self {
194 Self {
195 name: old.name,
196 entities: old.entities.into_iter().map(WireEntityV2::from).collect(),
197 types: old.types,
198 }
199 }
200}
201
202impl From<WireAttributeV2> for WireAttribute {
203 fn from(old: WireAttributeV2) -> Self {
204 Self {
205 name: old.name,
206 type_name: old.type_name,
207 optional: old.optional,
208 aggregate: old.aggregate,
209 aggregation: Vec::new(),
210 }
211 }
212}
213
214impl From<WireEntityV2> for WireEntity {
215 fn from(old: WireEntityV2) -> Self {
216 Self {
217 name: old.name,
218 supertype: old.supertype,
219 abstract_: old.abstract_,
220 attributes: old
221 .attributes
222 .into_iter()
223 .map(WireAttribute::from)
224 .collect(),
225 derived: old.derived,
226 where_rules: old.where_rules,
227 inverses: Vec::new(),
228 unique_rules: Vec::new(),
229 }
230 }
231}
232
233impl From<WireSchemaV2> for WireSchema {
234 fn from(old: WireSchemaV2) -> Self {
235 Self {
236 name: old.name,
237 entities: old.entities.into_iter().map(WireEntity::from).collect(),
238 types: old.types,
239 }
240 }
241}
242
243fn wire_aggregation(aggregation: &Aggregation) -> WireAggregation {
244 let bound = |bound: &Bound| match bound {
245 Bound::Integer(value) => WireBound::Integer(*value),
246 Bound::Unbounded => WireBound::Unbounded,
247 Bound::Expression(text) => WireBound::Expression(text.clone()),
248 };
249 WireAggregation {
250 kind: match aggregation.kind {
251 AggregateKind::List => WireAggregateKind::List,
252 AggregateKind::Set => WireAggregateKind::Set,
253 AggregateKind::Bag => WireAggregateKind::Bag,
254 AggregateKind::Array => WireAggregateKind::Array,
255 },
256 lower: bound(&aggregation.lower),
257 upper: bound(&aggregation.upper),
258 unique: aggregation.unique,
259 optional_elements: aggregation.optional_elements,
260 }
261}
262
263fn owned_aggregation(wire: WireAggregation) -> Aggregation {
264 let bound = |bound: WireBound| match bound {
265 WireBound::Integer(value) => Bound::Integer(value),
266 WireBound::Unbounded => Bound::Unbounded,
267 WireBound::Expression(text) => Bound::Expression(text),
268 };
269 let kind = match wire.kind {
270 WireAggregateKind::List => AggregateKind::List,
271 WireAggregateKind::Set => AggregateKind::Set,
272 WireAggregateKind::Bag => AggregateKind::Bag,
273 WireAggregateKind::Array => AggregateKind::Array,
274 };
275 let mut owned = Aggregation::new(kind, bound(wire.lower), bound(wire.upper));
276 owned.unique = wire.unique;
277 owned.optional_elements = wire.optional_elements;
278 owned
279}
280
281impl From<&Schema> for WireSchema {
282 fn from(schema: &Schema) -> Self {
283 Self {
284 name: schema.name().to_owned(),
285 entities: schema
286 .entities()
287 .map(|entity| WireEntity {
288 name: entity.name.clone(),
289 supertype: entity.supertype().map(str::to_owned),
290 abstract_: entity.abstract_,
291 attributes: entity
292 .attributes
293 .iter()
294 .map(|attribute| WireAttribute {
295 name: attribute.name.clone(),
296 type_name: attribute.type_name.clone(),
297 optional: attribute.optional,
298 aggregate: attribute.aggregate,
299 aggregation: attribute
300 .aggregation
301 .iter()
302 .map(wire_aggregation)
303 .collect(),
304 })
305 .collect(),
306 derived: entity.derived.clone(),
307 where_rules: entity
308 .where_rules
309 .iter()
310 .map(|rule| WireWhereRule {
311 label: rule.label.clone(),
312 expression: rule.expression.clone(),
313 })
314 .collect(),
315 inverses: entity
316 .inverses
317 .iter()
318 .map(|inverse| WireInverse {
319 name: inverse.name.clone(),
320 redeclares: inverse.redeclares.clone(),
321 entity: inverse.entity.clone(),
322 for_attribute: inverse.for_attribute.clone(),
323 aggregation: inverse.aggregation.as_ref().map(wire_aggregation),
324 })
325 .collect(),
326 unique_rules: entity
327 .unique_rules
328 .iter()
329 .map(|rule| WireUnique {
330 label: rule.label.clone(),
331 attributes: rule.attributes.clone(),
332 })
333 .collect(),
334 })
335 .collect(),
336 types: schema
337 .types()
338 .map(|type_def| WireType {
339 name: type_def.name.clone(),
340 kind: match &type_def.kind {
341 TypeKind::Defined(alias) => WireTypeKind::Defined(alias.clone()),
342 TypeKind::Enumeration(members) => {
343 WireTypeKind::Enumeration(members.clone())
344 }
345 TypeKind::Select(members) => WireTypeKind::Select(members.clone()),
346 },
347 })
348 .collect(),
349 }
350 }
351}
352
353impl From<WireSchema> for Schema {
354 fn from(wire: WireSchema) -> Self {
355 let entities = wire
356 .entities
357 .into_iter()
358 .map(|entity| {
359 let mut def = EntityDef::new(entity.name);
360 if let Some(supertype) = entity.supertype {
361 def = def.with_supertype(supertype);
362 }
363 def.abstract_ = entity.abstract_;
364 def.attributes = entity
365 .attributes
366 .into_iter()
367 .map(|attribute| {
368 let mut built = Attribute::new(attribute.name, attribute.type_name);
369 built.optional = attribute.optional;
370 built.aggregate = attribute.aggregate;
371 built.aggregation = attribute
372 .aggregation
373 .into_iter()
374 .map(owned_aggregation)
375 .collect();
376 built
377 })
378 .collect();
379 def.derived = entity.derived;
380 def.where_rules = entity
381 .where_rules
382 .into_iter()
383 .map(|rule| WhereRule::new(rule.label, rule.expression))
384 .collect();
385 def.inverses = entity
386 .inverses
387 .into_iter()
388 .map(|inverse| {
389 let mut built = InverseAttribute::new(
390 inverse.name,
391 inverse.entity,
392 inverse.for_attribute,
393 );
394 built.redeclares = inverse.redeclares;
395 built.aggregation = inverse.aggregation.map(owned_aggregation);
396 built
397 })
398 .collect();
399 def.unique_rules = entity
400 .unique_rules
401 .into_iter()
402 .map(|rule| {
403 let mut built = UniqueRule::unlabelled(rule.attributes);
404 built.label = rule.label;
405 built
406 })
407 .collect();
408 def
409 })
410 .collect();
411 let types = wire
412 .types
413 .into_iter()
414 .map(|type_def| {
415 let kind = match type_def.kind {
416 WireTypeKind::Defined(alias) => TypeKind::Defined(alias),
417 WireTypeKind::Enumeration(members) => TypeKind::Enumeration(members),
418 WireTypeKind::Select(members) => TypeKind::Select(members),
419 };
420 TypeDef::new(type_def.name, kind)
421 })
422 .collect();
423 Schema::new(wire.name, entities, types)
424 }
425}
426
427pub fn decode_schema(bytes: &[u8]) -> Result<Schema, BundledSchemaError> {
435 if bytes.len() > MAX_ARTIFACT_BYTES {
436 return Err(BundledSchemaError::TooLarge {
437 actual: bytes.len(),
438 limit: MAX_ARTIFACT_BYTES,
439 });
440 }
441 if !bytes.starts_with(&MAGIC) {
442 return Err(BundledSchemaError::BadMagic);
443 }
444 if bytes.len() < MIN_HEADER_BYTES {
445 return Err(BundledSchemaError::TruncatedHeader {
446 actual: bytes.len(),
447 required: MIN_HEADER_BYTES,
448 });
449 }
450 let header_config = bincode::config::standard().with_limit::<16>();
451 let (format_version, version_bytes): (u16, usize) =
452 bincode::decode_from_slice(&bytes[MAGIC.len()..], header_config)
453 .map_err(|error| BundledSchemaError::Decode(error.to_string()))?;
454 let payload_bytes = &bytes[MAGIC.len() + version_bytes..];
455 let config = bincode::config::standard().with_limit::<MAX_ARTIFACT_BYTES>();
456 let (wire, consumed): (WireSchema, usize) = match format_version {
459 1 => {
460 let (old, consumed): (WireSchemaV1, usize) =
461 bincode::decode_from_slice(payload_bytes, config)
462 .map_err(|error| BundledSchemaError::Decode(error.to_string()))?;
463 (WireSchema::from(WireSchemaV2::from(old)), consumed)
464 }
465 2 => {
466 let (old, consumed): (WireSchemaV2, usize) =
467 bincode::decode_from_slice(payload_bytes, config)
468 .map_err(|error| BundledSchemaError::Decode(error.to_string()))?;
469 (WireSchema::from(old), consumed)
470 }
471 FORMAT_VERSION => bincode::decode_from_slice(payload_bytes, config)
472 .map_err(|error| BundledSchemaError::Decode(error.to_string()))?,
473 other => return Err(BundledSchemaError::UnsupportedVersion(other)),
474 };
475 if consumed != payload_bytes.len() {
476 return Err(BundledSchemaError::TrailingBytes(
477 payload_bytes.len() - consumed,
478 ));
479 }
480 Ok(wire.into())
481}
482
483#[cfg(feature = "generation")]
489pub fn encode_schema(schema: &Schema) -> Result<Vec<u8>, bincode::error::EncodeError> {
490 let wire = WireSchema::from(schema);
491 let payload = bincode::encode_to_vec(wire, bincode::config::standard())?;
492 let version = bincode::encode_to_vec(FORMAT_VERSION, bincode::config::standard())?;
493 let mut bytes = Vec::with_capacity(MAGIC.len() + version.len() + payload.len());
494 bytes.extend_from_slice(&MAGIC);
495 bytes.extend_from_slice(&version);
496 bytes.extend_from_slice(&payload);
497 Ok(bytes)
498}
499
500#[derive(Debug, Clone, PartialEq, Eq, Error)]
502#[non_exhaustive]
503pub enum BundledSchemaError {
504 #[error("cannot decode schema artifact: {0}")]
506 Decode(String),
507 #[error("schema artifact is {actual} bytes; limit is {limit} bytes")]
509 TooLarge {
510 actual: usize,
512 limit: usize,
514 },
515 #[error("schema artifact header is {actual} bytes; at least {required} bytes are required")]
517 TruncatedHeader {
518 actual: usize,
520 required: usize,
522 },
523 #[error("schema artifact magic is invalid")]
525 BadMagic,
526 #[error("unsupported schema artifact format version {0}")]
528 UnsupportedVersion(u16),
529 #[error("schema artifact has {0} trailing bytes")]
531 TrailingBytes(usize),
532}
533
534#[cfg(all(test, feature = "generation"))]
535mod tests {
536 use super::*;
537
538 fn sample() -> Schema {
539 Schema::new(
540 "IFC4",
541 vec![
542 EntityDef::new("IfcRoot")
543 .abstract_entity()
544 .with_attribute(Attribute::new("GlobalId", "IfcGloballyUniqueId"))
545 .with_where_rule(WhereRule::new("WR1", "")),
546 EntityDef::new("IfcWall")
547 .with_supertype("IfcRoot")
548 .with_attribute(Attribute::new("Name", "IfcLabel").optional())
549 .with_attribute(Attribute::new("Tags", "IfcLabel").aggregate())
550 .with_attribute(
551 Attribute::new("Coords", "IfcLengthMeasure")
552 .with_aggregation(
553 Aggregation::new(
554 AggregateKind::List,
555 Bound::Integer(1),
556 Bound::Unbounded,
557 )
558 .unique(),
559 )
560 .with_aggregation(
561 Aggregation::new(
562 AggregateKind::Array,
563 Bound::Integer(1),
564 Bound::Expression("SELF\\IfcWall.Dim".into()),
565 )
566 .optional_elements(),
567 ),
568 )
569 .with_inverse(
570 InverseAttribute::new("Holes", "Voiding", "RelatingElement")
571 .with_aggregation(Aggregation::new(
572 AggregateKind::Set,
573 Bound::Integer(0),
574 Bound::Unbounded,
575 )),
576 )
577 .with_inverse(
578 InverseAttribute::new("Owner", "Owning", "Owned").redeclaring("IfcRoot"),
579 )
580 .with_unique_rule(UniqueRule::new("UR1", vec!["Name".into()]))
581 .with_unique_rule(UniqueRule::unlabelled(
582 vec!["SELF\\IfcRoot.GlobalId".into()],
583 ))
584 .with_derived("Dim"),
585 ],
586 vec![
587 TypeDef::new("IfcLabel", TypeKind::Defined("STRING".into())),
588 TypeDef::new("IfcSide", TypeKind::Enumeration(vec!["LEFT".into()])),
589 TypeDef::new("IfcValue", TypeKind::Select(vec!["IfcLabel".into()])),
590 ],
591 )
592 }
593
594 #[test]
595 fn round_trips_through_the_wire_format() {
596 let schema = sample();
597 let bytes = encode_schema(&schema).expect("encode");
598 let decoded = decode_schema(&bytes).expect("decode");
599 assert_eq!(decoded, schema);
600 }
601
602 #[test]
603 fn rejects_trailing_bytes() {
604 let schema = sample();
605 let mut bytes = encode_schema(&schema).expect("encode");
606 bytes.push(0);
607 assert!(matches!(
608 decode_schema(&bytes),
609 Err(BundledSchemaError::TrailingBytes(1))
610 ));
611 }
612}
613
614#[cfg(test)]
615mod header_tests {
616 use super::*;
617
618 fn framed(version: u16, payload: &[u8]) -> Vec<u8> {
619 let mut bytes = MAGIC.to_vec();
620 bytes.extend(bincode::encode_to_vec(version, bincode::config::standard()).unwrap());
621 bytes.extend_from_slice(payload);
622 bytes
623 }
624
625 fn v2_payload() -> Vec<u8> {
626 let old = WireSchemaV2 {
627 name: "IFC4".into(),
628 entities: vec![WireEntityV2 {
629 name: "IfcPolyline".into(),
630 supertype: None,
631 abstract_: false,
632 attributes: vec![WireAttributeV2 {
633 name: "Points".into(),
634 type_name: "IfcCartesianPoint".into(),
635 optional: false,
636 aggregate: true,
637 }],
638 derived: Vec::new(),
639 where_rules: vec![WireWhereRule {
640 label: "SameDim".into(),
641 expression: String::new(),
642 }],
643 }],
644 types: Vec::new(),
645 };
646 bincode::encode_to_vec(old, bincode::config::standard()).unwrap()
647 }
648
649 #[test]
652 fn a_format_2_artifact_still_decodes() {
653 let schema = decode_schema(&framed(2, &v2_payload())).expect("v2 decodes");
654 let entity = schema.entity("IfcPolyline").unwrap();
655 assert!(entity.attributes[0].aggregate);
656 assert!(entity.attributes[0].aggregation.is_empty());
657 assert!(entity.inverses.is_empty() && entity.unique_rules.is_empty());
658 assert_eq!(entity.where_rules[0].label, "SameDim");
659 assert!(decode_schema(&framed(FORMAT_VERSION, &v2_payload())).is_err());
661 }
662
663 #[test]
664 fn rejects_bad_magic() {
665 let bytes = vec![0u8; MIN_HEADER_BYTES + 1];
666 assert!(matches!(
667 decode_schema(&bytes),
668 Err(BundledSchemaError::BadMagic)
669 ));
670 }
671
672 #[test]
673 fn rejects_truncated_header() {
674 assert!(matches!(
675 decode_schema(&MAGIC),
676 Err(BundledSchemaError::TruncatedHeader { .. })
677 ));
678 }
679
680 #[test]
681 fn decode_rejects_input_above_resource_budget() {
682 let bytes = vec![0; MAX_ARTIFACT_BYTES + 1];
683 assert!(matches!(
684 decode_schema(&bytes),
685 Err(BundledSchemaError::TooLarge { .. })
686 ));
687 }
688}