Skip to main content

cedar_policy_symcc/symcc/
decoder.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 decoder, which is the inverse of the encoder
18//! that parses a subset of SMT-LIB terms and commands required for (get-model)
19
20use 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/// Errors during decoding, i.e., converting SMT terms
54/// to our internal [`Term`] representation.
55#[derive(Debug, Diagnostic, Error)]
56pub enum DecodeError {
57    /// Error parsing an s-expression
58    #[error(transparent)]
59    SExprParse(#[from] SExprParseError),
60    /// Failed to parse an SMT numeral.
61    #[error("Invalid numeric token: {0}")]
62    ParseIntError(#[from] std::num::ParseIntError),
63    /// Integer overflow.
64    #[error("Integer overflow")]
65    IntegerOverflow,
66    /// Model of an unexpected form returned by the solver.
67    #[error("Model of an unexpected form returned by the solver")]
68    UnexpectedModel,
69    /// Unknown SMT type.
70    #[error("Unknown SMT type: {0}")]
71    UnknownType(SExpr),
72    /// Unknown SMT literal.
73    #[error("Unknown SMT literal: {0}")]
74    UnknownLiteral(SExpr),
75    /// Unmatched types.
76    #[error("Unmatched type: expected {0:?}, found {1:?}")]
77    UnmatchedType(TermType, TermType),
78    /// Unmatched field type.
79    #[error("Unmatched field type: expected {0:?}, found {1:?}")]
80    UnmatchedFieldType(TermType, TermType),
81    /// Invalid set type.
82    #[error("Invalid set type: {0}")]
83    InvalidSetType(SExpr),
84    /// Invalid option type.
85    #[error("Invalid option type: {0}")]
86    InvalidOptionType(SExpr),
87    /// `set.union` applied to non-literals.
88    #[error("set.union applied to non-literals {0:?} and {1:?}")]
89    SetUnionNonLiterals(Term, Term),
90    /// Unmatched record type fields.
91    #[error("Unmatched record type fields")]
92    UnmatchedRecordType,
93    /// Unknown variable.
94    #[error("Unknown variable: {0}")]
95    UnknownVariable(String),
96    /// Unknown unary function.
97    #[error("Unknown unary function: {0}")]
98    UnknownUUF(String),
99    /// Unexpected form of unary function model.
100    #[error("Unexpected form of unary function model: {0}")]
101    UnexpectedUnaryFunctionForm(SExpr),
102    /// Bit-vector error.
103    #[error("Bit-vector error")]
104    BitVecError(#[from] BitVecError),
105    /// Bitvector of a zero width, which we do not support.
106    #[error("Bitvector of zero width")]
107    ZeroWidthBitVec,
108}
109
110/// Maps from SMT symbols their corresponding variables
111/// (principal, action, resource) and entity types.
112#[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    /// Extracts the reverse mapping from SMT symbols to
122    /// Term-level names from the encoder state.
123    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    /// Default literal of a type.
170    /// Used as placeholders for SMT partial applications.
171    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                // If the entity is an enum type, we return the first enum
179                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                        "" // This case should not happen on a well-formed `SymEnv`
188                    }
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    /// Similar to [`TermType::default_literal`], but for [`Uuf`].
232    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    /// Checks if the [`SExpr`] is the given symbol.
244    fn is_symbol(&self, s: &str) -> bool {
245        match self {
246            SExpr::Symbol(sym) => sym == s,
247            _ => false,
248        }
249    }
250
251    /// Checks if the [`SExpr`] is an `App` where the target function is the given symbol.
252    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    /// If this [`SExpr`] is an `App` applying the function named `func`, returns its arguments
257    /// (excluding the function symbol itself).
258    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    /// If this [`SExpr`] is an `App` applying the function named `func` to
269    /// exactly `N` arguments, returns those arguments (excluding the function symbol itself).
270    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    /// Decodes [`TermType`] from an [`SExpr`].
276    pub fn decode_type(&self, id_maps: &IdMaps<'_>) -> Result<TermType, DecodeError> {
277        match self {
278            // Atomic types
279            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                    // Entity or record type
297                    _ => id_maps
298                        .types
299                        .get(s)
300                        .copied()
301                        .cloned()
302                        .ok_or_else(|| DecodeError::UnknownType(self.clone())),
303                }
304            }
305
306            // Parametrized types
307            SExpr::App(args) => {
308                match args.as_slice() {
309                    // (_ BitVec n)
310                    [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                    // (Option x)
319                    [SExpr::Symbol(option), param] if option == "Option" => {
320                        let ty = param.decode_type(id_maps)?;
321                        Ok(TermType::option_of(ty))
322                    }
323
324                    // (Set x)
325                    [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    /// Decodes an [`SExpr`] as an entity UID or record.
339    /// Corresponds to `SExpr.decodeLit.constructEntityOrRecord` in Lean.
340    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            // Entity UID
348            (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            // Record
354            (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    /// Helper function to decode more complex applications as literals.
383    /// Corresponds to `SExpr.decodeLit.construct` in Lean.
384    ///
385    /// This function accepts an optional expected type which it uses to assign
386    /// a type to a `none` expression without explicit type annotation (as is
387    /// emitted by Z3), but it does not use this type to do any additional
388    /// typechecking.
389    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            // Sometimes cvc5 does not simplify the terms in the model,
397            // and having these custom interpreters alleviates such issues
398            // (e.g., https://github.com/cvc5/cvc5/issues/11928).
399
400            // (not <v>)
401            [SExpr::Symbol(not_tok), v] if not_tok == "not" => {
402                Ok(factory::not(v.decode_literal(id_maps)?))
403            }
404
405            // (or <v1> <v2>)
406            [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            // (= <v1> <v2>)
412            [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            // (ite <cond> <then> <else>)
418            [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            // (bvnego <v1>)
427            [SExpr::Symbol(bvnego_tok), v] if bvnego_tok == "bvnego" => {
428                Ok(factory::bvnego(v.decode_literal(id_maps)?))
429            }
430
431            // (bvsaddo <v1> <v2>)
432            [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            // (bvsmulo <v1> <v2>)
437            [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            // (as none <typ>)
442            [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            // ((as some <typ>) <val>)
452            #[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            // (some <val>) without type annotation (Z3 produces this)
477            [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            // (as set.empty <set_typ>)
488            [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            // (set.singleton <val>)
503            [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            // (set.union <set1> <set2>)
518            [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                    // Merge two set literals
530                    (
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            // Decimal
546            [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            // Datetime
559            [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            // Duration
570            [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            // IPv4/IPv6
581            [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            // (_ bvN W) bitvector literals.  Emitted by cvc5 if called with `--bv-print-consts-as-indexed-symbols`.
614            [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                // Check that `val` fits in declared width. If width is at least 128,
623                // then all 128 bit vals must fit.
624                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            // Entity UID or record
631            [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    /// Decodes a literal (with only SMT constants and no bound variables).
640    pub fn decode_literal(&self, id_maps: &IdMaps<'_>) -> Result<Term, DecodeError> {
641        self.decode_literal_expecting(id_maps, None)
642    }
643
644    /// Decodes a literal term.
645    ///
646    /// This function accepts an optional expected type which it uses to assign
647    /// a type to a `none` expression without explicit type annotation (as is
648    /// emitted by Z3), but it does not use this type to do any additional
649    /// typechecking.
650    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            // Bare `none` without type annotation (Z3 produces this)
663            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            // Empty record type
669            SExpr::Symbol(s) if id_maps.types.contains_key(s) => {
670                self.decode_entity_or_record(id_maps, s, &[])
671            }
672
673            // Entity enum
674            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            // More complex applications
682            SExpr::App(args) => self.decode_literal_app(id_maps, args, expected_ty),
683
684            _ => Err(DecodeError::UnknownLiteral(self.clone())),
685        }
686    }
687
688    /// Decodes a constant definition in the model.
689    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    /// Decodes a unary function with the forms:
715    /// * `(ite (= lit x) <lit> (ite (= <lit> x) default))`
716    /// * `(or (= <literal> arg) (= <literal> arg))`
717    /// * `(= <lit> arg)`
718    ///
719    /// TODO: generalize to other forms?
720    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        // First check if the SMT name actually corresponds to a UUF
729        let Some(&uuf) = id_maps.uufs.get(name) else {
730            return Err(DecodeError::UnknownUUF(name.to_string()));
731        };
732
733        // Check that argument type and return type match those of the UUF
734        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            // `ite` case also handles constant functions without any conditions
751            Self::decode_ite_table(uuf, id_maps, arg_name, &ret_ty, body)
752        }
753    }
754
755    /// Decode UDF table `(ite (= lit x) <lit> (ite (= <lit> x) default))`
756    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        // Decode the body as a nested ite term
764        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        // Next `App` isn't an `ite`, so decode it as the default value.
777        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    /// Decode UDF table with a disjunction `(or (= <literal> arg) (= <literal> arg))`
790    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    /// Decode UDF table with a single entry `(= <lit> arg)`
821    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    /// Get the literal in an s-expr with the shape `(= <lit> <arg>)` or `(= <arg> <lit>)`
843    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    /// Decodes the output of `(get-model)` to as [`Interpretation`].
861    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        // TODO: better error handling here
874        for cmd in cmds {
875            let SExpr::App(sub_exprs) = cmd else {
876                return Err(DecodeError::UnexpectedModel);
877            };
878
879            // sub_exprs should be of the form
880            // "define-fun" <name> (<args>) <ret_type> <body>
881            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                        // Decode unary function (skip if not a known UUF)
887                        [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                                // else: skip unknown unary functions (e.g., Z3 intermediate terms)
896                            }
897                            _ => return Err(DecodeError::UnexpectedModel),
898                        },
899
900                        // Decode SMT constant definition as interpretation to a Cedar variable
901                        // (skip if not a known variable — Z3 includes define-fun entries
902                        // for intermediate terms that aren't declare-const variables)
903                        [] => {
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
923/// Decodes the output of `(get-model)` to as [`Interpretation`].
924pub 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        // Hex bitvec literal (#xNN)
1024        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        // Indexed bitvec literal (_ bvN W)
1033        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    /// Z3 includes `define-fun` entries in its model for intermediate terms
1064    /// (e.g., terms introduced by `define-fun` in the input), not just
1065    /// `declare-const` variables. The decoder should skip these unknown
1066    /// symbols rather than failing with `UnknownVariable`.
1067    ///
1068    /// Reproduces the model format seen when using Z3 4.12.5 as the solver:
1069    /// ```text
1070    /// (define-fun t0 () E0 (E0 "!0!"))    <-- the actual declare-const var
1071    /// (define-fun t3 () Bool (not ...))    <-- intermediate, not in IdMaps
1072    /// (define-fun t1 () E0 (E0 "a"))      <-- intermediate
1073    /// (define-fun t2 () Bool (= t0 ...))  <-- intermediate
1074    /// ```
1075    #[test]
1076    fn decode_model_skips_unknown_define_funs() {
1077        // Z3-style model with extra define-funs for intermediate terms.
1078        // The decoder should skip unknown names and still decode known vars.
1079        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)]), // only t0 in consts
1093                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        // Z3-style model with an unknown unary function (not in IdMaps.uufs).
1105        // The decoder should skip it.
1106        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)]), // only t0 in consts
1120                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    /// Z3 puts the bound variable on the lhs of `=` in ite-chains (i.e., for a uuf).
1182    /// `(ite (= x!0 <literal>) ...)` instead of cvc5's `(ite (= <literal> x) ...)`
1183    #[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        // Z3 model for attrs[E] with two entities having different attrs
1207        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    /// Z3 produces bare `none` / `(some val)` without type annotations.
1299    #[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        // bare `none` in a record field
1308        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        // bare `(some "x")` in a record field
1333        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        // bare `none` as a direct constant
1354        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    /// Record constructor with wrong number of fields.
1509    #[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        // Record type has 2 fields but we only provide 1 argument
1524        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    /// Record field value has wrong type.
1538    #[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        // Field expects String but gets Bool
1550        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        // Entity expects (E0 "id") but gets (E0 true)
1576        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    /// Unary function arg has wrong form: the inner pair is not [Symbol, type].
1646    #[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}