Skip to main content

cedar_policy_symcc/symcc/
term.rs

1/*
2 * Copyright Cedar Contributors
3 *
4 * Licensed under the Apache License, Version 2.0 (the "License");
5 * you may not use this file except in compliance with the License.
6 * You may obtain a copy of the License at
7 *
8 *      https://www.apache.org/licenses/LICENSE-2.0
9 *
10 * Unless required by applicable law or agreed to in writing, software
11 * distributed under the License is distributed on an "AS IS" BASIS,
12 * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13 * See the License for the specific language governing permissions and
14 * limitations under the License.
15 */
16
17//! A simply typed IR to which we reduce Cedar expressions during symbolic compilation.
18//!
19//! The Term language has a straightforward translation to SMTLib. It is designed to
20//! reduce the semantic gap between Cedar and SMTLib, and to facilitate proofs of
21//! soundness and completeness of the Cedar symbolic compiler.
22//!
23//! Terms should _not_ be created directly using `Term` constructors. Instead, they
24//! should be created using the factory functions defined in `factory.rs`.
25//! The factory functions check the types of their arguments, perform optimizations,
26//! and ensure that applying them to well-formed terms results in well-formed terms.
27//!
28//! See `term_type.rs` and `op.rs` for definitions of Term types and operators.
29
30use smol_str::{format_smolstr, SmolStr};
31
32use super::bitvec::BitVec;
33use super::ext::Ext;
34use super::op::Op;
35use super::term_type::TermType;
36use super::type_abbrevs::*;
37use std::{
38    collections::{BTreeMap, BTreeSet},
39    ops::Deref,
40    sync::Arc,
41};
42
43/// A typed variable.
44#[derive(Clone, Debug, PartialEq, Eq, Ord, PartialOrd)]
45pub struct TermVar {
46    /// A unique identifier of the variable.
47    pub id: SmolStr,
48    /// Type of the variable.
49    pub ty: TermType,
50}
51
52/// Primitive terms.
53/// Variants must be defined in alphabetical order.
54#[derive(Clone, Debug, PartialEq, Eq, Ord, PartialOrd)]
55pub enum TermPrim {
56    /// Literal bitvec
57    Bitvec(BitVec),
58    /// Literal bool
59    Bool(bool),
60    /// Literal EntityUID
61    Entity(EntityUID),
62    /// Literal extension value
63    Ext(Ext),
64    /// Literal string
65    String(SmolStr),
66}
67
68/// Intermediate representation of [`Term`]s.
69/// Variants must be defined in alphabetical order.
70#[derive(Clone, Debug, PartialEq, Eq, Ord, PartialOrd)]
71pub enum Term {
72    /// Function calls
73    App {
74        /// Function being called
75        op: Op,
76        /// Arguments
77        args: Arc<Vec<Term>>,
78        /// Return type of the function
79        ret_ty: TermType,
80    },
81    /// None
82    None(TermType),
83    /// Literal
84    Prim(TermPrim),
85    /// Records
86    Record(Arc<BTreeMap<Attr, Term>>),
87    /// Sets
88    Set {
89        /// Elements of the set (as `Term`)
90        elts: Arc<BTreeSet<Term>>,
91        /// Type shared by all elements of the set
92        elts_ty: TermType,
93    },
94    /// Some
95    Some(Arc<Term>),
96    /// Variable
97    Var(TermVar),
98}
99
100// Corresponding to the `Coe` instances in Lean
101impl From<bool> for Term {
102    fn from(b: bool) -> Self {
103        if b {
104            Term::Prim(TermPrim::Bool(true))
105        } else {
106            Term::Prim(TermPrim::Bool(false))
107        }
108    }
109}
110
111impl From<i64> for Term {
112    fn from(i: i64) -> Self {
113        Term::Prim(TermPrim::Bitvec(BitVec::of_int(SIXTY_FOUR, i.into())))
114    }
115}
116
117impl From<BitVec> for Term {
118    fn from(bv: BitVec) -> Self {
119        Term::Prim(TermPrim::Bitvec(bv))
120    }
121}
122
123impl From<SmolStr> for Term {
124    fn from(s: SmolStr) -> Self {
125        Term::Prim(TermPrim::String(s))
126    }
127}
128
129impl From<EntityUID> for Term {
130    fn from(uid: EntityUID) -> Self {
131        Term::Prim(TermPrim::Entity(uid))
132    }
133}
134
135impl From<Ext> for Term {
136    fn from(ext: Ext) -> Self {
137        Term::Prim(TermPrim::Ext(ext))
138    }
139}
140
141impl From<TermVar> for Term {
142    fn from(v: TermVar) -> Self {
143        Term::Var(v)
144    }
145}
146
147impl TermPrim {
148    /// Returns the type of the primitive term.
149    pub fn type_of(&self) -> TermType {
150        match self {
151            TermPrim::Bool(_) => TermType::Bool,
152            TermPrim::Bitvec(v) => TermType::Bitvec { n: v.width() },
153            TermPrim::String(_) => TermType::String,
154            TermPrim::Entity(e) => TermType::Entity {
155                ety: e.type_name().clone(),
156            },
157            TermPrim::Ext(Ext::Decimal { .. }) => TermType::Ext {
158                xty: ExtType::Decimal,
159            },
160            TermPrim::Ext(Ext::Ipaddr { .. }) => TermType::Ext {
161                xty: ExtType::IpAddr,
162            },
163            TermPrim::Ext(Ext::Duration { .. }) => TermType::Ext {
164                xty: ExtType::Duration,
165            },
166            TermPrim::Ext(Ext::Datetime { .. }) => TermType::Ext {
167                xty: ExtType::DateTime,
168            },
169        }
170    }
171}
172
173impl Term {
174    /// Computes the type of a term.
175    pub fn type_of(&self) -> TermType {
176        match self {
177            Term::Prim(l) => l.type_of(),
178            Term::Var(v) => v.ty.clone(),
179            Term::None(ty) => TermType::option_of(ty.clone()),
180            Term::Some(t) => TermType::option_of(t.type_of()),
181            Term::Set { elts_ty, .. } => TermType::set_of(elts_ty.clone()),
182            Term::Record(m) => {
183                let rty = Arc::new(m.iter().map(|(k, v)| (k.clone(), v.type_of())).collect());
184                TermType::Record { rty }
185            }
186            Term::App { ret_ty, .. } => ret_ty.clone(),
187        }
188    }
189
190    /// Checks if the term is a literal, i.e., contains no variables or applications.
191    pub fn is_literal(&self) -> bool {
192        match self {
193            Term::Prim(_) => true,
194            Term::None(_) => true,
195            Term::Some(t) => t.is_literal(),
196            Term::Set { elts, .. } => elts.iter().all(Term::is_literal),
197            Term::Record(m) => m.values().all(Term::is_literal),
198            _ => false,
199        }
200    }
201}
202
203impl std::fmt::Display for Term {
204    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
205        match self {
206            Term::Prim(prim) => write!(f, "{prim}"),
207            Term::Var(var) => write!(f, "{}", var.id),
208            Term::None(_) => write!(f, "None"),
209            Term::Some(t) => write!(f, "Some({t})"),
210            Term::Set { elts, .. } => {
211                write!(f, "[")?;
212                let mut first = true;
213                for elt in elts.iter() {
214                    if !first {
215                        write!(f, ", ")?;
216                    }
217                    write!(f, "{elt}")?;
218                    first = false;
219                }
220                write!(f, "]")
221            }
222            Term::Record(map) => {
223                if map.is_empty() {
224                    write!(f, "{{}}")
225                } else {
226                    write!(f, "{{ ")?;
227                    let mut first = true;
228                    for (k, v) in map.iter() {
229                        if !first {
230                            write!(f, ", ")?;
231                        }
232                        write!(f, "{k}: {v}")?;
233                        first = false;
234                    }
235                    write!(f, " }}")
236                }
237            }
238            Term::App { op, args, .. } => {
239                write!(
240                    f,
241                    "{op}(",
242                    op = match op {
243                        Op::Ext(ext) => SmolStr::new(ext.mk_name()),
244                        Op::Uuf(uuf) => uuf.id.clone(),
245                        Op::RecordGet(attr) => format_smolstr!("getattr[\"{attr}\"]"),
246                        Op::StringLike(pat) =>
247                            format_smolstr!("like[\"{pat}\"]", pat = pat.deref()),
248                        _ => SmolStr::new(op.mk_name()),
249                    }
250                )?;
251                let mut first = true;
252                for arg in args.iter() {
253                    if !first {
254                        write!(f, ", ")?;
255                    }
256                    write!(f, "{arg}")?;
257                    first = false;
258                }
259                write!(f, ")")
260            }
261        }
262    }
263}
264
265impl std::fmt::Display for TermPrim {
266    #[expect(
267        clippy::unwrap_used,
268        reason = "for now, allowing panics in this Display impl intended for debugging"
269    )]
270    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
271        match self {
272            TermPrim::Bool(b) => write!(f, "{b}"),
273            TermPrim::Bitvec(bv) => write!(f, "{bv}"),
274            TermPrim::String(s) => write!(f, "\"{s}\""),
275            TermPrim::Entity(e) => write!(f, "{e}"),
276            TermPrim::Ext(ext) => write!(
277                f,
278                "{}",
279                cedar_policy_core::ast::Value::try_from(ext).unwrap()
280            ),
281        }
282    }
283}
284
285#[cfg(test)]
286mod test {
287    use super::super::factory;
288    use super::*;
289
290    use cedar_policy::EntityTypeName;
291    use std::str::FromStr;
292
293    #[test]
294    fn term_display() {
295        let term = Term::from(false);
296        insta::with_settings!({ description => format!("{term:?}") }, {
297            insta::assert_snapshot!(term.to_string(), @"false");
298        });
299
300        let term = Term::from(334);
301        insta::with_settings!({ description => format!("{term:?}") }, {
302            insta::assert_snapshot!(term.to_string(), @"(bv64 334)");
303        });
304
305        let term = Term::from(SmolStr::new_static("hello I am a string"));
306        insta::with_settings!({ description => format!("{term:?}") }, {
307            insta::assert_snapshot!(term.to_string(), @r#""hello I am a string""#);
308        });
309
310        let term = Term::from(EntityUID::from_str("App::Domain::\"Component\"").unwrap());
311        insta::with_settings!({ description => format!("{term:?}") }, {
312            insta::assert_snapshot!(term.to_string(), @r#"App::Domain::"Component""#);
313        });
314
315        let term = Term::from(TermVar {
316            id: SmolStr::new_static("principal"),
317            ty: TermType::Entity {
318                ety: EntityTypeName::from_str("A::B::CDEFG").unwrap(),
319            },
320        });
321        insta::with_settings!({ description => format!("{term:?}") }, {
322            insta::assert_snapshot!(term.to_string(), @"principal");
323        });
324
325        let term = Term::from(Ext::parse_decimal("-0.11").unwrap());
326        insta::with_settings!({ description => format!("{term:?}") }, {
327            insta::assert_snapshot!(term.to_string(), @r#"decimal("-0.1100")"#);
328        });
329
330        let term = Term::from(Ext::parse_decimal("34567.8901").unwrap());
331        insta::with_settings!({ description => format!("{term:?}") }, {
332            insta::assert_snapshot!(term.to_string(), @r#"decimal("34567.8901")"#);
333        });
334
335        let term = Term::from(Ext::parse_ip("192.168.0.0/24").unwrap());
336        insta::with_settings!({ description => format!("{term:?}") }, {
337            insta::assert_snapshot!(term.to_string(), @r#"ip("192.168.0.0/24")"#);
338        });
339
340        let term = Term::from(Ext::parse_ip("ffee::1").unwrap());
341        insta::with_settings!({ description => format!("{term:?}") }, {
342            insta::assert_snapshot!(term.to_string(), @r#"ip("ffee:0000:0000:0000:0000:0000:0000:0001/128")"#);
343        });
344
345        let term = Term::from(Ext::parse_duration("3m7s").unwrap());
346        insta::with_settings!({ description => format!("{term:?}") }, {
347            // TODO: this one isn't the prettiest, but could be that this
348            // representation is helpful for someone debugging at the Term
349            // level; not sure what's optimal here
350            insta::assert_snapshot!(term.to_string(), @r#"duration("187000ms")"#);
351        });
352
353        let term = Term::from(Ext::parse_duration("1d0m76s111ms").unwrap());
354        insta::with_settings!({ description => format!("{term:?}") }, {
355            insta::assert_snapshot!(term.to_string(), @r#"duration("86476111ms")"#);
356        });
357
358        let term = Term::from(Ext::parse_datetime("2001-07-07").unwrap());
359        insta::with_settings!({ description => format!("{term:?}") }, {
360            // TODO: not pretty
361            insta::assert_snapshot!(term.to_string(), @r#"datetime("1970-01-01").offset(duration("994464000000ms"))"#);
362        });
363
364        let term = Term::from(Ext::parse_datetime("2010-12-31T11:59:59Z").unwrap());
365        insta::with_settings!({ description => format!("{term:?}") }, {
366            // TODO: not pretty
367            insta::assert_snapshot!(term.to_string(), @r#"datetime("1970-01-01").offset(duration("1293796799000ms"))"#);
368        });
369
370        let term = Term::from(Ext::parse_datetime("2010-12-31T11:59:59.777Z").unwrap());
371        insta::with_settings!({ description => format!("{term:?}") }, {
372            // TODO: not pretty
373            insta::assert_snapshot!(term.to_string(), @r#"datetime("1970-01-01").offset(duration("1293796799777ms"))"#);
374        });
375
376        let term = Term::from(Ext::parse_datetime("2010-12-31T11:59:59.777+1134").unwrap());
377        insta::with_settings!({ description => format!("{term:?}") }, {
378            // TODO: not pretty
379            insta::assert_snapshot!(term.to_string(), @r#"datetime("1970-01-01").offset(duration("1293755159777ms"))"#);
380        });
381
382        let term = Term::Some(Arc::new(factory::set_of(
383            [Term::from(36), Term::from(-1240)],
384            TermType::Bitvec { n: SIXTY_FOUR },
385        )));
386        insta::with_settings!({ description => format!("{term:?}") }, {
387            insta::assert_snapshot!(term.to_string(), @"Some([(bv64 36), (bv64 18446744073709550376)])");
388        });
389
390        let term = factory::record_of([
391            ("foo".into(), Term::from(-321)),
392            ("bar".into(), Term::from(SmolStr::new_static("a string"))),
393            (
394                "weird key!".into(),
395                Term::from(Ext::parse_decimal("2.222").unwrap()),
396            ),
397        ]);
398        insta::with_settings!({ description => format!("{term:?}") }, {
399            insta::assert_snapshot!(term.to_string(), @r#"{ bar: "a string", foo: (bv64 18446744073709551295), weird key!: decimal("2.2220") }"#);
400        });
401
402        let context = Term::from(TermVar {
403            id: SmolStr::new_static("context"),
404            ty: TermType::Record {
405                rty: Arc::new(
406                    [
407                        (SmolStr::new("foo"), TermType::Bitvec { n: SIXTY_FOUR }),
408                        (SmolStr::new("abc"), TermType::Bool),
409                        (SmolStr::new("def"), TermType::Bool),
410                        (SmolStr::new("zyx"), TermType::Bool),
411                        (SmolStr::new("path"), TermType::String),
412                    ]
413                    .into_iter()
414                    .collect(),
415                ),
416            },
417        });
418        let term = factory::bvslt(
419            factory::record_get(context.clone(), &SmolStr::new("foo")),
420            Term::from(12),
421        );
422        insta::with_settings!({ description => format!("{term:?}") }, {
423            insta::assert_snapshot!(term.to_string(), @r#"bvslt(getattr["foo"](context), (bv64 12))"#);
424        });
425
426        let term = factory::and(
427            factory::or(
428                factory::record_get(context.clone(), &SmolStr::new("abc")),
429                factory::record_get(context.clone(), &SmolStr::new("def")),
430            ),
431            factory::record_get(context.clone(), &SmolStr::new("zyx")),
432        );
433        insta::with_settings!({ description => format!("{term:?}") }, {
434            insta::assert_snapshot!(term.to_string(), @r#"and(or(getattr["abc"](context), getattr["def"](context)), getattr["zyx"](context))"#);
435        });
436
437        let term = factory::ite(
438            factory::record_get(context.clone(), &SmolStr::new("abc")),
439            Term::from(777),
440            Term::from(888),
441        );
442        insta::with_settings!({ description => format!("{term:?}") }, {
443            insta::assert_snapshot!(term.to_string(), @r#"ite(getattr["abc"](context), (bv64 777), (bv64 888))"#);
444        });
445
446        let term = factory::string_like(
447            factory::record_get(context, &SmolStr::new("path")),
448            cedar_policy_core::ast::Pattern::from_iter([
449                cedar_policy_core::ast::PatternElem::Char('a'),
450                cedar_policy_core::ast::PatternElem::Wildcard,
451                cedar_policy_core::ast::PatternElem::Char('z'),
452            ])
453            .into(),
454        );
455        insta::with_settings!({ description => format!("{term:?}") }, {
456            insta::assert_snapshot!(term.to_string(), @r#"like["a*z"](getattr["path"](context))"#);
457        });
458    }
459}