Skip to main content

cedar_policy_symcc/symcc/
concretizer.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 includes functions to convert
18//! literal Term/SymRequest/SymEntities to their
19//! concrete versions
20
21use std::borrow::Borrow;
22use std::collections::{BTreeMap, BTreeSet, HashSet};
23use std::sync::Arc;
24
25use cedar_policy::{Entities, EntityId, EntityTypeName, EntityUid, EvalResult, Request};
26use cedar_policy_core::ast::{
27    is_normalized_ident, Context, Entity, EntityAttrEvaluationError, Expr, ExprVisitor, Literal,
28    Set, Value, ValueKind,
29};
30use cedar_policy_core::entities::{NoEntitiesSchema, TCComputation};
31use cedar_policy_core::extensions::Extensions;
32use cedar_policy_core::parser::Loc;
33use miette::Diagnostic;
34use num_bigint::{BigInt, TryFromBigIntError};
35use smol_str::SmolStr;
36use thiserror::Error;
37
38use crate::symcc::factory;
39
40use super::env::{SymEntities, SymEntityData, SymRequest};
41use super::ext::ExtError;
42use super::function::{Udf, UnaryFunction};
43use super::term::{Term, TermPrim};
44use super::SymEnv;
45
46/// Errors that happen during concretization, i.e., the process
47/// of converting literal [`Term`]s back to representatinos in
48/// [`cedar_policy`]/[`cedar_policy_core`].
49#[derive(Debug, Diagnostic, Error)]
50pub enum ConcretizeError {
51    /// Got no policies, expected to have at least one.
52    #[error("expected to have at least one policy")]
53    NoPolicies,
54    /// Expecting a literal entity.
55    #[error("Not a literal entity: {0:?}")]
56    NotLiteralEntity(Term),
57    /// Expecting a literal string.
58    #[error("Not a literal string: {0:?}")]
59    NotLiteralString(Term),
60    /// Cannot convert a term to a value.
61    #[error("Unable to convert {0:?} to a value")]
62    UnableToConvertToValue(Term),
63    /// Cannot convert a term to a context.
64    #[error("Unable to convert {0:?} to a context")]
65    UnableToConvertToContext(Term),
66    /// Unable to construct entity.
67    #[error("Unable to construct a valid entity")]
68    UnableToConstructEntity(#[from] EntityAttrEvaluationError),
69    /// Entity type not found.
70    #[error("Entity type not found: {0}")]
71    EntityTypeNotFound(EntityTypeName),
72    /// Unable to falidate request.
73    #[error("Request validation error")]
74    RequestValidationError(#[from] cedar_policy::RequestValidationError),
75    /// Errors when constructing entity.
76    #[error("Unable to construct entities")]
77    EntitiesError(#[from] cedar_policy::entities_errors::EntitiesError),
78    /// Fail to convert from big integer.
79    #[error("Unable to convert BitVec to integer")]
80    TryFromBigIntError(#[from] TryFromBigIntError<BigInt>),
81    /// Extension error.
82    #[error("extension error")]
83    ExtError(#[from] ExtError),
84    /// Unsupported expression.
85    #[error("unsupported expression: {0}")]
86    UnsupportedExpr(Expr),
87}
88
89/// A concrete environment recovered from a [`SymEnv`].
90#[derive(Debug, Clone, PartialEq)]
91pub struct Env {
92    /// Concrete request
93    pub request: Request,
94    /// Concrete entities
95    pub entities: Entities,
96}
97
98/// Write an entity's attributes or tags
99fn fmt_attrs_or_tags<'a, E>(
100    f: &mut std::fmt::Formatter<'_>,
101    label: Option<&str>,
102    values: impl IntoIterator<Item = (&'a str, Result<EvalResult, E>)>,
103) -> std::fmt::Result {
104    let mut values = values.into_iter().peekable();
105    if values.peek().is_none() {
106        return Ok(());
107    }
108    if let Some(label) = label {
109        write!(f, " {label}")?;
110    }
111    writeln!(f, " {{")?;
112    for (k, v) in values {
113        write!(f, "    ")?;
114        if is_normalized_ident(k) {
115            write!(f, "{k}")?;
116        } else {
117            write!(f, "\"{}\"", k.escape_debug())?;
118        }
119        match v {
120            Ok(val) => writeln!(f, ": {val},")?,
121            Err(_) => writeln!(f, ": <unknown>,")?,
122        }
123    }
124    write!(f, "  }}")
125}
126
127impl std::fmt::Display for Env {
128    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
129        let req = &self.request;
130        let principal = req
131            .principal()
132            .map_or_else(|| "unknown".to_string(), |p| p.to_string());
133        let action = req
134            .action()
135            .map_or_else(|| "unknown".to_string(), |a| a.to_string());
136        let resource = req
137            .resource()
138            .map_or_else(|| "unknown".to_string(), |r| r.to_string());
139        write!(
140            f,
141            "principal: {principal}, action: {action}, resource: {resource}"
142        )?;
143        if let Some(ctx) = req.context() {
144            write!(f, "\ncontext: {ctx}")?;
145        }
146        let entities = &self.entities;
147        if !entities.is_empty() {
148            writeln!(f, "\nentities: [")?;
149            for entity in entities.iter() {
150                let uid = entity.uid();
151                write!(f, "  {uid}")?;
152                // ancestors
153                if let Some(ancestors) = entities.ancestors(&uid) {
154                    let ancs: Vec<_> = ancestors.map(|a| a.to_string()).collect();
155                    if !ancs.is_empty() {
156                        write!(f, " in [{}]", ancs.join(", "))?;
157                    }
158                }
159                fmt_attrs_or_tags(f, None, entity.attrs())?;
160                fmt_attrs_or_tags(f, Some("tags"), entity.tags())?;
161                writeln!(f, ",")?;
162            }
163            write!(f, "]")?;
164        }
165        Ok(())
166    }
167}
168
169/// Tries to extract an `EntityUid` from a `Term`.
170/// Corresponds to `Term.entityUID?` in `Concretize.lean`
171impl TryFrom<&Term> for EntityUid {
172    type Error = ConcretizeError;
173
174    fn try_from(term: &Term) -> Result<Self, Self::Error> {
175        if let Term::Prim(TermPrim::Entity(uid)) = term {
176            Ok(uid.clone())
177        } else {
178            Err(ConcretizeError::NotLiteralEntity(term.clone()))
179        }
180    }
181}
182
183/// Tries to extract a set of `EntityUid`'s from a `Term`.
184/// Corresponds `Term.setOfEntityUIDs?` in `Concretize.lean`
185impl TryFrom<&Term> for BTreeSet<EntityUid> {
186    type Error = ConcretizeError;
187
188    fn try_from(term: &Term) -> Result<Self, Self::Error> {
189        if let Term::Set { elts, .. } = term {
190            Ok(elts
191                .iter()
192                .map(|t| t.try_into())
193                .collect::<Result<_, _>>()?)
194        } else {
195            Err(ConcretizeError::NotLiteralEntity(term.clone()))
196        }
197    }
198}
199
200/// Tries to convert a `Term` to a string.
201impl TryFrom<&Term> for SmolStr {
202    type Error = ConcretizeError;
203
204    fn try_from(term: &Term) -> Result<Self, Self::Error> {
205        if let Term::Prim(TermPrim::String(s)) = term {
206            Ok(s.clone())
207        } else {
208            Err(ConcretizeError::NotLiteralString(term.clone()))
209        }
210    }
211}
212
213/// Tries to extract a set of `Strings`'s from a `Term`.
214impl TryFrom<&Term> for BTreeSet<SmolStr> {
215    type Error = ConcretizeError;
216
217    fn try_from(term: &Term) -> Result<Self, Self::Error> {
218        if let Term::Set { elts, .. } = term {
219            Ok(elts
220                .iter()
221                .map(|t| t.try_into())
222                .collect::<Result<_, _>>()?)
223        } else {
224            Err(ConcretizeError::NotLiteralEntity(term.clone()))
225        }
226    }
227}
228
229impl TryFrom<&Term> for Value {
230    type Error = ConcretizeError;
231
232    fn try_from(term: &Term) -> Result<Self, Self::Error> {
233        match term {
234            Term::Prim(TermPrim::Bool(b)) => {
235                Ok(Value::new(ValueKind::Lit(Literal::Bool(*b)), None))
236            }
237
238            Term::Prim(TermPrim::Bitvec(v)) => Ok(Value::new(
239                ValueKind::Lit(Literal::Long(v.to_int().try_into()?)),
240                None,
241            )),
242
243            Term::Prim(TermPrim::String(s)) => {
244                Ok(Value::new(ValueKind::Lit(Literal::String(s.clone())), None))
245            }
246
247            Term::Prim(TermPrim::Entity(uid)) => Ok(Value::new(
248                ValueKind::Lit(Literal::EntityUID(Arc::new(uid.clone().into()))),
249                None,
250            )),
251
252            Term::Prim(TermPrim::Ext(ext)) => Ok(Self::try_from(ext)?),
253
254            Term::Set { elts, .. } => Ok(Value::new(
255                ValueKind::Set(Set::new(
256                    elts.iter()
257                        .map(|t| t.try_into())
258                        .collect::<Result<Vec<_>, _>>()?,
259                )),
260                None,
261            )),
262
263            Term::Record(rec) => Ok(Value::new(
264                ValueKind::Record(Arc::new(
265                    rec.iter()
266                        .map(|(k, v)| {
267                            if let Term::Some(t) = v {
268                                Ok(Some((k.clone(), t.as_ref().try_into()?)))
269                            } else if let Term::None(_) = v {
270                                // None fields are simply ignored
271                                Ok(None)
272                            } else {
273                                Ok(Some((k.clone(), v.try_into()?)))
274                            }
275                        })
276                        .collect::<Result<Vec<Option<_>>, ConcretizeError>>()?
277                        .into_iter()
278                        .flatten()
279                        .collect(),
280                )),
281                None,
282            )),
283
284            // Otherwise it's not convertable
285            _ => Err(ConcretizeError::UnableToConvertToValue(term.clone())),
286        }
287    }
288}
289
290impl SymRequest {
291    pub fn concretize(&self) -> Result<Request, ConcretizeError> {
292        Ok(Request::new(
293            (&self.principal).try_into()?,
294            (&self.action).try_into()?,
295            (&self.resource).try_into()?,
296            Context::Value(self.context.try_into_record()?).into(),
297            None, // TODO: schema == None disables request validation
298        )?)
299    }
300
301    fn get_all_entity_uids(&self, uids: &mut BTreeSet<EntityUid>) {
302        self.context.get_all_entity_uids(uids);
303        self.principal.get_all_entity_uids(uids);
304        self.action.get_all_entity_uids(uids);
305        self.resource.get_all_entity_uids(uids);
306    }
307}
308
309impl Term {
310    /// Tries to convert a term into a record
311    ///
312    /// Corresponds to `Term.recordValue?` in `Concretize.lean`
313    fn try_into_record(&self) -> Result<Arc<BTreeMap<SmolStr, Value>>, ConcretizeError> {
314        if let Value {
315            value: ValueKind::Record(record),
316            ..
317        } = self.try_into()?
318        {
319            Ok(record)
320        } else {
321            Err(ConcretizeError::UnableToConvertToContext(self.clone()))
322        }
323    }
324
325    /// Collect all entity UIDs occurring in the term
326    ///
327    /// Corresponds to `Term.entityUIDs` in `Concretizer.lean`
328    pub(crate) fn get_all_entity_uids(&self, uids: &mut BTreeSet<EntityUid>) {
329        match self {
330            Term::Prim(TermPrim::Entity(uid)) => {
331                uids.insert(uid.clone());
332            }
333
334            Term::Some(t) => {
335                t.get_all_entity_uids(uids);
336            }
337
338            Term::Set { elts, .. } => {
339                for t in elts.iter() {
340                    t.get_all_entity_uids(uids);
341                }
342            }
343
344            Term::Record(rec) => {
345                for t in rec.values() {
346                    t.get_all_entity_uids(uids);
347                }
348            }
349
350            Term::App { args, .. } => {
351                for t in args.iter() {
352                    t.get_all_entity_uids(uids);
353                }
354            }
355
356            _ => {}
357        }
358    }
359}
360
361impl Udf {
362    fn get_all_entity_uids(&self, uids: &mut BTreeSet<EntityUid>) {
363        self.default.get_all_entity_uids(uids);
364        for (k, v) in self.table.iter() {
365            k.get_all_entity_uids(uids);
366            v.get_all_entity_uids(uids);
367        }
368    }
369}
370
371impl UnaryFunction {
372    /// Corresponds to `UnaryFunction.entityUIDs` in `Concretize.lean`
373    fn get_all_entity_uids(&self, uids: &mut BTreeSet<EntityUid>) {
374        match self {
375            UnaryFunction::Udf(udf) => udf.get_all_entity_uids(uids),
376            UnaryFunction::Uuf(_) => {}
377        }
378    }
379}
380
381impl SymEntityData {
382    /// Concretizes a particular entity.
383    pub fn concretize(&self, euid: &EntityUid) -> Result<Entity, ConcretizeError> {
384        let tuid = Term::Prim(TermPrim::Entity(euid.clone()));
385
386        let concrete_attrs = factory::app(self.attrs.clone(), tuid.clone()).try_into_record()?;
387
388        // For each ancestor entity type, apply the suitable ancestor function
389        // to obtain a concrete set of ancestor EUIDs
390        let concrete_ancestors = self
391            .ancestors
392            .values()
393            .map(|ancestor| {
394                let euids: BTreeSet<EntityUid> =
395                    (&factory::app(ancestor.clone(), tuid.clone())).try_into()?;
396
397                Ok(euids.into_iter().map(|euid| euid.as_ref().clone()))
398            })
399            .collect::<Result<Vec<_>, ConcretizeError>>()?
400            .into_iter()
401            .flatten()
402            .collect::<HashSet<_>>();
403
404        // Read tags from the model
405        let tags = if let Some(tags) = &self.tags {
406            // Get all valid tag keys first
407            let keys: BTreeSet<SmolStr> =
408                (&factory::app(tags.keys.clone(), tuid.clone())).try_into()?;
409
410            keys.into_iter()
411                .map(|k| {
412                    // Using get_tag_unchecked here since we know already that k is in the key set
413                    let val: Value = (&tags
414                        .get_tag_unchecked(tuid.clone(), Term::Prim(TermPrim::String(k.clone()))))
415                        .try_into()?;
416
417                    Ok((k, val.into()))
418                })
419                .collect::<Result<_, ConcretizeError>>()?
420        } else {
421            BTreeMap::new()
422        };
423
424        Ok(Entity::new(
425            euid.as_ref().clone(),
426            concrete_attrs
427                .as_ref()
428                .clone()
429                .into_iter()
430                .map(|(k, v)| (k, v.into())),
431            HashSet::new(),
432            concrete_ancestors,
433            tags,
434            Extensions::all_available(),
435        )?)
436    }
437
438    /// Corresponds to `SymEntityData.entityUIDs` in `Concretize.lean`
439    fn get_all_entity_uids(&self, ety: &EntityTypeName, uids: &mut BTreeSet<EntityUid>) {
440        if let Some(members) = &self.members {
441            for member in members {
442                uids.insert(EntityUid::from_type_name_and_id(
443                    ety.clone(),
444                    EntityId::new(member),
445                ));
446            }
447        }
448
449        self.attrs.get_all_entity_uids(uids);
450
451        for ancestor in self.ancestors.values() {
452            ancestor.get_all_entity_uids(uids);
453        }
454
455        if let Some(tags) = &self.tags {
456            // tags.keys.get_all_entity_uids(uids);
457            tags.vals.get_all_entity_uids(uids);
458        }
459    }
460}
461
462impl SymEntities {
463    /// Concretizes a literal SymEntities to Entities
464    pub fn concretize(&self, all_euids: &BTreeSet<EntityUid>) -> Result<Entities, ConcretizeError> {
465        let mut entities = Vec::with_capacity(all_euids.len());
466
467        for euid in all_euids {
468            let sym_entity_data =
469                self.0
470                    .get(euid.type_name())
471                    .ok_or(ConcretizeError::EntityTypeNotFound(
472                        euid.type_name().clone(),
473                    ))?;
474
475            entities.push(sym_entity_data.concretize(euid)?);
476        }
477
478        // As the internal cedar_policy_core::entities::Entities
479        let internal_entities = cedar_policy_core::entities::Entities::from_entities(
480            entities,
481            None::<&NoEntitiesSchema>,
482            #[cfg(debug_assertions)]
483            TCComputation::EnforceAlreadyComputed,
484            #[cfg(not(debug_assertions))]
485            TCComputation::AssumeAlreadyComputed,
486            Extensions::all_available(),
487        )?;
488
489        Ok(Entities::from(internal_entities))
490    }
491
492    /// Corresponds to `SymEntities.entityUIDs` in `Concretize.lean`
493    fn get_all_entity_uids(&self, uids: &mut BTreeSet<EntityUid>) {
494        for (ety, data) in self.0.iter() {
495            data.get_all_entity_uids(ety, uids);
496        }
497    }
498}
499
500/// An [`ExprVisitor`] to collect all entity UIDs occurring in an expression.
501///
502/// Corresponds to `Expr.entityUIDs` in `Concretize.lean`.
503struct EntityUIDCollector<'a>(&'a mut BTreeSet<EntityUid>);
504
505impl ExprVisitor for EntityUIDCollector<'_> {
506    type Output = ();
507
508    fn visit_literal(&mut self, lit: &Literal, _: Option<&Loc>) -> Option<Self::Output> {
509        if let Literal::EntityUID(euid) = lit {
510            self.0.insert(euid.as_ref().clone().into());
511        }
512        None
513    }
514}
515
516impl SymEnv {
517    /// Concretizes a literal [`SymEnv`] to a concrete [`Env`].
518    ///
519    /// In most cases, one should use [`SymEnv::extract`] instead
520    /// to ensure well-formed output [`Env`].
521    pub(crate) fn concretize<E: Borrow<Expr>>(
522        &self,
523        exprs: impl IntoIterator<Item = E>,
524    ) -> Result<Env, ConcretizeError> {
525        let mut uids = BTreeSet::new();
526        self.request.get_all_entity_uids(&mut uids);
527        self.entities.get_all_entity_uids(&mut uids);
528
529        // Instead of using `footprint` and `Term::get_all_entity_uids`,
530        // we collect EUIDs in expressions directly to avoid incorrect
531        // short-circuiting in an incomplete entity store.
532        let mut visitor = EntityUIDCollector(&mut uids);
533        for expr in exprs {
534            visitor.visit_expr(expr.borrow());
535        }
536
537        Ok(Env {
538            request: self.request.concretize()?,
539            entities: self.entities.concretize(&uids)?,
540        })
541    }
542}
543
544#[cfg(test)]
545mod test {
546    use super::Env;
547    use cedar_policy::{Context, Entities, Entity, Request, RestrictedExpression};
548    use insta::assert_snapshot;
549
550    fn display_env(
551        attrs: impl IntoIterator<Item = (String, RestrictedExpression)>,
552        tags: impl IntoIterator<Item = (String, RestrictedExpression)>,
553    ) -> String {
554        let entity =
555            Entity::new_with_tags(r#"User::"alice""#.parse().unwrap(), attrs, [], tags).unwrap();
556        let request = Request::new(
557            r#"User::"alice""#.parse().unwrap(),
558            r#"Action::"view""#.parse().unwrap(),
559            r#"Photo::"vacation""#.parse().unwrap(),
560            Context::empty(),
561            None,
562        )
563        .unwrap();
564        Env {
565            request,
566            entities: Entities::from_entities([entity], None).unwrap(),
567        }
568        .to_string()
569    }
570
571    #[test]
572    fn test_display_env() {
573        assert_snapshot!(display_env(
574            [("a".to_string(), RestrictedExpression::new_long(0))],
575            [("b".to_string(), RestrictedExpression::new_long(1))],
576        ), @r#"
577        principal: User::"alice", action: Action::"view", resource: Photo::"vacation"
578        context: {}
579        entities: [
580          User::"alice" {
581            a: 0,
582          } tags {
583            b: 1,
584          },
585        ]
586        "#);
587
588        assert_snapshot!(display_env([("a".to_string(), RestrictedExpression::new_long(0))], [],), @r#"
589        principal: User::"alice", action: Action::"view", resource: Photo::"vacation"
590        context: {}
591        entities: [
592          User::"alice" {
593            a: 0,
594          },
595        ]
596        "#);
597
598        assert_snapshot!(display_env([], [("b".to_string(), RestrictedExpression::new_long(1))],), @r#"
599        principal: User::"alice", action: Action::"view", resource: Photo::"vacation"
600        context: {}
601        entities: [
602          User::"alice" tags {
603            b: 1,
604          },
605        ]
606        "#);
607
608        assert_snapshot!(display_env([], []), @r#"
609        principal: User::"alice", action: Action::"view", resource: Photo::"vacation"
610        context: {}
611        entities: [
612          User::"alice",
613        ]
614        "#);
615    }
616
617    #[test]
618    fn display_attr_and_tag_names() {
619        // the whole `Env`, for context on the per-name snapshots below
620        let display_attr = |attr: &str| {
621            display_env([(attr.to_string(), RestrictedExpression::new_long(0))], [])
622                .lines()
623                .find(|l| l.starts_with("    "))
624                .expect("no attr line")
625                .trim()
626                .to_string()
627        };
628        assert_snapshot!(display_attr("plain"), @"plain: 0,");
629        assert_snapshot!(display_attr("new\nline"), @r#""new\nline": 0,"#);
630
631        let display_tag = |tag: &str| {
632            display_env([], [(tag.to_string(), RestrictedExpression::new_long(0))])
633                .lines()
634                .find(|l| l.starts_with("    "))
635                .expect("no tag line")
636                .trim()
637                .to_string()
638        };
639        assert_snapshot!(display_tag("plain"), @"plain: 0,");
640        assert_snapshot!(display_tag("spa ce"), @r#""spa ce": 0,"#);
641    }
642}