Skip to main content

cedar_policy_symcc/symcc/
encoder.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//! This module defines the Cedar encoder, which translates a list of boolean Terms
18//! into a list of SMT assertions. Term encoding is trusted.
19//!
20//!  We use the following type representations for primitive types:
21//!  * `TermType.bool`:     builtin SMT `Bool` type
22//!  * `TermType.string`:   builtin SMT `String` type
23//!  * `TermType.bitvec n`: builtin SMT `(_ BitVec n)` type
24//!
25//!  We will represent non-primitive types as SMT algebraic data types:
26//!  * `TermType.option T`: a parameterized SMT algebraic datatype of the same name,
27//!    and with the constructors `(some (val T))` and `(none)`. For each constructor
28//!    argument, SMTLib introduces a corresponding (total) selector function. We
29//!    will translate `Term.some` nodes in the Term language as applications of the
30//!    `val` selector function.
31//!  * `TermType.entity E`: we represent Cedar entities of entity type E as values
32//!    of the SMT algebraic datatype E with a single constructor, `(E (E_eid String))`.
33//!    The selector is named `E_eid`, after the entity type, since SMT-LIB requires
34//!    unique selector and constructor names across all datatypes.
35//!    Each entity type E gets an uninterpreted function `f: E → Record_E` that maps
36//!    instances of E to their attributes.  Similarly, each E
37//!    gets N uninterpreted functions `g₁: E → Set E₁, ..., gₙ: E → Set Eₙ` that map
38//!    each instance of E to its ancestor sets of the given types, as specified by
39//!    the `memberOf` relation in the schema.
40//!  * `TermType.Record (Map Attr TermType)`: we represent each record term type as
41//!    an SMT algebraic datatype with a single constructor. The order of arguments
42//!    to the constructor (the attributes) is important, so we fix that to be the
43//!    lexicographic order on the attribute names of the underlying record type. We
44//!    use the argument selector functions to translate `record.get` applications.
45//!    We can't use raw Cedar attribute names for argument names because they may
46//!    not be valid SMT identifiers. So, we'll keep a mapping from the attribute
47//!    names to their unique SMT ids. In general, we'll name SMT record types as
48//!    "R`<i>`" where `<i>` is a natural number and attributes within the record as
49//!    "R`<i>`a`<j>`", where `<j>` is the attribute's position in the constructor argument
50//!    list.
51//!
52//!  Similarly to types and attributes, all uninterpreted functions, variables, and
53//!  Terms are mapped to their SMT encoding that conforms to the SMTLib syntax. We
54//!  keep track of these mappings to ensure that each Term construct is translated
55//!  to its SMT encoding exactly once.  This translation invariant is necessary for
56//!  correctness in the case of record type declarations, UF names, and variable
57//!  names; and it is necessary for compactness in the case of terms. In
58//!  particular, the resulting SMT encoding will be in A-normal form (ANF): the body
59//!  of every s-expression in the encoding consists of atomic subterms (identifiers
60//!  or literals).
61
62use 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/// Errors during encoding, i.e., converting [`Term`]
87/// to SMT-LIB 2 format.
88#[derive(Debug, Diagnostic, Error)]
89pub enum EncodeError {
90    /// IO error.
91    #[error("IO error during SMT encoding")]
92    Io(#[from] std::io::Error),
93    /// Missing member in enum entity.
94    #[error("missing member {0} in enum entity")]
95    EnumMissingMember(EntityUID),
96    /// Record missing attribute.
97    #[error("record missing attribute {0}")]
98    RecordMissingAttr(Attr),
99    /// Expecting a record type.
100    #[error("expecting a record type, got {0:?}")]
101    ExpectRecord(TermType),
102    /// Missing type encoding.
103    #[error("missing type encoding for {0:?}")]
104    MissingTypeEncoding(TermType),
105    /// Malformed record get.
106    #[error("malformed record get")]
107    MalformedRecordGet,
108    /// Unable to encode string.
109    #[error("unable to encode string \"{0}\" in SMT as it exceeds the max supported code point")]
110    EncodeStringFailed(SmolStr),
111    /// Unable to encode pattern.
112    #[error("unable to encode pattern {0:?} in SMT as it exceeds the max supported code point")]
113    EncodePatternFailed(OrdPattern),
114    /// Bit-vector error.
115    #[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
154// We don't need these
155// def typeNum : EncoderM Nat := do return (← get).types.size
156// def termNum : EncoderM Nat := do return (← get).terms.size
157// def uufNum  : EncoderM Nat := do return (← get).uufs.size
158
159impl<'a, S> Encoder<'a, S> {
160    /// Corresponds to `EncoderState.init` in Lean
161    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    /// "Finalize" the encoder, removing the ability to write to the `script`, and thus also dropping all borrows inherent in the `S` type.
176    /// The resulting encoder can have its state inspected (e.g., by the decoder), but can no longer encode anything new.
177    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    /// Returns `id` to match the Lean
190    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        // Getting the index of a key in `BTreeMap` should be ok
409        // (it wouldn't be for `HashMap`)
410        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]; // guaranteed to exist because we already checked that `args.len() == 1`
543
544                // don't encode bvnego itself, for compatibility with older CVC5 (bvnego was
545                // introduced in CVC5 1.1.2)
546                // this rewrite is done in the encoder and is thus trusted; see notes here in
547                // the Lean
548                match t.type_of() {
549                    TermType::Bitvec { n } => {
550                        // more fancy and possibly more optimized, but hard to prove termination in Lean:
551                        // self.encode_term(&factory::eq(t, &BitVec::int_min(n))).await?
552                        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                        // we could put anything here and be sound, because `Bvnego` should only be
564                        // applied to Terms of type `Bitvec`
565                        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    /// Once you've generated `Asserts` with one of the functions in verifier.rs, you
583    /// can use this function to encode them as SMTLib assertions.
584    ///
585    /// Note that `encode()` itself first resets the solver in order to define datatypes
586    /// etc.
587    ///
588    /// In Lean, this is a standalone function which takes a `SymEnv`, uses that to
589    /// construct an `Encoder` (`EncoderState` in Lean), and then does the encoding.
590    /// Here in Rust, we have this as a method on `Encoder`, so the caller first
591    /// constructs an `Encoder` themselves with the `SymEnv`, then calls this.
592    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
608/// The maximum Unicode code point supported in SMT-LIB 2.7.
609/// Also see `num_codes` in cvc5:
610/// https://github.com/cvc5/cvc5/blob/b78e7ed23348659db52a32765ad181ae0c26bbd5/src/util/string.h#L53
611pub const SMT_LIB_MAX_CODE_POINT: u32 = 196607;
612
613/// This function needs to encode unicode strings with two levels of
614/// escape sequences:
615/// - At the string theory level, we need to encode all non-printable
616///   unicode characters as `\u{xxxx}`, where a character is printable
617///   if its code point is within [32, 126] (see also the note on string
618///   literals in https://smt-lib.org/theories-UnicodeStrings.shtml).
619/// - At the parser level, we need to replace any single `"` character
620///   with `""`, according to the SMT-LIB 2.7 standard on string literals:
621///   https://smt-lib.org/papers/smt-lib-reference-v2.7-r2025-07-07.pdf
622///
623/// Note in particular that `\\` is NOT an escape sequence,
624/// so cvc5 will read `\\u{0}` as a two-character string with
625/// characters `\u{5c}` and `\0`.
626pub(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            // This is to avoid unexpectedly escape some characters
633            out.push_str("\\u{5c}");
634        } else if 32 as char <= c && c <= 126 as char {
635            out.push(c);
636        } else {
637            // Encode non-printable character
638            if c as u32 > SMT_LIB_MAX_CODE_POINT {
639                return None; // Invalid code point for SMT-LIB
640            }
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    /// Compiles `expr` against the schema shared with
877    /// `compiler::ext_has_attr_tests` and returns the SMT text the encoder
878    /// emits for it.
879    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    // entity base, optional then present
900    #[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    // entity base, present then optional
935    #[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    // record base, present then optional
974    #[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    // record base, optional then present
1014    #[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        // Size increases linearly, not exponentially
1076        assert_eq!(n_3 - n_2, n_4 - n_3);
1077    }
1078}