1use async_recursion::async_recursion;
63use itertools::Itertools;
64use miette::Diagnostic;
65use smol_str::{format_smolstr, SmolStr, ToSmolStr};
66use std::collections::{BTreeMap, BTreeSet};
67use std::fmt::Write;
68use thiserror::Error;
69
70use cedar_policy_core::ast::PatternElem;
71
72use super::{
73 bitvec::{BitVec, BitVecError},
74 env::SymEnv,
75 ext::Ext,
76 extension_types::ipaddr::{CIDRv4, CIDRv6, IPNet, IPv4Prefix, IPv6Prefix},
77 op::{ExtOp, Op, Uuf},
78 smtlib_script::SmtLibScript,
79 term::{Term, TermPrim, TermVar},
80 term_type::TermType,
81 type_abbrevs::*,
82};
83
84use super::extension_types::ipaddr::{V4_WIDTH, V6_WIDTH};
85
86#[derive(Debug, Diagnostic, Error)]
89pub enum EncodeError {
90 #[error("IO error during SMT encoding")]
92 Io(#[from] std::io::Error),
93 #[error("missing member {0} in enum entity")]
95 EnumMissingMember(EntityUID),
96 #[error("record missing attribute {0}")]
98 RecordMissingAttr(Attr),
99 #[error("expecting a record type, got {0:?}")]
101 ExpectRecord(TermType),
102 #[error("missing type encoding for {0:?}")]
104 MissingTypeEncoding(TermType),
105 #[error("malformed record get")]
107 MalformedRecordGet,
108 #[error("unable to encode string \"{0}\" in SMT as it exceeds the max supported code point")]
110 EncodeStringFailed(SmolStr),
111 #[error("unable to encode pattern {0:?} in SMT as it exceeds the max supported code point")]
113 EncodePatternFailed(OrdPattern),
114 #[error("bit-vector error")]
116 BitVecError(#[from] BitVecError),
117}
118
119type Result<T> = std::result::Result<T, EncodeError>;
120
121#[derive(Debug)]
122pub struct Encoder<'a, S> {
123 pub(super) terms: BTreeMap<Term, SmolStr>,
124 pub(super) types: BTreeMap<TermType, SmolStr>,
125 pub(super) uufs: BTreeMap<Uuf, SmolStr>,
126 pub(super) enums: BTreeMap<&'a EntityType, &'a BTreeSet<SmolStr>>,
127 script: S,
128}
129
130fn term_id(n: usize) -> SmolStr {
131 format_smolstr!("t{n}")
132}
133
134fn uuf_id(n: usize) -> SmolStr {
135 format_smolstr!("f{n}")
136}
137
138fn entity_type_id(n: usize) -> SmolStr {
139 format_smolstr!("E{n}")
140}
141
142pub(super) fn enum_id(e: &str, n: usize) -> SmolStr {
143 format_smolstr!("{e}_m{n}")
144}
145
146fn record_type_id(n: usize) -> SmolStr {
147 format_smolstr!("R{n}")
148}
149
150fn record_attr_id(r: &str, n: usize) -> SmolStr {
151 format_smolstr!("{r}_a{n}")
152}
153
154impl<'a, S> Encoder<'a, S> {
160 pub fn new(env: &'a SymEnv, script: S) -> Result<Self> {
162 Ok(Encoder {
163 terms: BTreeMap::new(),
164 types: BTreeMap::new(),
165 uufs: BTreeMap::new(),
166 enums: env
167 .entities
168 .iter()
169 .filter_map(|(ety, d)| Some((ety, d.members.as_ref()?)))
170 .collect(),
171 script,
172 })
173 }
174
175 pub fn finalize(self) -> Encoder<'a, ()> {
178 Encoder {
179 terms: self.terms,
180 types: self.types,
181 uufs: self.uufs,
182 enums: self.enums,
183 script: (),
184 }
185 }
186}
187
188impl<S: tokio::io::AsyncWrite + Unpin + Send> Encoder<'_, S> {
189 pub async fn declare_type<T: AsRef<str>>(
191 &mut self,
192 id: T,
193 mks: impl IntoIterator<Item = &str>,
194 ) -> Result<T> {
195 self.script
196 .declare_datatype(id.as_ref(), vec![], mks)
197 .await?;
198 Ok(id)
199 }
200
201 pub async fn declare_entity_type(&mut self, ety: &EntityType) -> Result<SmolStr> {
202 let ety_id = entity_type_id(self.types.len());
203 match self.enums.get(ety) {
204 Some(members) => {
205 self.script
206 .comment(&format_smolstr!("{ety}::[{}]", members.iter().join(", ")))
207 .await?;
208 let mks: Vec<_> = members
209 .iter()
210 .enumerate()
211 .map(|(i, _)| format_smolstr!("({})", enum_id(&ety_id, i)))
212 .collect();
213 self.declare_type(ety_id, mks.iter().map(|s| s.as_str()))
214 .await
215 }
216 None => {
217 self.script.comment(&ety.to_string()).await?;
218 self.declare_type(
219 ety_id.clone(),
220 [format_smolstr!("({ety_id} ({ety_id}_eid String))").as_str()],
221 )
222 .await
223 }
224 }
225 }
226
227 pub async fn declare_ext_type(&mut self, ext_ty: ExtType) -> Result<&'static str> {
228 match ext_ty {
229 ExtType::Decimal => {
230 self.declare_type("Decimal", ["(Decimal (decimalVal (_ BitVec 64)))"])
231 .await
232 }
233 ExtType::IpAddr => {
234 self.declare_type(
235 "IPAddr",
236 [
237 "(V4 (addrV4 (_ BitVec 32)) (prefixV4 (Option (_ BitVec 5))))",
238 "(V6 (addrV6 (_ BitVec 128)) (prefixV6 (Option (_ BitVec 7))))",
239 ],
240 )
241 .await
242 }
243 ExtType::Duration => {
244 self.declare_type("Duration", ["(Duration (durationVal (_ BitVec 64)))"])
245 .await
246 }
247 ExtType::DateTime => {
248 self.declare_type("Datetime", ["(Datetime (datetimeVal (_ BitVec 64)))"])
249 .await
250 }
251 }
252 }
253
254 pub async fn declare_record_type<'r>(
255 &mut self,
256 rty: impl IntoIterator<Item = &'r (Attr, SmolStr)> + Clone,
257 ) -> Result<SmolStr> {
258 let rty_id = record_type_id(self.types.len());
259 let mut attrs = rty
260 .clone()
261 .into_iter()
262 .enumerate()
263 .map(|(i, (_, ty))| format_smolstr!("({} {})", record_attr_id(&rty_id, i), ty));
264 self.script
265 .comment(&format_smolstr!(
266 "{{{}}}",
267 rty.into_iter().map(|(k, _)| k).join(", ")
268 ))
269 .await?;
270 self.declare_type(
271 rty_id.clone(),
272 [format_smolstr!("({} {})", rty_id, attrs.join(" ")).as_str()],
273 )
274 .await
275 }
276
277 #[async_recursion]
278 pub async fn encode_type(&mut self, ty: &TermType) -> Result<SmolStr> {
279 match self.types.get(ty) {
280 Some(enc) => Ok(enc.clone()),
281 None => {
282 let enc = match ty {
283 TermType::Bool => {
284 return Ok(SmolStr::new_static("Bool"));
285 }
286 TermType::String => {
287 return Ok(SmolStr::new_static("String"));
288 }
289 TermType::Bitvec { n } => {
290 return Ok(format_smolstr!("(_ BitVec {n})"));
291 }
292 TermType::Option { ref ty } => {
293 return Ok(format_smolstr!("(Option {})", self.encode_type(ty).await?));
294 }
295 TermType::Set { ty } => {
296 return Ok(format_smolstr!("(Set {})", self.encode_type(ty).await?));
297 }
298 TermType::Entity { ety } => self.declare_entity_type(ety).await?,
299 TermType::Ext { xty } => {
300 SmolStr::new_static(self.declare_ext_type(*xty).await?)
301 }
302 TermType::Record { rty } => {
303 let mut record_type = Vec::with_capacity(rty.len());
304 for (k, v) in rty.iter() {
305 record_type.push((k.clone(), self.encode_type(v).await?));
306 }
307 self.declare_record_type(record_type.iter()).await?
308 }
309 };
310 self.types.insert(ty.clone(), enc.clone());
311 Ok(enc)
312 }
313 }
314 }
315
316 pub async fn declare_var(&mut self, v: &TermVar, ty_enc: &str) -> Result<SmolStr> {
317 let id = term_id(self.terms.len());
318 self.script.comment(&format_smolstr!("{:?}", v.id)).await?;
319 self.script.declare_const(&id, ty_enc).await?;
320 Ok(id)
321 }
322
323 pub async fn define_term(&mut self, ty_enc: &str, t_enc: &str) -> Result<SmolStr> {
324 let id = term_id(self.terms.len());
325 self.script.define_fun(&id, [], ty_enc, t_enc).await?;
326 Ok(id)
327 }
328
329 pub async fn define_set<'s>(
330 &mut self,
331 ty_enc: &str,
332 t_encs: impl ExactSizeIterator<Item = &'s str>,
333 ) -> Result<SmolStr> {
334 let set_term = if t_encs.len() == 0 {
335 format!("(as set.empty {ty_enc})")
336 } else {
337 format!(
338 "(set.insert {} (as set.empty {}))",
339 t_encs.format(" "),
340 ty_enc
341 )
342 };
343 self.define_term(ty_enc, &set_term).await
344 }
345
346 pub async fn define_record<'s>(
347 &mut self,
348 ty_enc: &str,
349 t_encs: impl IntoIterator<Item = &'s str>,
350 ) -> Result<SmolStr> {
351 let t_encs = t_encs.into_iter().join(" ");
352 let t_enc = if t_encs.is_empty() {
353 ty_enc
354 } else {
355 &format_smolstr!("({ty_enc} {})", t_encs)
356 };
357 self.define_term(ty_enc, t_enc).await
358 }
359
360 pub async fn encode_uuf(&mut self, uuf: &Uuf) -> Result<SmolStr> {
361 match self.uufs.get(uuf) {
362 Some(enc) => Ok(enc.clone()),
363 None => {
364 let id = uuf_id(self.uufs.len());
365 self.script.comment(&uuf.id).await?;
366 let encoded_arg_type = self.encode_type(&uuf.arg).await?;
367 let encoded_out_type = self.encode_type(&uuf.out).await?;
368 self.script
369 .declare_fun(&id, [encoded_arg_type.as_str()], &encoded_out_type)
370 .await?;
371 self.uufs.insert(uuf.clone(), id.clone());
372 Ok(id)
373 }
374 }
375 }
376
377 pub async fn define_entity(&mut self, ty_enc: &str, entity: &EntityUID) -> Result<SmolStr> {
378 match self.enums.get(entity.type_name()) {
379 Some(members) => {
380 let entity_ind = match members
381 .iter()
382 .position(|s| s == <EntityID as AsRef<str>>::as_ref(entity.id()))
383 {
384 Some(ind) => ind,
385 None => return Err(EncodeError::EnumMissingMember(entity.clone())),
386 };
387 Ok(enum_id(ty_enc, entity_ind))
388 }
389 None => {
390 self.define_term(
391 ty_enc,
392 &format_smolstr!(
393 "({ty_enc} \"{}\")",
394 encode_string(<EntityID as AsRef<str>>::as_ref(entity.id())).ok_or_else(
395 || EncodeError::EncodeStringFailed(format_smolstr!(
396 "{:?}",
397 entity.id()
398 ))
399 )?
400 ),
401 )
402 .await
403 }
404 }
405 }
406
407 fn index_of_attr(a: &Attr, t_ty: &TermType) -> Result<usize> {
408 match t_ty {
411 TermType::Record { rty } => match rty.keys().position(|k| k == a) {
412 Some(ind) => Ok(ind),
413 None => Err(EncodeError::RecordMissingAttr(a.clone())),
414 },
415 _ => Err(EncodeError::ExpectRecord(t_ty.clone())),
416 }
417 }
418
419 pub async fn define_record_get(
420 &mut self,
421 ty_enc: &str,
422 a: &Attr,
423 t_enc: &str,
424 ty: &TermType,
425 ) -> Result<SmolStr> {
426 let r_id = match self.types.get(ty) {
427 Some(t) => t,
428 None => return Err(EncodeError::MissingTypeEncoding(ty.clone())),
429 };
430 let a_id = Self::index_of_attr(a, ty)?;
431 self.define_term(
432 ty_enc,
433 &format_smolstr!("({} {t_enc})", record_attr_id(r_id, a_id)),
434 )
435 .await
436 }
437
438 pub async fn define_app<'b>(
439 &mut self,
440 ty_enc: &str,
441 op: &Op,
442 t_encs: impl IntoIterator<Item = SmolStr>,
443 ts: impl IntoIterator<Item = &'b Term>,
444 ) -> Result<SmolStr> {
445 let args = t_encs.into_iter().join(" ");
446 match op {
447 Op::RecordGet(a) => {
448 let ty = match ts.into_iter().next() {
449 Some(t) => t.type_of(),
450 None => return Err(EncodeError::MalformedRecordGet),
451 };
452 self.define_record_get(ty_enc, a, &args, &ty).await
453 }
454 Op::StringLike(p) => {
455 self.define_term(
456 ty_enc,
457 &format_smolstr!(
458 "(str.in_re {args} {})",
459 encode_pattern(p)
460 .ok_or_else(|| EncodeError::EncodePatternFailed(p.clone()))?
461 ),
462 )
463 .await
464 }
465 Op::Uuf(f) => {
466 let encoded_uuf = self.encode_uuf(f).await?;
467 self.define_term(ty_enc, &format_smolstr!("({} {args})", encoded_uuf))
468 .await
469 }
470 _ => {
471 self.define_term(ty_enc, &format_smolstr!("({} {args})", encode_op(op)))
472 .await
473 }
474 }
475 }
476
477 #[async_recursion]
478 pub async fn encode_term(&mut self, t: &Term) -> Result<SmolStr> {
479 if let Some(enc) = self.terms.get(t) {
480 return Ok(enc.clone());
481 }
482 let ty_enc = self.encode_type(&t.type_of()).await?;
483 let enc = match &t {
484 Term::Var(v) => self.declare_var(v, &ty_enc).await?,
485 Term::Prim(p) => match p {
486 TermPrim::Bool(b) => {
487 return Ok({
488 if *b {
489 SmolStr::new_static("true")
490 } else {
491 SmolStr::new_static("false")
492 }
493 });
494 }
495 TermPrim::Bitvec(bv) => {
496 return Ok(encode_bitvec(bv));
497 }
498 TermPrim::String(s) => {
499 return Ok(format_smolstr!(
500 "\"{}\"",
501 encode_string(s)
502 .ok_or_else(|| EncodeError::EncodeStringFailed(s.clone()))?
503 ));
504 }
505 TermPrim::Entity(e) => self.define_entity(&ty_enc, e).await?,
506 TermPrim::Ext(x) => self.define_term(&ty_enc, &encode_ext(x)).await?,
507 },
508 Term::None(_) => {
509 self.define_term(&ty_enc, &format_smolstr!("(as none {ty_enc})"))
510 .await?
511 }
512 Term::Some(t1) => {
513 let encoded_term = self.encode_term(t1).await?;
514 self.define_term(&ty_enc, &format_smolstr!("(some {encoded_term})"))
515 .await?
516 }
517 Term::Set { elts, .. } => {
518 let mut encoded_terms = Vec::with_capacity(elts.len());
519 for elt in elts.iter() {
520 encoded_terms.push(self.encode_term(elt).await?);
521 }
522 self.define_set(&ty_enc, encoded_terms.iter().map(|s| s.as_str()))
523 .await?
524 }
525 Term::Record(ats) => {
526 let mut encoded_terms = Vec::with_capacity(ats.len());
527 for t in ats.values() {
528 encoded_terms.push(self.encode_term(t).await?);
529 }
530 self.define_record(&ty_enc, encoded_terms.iter().map(|s| s.as_str()))
531 .await?
532 }
533 Term::App {
534 op: Op::Bvnego,
535 args,
536 ret_ty: TermType::Bool,
537 } if args.len() == 1 => {
538 #[expect(
539 clippy::indexing_slicing,
540 reason = "Slice of length 1 can be indexed by 0"
541 )]
542 let t = &args[0]; match t.type_of() {
549 TermType::Bitvec { n } => {
550 let t_enc = self.encode_term(t).await?;
553 self.define_app(
554 &ty_enc,
555 &Op::Eq,
556 [t_enc, encode_bitvec(&BitVec::int_min(n))],
557 [t, &BitVec::int_min(n).into()],
558 )
559 .await?
560 }
561 _ => {
562 debug_assert!(false, "`Bvnego` should only be applied to `Bitvec`");
563 SmolStr::new_static("false")
566 }
567 }
568 }
569 Term::App { op, args, .. } => {
570 let mut encoded_terms = Vec::with_capacity(args.len());
571 for arg in args.iter() {
572 encoded_terms.push(self.encode_term(arg).await?);
573 }
574 self.define_app(&ty_enc, op, encoded_terms, args.iter())
575 .await?
576 }
577 };
578 self.terms.insert(t.clone(), enc.clone());
579 Ok(enc)
580 }
581
582 pub async fn encode(&mut self, ts: impl ExactSizeIterator<Item = &Term>) -> Result<()> {
593 self.script
594 .declare_datatype("Option", ["X"], ["(none)", "(some (val X))"])
595 .await?;
596 let mut ids: Vec<_> = Vec::with_capacity(ts.len());
597 for t in ts {
598 let id = self.encode_term(t).await?;
599 ids.push(id);
600 }
601 for id in ids {
602 self.script.assert(&id).await?;
603 }
604 Ok(())
605 }
606}
607
608pub const SMT_LIB_MAX_CODE_POINT: u32 = 196607;
612
613pub(super) fn encode_string(s: &str) -> Option<String> {
627 let mut out = String::with_capacity(s.len());
628 for c in s.chars() {
629 if c == '"' {
630 out.push_str("\"\"");
631 } else if c == '\\' {
632 out.push_str("\\u{5c}");
634 } else if 32 as char <= c && c <= 126 as char {
635 out.push(c);
636 } else {
637 if c as u32 > SMT_LIB_MAX_CODE_POINT {
639 return None; }
641 #[expect(clippy::unwrap_used, reason = "writing string cannot fail")]
642 write!(out, "\\u{{{:x}}}", c as u32).unwrap();
643 }
644 }
645 Some(out)
646}
647
648fn encode_bitvec(bv: &BitVec) -> SmolStr {
649 format_smolstr!("(_ bv{} {})", bv.as_nat(), bv.width())
650}
651
652fn encode_ipaddr_prefix_v4(pre: &IPv4Prefix) -> SmolStr {
653 match pre.as_bitvec() {
654 Some(pre) => format_smolstr!("(some {})", encode_bitvec(pre)),
655 None => format_smolstr!("(as none (Option (_ BitVec {V4_WIDTH})))"),
656 }
657}
658
659fn encode_ipaddr_prefix_v6(pre: &IPv6Prefix) -> SmolStr {
660 match pre.as_bitvec() {
661 Some(pre) => format_smolstr!("(some {})", encode_bitvec(pre)),
662 None => format_smolstr!("(as none (Option (_ BitVec {V6_WIDTH})))"),
663 }
664}
665
666fn encode_ext(e: &Ext) -> SmolStr {
667 match e {
668 Ext::Decimal { d } => {
669 let bv_enc = encode_bitvec(&BitVec::of_int(SIXTY_FOUR, d.0.into()));
670 format_smolstr!("(Decimal {bv_enc})")
671 }
672 Ext::Ipaddr {
673 ip: IPNet::V4(CIDRv4 { addr, prefix }),
674 } => {
675 let addr = encode_bitvec(addr.as_bitvec());
676 let pre = encode_ipaddr_prefix_v4(prefix);
677 format_smolstr!("(V4 {addr} {pre})")
678 }
679 Ext::Ipaddr {
680 ip: IPNet::V6(CIDRv6 { addr, prefix }),
681 } => {
682 let addr = encode_bitvec(addr.as_bitvec());
683 let pre = encode_ipaddr_prefix_v6(prefix);
684 format_smolstr!("(V6 {addr} {pre})")
685 }
686 Ext::Duration { d } => {
687 let bv_enc = encode_bitvec(&BitVec::of_int(SIXTY_FOUR, d.to_milliseconds().into()));
688 format_smolstr!("(Duration {bv_enc})")
689 }
690 Ext::Datetime { dt } => {
691 let bv_enc = encode_bitvec(&BitVec::of_i128(SIXTY_FOUR, i64::from(dt).into()));
692 format_smolstr!("(Datetime {bv_enc})")
693 }
694 }
695}
696
697fn encode_ext_op(ext_op: &ExtOp) -> &'static str {
698 match ext_op {
699 ExtOp::DecimalVal => "decimalVal",
700 ExtOp::IpaddrIsV4 => "(_ is V4)",
701 ExtOp::IpaddrAddrV4 => "addrV4",
702 ExtOp::IpaddrPrefixV4 => "prefixV4",
703 ExtOp::IpaddrAddrV6 => "addrV6",
704 ExtOp::IpaddrPrefixV6 => "prefixV6",
705 ExtOp::DatetimeVal => "datetimeVal",
706 ExtOp::DatetimeOfBitVec => "Datetime",
707 ExtOp::DurationVal => "durationVal",
708 ExtOp::DurationOfBitVec => "Duration",
709 }
710}
711
712fn encode_op(op: &Op) -> SmolStr {
713 match op {
714 Op::Eq => SmolStr::new_static("="),
715 Op::ZeroExtend(n) => format_smolstr!("(_ zero_extend {n})"),
716 Op::OptionGet => SmolStr::new_static("val"),
717 Op::Ext(xop) => SmolStr::new_static(encode_ext_op(xop)),
718 _ => SmolStr::new_static(op.mk_name()),
719 }
720}
721
722fn encode_pat_elem(pat_elem: PatternElem) -> Option<SmolStr> {
723 Some(match pat_elem {
724 PatternElem::Wildcard => SmolStr::new_static("(re.* re.allchar)"),
725 PatternElem::Char(c) => {
726 format_smolstr!("(str.to_re \"{}\")", encode_string(&c.to_smolstr())?)
727 }
728 })
729}
730
731fn encode_pattern(pattern: &OrdPattern) -> Option<SmolStr> {
732 if pattern.get_elems().is_empty() {
733 Some(SmolStr::new_static("(str.to_re \"\")"))
734 } else if pattern.get_elems().len() == 1 {
735 #[expect(
736 clippy::indexing_slicing,
737 reason = "Slice of length 1 can be indexed by 0"
738 )]
739 encode_pat_elem(pattern.get_elems()[0])
740 } else {
741 Some(format_smolstr!(
742 "(re.++ {})",
743 pattern
744 .iter()
745 .copied()
746 .map(encode_pat_elem)
747 .collect::<Option<Vec<_>>>()?
748 .into_iter()
749 .join(" ")
750 ))
751 }
752}
753
754#[cfg(test)]
755mod unit_tests {
756 use std::{collections::BTreeSet, str::FromStr};
757
758 use crate::symcc::env::{SymEntities, SymEnv, SymRequest};
759 use cedar_policy::EntityTypeName;
760 use smol_str::SmolStr;
761
762 use super::Encoder;
763 use crate::symcc::term_type::TermType;
764 use std::collections::BTreeMap;
765 use std::sync::Arc;
766
767 #[tokio::test]
768 async fn declare_type() {
769 let symenv = SymEnv {
770 request: SymRequest::empty_sym_req(),
771 entities: Arc::new(SymEntities(BTreeMap::new())),
772 };
773 let mut encoder = Encoder::new(&symenv, Vec::<u8>::new()).unwrap();
774 encoder
775 .declare_type("foo", ["(Bar1 (baz String))"])
776 .await
777 .unwrap();
778 }
779
780 #[tokio::test]
781 async fn declare_entity_type() {
782 let symenv = SymEnv {
783 request: SymRequest::empty_sym_req(),
784 entities: Arc::new(SymEntities(BTreeMap::new())),
785 };
786 let mut encoder = Encoder::new(&symenv, Vec::<u8>::new()).unwrap();
787 let ety = cedar_policy::EntityTypeName::from_str("User").unwrap();
788 let empty_set = BTreeSet::new();
789 encoder.enums.insert(&ety, &empty_set);
790 encoder.declare_entity_type(&ety).await.unwrap();
791 }
792
793 #[tokio::test]
794 async fn declare_empty_record_type() {
795 let symenv = SymEnv {
796 request: SymRequest::empty_sym_req(),
797 entities: Arc::new(SymEntities(BTreeMap::new())),
798 };
799 let mut encoder = Encoder::new(&symenv, Vec::<u8>::new()).unwrap();
800 encoder.declare_record_type(vec![]).await.unwrap();
801 }
802
803 #[tokio::test]
804 async fn declare_record_type() {
805 let symenv = SymEnv {
806 request: SymRequest::empty_sym_req(),
807 entities: Arc::new(SymEntities(BTreeMap::new())),
808 };
809 let mut encoder = Encoder::new(&symenv, Vec::<u8>::new()).unwrap();
810 encoder
811 .declare_record_type(std::iter::once(&("foo".into(), SmolStr::new_static("bar"))))
812 .await
813 .unwrap();
814 }
815
816 #[tokio::test]
817 async fn encode_bool_type() {
818 let symenv = SymEnv {
819 request: SymRequest::empty_sym_req(),
820 entities: Arc::new(SymEntities(BTreeMap::new())),
821 };
822 let mut encoder = Encoder::new(&symenv, Vec::<u8>::new()).unwrap();
823 encoder.encode_type(&TermType::Bool).await.unwrap();
824 }
825
826 #[tokio::test]
827 async fn encode_string_type() {
828 let symenv = SymEnv {
829 request: SymRequest::empty_sym_req(),
830 entities: Arc::new(SymEntities(BTreeMap::new())),
831 };
832 let mut encoder = Encoder::new(&symenv, Vec::<u8>::new()).unwrap();
833 encoder.encode_type(&TermType::String).await.unwrap();
834 }
835
836 #[tokio::test]
837 async fn encode_uuf() {
838 let symenv = SymEnv {
839 request: SymRequest::empty_sym_req(),
840 entities: Arc::new(SymEntities(BTreeMap::new())),
841 };
842 let mut encoder = Encoder::new(&symenv, Vec::<u8>::new()).unwrap();
843 let my_uuf = crate::symcc::op::Uuf {
844 id: "my_fun".into(),
845 arg: TermType::Bool,
846 out: TermType::Bool,
847 };
848 encoder.encode_uuf(&my_uuf).await.unwrap();
849 }
850
851 #[tokio::test]
852 async fn define_entity() {
853 use cedar_policy::EntityUid;
854 let symenv = SymEnv {
855 request: SymRequest::empty_sym_req(),
856 entities: Arc::new(SymEntities(BTreeMap::new())),
857 };
858 let mut encoder = Encoder::new(&symenv, Vec::<u8>::new()).unwrap();
859 let entity_type_name = EntityTypeName::from_str("User").unwrap();
860 let entity = EntityUid::from_type_name_and_id(
861 entity_type_name.clone(),
862 cedar_policy::EntityId::from_str("alice").unwrap(),
863 );
864 let entity_ty_enc = encoder
865 .encode_type(&TermType::Entity {
866 ety: entity_type_name,
867 })
868 .await
869 .unwrap();
870 encoder
871 .define_entity(&entity_ty_enc, &entity)
872 .await
873 .unwrap();
874 }
875
876 async fn compile_and_encode(expr: &str) -> String {
880 use crate::symcc::compiler::{
881 compile,
882 ext_has_attr_tests::{parse_expr, sym_env},
883 };
884
885 let symenv = sym_env();
886 let term = compile(&parse_expr(expr), &symenv).unwrap();
887
888 let mut encoder = Encoder::new(&symenv, Vec::<u8>::new()).unwrap();
889 encoder.encode_term(&term).await.unwrap();
890
891 String::from_utf8(encoder.script).unwrap()
892 }
893
894 #[tokio::test]
895 async fn ext_has_attr_compiles_to_expected_smt() {
896 insta::assert_snapshot!(compile_and_encode("context has rec.x").await, @"(define-fun t0 () (Option Bool) (some true))");
897 }
898
899 #[tokio::test]
901 async fn ext_has_attr_entity_optional_then_present_smt() {
902 insta::assert_snapshot!(compile_and_encode("principal has thing1.id").await, @r#"
903 ; Thing
904 (declare-datatype E0 (
905 (E0 (E0_eid String))))
906 ; Thing2
907 (declare-datatype E1 (
908 (E1 (E1_eid String))))
909 ; {id, thing2, thing2bis}
910 (declare-datatype R2 (
911 (R2 (R2_a0 String) (R2_a1 E1) (R2_a2 (Option E1)))))
912 ; {name, thing1, thing2, x, xopt}
913 (declare-datatype R3 (
914 (R3 (R3_a0 String) (R3_a1 (Option E0)) (R3_a2 E1) (R3_a3 R2) (R3_a4 (Option R2)))))
915 ; User
916 (declare-datatype E4 (
917 (E4 (E4_eid String))))
918 ; "principal"
919 (declare-const t0 E4)
920 ; attrs[User]
921 (declare-fun f0 (E4) R3)
922 (define-fun t1 () R3 (f0 t0))
923 (define-fun t2 () (Option E0) (R3_a1 t1))
924 (define-fun t3 () (Option E0) (as none (Option E0)))
925 (define-fun t4 () Bool (= t2 t3))
926 (define-fun t5 () Bool (not t4))
927 (define-fun t6 () (Option Bool) (as none (Option Bool)))
928 (define-fun t7 () (Option Bool) (some false))
929 (define-fun t8 () (Option Bool) (ite t4 t6 t7))
930 (define-fun t9 () (Option Bool) (ite t5 t8 t7))
931 "#);
932 }
933
934 #[tokio::test]
936 async fn ext_has_attr_entity_present_then_optional_smt() {
937 insta::assert_snapshot!(compile_and_encode("principal has thing2.opt").await, @r#"
938 ; {id, opt}
939 (declare-datatype R0 (
940 (R0 (R0_a0 String) (R0_a1 (Option (_ BitVec 64))))))
941 ; Thing2
942 (declare-datatype E1 (
943 (E1 (E1_eid String))))
944 ; Thing
945 (declare-datatype E2 (
946 (E2 (E2_eid String))))
947 ; {id, thing2, thing2bis}
948 (declare-datatype R3 (
949 (R3 (R3_a0 String) (R3_a1 E1) (R3_a2 (Option E1)))))
950 ; {name, thing1, thing2, x, xopt}
951 (declare-datatype R4 (
952 (R4 (R4_a0 String) (R4_a1 (Option E2)) (R4_a2 E1) (R4_a3 R3) (R4_a4 (Option R3)))))
953 ; User
954 (declare-datatype E5 (
955 (E5 (E5_eid String))))
956 ; "principal"
957 (declare-const t0 E5)
958 ; attrs[User]
959 (declare-fun f0 (E5) R4)
960 (define-fun t1 () R4 (f0 t0))
961 (define-fun t2 () E1 (R4_a2 t1))
962 ; attrs[Thing2]
963 (declare-fun f1 (E1) R0)
964 (define-fun t3 () R0 (f1 t2))
965 (define-fun t4 () (Option (_ BitVec 64)) (R0_a1 t3))
966 (define-fun t5 () (Option (_ BitVec 64)) (as none (Option (_ BitVec 64))))
967 (define-fun t6 () Bool (= t4 t5))
968 (define-fun t7 () Bool (not t6))
969 (define-fun t8 () (Option Bool) (some t7))
970 "#);
971 }
972
973 #[tokio::test]
975 async fn ext_has_attr_record_present_then_optional_smt() {
976 insta::assert_snapshot!(compile_and_encode("principal.x has thing2.opt").await, @r#"
977 ; {id, opt}
978 (declare-datatype R0 (
979 (R0 (R0_a0 String) (R0_a1 (Option (_ BitVec 64))))))
980 ; Thing2
981 (declare-datatype E1 (
982 (E1 (E1_eid String))))
983 ; {id, thing2, thing2bis}
984 (declare-datatype R2 (
985 (R2 (R2_a0 String) (R2_a1 E1) (R2_a2 (Option E1)))))
986 ; Thing
987 (declare-datatype E3 (
988 (E3 (E3_eid String))))
989 ; {name, thing1, thing2, x, xopt}
990 (declare-datatype R4 (
991 (R4 (R4_a0 String) (R4_a1 (Option E3)) (R4_a2 E1) (R4_a3 R2) (R4_a4 (Option R2)))))
992 ; User
993 (declare-datatype E5 (
994 (E5 (E5_eid String))))
995 ; "principal"
996 (declare-const t0 E5)
997 ; attrs[User]
998 (declare-fun f0 (E5) R4)
999 (define-fun t1 () R4 (f0 t0))
1000 (define-fun t2 () R2 (R4_a3 t1))
1001 (define-fun t3 () E1 (R2_a1 t2))
1002 ; attrs[Thing2]
1003 (declare-fun f1 (E1) R0)
1004 (define-fun t4 () R0 (f1 t3))
1005 (define-fun t5 () (Option (_ BitVec 64)) (R0_a1 t4))
1006 (define-fun t6 () (Option (_ BitVec 64)) (as none (Option (_ BitVec 64))))
1007 (define-fun t7 () Bool (= t5 t6))
1008 (define-fun t8 () Bool (not t7))
1009 (define-fun t9 () (Option Bool) (some t8))
1010 "#);
1011 }
1012
1013 #[tokio::test]
1015 async fn ext_has_attr_record_optional_then_present_smt() {
1016 insta::assert_snapshot!(compile_and_encode("principal.x has thing2bis.id").await, @r#"
1017 ; Thing2
1018 (declare-datatype E0 (
1019 (E0 (E0_eid String))))
1020 ; {id, thing2, thing2bis}
1021 (declare-datatype R1 (
1022 (R1 (R1_a0 String) (R1_a1 E0) (R1_a2 (Option E0)))))
1023 ; Thing
1024 (declare-datatype E2 (
1025 (E2 (E2_eid String))))
1026 ; {name, thing1, thing2, x, xopt}
1027 (declare-datatype R3 (
1028 (R3 (R3_a0 String) (R3_a1 (Option E2)) (R3_a2 E0) (R3_a3 R1) (R3_a4 (Option R1)))))
1029 ; User
1030 (declare-datatype E4 (
1031 (E4 (E4_eid String))))
1032 ; "principal"
1033 (declare-const t0 E4)
1034 ; attrs[User]
1035 (declare-fun f0 (E4) R3)
1036 (define-fun t1 () R3 (f0 t0))
1037 (define-fun t2 () R1 (R3_a3 t1))
1038 (define-fun t3 () (Option E0) (R1_a2 t2))
1039 (define-fun t4 () (Option E0) (as none (Option E0)))
1040 (define-fun t5 () Bool (= t3 t4))
1041 (define-fun t6 () Bool (not t5))
1042 (define-fun t7 () (Option Bool) (as none (Option Bool)))
1043 (define-fun t8 () (Option Bool) (some true))
1044 (define-fun t9 () (Option Bool) (ite t5 t7 t8))
1045 (define-fun t10 () (Option Bool) (some false))
1046 (define-fun t11 () (Option Bool) (ite t6 t9 t10))
1047 "#);
1048 }
1049}
1050
1051#[cfg(test)]
1052mod deep_extended_has_chain_tests {
1053 use crate::symcc::compiler::compile;
1054 use crate::symcc::test_utils::{deep_chain_sym_env, deep_has_chain_expr};
1055
1056 use super::Encoder;
1057
1058 async fn compile_encde_at_depth(depth: usize) -> String {
1059 let symenv = deep_chain_sym_env(depth);
1060 let term =
1061 compile(&deep_has_chain_expr(depth), &symenv).expect("expression should compile");
1062 let mut encoder = Encoder::new(&symenv, Vec::<u8>::new()).unwrap();
1063 encoder.encode_term(&term).await.unwrap();
1064 String::from_utf8(encoder.script).unwrap()
1065 }
1066
1067 #[tokio::test]
1068 async fn nested_has_chain_encodes_linearly_not_exponentially() {
1069 let smt_at_2 = compile_encde_at_depth(2).await;
1070 let smt_at_3 = compile_encde_at_depth(3).await;
1071 let smt_at_4 = compile_encde_at_depth(4).await;
1072 let n_2 = smt_at_2.matches("define-fun").count();
1073 let n_3 = smt_at_3.matches("define-fun").count();
1074 let n_4 = smt_at_4.matches("define-fun").count();
1075 assert_eq!(n_3 - n_2, n_4 - n_3);
1077 }
1078}