1use std::collections::{BTreeMap, BTreeSet};
21use std::sync::Arc;
22
23use cedar_policy::{EntityId, EntityUid};
24use miette::Diagnostic;
25use smol_str::SmolStr;
26use thiserror::Error;
27
28use crate::symcc::bitvec::BitVecError;
29use crate::symcc::decoder::sexpr::SExprParseError;
30use crate::symcc::env::SymEntityData;
31use crate::symcc::extension_types::ipaddr::{
32 CIDRv4, CIDRv6, IPv4Addr, IPv4Prefix, IPv6Addr, IPv6Prefix,
33};
34use crate::symcc::type_abbrevs::{ExtType, Width, SIXTY_FOUR};
35use crate::SymEnv;
36
37use super::bitvec::BitVec;
38use super::encoder::Encoder;
39use super::ext::Ext;
40use super::extension_types::datetime::{Datetime, Duration};
41use super::extension_types::decimal::Decimal;
42use super::extension_types::ipaddr::IPNet;
43use super::factory;
44use super::function::Udf;
45use super::interpretation::Interpretation;
46use super::op::Uuf;
47use super::term::{Term, TermPrim, TermVar};
48use super::term_type::TermType;
49
50mod sexpr;
51use sexpr::{parse_sexpr, SExpr};
52
53#[derive(Debug, Diagnostic, Error)]
56pub enum DecodeError {
57 #[error(transparent)]
59 SExprParse(#[from] SExprParseError),
60 #[error("Invalid numeric token: {0}")]
62 ParseIntError(#[from] std::num::ParseIntError),
63 #[error("Integer overflow")]
65 IntegerOverflow,
66 #[error("Model of an unexpected form returned by the solver")]
68 UnexpectedModel,
69 #[error("Unknown SMT type: {0}")]
71 UnknownType(SExpr),
72 #[error("Unknown SMT literal: {0}")]
74 UnknownLiteral(SExpr),
75 #[error("Unmatched type: expected {0:?}, found {1:?}")]
77 UnmatchedType(TermType, TermType),
78 #[error("Unmatched field type: expected {0:?}, found {1:?}")]
80 UnmatchedFieldType(TermType, TermType),
81 #[error("Invalid set type: {0}")]
83 InvalidSetType(SExpr),
84 #[error("Invalid option type: {0}")]
86 InvalidOptionType(SExpr),
87 #[error("set.union applied to non-literals {0:?} and {1:?}")]
89 SetUnionNonLiterals(Term, Term),
90 #[error("Unmatched record type fields")]
92 UnmatchedRecordType,
93 #[error("Unknown variable: {0}")]
95 UnknownVariable(String),
96 #[error("Unknown unary function: {0}")]
98 UnknownUUF(String),
99 #[error("Unexpected form of unary function model: {0}")]
101 UnexpectedUnaryFunctionForm(SExpr),
102 #[error("Bit-vector error")]
104 BitVecError(#[from] BitVecError),
105 #[error("Bitvector of zero width")]
107 ZeroWidthBitVec,
108}
109
110#[derive(Debug)]
113pub struct IdMaps<'a> {
114 types: BTreeMap<&'a SmolStr, &'a TermType>,
115 vars: BTreeMap<&'a SmolStr, &'a TermVar>,
116 uufs: BTreeMap<&'a SmolStr, &'a Uuf>,
117 enums: BTreeMap<SmolStr, EntityUid>,
118}
119
120impl<'a> IdMaps<'a> {
121 pub fn from_encoder<S>(encoder: &'a Encoder<'_, S>) -> Self {
124 let mut types = BTreeMap::new();
125 let mut vars = BTreeMap::new();
126 let mut uufs = BTreeMap::new();
127 let mut enums = BTreeMap::new();
128
129 for (term, enc) in &encoder.types {
130 types.insert(enc, term);
131 }
132
133 for (term, enc) in &encoder.terms {
134 if let Term::Var(var) = term {
135 vars.insert(enc, var);
136 }
137 }
138
139 for (uuf, id) in &encoder.uufs {
140 uufs.insert(id, uuf);
141 }
142
143 for (&entity_type, &enum_ids) in &encoder.enums {
144 if let Some(entity_type_id) = encoder.types.get(&TermType::Entity {
145 ety: entity_type.clone(),
146 }) {
147 for (i, enum_id) in enum_ids.iter().enumerate() {
148 enums.insert(
149 super::encoder::enum_id(entity_type_id, i),
150 EntityUid::from_type_name_and_id(
151 entity_type.clone(),
152 EntityId::new(enum_id),
153 ),
154 );
155 }
156 }
157 }
158
159 Self {
160 types,
161 vars,
162 uufs,
163 enums,
164 }
165 }
166}
167
168impl TermType {
169 pub fn default_literal(&self, env: &SymEnv) -> Term {
172 match self {
173 TermType::Bool => Term::Prim(TermPrim::Bool(false)),
174 TermType::Bitvec { n } => Term::Prim(TermPrim::Bitvec(BitVec::of_u128(*n, 0))),
175 TermType::String => Term::Prim(TermPrim::String(SmolStr::new_static(""))),
176
177 TermType::Entity { ety } => {
178 let eid = if let Some(SymEntityData {
180 members: Some(eids),
181 ..
182 }) = env.entities.get(ety)
183 {
184 if let Some(eid) = eids.first() {
185 eid
186 } else {
187 "" }
189 } else {
190 ""
191 };
192 Term::Prim(TermPrim::Entity(EntityUid::from_type_name_and_id(
193 ety.clone(),
194 EntityId::new(eid),
195 )))
196 }
197
198 TermType::Ext { xty } => match xty {
199 ExtType::Decimal => Term::Prim(TermPrim::Ext(Ext::Decimal { d: Decimal(0) })),
200
201 ExtType::DateTime => Term::Prim(TermPrim::Ext(Ext::Datetime {
202 dt: Datetime::default(),
203 })),
204
205 ExtType::Duration => Term::Prim(TermPrim::Ext(Ext::Duration {
206 d: Duration::default(),
207 })),
208
209 ExtType::IpAddr => Term::Prim(TermPrim::Ext(Ext::Ipaddr {
210 ip: IPNet::default(),
211 })),
212 },
213
214 TermType::Option { ty } => Term::None(ty.as_ref().clone()),
215
216 TermType::Set { ty } => Term::Set {
217 elts: Arc::new(BTreeSet::new()),
218 elts_ty: ty.as_ref().clone(),
219 },
220
221 TermType::Record { rty } => Term::Record(Arc::new(
222 rty.iter()
223 .map(|(k, v)| (k.clone(), v.default_literal(env)))
224 .collect(),
225 )),
226 }
227 }
228}
229
230impl Uuf {
231 pub fn default_udf(&self, env: &SymEnv) -> Udf {
233 Udf {
234 arg: self.arg.clone(),
235 out: self.out.clone(),
236 table: Arc::new(BTreeMap::new()),
237 default: self.out.default_literal(env),
238 }
239 }
240}
241
242impl SExpr {
243 fn is_symbol(&self, s: &str) -> bool {
245 match self {
246 SExpr::Symbol(sym) => sym == s,
247 _ => false,
248 }
249 }
250
251 fn is_app_of(&self, s: &str) -> bool {
253 matches!(self, SExpr::App(sexprs) if sexprs.first().is_some_and(|e| e.is_symbol(s)))
254 }
255
256 fn as_app(&self, func: &str) -> Option<&[SExpr]> {
259 match self {
260 SExpr::App(sexprs) => match sexprs.as_slice() {
261 [SExpr::Symbol(f), args @ ..] if f == func => Some(args),
262 _ => None,
263 },
264 _ => None,
265 }
266 }
267
268 fn as_app_n<const N: usize>(&self, func: &str) -> Option<&[SExpr; N]> {
271 self.as_app(func)
272 .and_then(|args| <&[SExpr; N]>::try_from(args).ok())
273 }
274
275 pub fn decode_type(&self, id_maps: &IdMaps<'_>) -> Result<TermType, DecodeError> {
277 match self {
278 SExpr::Symbol(s) => {
280 match s.as_str() {
281 "Bool" => Ok(TermType::Bool),
282 "String" => Ok(TermType::String),
283 "Decimal" => Ok(TermType::Ext {
284 xty: ExtType::Decimal,
285 }),
286 "IPAddr" => Ok(TermType::Ext {
287 xty: ExtType::IpAddr,
288 }),
289 "Duration" => Ok(TermType::Ext {
290 xty: ExtType::Duration,
291 }),
292 "Datetime" => Ok(TermType::Ext {
293 xty: ExtType::DateTime,
294 }),
295
296 _ => id_maps
298 .types
299 .get(s)
300 .copied()
301 .cloned()
302 .ok_or_else(|| DecodeError::UnknownType(self.clone())),
303 }
304 }
305
306 SExpr::App(args) => {
308 match args.as_slice() {
309 [SExpr::Symbol(app), SExpr::Symbol(bit_vec), SExpr::Numeral(n)]
311 if app == "_" && bit_vec == "BitVec" =>
312 {
313 let n = u32::try_from(*n).map_err(|_| DecodeError::IntegerOverflow)?;
314 let n = Width::new(n).ok_or(DecodeError::ZeroWidthBitVec)?;
315 Ok(TermType::Bitvec { n })
316 }
317
318 [SExpr::Symbol(option), param] if option == "Option" => {
320 let ty = param.decode_type(id_maps)?;
321 Ok(TermType::option_of(ty))
322 }
323
324 [SExpr::Symbol(set), param] if set == "Set" => {
326 let ty = param.decode_type(id_maps)?;
327 Ok(TermType::set_of(ty))
328 }
329
330 _ => Err(DecodeError::UnknownType(self.clone())),
331 }
332 }
333
334 _ => Err(DecodeError::UnknownType(self.clone())),
335 }
336 }
337
338 fn decode_entity_or_record(
341 &self,
342 id_maps: &IdMaps<'_>,
343 name: &SmolStr,
344 args: &[SExpr],
345 ) -> Result<Term, DecodeError> {
346 match (id_maps.types.get(name), args) {
347 (Some(TermType::Entity { ety }), [SExpr::String(e)]) => {
349 let uid = EntityUid::from_type_name_and_id(ety.clone(), EntityId::new(e));
350 Ok(Term::Prim(TermPrim::Entity(uid)))
351 }
352
353 (Some(TermType::Record { rty }), fields) => {
355 if fields.len() != rty.len() {
356 return Err(DecodeError::UnmatchedRecordType);
357 }
358
359 let mut record = BTreeMap::new();
360
361 for (field, (field_name, field_ty)) in fields.iter().zip(rty.iter()) {
362 let decoded_field = field.decode_literal_expecting(id_maps, Some(field_ty))?;
363 let decoded_field_ty = decoded_field.type_of();
364
365 if &decoded_field_ty != field_ty {
366 return Err(DecodeError::UnmatchedFieldType(
367 decoded_field_ty,
368 field_ty.clone(),
369 ));
370 }
371
372 record.insert(field_name.clone(), decoded_field);
373 }
374
375 Ok(Term::Record(Arc::new(record)))
376 }
377
378 _ => Err(DecodeError::UnknownLiteral(self.clone())),
379 }
380 }
381
382 fn decode_literal_app(
390 &self,
391 id_maps: &IdMaps<'_>,
392 args: &[SExpr],
393 expected_ty: Option<&TermType>,
394 ) -> Result<Term, DecodeError> {
395 match args {
396 [SExpr::Symbol(not_tok), v] if not_tok == "not" => {
402 Ok(factory::not(v.decode_literal(id_maps)?))
403 }
404
405 [SExpr::Symbol(or_tok), v1, v2] if or_tok == "or" => Ok(factory::or(
407 v1.decode_literal(id_maps)?,
408 v2.decode_literal(id_maps)?,
409 )),
410
411 [SExpr::Symbol(eq_tok), v1, v2] if eq_tok == "=" => Ok(factory::eq(
413 v1.decode_literal(id_maps)?,
414 v2.decode_literal(id_maps)?,
415 )),
416
417 [SExpr::Symbol(ite_tok), cond, true_branch, false_branch] if ite_tok == "ite" => {
419 Ok(factory::ite(
420 cond.decode_literal(id_maps)?,
421 true_branch.decode_literal_expecting(id_maps, expected_ty)?,
422 false_branch.decode_literal_expecting(id_maps, expected_ty)?,
423 ))
424 }
425
426 [SExpr::Symbol(bvnego_tok), v] if bvnego_tok == "bvnego" => {
428 Ok(factory::bvnego(v.decode_literal(id_maps)?))
429 }
430
431 [SExpr::Symbol(bvsaddo_tok), v1, v2] if bvsaddo_tok == "bvsaddo" => Ok(
433 factory::bvsaddo(v1.decode_literal(id_maps)?, v2.decode_literal(id_maps)?),
434 ),
435
436 [SExpr::Symbol(bvsmulo_tok), v1, v2] if bvsmulo_tok == "bvsmulo" => Ok(
438 factory::bvsmulo(v1.decode_literal(id_maps)?, v2.decode_literal(id_maps)?),
439 ),
440
441 [SExpr::Symbol(as_tok), SExpr::Symbol(none), typ]
443 if as_tok == "as" && none == "none" =>
444 {
445 match typ.decode_type(id_maps)? {
446 TermType::Option { ty } => Ok(Term::None(Arc::unwrap_or_clone(ty))),
447 _ => Err(DecodeError::InvalidOptionType(typ.clone())),
448 }
449 }
450
451 #[expect(
453 clippy::indexing_slicing,
454 reason = "Slice of length 3 can be indexed by 0-2"
455 )]
456 [SExpr::App(as_some_typ), val]
457 if as_some_typ.len() == 3
458 && as_some_typ[0].is_symbol("as")
459 && as_some_typ[1].is_symbol("some") =>
460 {
461 let ty = as_some_typ[2].decode_type(id_maps)?;
462 let inner_ty = match &ty {
463 TermType::Option { ty } => Some(ty.as_ref()),
464 _ => None,
465 };
466 let val = Term::Some(Arc::new(val.decode_literal_expecting(id_maps, inner_ty)?));
467 let val_ty = val.type_of();
468
469 if val_ty != ty {
470 return Err(DecodeError::UnmatchedType(val_ty, ty));
471 }
472
473 Ok(val)
474 }
475
476 [SExpr::Symbol(some), val] if some == "some" => {
478 let inner_ty = match expected_ty {
479 None => None,
480 Some(TermType::Option { ty }) => Some(ty.as_ref()),
481 Some(_) => return Err(DecodeError::UnknownLiteral(self.clone())),
482 };
483 let val = val.decode_literal_expecting(id_maps, inner_ty)?;
484 Ok(Term::Some(Arc::new(val)))
485 }
486
487 [SExpr::Symbol(as_tok), SExpr::Symbol(set_empty), typ]
489 if as_tok == "as" && set_empty == "set.empty" =>
490 {
491 let ty = typ.decode_type(id_maps)?;
492
493 match ty {
494 TermType::Set { ty } => Ok(Term::Set {
495 elts: Arc::new(BTreeSet::new()),
496 elts_ty: Arc::unwrap_or_clone(ty),
497 }),
498 _ => Err(DecodeError::InvalidSetType(typ.clone())),
499 }
500 }
501
502 [SExpr::Symbol(set_singleton), val] if set_singleton == "set.singleton" => {
504 let elt_ty = match expected_ty {
505 None => None,
506 Some(TermType::Set { ty }) => Some(ty.as_ref()),
507 Some(_) => return Err(DecodeError::UnknownLiteral(self.clone())),
508 };
509 let val = val.decode_literal_expecting(id_maps, elt_ty)?;
510 let val_ty = val.type_of();
511 Ok(Term::Set {
512 elts: Arc::new(BTreeSet::from([val])),
513 elts_ty: val_ty,
514 })
515 }
516
517 [SExpr::Symbol(set_union), set1, set2] if set_union == "set.union" => {
519 let set1 = set1.decode_literal_expecting(id_maps, expected_ty)?;
520 let set2 = set2.decode_literal_expecting(id_maps, expected_ty)?;
521 let set1_ty = set1.type_of();
522 let set2_ty = set2.type_of();
523
524 if set1_ty != set2_ty {
525 return Err(DecodeError::UnmatchedType(set1_ty, set2_ty));
526 }
527
528 match (set1, set2) {
529 (
531 Term::Set {
532 elts: elts1,
533 elts_ty,
534 },
535 Term::Set { elts: elts2, .. },
536 ) => Ok(Term::Set {
537 elts: Arc::new(elts1.union(&elts2).cloned().collect()),
538 elts_ty,
539 }),
540
541 (set1, set2) => Err(DecodeError::SetUnionNonLiterals(set1, set2)),
542 }
543 }
544
545 [SExpr::Symbol(decimal), SExpr::BitVec(bv)]
547 if decimal == "Decimal" && bv.width() == SIXTY_FOUR =>
548 {
549 Ok(Term::Prim(TermPrim::Ext(Ext::Decimal {
550 d: Decimal(
551 bv.to_int()
552 .try_into()
553 .map_err(|_| DecodeError::UnknownLiteral(self.clone()))?,
554 ),
555 })))
556 }
557
558 [SExpr::Symbol(datetime), SExpr::BitVec(bv)]
560 if datetime == "Datetime" && bv.width() == SIXTY_FOUR =>
561 {
562 let dt: i64 = bv
563 .to_int()
564 .try_into()
565 .map_err(|_| DecodeError::IntegerOverflow)?;
566 Ok(Term::Prim(TermPrim::Ext(Ext::Datetime { dt: dt.into() })))
567 }
568
569 [SExpr::Symbol(duration), SExpr::BitVec(bv)]
571 if duration == "Duration" && bv.width() == SIXTY_FOUR =>
572 {
573 let d: i64 = bv
574 .to_int()
575 .try_into()
576 .map_err(|_| DecodeError::IntegerOverflow)?;
577 Ok(Term::Prim(TermPrim::Ext(Ext::Duration { d: d.into() })))
578 }
579
580 [SExpr::Symbol(ip), addr, prefix] if (ip == "V4" || ip == "V6") => {
582 let addr = match addr.decode_literal(id_maps)? {
583 Term::Prim(TermPrim::Bitvec(bv)) => bv,
584 _ => Err(DecodeError::UnknownLiteral(self.clone()))?,
585 };
586 let prefix = match prefix.decode_literal(id_maps)? {
587 Term::Some(t) => match Arc::unwrap_or_clone(t) {
588 Term::Prim(TermPrim::Bitvec(bv)) => Some(bv),
589 _ => Err(DecodeError::UnknownLiteral(self.clone()))?,
590 },
591 Term::None(..) => None,
592 _ => Err(DecodeError::UnknownLiteral(self.clone()))?,
593 };
594 Ok(Term::Prim(TermPrim::Ext(Ext::Ipaddr {
595 ip: if ip == "V4" {
596 IPNet::V4(CIDRv4 {
597 addr: IPv4Addr::try_from_bitvec(addr)
598 .ok_or_else(|| DecodeError::UnknownLiteral(self.clone()))?,
599 prefix: IPv4Prefix::try_from_bitvec(prefix)
600 .ok_or_else(|| DecodeError::UnknownLiteral(self.clone()))?,
601 })
602 } else {
603 IPNet::V6(CIDRv6 {
604 addr: IPv6Addr::try_from_bitvec(addr)
605 .ok_or_else(|| DecodeError::UnknownLiteral(self.clone()))?,
606 prefix: IPv6Prefix::try_from_bitvec(prefix)
607 .ok_or_else(|| DecodeError::UnknownLiteral(self.clone()))?,
608 })
609 },
610 })))
611 }
612
613 [SExpr::Symbol(underscore), SExpr::Symbol(bv_val), SExpr::Numeral(w)]
615 if underscore == "_" && bv_val.starts_with("bv") =>
616 {
617 #[expect(clippy::string_slice, reason = "starts_with guarantees len >= 2")]
618 let val_str = &bv_val[2..];
619 let val: u128 = val_str.parse().map_err(DecodeError::ParseIntError)?;
620 let width = u32::try_from(*w).map_err(|_| DecodeError::IntegerOverflow)?;
621 let width = Width::new(width).ok_or(DecodeError::ZeroWidthBitVec)?;
622 if width.get() < 128 && val >= (1u128 << width.get()) {
625 return Err(DecodeError::IntegerOverflow);
626 }
627 Ok(Term::Prim(TermPrim::Bitvec(BitVec::of_u128(width, val))))
628 }
629
630 [SExpr::Symbol(name), rest_args @ ..] => {
632 self.decode_entity_or_record(id_maps, name, rest_args)
633 }
634
635 _ => Err(DecodeError::UnknownLiteral(self.clone())),
636 }
637 }
638
639 pub fn decode_literal(&self, id_maps: &IdMaps<'_>) -> Result<Term, DecodeError> {
641 self.decode_literal_expecting(id_maps, None)
642 }
643
644 fn decode_literal_expecting(
651 &self,
652 id_maps: &IdMaps<'_>,
653 expected_ty: Option<&TermType>,
654 ) -> Result<Term, DecodeError> {
655 match self {
656 SExpr::BitVec(bv) => Ok(Term::Prim(TermPrim::Bitvec(bv.clone()))),
657 SExpr::String(s) => Ok(Term::Prim(TermPrim::String(SmolStr::new(s)))),
658
659 SExpr::Symbol(s) if s == "true" => Ok(Term::Prim(TermPrim::Bool(true))),
660 SExpr::Symbol(s) if s == "false" => Ok(Term::Prim(TermPrim::Bool(false))),
661
662 SExpr::Symbol(s) if s == "none" => match expected_ty {
664 Some(TermType::Option { ty }) => Ok(Term::None(ty.as_ref().clone())),
665 _ => Err(DecodeError::UnknownLiteral(self.clone())),
666 },
667
668 SExpr::Symbol(s) if id_maps.types.contains_key(s) => {
670 self.decode_entity_or_record(id_maps, s, &[])
671 }
672
673 SExpr::Symbol(e) => id_maps
675 .enums
676 .get(e.as_str())
677 .cloned()
678 .map(|uid| Term::Prim(TermPrim::Entity(uid)))
679 .ok_or_else(|| DecodeError::UnknownLiteral(self.clone())),
680
681 SExpr::App(args) => self.decode_literal_app(id_maps, args, expected_ty),
683
684 _ => Err(DecodeError::UnknownLiteral(self.clone())),
685 }
686 }
687
688 pub fn decode_var(
690 id_maps: &IdMaps<'_>,
691 name: &SmolStr,
692 typ: &SExpr,
693 value: &SExpr,
694 ) -> Result<(TermVar, Term), DecodeError> {
695 let Some(&term_var) = id_maps.vars.get(name) else {
696 return Err(DecodeError::UnknownVariable(name.to_string()));
697 };
698
699 let ty = typ.decode_type(id_maps)?;
700 let val = value.decode_literal_expecting(id_maps, Some(&ty))?;
701 let val_ty = val.type_of();
702
703 if val_ty != ty {
704 return Err(DecodeError::UnmatchedType(val_ty, ty));
705 }
706
707 if term_var.ty != ty {
708 return Err(DecodeError::UnmatchedType(term_var.ty.clone(), ty));
709 }
710
711 Ok((term_var.clone(), val))
712 }
713
714 pub fn decode_unary_function(
721 id_maps: &IdMaps<'_>,
722 name: &SmolStr,
723 arg_name: &str,
724 arg_typ: &SExpr,
725 ret_typ: &SExpr,
726 body: &SExpr,
727 ) -> Result<(Uuf, Udf), DecodeError> {
728 let Some(&uuf) = id_maps.uufs.get(name) else {
730 return Err(DecodeError::UnknownUUF(name.to_string()));
731 };
732
733 let arg_ty = arg_typ.decode_type(id_maps)?;
735 let ret_ty = ret_typ.decode_type(id_maps)?;
736
737 if arg_ty != uuf.arg {
738 return Err(DecodeError::UnmatchedType(arg_ty, uuf.arg.clone()));
739 }
740
741 if ret_ty != uuf.out {
742 return Err(DecodeError::UnmatchedType(ret_ty, uuf.out.clone()));
743 }
744
745 if body.is_app_of("or") {
746 Self::decode_or_table(uuf, id_maps, arg_name, body)
747 } else if body.is_app_of("=") {
748 Self::decode_eq_table(uuf, id_maps, arg_name, body)
749 } else {
750 Self::decode_ite_table(uuf, id_maps, arg_name, &ret_ty, body)
752 }
753 }
754
755 fn decode_ite_table(
757 uuf: &Uuf,
758 id_maps: &IdMaps<'_>,
759 arg_name: &str,
760 ret_ty: &TermType,
761 body: &SExpr,
762 ) -> Result<(Uuf, Udf), DecodeError> {
763 let mut table = BTreeMap::new();
765
766 let mut cur_body = body;
767
768 while let Some([cond, then_expr, else_expr]) = cur_body.as_app_n("ite") {
769 table.insert(
770 Self::decode_eq_operand(arg_name, id_maps, cond)?,
771 then_expr.decode_literal_expecting(id_maps, Some(ret_ty))?,
772 );
773 cur_body = else_expr;
774 }
775
776 let default = cur_body.decode_literal_expecting(id_maps, Some(ret_ty))?;
778 Ok((
779 uuf.clone(),
780 Udf {
781 arg: uuf.arg.clone(),
782 out: uuf.out.clone(),
783 table: Arc::new(table),
784 default,
785 },
786 ))
787 }
788
789 fn decode_or_table(
791 uuf: &Uuf,
792 id_maps: &IdMaps<'_>,
793 arg_name: &str,
794 body: &SExpr,
795 ) -> Result<(Uuf, Udf), DecodeError> {
796 let disjuncts = body
797 .as_app("or")
798 .ok_or_else(|| DecodeError::UnexpectedUnaryFunctionForm(body.clone()))?;
799
800 let mut table = BTreeMap::new();
801
802 for expr in disjuncts {
803 table.insert(
804 Self::decode_eq_operand(arg_name, id_maps, expr)?,
805 Term::Prim(TermPrim::Bool(true)),
806 );
807 }
808
809 Ok((
810 uuf.clone(),
811 Udf {
812 arg: uuf.arg.clone(),
813 out: uuf.out.clone(),
814 table: Arc::new(table),
815 default: Term::Prim(TermPrim::Bool(false)),
816 },
817 ))
818 }
819
820 fn decode_eq_table(
822 uuf: &Uuf,
823 id_maps: &IdMaps<'_>,
824 arg_name: &str,
825 body: &SExpr,
826 ) -> Result<(Uuf, Udf), DecodeError> {
827 let cond_lit_term = Self::decode_eq_operand(arg_name, id_maps, body)?;
828 Ok((
829 uuf.clone(),
830 Udf {
831 arg: uuf.arg.clone(),
832 out: uuf.out.clone(),
833 table: Arc::new(BTreeMap::from([(
834 cond_lit_term,
835 Term::Prim(TermPrim::Bool(true)),
836 )])),
837 default: Term::Prim(TermPrim::Bool(false)),
838 },
839 ))
840 }
841
842 fn decode_eq_operand(
844 arg_name: &str,
845 id_maps: &IdMaps<'_>,
846 eq: &SExpr,
847 ) -> Result<Term, DecodeError> {
848 let [lhs, rhs] = eq.as_app_n("=").ok_or(DecodeError::UnexpectedModel)?;
849
850 if rhs.is_symbol(arg_name) {
851 lhs
852 } else if lhs.is_symbol(arg_name) {
853 rhs
854 } else {
855 return Err(DecodeError::UnexpectedModel);
856 }
857 .decode_literal(id_maps)
858 }
859
860 fn decode_model<'a>(
862 &self,
863 env: &'a SymEnv,
864 id_maps: &IdMaps<'_>,
865 ) -> Result<Interpretation<'a>, DecodeError> {
866 let SExpr::App(cmds) = self else {
867 return Err(DecodeError::UnexpectedModel);
868 };
869
870 let mut vars = BTreeMap::new();
871 let mut funs = BTreeMap::new();
872
873 for cmd in cmds {
875 let SExpr::App(sub_exprs) = cmd else {
876 return Err(DecodeError::UnexpectedModel);
877 };
878
879 match sub_exprs.as_slice() {
882 [SExpr::Symbol(define_fun), SExpr::Symbol(name), SExpr::App(args), ret_ty, body]
883 if define_fun == "define-fun" =>
884 {
885 match args.as_slice() {
886 [SExpr::App(arg)] if arg.len() == 2 => match arg.as_slice() {
888 [SExpr::Symbol(arg_name), arg_ty] => {
889 if id_maps.uufs.contains_key(name) {
890 let (uuf, udf) = Self::decode_unary_function(
891 id_maps, name, arg_name, arg_ty, ret_ty, body,
892 )?;
893 funs.insert(uuf, udf);
894 }
895 }
897 _ => return Err(DecodeError::UnexpectedModel),
898 },
899
900 [] => {
904 if id_maps.vars.contains_key(name) {
905 let (term_var, term) =
906 Self::decode_var(id_maps, name, ret_ty, body)?;
907 vars.insert(term_var, term);
908 }
909 }
910
911 _ => return Err(DecodeError::UnexpectedModel),
912 }
913 }
914
915 _ => return Err(DecodeError::UnexpectedModel),
916 }
917 }
918
919 Ok(Interpretation { vars, funs, env })
920 }
921}
922
923pub fn decode_model<'a>(
925 model: &str,
926 env: &'a SymEnv,
927 id_maps: &IdMaps<'_>,
928) -> Result<Interpretation<'a>, DecodeError> {
929 let model_sexpr = parse_sexpr(model.as_bytes())?;
930 model_sexpr.decode_model(env, id_maps)
931}
932
933#[cfg(test)]
934mod test_decode {
935 use std::{
936 collections::BTreeMap,
937 num::NonZeroU32,
938 str::FromStr,
939 sync::{Arc, LazyLock},
940 };
941
942 use cedar_policy::{EntityId, EntityTypeName, EntityUid, RequestEnv, Schema};
943 use smol_str::SmolStr;
944
945 use cool_asserts::assert_matches;
946
947 use crate::{
948 bitvec::BitVec,
949 err::Term,
950 op::Uuf,
951 symcc::decoder::{sexpr::parse_sexpr, DecodeError, IdMaps},
952 term::{TermPrim, TermVar},
953 term_type::TermType,
954 SymEnv,
955 };
956
957 static TEST_ENV: LazyLock<SymEnv> = LazyLock::new(|| {
958 SymEnv::new(
959 &Schema::from_cedarschema_str(
960 "entity E; action A appliesTo { principal: [E], resource: [E] };",
961 )
962 .unwrap()
963 .0,
964 &RequestEnv::new(
965 "E".parse().unwrap(),
966 "Action::\"A\"".parse().unwrap(),
967 "E".parse().unwrap(),
968 ),
969 )
970 .expect("Malformed sym env.")
971 });
972
973 #[track_caller]
974 fn assert_decode_var(model: &str, var: SmolStr, ty: TermType, expected: impl Into<Term>) {
975 let sexpr = parse_sexpr(model.as_bytes()).expect("failed to parse model sexpr");
976 let var = TermVar { id: var, ty };
977 let actual = sexpr
978 .decode_model(
979 &TEST_ENV,
980 &IdMaps {
981 types: BTreeMap::new(),
982 vars: BTreeMap::from([(&var.id, &var)]),
983 uufs: BTreeMap::new(),
984 enums: BTreeMap::new(),
985 },
986 )
987 .expect("failed to decode model")
988 .vars
989 .get(&var)
990 .expect("could not find expected var in model")
991 .clone();
992 assert_eq!(actual, expected.into());
993 }
994
995 #[test]
996 fn decode_literals() {
997 assert_decode_var(
998 "((define-fun x () Bool true))",
999 "x".into(),
1000 TermType::Bool,
1001 true,
1002 );
1003 assert_decode_var(
1004 "((define-fun x () Bool false))",
1005 "x".into(),
1006 TermType::Bool,
1007 false,
1008 );
1009 assert_decode_var(
1010 "((define-fun x () (_ BitVec 2) #b11))",
1011 "x".into(),
1012 TermType::Bitvec {
1013 n: NonZeroU32::new(2).unwrap(),
1014 },
1015 BitVec::of_i128(NonZeroU32::new(2).unwrap(), 3),
1016 );
1017 assert_decode_var(
1018 r#"((define-fun x () String "foo"))"#,
1019 "x".into(),
1020 TermType::String,
1021 SmolStr::new_static("foo"),
1022 );
1023 assert_decode_var(
1025 "((define-fun x () (_ BitVec 8) #xFF))",
1026 "x".into(),
1027 TermType::Bitvec {
1028 n: NonZeroU32::new(8).unwrap(),
1029 },
1030 BitVec::of_i128(NonZeroU32::new(8).unwrap(), -1),
1031 );
1032 assert_decode_var(
1034 "((define-fun x () (_ BitVec 8) (_ bv42 8)))",
1035 "x".into(),
1036 TermType::Bitvec {
1037 n: NonZeroU32::new(8).unwrap(),
1038 },
1039 BitVec::of_u128(NonZeroU32::new(8).unwrap(), 42),
1040 );
1041 }
1042
1043 #[test]
1044 fn decode_indexed_bv_err() {
1045 let id_maps = IdMaps {
1046 types: BTreeMap::new(),
1047 vars: BTreeMap::new(),
1048 uufs: BTreeMap::new(),
1049 enums: BTreeMap::new(),
1050 };
1051 assert_matches!(
1052 parse_sexpr(b"(_ bv0 0)").unwrap().decode_literal(&id_maps),
1053 Err(DecodeError::ZeroWidthBitVec)
1054 );
1055 assert_matches!(
1056 parse_sexpr(b"(_ bv256 8)")
1057 .unwrap()
1058 .decode_literal(&id_maps),
1059 Err(DecodeError::IntegerOverflow)
1060 );
1061 }
1062
1063 #[test]
1076 fn decode_model_skips_unknown_define_funs() {
1077 let z3_model = r#"(
1080 (define-fun t0 () Bool true)
1081 (define-fun t3 () Bool (not true))
1082 )"#;
1083 let sexpr = parse_sexpr(z3_model.as_bytes()).expect("failed to parse");
1084 let var = TermVar {
1085 id: "t0".into(),
1086 ty: TermType::Bool,
1087 };
1088 let result = sexpr.decode_model(
1089 &TEST_ENV,
1090 &IdMaps {
1091 types: BTreeMap::new(),
1092 vars: BTreeMap::from([(&var.id, &var)]), uufs: BTreeMap::new(),
1094 enums: BTreeMap::new(),
1095 },
1096 );
1097 let interp = result.expect("decode_model should skip unknown define-funs");
1098 let val = interp.vars.get(&var).expect("t0 should be in the model");
1099 assert_eq!(*val, Term::Prim(crate::symcc::term::TermPrim::Bool(true)));
1100 }
1101
1102 #[test]
1103 fn decode_model_skips_unknown_unary_funs() {
1104 let z3_model = r#"(
1107 (define-fun t0 () Bool true)
1108 (define-fun unknown_fn ((x Bool)) Bool (ite (= true x) false true))
1109 )"#;
1110 let sexpr = parse_sexpr(z3_model.as_bytes()).expect("failed to parse");
1111 let var = TermVar {
1112 id: "t0".into(),
1113 ty: TermType::Bool,
1114 };
1115 let result = sexpr.decode_model(
1116 &TEST_ENV,
1117 &IdMaps {
1118 types: BTreeMap::new(),
1119 vars: BTreeMap::from([(&var.id, &var)]), uufs: BTreeMap::new(),
1121 enums: BTreeMap::new(),
1122 },
1123 );
1124 let interp = result.expect("decode_model should skip unknown unary funs");
1125 let val = interp.vars.get(&var).expect("t0 should be in the model");
1126 assert_eq!(*val, Term::Prim(crate::symcc::term::TermPrim::Bool(true)));
1127 }
1128
1129 #[test]
1130 fn decode_application() {
1131 assert_decode_var(
1132 "((define-fun x () Bool (not true)))",
1133 "x".into(),
1134 TermType::Bool,
1135 false,
1136 );
1137 assert_decode_var(
1138 "((define-fun x () Bool (not (not true))))",
1139 "x".into(),
1140 TermType::Bool,
1141 true,
1142 );
1143 assert_decode_var(
1144 "((define-fun x () Bool (or false true)))",
1145 "x".into(),
1146 TermType::Bool,
1147 true,
1148 );
1149 assert_decode_var(
1150 "((define-fun x () Bool (= true false)))",
1151 "x".into(),
1152 TermType::Bool,
1153 false,
1154 );
1155 assert_decode_var(
1156 r#"((define-fun x () String (ite false "foo" "bar")))"#,
1157 "x".into(),
1158 TermType::String,
1159 SmolStr::new_static("bar"),
1160 );
1161 assert_decode_var(
1162 "((define-fun x () Bool (bvnego #b10)))",
1163 "x".into(),
1164 TermType::Bool,
1165 true,
1166 );
1167 assert_decode_var(
1168 "((define-fun x () Bool (bvsaddo #b01 #b01)))",
1169 "x".into(),
1170 TermType::Bool,
1171 true,
1172 );
1173 assert_decode_var(
1174 "((define-fun x () Bool (bvsmulo #b010 #b010)))",
1175 "x".into(),
1176 TermType::Bool,
1177 true,
1178 );
1179 }
1180
1181 #[test]
1184 fn decode_z3_uuf_arg_on_left() {
1185 let entity_ty = TermType::Entity {
1186 ety: EntityTypeName::from_str("E0").unwrap().clone(),
1187 };
1188 let record_ty = TermType::Record {
1189 rty: Arc::new(BTreeMap::from([("admin".into(), TermType::Bool)])),
1190 };
1191 let uuf = Uuf {
1192 id: "attrs".into(),
1193 arg: entity_ty.clone(),
1194 out: record_ty.clone(),
1195 };
1196 let ety_id: SmolStr = "E0".into();
1197 let rty_id: SmolStr = "R0".into();
1198 let uuf_id: SmolStr = "f0".into();
1199 let id_maps = IdMaps {
1200 types: BTreeMap::from([(&ety_id, &entity_ty), (&rty_id, &record_ty)]),
1201 vars: BTreeMap::new(),
1202 uufs: BTreeMap::from([(&uuf_id, &uuf)]),
1203 enums: BTreeMap::new(),
1204 };
1205
1206 let sexpr = parse_sexpr(
1208 br#"((define-fun f0 ((x!0 E0)) R0 (ite (= x!0 (E0 "bob")) (R0 false) (R0 true))))"#,
1209 )
1210 .unwrap();
1211 let udf = sexpr
1212 .decode_model(&TEST_ENV, &id_maps)
1213 .unwrap()
1214 .funs
1215 .remove(&uuf)
1216 .unwrap();
1217 let bob_key = Term::Prim(TermPrim::Entity(EntityUid::from_type_name_and_id(
1218 EntityTypeName::from_str("E0").unwrap(),
1219 EntityId::new("bob"),
1220 )));
1221 let rec = |b| Term::Record(Arc::new(BTreeMap::from([("admin".into(), Term::from(b))])));
1222 assert_eq!(udf.table.get(&bob_key), Some(&rec(false)));
1223 assert_eq!(udf.default, rec(true));
1224 }
1225
1226 #[test]
1227 fn decode_bool_uuf_eq() {
1228 let entity_ty = TermType::Entity {
1229 ety: EntityTypeName::from_str("E0").unwrap(),
1230 };
1231 let uuf = Uuf {
1232 id: "f".into(),
1233 arg: entity_ty.clone(),
1234 out: TermType::Bool,
1235 };
1236 let ety_id: SmolStr = "E0".into();
1237 let uuf_id: SmolStr = "f0".into();
1238 let id_maps = IdMaps {
1239 types: BTreeMap::from([(&ety_id, &entity_ty)]),
1240 vars: BTreeMap::new(),
1241 uufs: BTreeMap::from([(&uuf_id, &uuf)]),
1242 enums: BTreeMap::new(),
1243 };
1244
1245 let sexpr =
1246 parse_sexpr(br#"((define-fun f0 ((_arg_1 E0)) Bool (= (E0 "") _arg_1)))"#).unwrap();
1247 let interp = sexpr
1248 .decode_model(&TEST_ENV, &id_maps)
1249 .expect("Bool-codomain UF model with `=` body should decode");
1250 let udf = interp.funs.get(&uuf).expect("f0 should be in the model");
1251 let empty_key = Term::Prim(TermPrim::Entity(EntityUid::from_type_name_and_id(
1252 EntityTypeName::from_str("E0").unwrap(),
1253 EntityId::new(""),
1254 )));
1255 assert_eq!(udf.table.get(&empty_key), Some(&Term::from(true)));
1256 assert_eq!(udf.default, Term::from(false));
1257 }
1258
1259 #[test]
1260 fn decode_bool_uuf_or() {
1261 let entity_ty = TermType::Entity {
1262 ety: EntityTypeName::from_str("E0").unwrap(),
1263 };
1264 let uuf = Uuf {
1265 id: "f".into(),
1266 arg: entity_ty.clone(),
1267 out: TermType::Bool,
1268 };
1269 let ety_id: SmolStr = "E0".into();
1270 let uuf_id: SmolStr = "f0".into();
1271 let id_maps = IdMaps {
1272 types: BTreeMap::from([(&ety_id, &entity_ty)]),
1273 vars: BTreeMap::new(),
1274 uufs: BTreeMap::from([(&uuf_id, &uuf)]),
1275 enums: BTreeMap::new(),
1276 };
1277 let key = |eid: &str| {
1278 Term::Prim(TermPrim::Entity(EntityUid::from_type_name_and_id(
1279 EntityTypeName::from_str("E0").unwrap(),
1280 EntityId::new(eid),
1281 )))
1282 };
1283
1284 let sexpr = parse_sexpr(
1285 br#"((define-fun f0 ((_arg_1 E0)) Bool (or (= (E0 "a") _arg_1) (= _arg_1 (E0 "b")) (= _arg_1 (E0 "c")))))"#,
1286 )
1287 .unwrap();
1288 let interp = sexpr
1289 .decode_model(&TEST_ENV, &id_maps)
1290 .expect("cvc5 Bool-codomain UF model with `or` body should decode");
1291 let udf = interp.funs.get(&uuf).expect("f0 should be in the model");
1292 assert_eq!(udf.table.get(&key("a")), Some(&Term::from(true)));
1293 assert_eq!(udf.table.get(&key("b")), Some(&Term::from(true)));
1294 assert_eq!(udf.table.get(&key("c")), Some(&Term::from(true)));
1295 assert_eq!(udf.default, Term::from(false));
1296 }
1297
1298 #[test]
1300 fn decode_z3_bare_none_and_some() {
1301 let opt_str = TermType::option_of(TermType::String);
1302 let rty = TermType::Record {
1303 rty: Arc::new(BTreeMap::from([("a".into(), opt_str.clone())])),
1304 };
1305 let type_id: SmolStr = "R0".into();
1306
1307 let var = TermVar {
1309 id: "t0".into(),
1310 ty: rty.clone(),
1311 };
1312 let sexpr = parse_sexpr(b"((define-fun t0 () R0 (R0 none)))").unwrap();
1313 let interp = sexpr
1314 .decode_model(
1315 &TEST_ENV,
1316 &IdMaps {
1317 types: BTreeMap::from([(&type_id, &rty)]),
1318 vars: BTreeMap::from([(&var.id, &var)]),
1319 uufs: BTreeMap::new(),
1320 enums: BTreeMap::new(),
1321 },
1322 )
1323 .expect("bare none in record");
1324 assert_eq!(
1325 *interp.vars.get(&var).unwrap(),
1326 Term::Record(Arc::new(BTreeMap::from([(
1327 "a".into(),
1328 Term::None(TermType::String)
1329 )])))
1330 );
1331
1332 let sexpr = parse_sexpr(br#"((define-fun t0 () R0 (R0 (some "x"))))"#).unwrap();
1334 let interp = sexpr
1335 .decode_model(
1336 &TEST_ENV,
1337 &IdMaps {
1338 types: BTreeMap::from([(&type_id, &rty)]),
1339 vars: BTreeMap::from([(&var.id, &var)]),
1340 uufs: BTreeMap::new(),
1341 enums: BTreeMap::new(),
1342 },
1343 )
1344 .expect("bare some in record");
1345 assert_eq!(
1346 *interp.vars.get(&var).unwrap(),
1347 Term::Record(Arc::new(BTreeMap::from([(
1348 "a".into(),
1349 Term::Some(Arc::new(Term::Prim(TermPrim::String("x".into()))))
1350 )])))
1351 );
1352
1353 let var2 = TermVar {
1355 id: "t0".into(),
1356 ty: opt_str.clone(),
1357 };
1358 let sexpr = parse_sexpr(b"((define-fun t0 () (Option String) none))").unwrap();
1359 let interp = sexpr
1360 .decode_model(
1361 &TEST_ENV,
1362 &IdMaps {
1363 types: BTreeMap::new(),
1364 vars: BTreeMap::from([(&var2.id, &var2)]),
1365 uufs: BTreeMap::new(),
1366 enums: BTreeMap::new(),
1367 },
1368 )
1369 .expect("bare none as constant");
1370 assert_eq!(
1371 *interp.vars.get(&var2).unwrap(),
1372 Term::None(TermType::String)
1373 );
1374 }
1375}
1376
1377#[cfg(test)]
1378mod test_decode_type_mismatch {
1379 use std::{collections::BTreeMap, num::NonZeroU32, sync::LazyLock};
1380
1381 use cedar_policy::{RequestEnv, Schema};
1382 use cool_asserts::assert_matches;
1383 use smol_str::SmolStr;
1384
1385 use crate::{
1386 op::Uuf,
1387 symcc::decoder::{sexpr::parse_sexpr, DecodeError, IdMaps},
1388 term::TermVar,
1389 term_type::TermType,
1390 SymEnv,
1391 };
1392
1393 static TEST_ENV: LazyLock<SymEnv> = LazyLock::new(|| {
1394 SymEnv::new(
1395 &Schema::from_cedarschema_str(
1396 "entity E; action A appliesTo { principal: [E], resource: [E] };",
1397 )
1398 .unwrap()
1399 .0,
1400 &RequestEnv::new(
1401 "E".parse().unwrap(),
1402 "Action::\"A\"".parse().unwrap(),
1403 "E".parse().unwrap(),
1404 ),
1405 )
1406 .expect("Malformed sym env.")
1407 });
1408
1409 #[test]
1410 fn env_var_model_mismatch() {
1411 let var = TermVar {
1412 id: "x".into(),
1413 ty: TermType::Bool,
1414 };
1415 let sexpr = parse_sexpr(br#"((define-fun x () String "hello"))"#).unwrap();
1416 let result = sexpr.decode_model(
1417 &TEST_ENV,
1418 &IdMaps {
1419 types: BTreeMap::new(),
1420 vars: BTreeMap::from([(&var.id, &var)]),
1421 uufs: BTreeMap::new(),
1422 enums: BTreeMap::new(),
1423 },
1424 );
1425 assert_matches!(
1426 result,
1427 Err(DecodeError::UnmatchedType(TermType::Bool, TermType::String))
1428 );
1429 }
1430
1431 #[test]
1432 fn env_arg_model_mismatch() {
1433 let uuf = Uuf {
1434 id: "f".into(),
1435 arg: TermType::Bool,
1436 out: TermType::Bool,
1437 };
1438 let uuf_id: SmolStr = "f0".into();
1439 let sexpr =
1440 parse_sexpr(br#"((define-fun f0 ((x String)) Bool (ite (= x "a") true false)))"#)
1441 .unwrap();
1442 let result = sexpr.decode_model(
1443 &TEST_ENV,
1444 &IdMaps {
1445 types: BTreeMap::new(),
1446 vars: BTreeMap::new(),
1447 uufs: BTreeMap::from([(&uuf_id, &uuf)]),
1448 enums: BTreeMap::new(),
1449 },
1450 );
1451 assert_matches!(
1452 result,
1453 Err(DecodeError::UnmatchedType(TermType::String, TermType::Bool))
1454 );
1455 }
1456
1457 #[test]
1458 fn env_ret_model_mismatch() {
1459 let uuf = Uuf {
1460 id: "f".into(),
1461 arg: TermType::Bool,
1462 out: TermType::Bool,
1463 };
1464 let uuf_id: SmolStr = "f0".into();
1465 let sexpr = parse_sexpr(br#"((define-fun f0 ((x Bool)) String (ite (= x true) "a" "b")))"#)
1466 .unwrap();
1467 let result = sexpr.decode_model(
1468 &TEST_ENV,
1469 &IdMaps {
1470 types: BTreeMap::new(),
1471 vars: BTreeMap::new(),
1472 uufs: BTreeMap::from([(&uuf_id, &uuf)]),
1473 enums: BTreeMap::new(),
1474 },
1475 );
1476 assert_matches!(
1477 result,
1478 Err(DecodeError::UnmatchedType(TermType::String, TermType::Bool))
1479 );
1480 }
1481
1482 #[test]
1483 fn illtyped_model() {
1484 let var = TermVar {
1485 id: "x".into(),
1486 ty: TermType::Bitvec {
1487 n: NonZeroU32::new(8).unwrap(),
1488 },
1489 };
1490 let sexpr = parse_sexpr(b"((define-fun x () (_ BitVec 8) #b01))").unwrap();
1491 let result = sexpr.decode_model(
1492 &TEST_ENV,
1493 &IdMaps {
1494 types: BTreeMap::new(),
1495 vars: BTreeMap::from([(&var.id, &var)]),
1496 uufs: BTreeMap::new(),
1497 enums: BTreeMap::new(),
1498 },
1499 );
1500 assert_matches!(
1501 result,
1502 Err(DecodeError::UnmatchedType(val_ty, declared_ty))
1503 if val_ty == TermType::Bitvec { n: NonZeroU32::new(2).unwrap() }
1504 && declared_ty == TermType::Bitvec { n: NonZeroU32::new(8).unwrap() }
1505 );
1506 }
1507
1508 #[test]
1510 fn record_field_count_mismatch() {
1511 use std::sync::Arc;
1512 let rty = TermType::Record {
1513 rty: Arc::new(BTreeMap::from([
1514 ("a".into(), TermType::Bool),
1515 ("b".into(), TermType::Bool),
1516 ])),
1517 };
1518 let rty_id: SmolStr = "R0".into();
1519 let var = TermVar {
1520 id: "x".into(),
1521 ty: rty.clone(),
1522 };
1523 let sexpr = parse_sexpr(b"((define-fun x () R0 (R0 true)))").unwrap();
1525 let result = sexpr.decode_model(
1526 &TEST_ENV,
1527 &IdMaps {
1528 types: BTreeMap::from([(&rty_id, &rty)]),
1529 vars: BTreeMap::from([(&var.id, &var)]),
1530 uufs: BTreeMap::new(),
1531 enums: BTreeMap::new(),
1532 },
1533 );
1534 assert_matches!(result, Err(DecodeError::UnmatchedRecordType));
1535 }
1536
1537 #[test]
1539 fn record_field_type_mismatch() {
1540 use std::sync::Arc;
1541 let rty = TermType::Record {
1542 rty: Arc::new(BTreeMap::from([("name".into(), TermType::String)])),
1543 };
1544 let rty_id: SmolStr = "R0".into();
1545 let var = TermVar {
1546 id: "x".into(),
1547 ty: rty.clone(),
1548 };
1549 let sexpr = parse_sexpr(b"((define-fun x () R0 (R0 true)))").unwrap();
1551 let result = sexpr.decode_model(
1552 &TEST_ENV,
1553 &IdMaps {
1554 types: BTreeMap::from([(&rty_id, &rty)]),
1555 vars: BTreeMap::from([(&var.id, &var)]),
1556 uufs: BTreeMap::new(),
1557 enums: BTreeMap::new(),
1558 },
1559 );
1560 assert_matches!(result, Err(DecodeError::UnmatchedFieldType(..)));
1561 }
1562
1563 #[test]
1564 fn entity_non_string_arg() {
1565 use cedar_policy::EntityTypeName;
1566 use std::str::FromStr;
1567 let ety = TermType::Entity {
1568 ety: EntityTypeName::from_str("E0").unwrap(),
1569 };
1570 let ety_id: SmolStr = "E0".into();
1571 let var = TermVar {
1572 id: "x".into(),
1573 ty: ety.clone(),
1574 };
1575 let sexpr = parse_sexpr(b"((define-fun x () E0 (E0 true)))").unwrap();
1577 let result = sexpr.decode_model(
1578 &TEST_ENV,
1579 &IdMaps {
1580 types: BTreeMap::from([(&ety_id, &ety)]),
1581 vars: BTreeMap::from([(&var.id, &var)]),
1582 uufs: BTreeMap::new(),
1583 enums: BTreeMap::new(),
1584 },
1585 );
1586 assert_matches!(result, Err(DecodeError::UnknownLiteral(..)));
1587 }
1588}
1589
1590#[cfg(test)]
1591mod test_decode_unexpected_model {
1592 use std::{collections::BTreeMap, sync::LazyLock};
1593
1594 use cedar_policy::{RequestEnv, Schema};
1595 use cool_asserts::assert_matches;
1596 use smol_str::SmolStr;
1597
1598 use crate::{
1599 op::Uuf,
1600 symcc::decoder::{sexpr::parse_sexpr, DecodeError, IdMaps},
1601 term_type::TermType,
1602 SymEnv,
1603 };
1604
1605 static TEST_ENV: LazyLock<SymEnv> = LazyLock::new(|| {
1606 SymEnv::new(
1607 &Schema::from_cedarschema_str(
1608 "entity E; action A appliesTo { principal: [E], resource: [E] };",
1609 )
1610 .unwrap()
1611 .0,
1612 &RequestEnv::new(
1613 "E".parse().unwrap(),
1614 "Action::\"A\"".parse().unwrap(),
1615 "E".parse().unwrap(),
1616 ),
1617 )
1618 .expect("Malformed sym env.")
1619 });
1620
1621 fn empty_id_maps() -> IdMaps<'static> {
1622 IdMaps {
1623 types: BTreeMap::new(),
1624 vars: BTreeMap::new(),
1625 uufs: BTreeMap::new(),
1626 enums: BTreeMap::new(),
1627 }
1628 }
1629
1630 #[rstest::rstest]
1631 #[case::top_level_not_app(b"true")]
1632 #[case::command_not_app(b"(42)")]
1633 #[case::not_define_fun(b"((declare-const x Bool))")]
1634 #[case::define_fun_too_few_parts(b"((define-fun x () Bool))")]
1635 #[case::define_fun_name_not_symbol(b"((define-fun 123 () Bool true))")]
1636 #[case::define_fun_args_not_app(b"((define-fun x foo Bool true))")]
1637 #[case::multi_arg_function(b"((define-fun f ((x Bool) (y Bool)) Bool true))")]
1638 fn unexpected_model(#[case] input: &[u8]) {
1639 let result = parse_sexpr(input)
1640 .unwrap()
1641 .decode_model(&TEST_ENV, &empty_id_maps());
1642 assert_matches!(result, Err(DecodeError::UnexpectedModel));
1643 }
1644
1645 #[test]
1647 fn unary_fun_arg_not_symbol() {
1648 let uuf = Uuf {
1649 id: "f".into(),
1650 arg: TermType::Bool,
1651 out: TermType::Bool,
1652 };
1653 let uuf_id: SmolStr = "f0".into();
1654 let sexpr = parse_sexpr(b"((define-fun f0 ((42 Bool)) Bool true))").unwrap();
1655 let result = sexpr.decode_model(
1656 &TEST_ENV,
1657 &IdMaps {
1658 types: BTreeMap::new(),
1659 vars: BTreeMap::new(),
1660 uufs: BTreeMap::from([(&uuf_id, &uuf)]),
1661 enums: BTreeMap::new(),
1662 },
1663 );
1664 assert_matches!(result, Err(DecodeError::UnexpectedModel));
1665 }
1666
1667 #[rstest::rstest]
1668 #[case::eq_missing_operand(br#"((define-fun f0 ((x E0)) Bool (= x)))"#)]
1669 #[case::eq_wrong_var(br#"((define-fun f0 ((x E0)) Bool (= (E0 "a") y)))"#)]
1670 #[case::eq_extra_operand(br#"((define-fun f0 ((x E0)) Bool (= (E0 "a") x x)))"#)]
1671 #[case::ite_missing_else(br#"((define-fun f0 ((x E0)) Bool (ite (= (E0 "a") x) true)))"#)]
1672 #[case::ite_extra_arg(
1673 br#"((define-fun f0 ((x E0)) Bool (ite (= (E0 "a") x) true false false)))"#
1674 )]
1675 fn malformed_uuf_table(#[case] input: &[u8]) {
1676 let entity_ty = TermType::Entity {
1677 ety: "E0".parse().unwrap(),
1678 };
1679 let uuf = Uuf {
1680 id: "f".into(),
1681 arg: entity_ty.clone(),
1682 out: TermType::Bool,
1683 };
1684 let ety_id: SmolStr = "E0".into();
1685 let uuf_id: SmolStr = "f0".into();
1686 let id_maps = IdMaps {
1687 types: BTreeMap::from([(&ety_id, &entity_ty)]),
1688 vars: BTreeMap::new(),
1689 uufs: BTreeMap::from([(&uuf_id, &uuf)]),
1690 enums: BTreeMap::new(),
1691 };
1692 let err = parse_sexpr(input)
1693 .unwrap()
1694 .decode_model(&TEST_ENV, &id_maps);
1695 assert_matches!(
1696 err,
1697 Err(DecodeError::UnexpectedModel
1698 | DecodeError::UnknownLiteral(_)
1699 | DecodeError::UnexpectedUnaryFunctionForm(_))
1700 );
1701 }
1702
1703 #[test]
1704 fn unknown_variable() {
1705 use crate::symcc::decoder::SExpr;
1706 let name: SmolStr = "unknown".into();
1707 let typ = SExpr::Symbol("Bool".into());
1708 let value = SExpr::Symbol("true".into());
1709 let result = SExpr::decode_var(&empty_id_maps(), &name, &typ, &value);
1710 assert_matches!(result, Err(DecodeError::UnknownVariable(v)) if v == "unknown");
1711 }
1712
1713 #[test]
1714 fn unknown_uuf() {
1715 use crate::symcc::decoder::SExpr;
1716 let name: SmolStr = "unknown_f".into();
1717 let arg_typ = SExpr::Symbol("Bool".into());
1718 let ret_typ = SExpr::Symbol("Bool".into());
1719 let body = SExpr::Symbol("true".into());
1720 let result =
1721 SExpr::decode_unary_function(&empty_id_maps(), &name, "x", &arg_typ, &ret_typ, &body);
1722 assert_matches!(result, Err(DecodeError::UnknownUUF(v)) if v == "unknown_f");
1723 }
1724}