Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
31 changes: 20 additions & 11 deletions cedar-lean-ffi/src/lean_ffi/tpe.rs
Original file line number Diff line number Diff line change
Expand Up @@ -189,9 +189,9 @@ mod test {
unsafe extern "C" {}

use cedar_policy::{
Context, Entities, EntityTypeName, EntityUid, PartialEntities, PartialEntity,
PartialEntityUid, PartialRequest, Policy, PolicyId, PolicySet, RestrictedExpression,
Schema,
Context, Entities, EntityTypeName, EntityUid, PartialAttribute, PartialEntities,
PartialEntity, PartialEntityUid, PartialRequest, Policy, PolicyId, PolicySet,
RestrictedExpression, Schema,
};
use cool_asserts::assert_matches;

Expand All @@ -203,6 +203,15 @@ mod test {
};
use cedar_policy_core::ast::{EntityUID, Value, Var};
use cedar_policy_core::tpe::residual::{Residual, ResidualKind};

fn partial_attrs<const N: usize>(
attrs: [(smol_str::SmolStr, RestrictedExpression); N],
) -> BTreeMap<smol_str::SmolStr, PartialAttribute> {
attrs
.into_iter()
.map(|(key, value)| (key, PartialAttribute::value(value)))
.collect()
}
use cedar_policy_core::validator::types::{EntityKind, EntityLUB, Type};

/// Helper to compare Rust and Lean TPE responses: decision, policy categorizations,
Expand Down Expand Up @@ -490,7 +499,7 @@ mod test {

let account = PartialEntity::new(
EntityUid::from_str(r#"Account::"checking""#).unwrap(),
Some(BTreeMap::from([
Some(partial_attrs([
("balance".into(), RestrictedExpression::new_long(10000)),
(
"owner".into(),
Expand All @@ -504,7 +513,7 @@ mod test {
.unwrap();
let user = PartialEntity::new(
EntityUid::from_str(r#"User::"alice""#).unwrap(),
Some(BTreeMap::from([
Some(partial_attrs([
(
"name".into(),
RestrictedExpression::new_string("Alice".into()),
Expand Down Expand Up @@ -572,7 +581,7 @@ mod test {

let user = PartialEntity::new(
EntityUid::from_str(r#"User::"alice""#).unwrap(),
Some(BTreeMap::from([
Some(partial_attrs([
(
"name".into(),
RestrictedExpression::new_string("Alice".into()),
Expand Down Expand Up @@ -625,7 +634,7 @@ mod test {

let account = PartialEntity::new(
EntityUid::from_str(r#"Account::"checking""#).unwrap(),
Some(BTreeMap::from([
Some(partial_attrs([
("balance".into(), RestrictedExpression::new_long(10000)),
(
"owner".into(),
Expand All @@ -639,7 +648,7 @@ mod test {
.unwrap();
let user = PartialEntity::new(
EntityUid::from_str(r#"User::"alice""#).unwrap(),
Some(BTreeMap::from([
Some(partial_attrs([
(
"name".into(),
RestrictedExpression::new_string("Alice".into()),
Expand Down Expand Up @@ -703,7 +712,7 @@ mod test {

let user = PartialEntity::new(
EntityUid::from_str(r#"User::"alice""#).unwrap(),
Some(BTreeMap::from([
Some(partial_attrs([
(
"name".into(),
RestrictedExpression::new_string("Alice".into()),
Expand Down Expand Up @@ -797,7 +806,7 @@ mod test {

let account = PartialEntity::new(
EntityUid::from_str(r#"Account::"checking""#).unwrap(),
Some(BTreeMap::from([
Some(partial_attrs([
("balance".into(), RestrictedExpression::new_long(10000)),
(
"owner".into(),
Expand All @@ -811,7 +820,7 @@ mod test {
.unwrap();
let user = PartialEntity::new(
EntityUid::from_str(r#"User::"alice""#).unwrap(),
Some(BTreeMap::from([
Some(partial_attrs([
(
"name".into(),
RestrictedExpression::new_string("Alice".into()),
Expand Down
126 changes: 71 additions & 55 deletions cedar-lean-ffi/src/messages.rs
Original file line number Diff line number Diff line change
Expand Up @@ -235,6 +235,11 @@ pub mod tpe {
Entities, EntityUid, PartialEntities, PartialEntityUid, PartialRequest, PolicySet, Request,
RestrictedExpression, Schema, proto::models as cedar_proto,
};
use cedar_policy_core::ast::{Expr, Value, ValueKind};
use cedar_policy_core::tpe::value::{
PartialAttribute as CorePartialAttribute, PartialRecord as CorePartialRecord,
PartialValue as CorePartialValue,
};
use smol_str::SmolStr;
use std::collections::{BTreeMap, HashMap, HashSet};

Expand Down Expand Up @@ -404,26 +409,58 @@ pub mod tpe {
}
}

fn expr_from_partial_value(value: &CorePartialValue) -> Option<Expr> {
match value {
CorePartialValue::Lit(lit) => Some(Expr::from(Value::new(lit.clone(), None))),
CorePartialValue::Set(set) => {
Some(Expr::from(Value::new(ValueKind::Set(set.clone()), None)))
}
CorePartialValue::ExtensionValue(ext) => Some(Expr::from(Value::new(
ValueKind::ExtensionValue(ext.clone()),
None,
))),
CorePartialValue::Record(record) => {
let fields = partial_record_exprs(record)?;
Expr::record(fields).ok()
}
}
}

/// Encode only records representable by the legacy concrete-value wire format.
/// Returning `None` makes the entire component unknown, which is conservative.
fn partial_record_exprs(record: &CorePartialRecord) -> Option<BTreeMap<SmolStr, Expr>> {
record
.attrs()
.filter_map(|(key, state)| match state {
CorePartialAttribute::Value(value) => {
Some(expr_from_partial_value(value).map(|expr| (key.clone(), expr)))
}
CorePartialAttribute::Absent => None,
CorePartialAttribute::Exists | CorePartialAttribute::Unknown => Some(None),
})
.collect()
}

fn partial_record_to_proto(
record: &CorePartialRecord,
) -> Option<HashMap<String, cedar_proto::Expr>> {
Some(
partial_record_exprs(record)?
.into_iter()
.map(|(key, expr)| (key.to_string(), cedar_proto::Expr::from(&expr)))
.collect(),
)
}

impl proto::PartialRequest {
fn from_inner(req: &cedar_policy_core::tpe::request::PartialRequest) -> Self {
use cedar_policy_core::ast::Expr;
let (context, has_context) = match req.context_attrs() {
Some(ctx) => (
ctx.iter()
.map(|(k, v)| {
(
k.to_string(),
cedar_policy::proto::models::Expr::from(&Expr::from(v.clone())),
)
})
.collect(),
true,
),
None => (Default::default(), false),
};
let (context, has_context) = req
.context()
.and_then(partial_record_to_proto)
.map_or_else(|| (Default::default(), false), |context| (context, true));
Self {
principal: Some(proto::PartialEntityUid::from_inner(req.principal())),
action: Some(cedar_policy::proto::models::EntityUid::from(req.action())),
action: Some(cedar_proto::EntityUid::from(req.action())),
resource: Some(proto::PartialEntityUid::from_inner(req.resource())),
context,
has_context,
Expand All @@ -444,46 +481,25 @@ pub mod tpe {

impl proto::PartialEntity {
fn from_inner(entity: &cedar_policy_core::tpe::entities::PartialEntity) -> Self {
use cedar_policy_core::ast::Expr;
let (attrs, has_attrs) = match entity.attrs() {
Some(a) => (
a.iter()
.map(|(k, v)| {
(
k.to_string(),
cedar_policy::proto::models::Expr::from(&Expr::from(v.clone())),
)
})
.collect(),
true,
),
None => (Default::default(), false),
};
let (ancestors, has_ancestors) = match entity.ancestors() {
Some(a) => (
a.iter()
.map(cedar_policy::proto::models::EntityUid::from)
.collect(),
true,
),
None => (Default::default(), false),
};
let (tags, has_tags) = match entity.tags() {
Some(t) => (
t.iter()
.map(|(k, v)| {
(
k.to_string(),
cedar_policy::proto::models::Expr::from(&Expr::from(v.clone())),
)
})
.collect(),
true,
),
None => (Default::default(), false),
};
let (attrs, has_attrs) = entity
.attrs()
.and_then(partial_record_to_proto)
.map_or_else(|| (Default::default(), false), |attrs| (attrs, true));
let (ancestors, has_ancestors) = entity.ancestors().map_or_else(
|| (Default::default(), false),
|ancestors| {
(
ancestors.iter().map(cedar_proto::EntityUid::from).collect(),
true,
)
},
);
let (tags, has_tags) = entity
.tags()
.and_then(partial_record_to_proto)
.map_or_else(|| (Default::default(), false), |tags| (tags, true));
Self {
uid: Some(cedar_policy::proto::models::EntityUid::from(entity.uid())),
uid: Some(cedar_proto::EntityUid::from(entity.uid())),
attrs,
ancestors,
tags,
Expand Down
33 changes: 20 additions & 13 deletions cedar-lean/Cedar/TPE/BatchedEvaluator.lean
Original file line number Diff line number Diff line change
Expand Up @@ -44,6 +44,7 @@ The batched evaluation loop for a single residual expression.
3. Exits if a value has been found or it hits the maximum iteration limit
-/
def batchedEvaluateLoop
(env : TypeEnv)
(residual : Residual)
(req : Request)
(loader : EntityLoader)
Expand All @@ -52,12 +53,12 @@ def batchedEvaluateLoop
| 0 => residual
| n + 1 =>
let toLoad := residual.allLiteralUIDs.filter (λ uid => (store.find? uid).isNone)
let newEntities := ((loader toLoad).mapOnValues MaybeEntityData.asPartial)
let newEntities := SlicedEntities.asPartial env.ets (loader toLoad)
let newStore := newEntities ++ store

match Cedar.TPE.evaluate residual req.asPartialRequest newStore with
match Cedar.TPE.evaluate env residual (req.asPartialRequest env.reqty.context) newStore with
| .val v _ty => .val v _ty
| newRes => batchedEvaluateLoop newRes req loader newStore n
| newRes => batchedEvaluateLoop env newRes req loader newStore n

def actionEntities (acts : ActionSchema) : PartialEntities :=
Map.make (acts.toList.map λ (uid, entry) =>
Expand All @@ -70,19 +71,21 @@ Performs a maximum of `iter` number of calls to `loader`,
but may perform fewer when a value is found.
-/
def batchedEvaluate
(acts : ActionSchema)
(env : TypeEnv)
(x : TypedExpr)
(req : Request)
(loader : EntityLoader)
(iters : Nat)
: Residual :=
let residual := Cedar.TPE.evaluate x.toResidual req.asPartialRequest (actionEntities acts)
batchedEvaluateLoop residual req loader (actionEntities acts) iters
let residual := Cedar.TPE.evaluate env x.toResidual
(req.asPartialRequest env.reqty.context) (actionEntities env.acts)
batchedEvaluateLoop env residual req loader (actionEntities env.acts) iters

/--
The batched authorization loop for authorization over a list of policies.
-/
def batchedAuthorizeLoop
(env : TypeEnv)
(residuals : List ResidualPolicy) (req : Request) (loader : EntityLoader)
(store : PartialEntities) (n : Nat)
: Response
Expand All @@ -94,12 +97,12 @@ def batchedAuthorizeLoop
| 0 => resp
| n + 1 =>
let toLoad := residuals.mapUnion (λ rp : ResidualPolicy => rp.residual.allLiteralUIDs)|>.filter (λ uid => (store.find? uid).isNone)
let newEntities := ((loader toLoad).mapOnValues MaybeEntityData.asPartial)
let newEntities := SlicedEntities.asPartial env.ets (loader toLoad)
let newStore := newEntities ++ store

let residuals : List ResidualPolicy := residuals.map λ rp =>
⟨rp.id, rp.effect, Cedar.TPE.evaluate rp.residual req.asPartialRequest newStore⟩
batchedAuthorizeLoop residuals req loader newStore n
⟨rp.id, rp.effect, Cedar.TPE.evaluate env rp.residual (req.asPartialRequest env.reqty.context) newStore⟩
batchedAuthorizeLoop env residuals req loader newStore n

/--
Evaluate an authorization request using an EntityLoader instead of a full Entities store.
Expand All @@ -113,10 +116,14 @@ def batchedAuthorize
(loader : EntityLoader)
(iters : Nat)
: Except Error Response := do
let residualPolicies ← policies.mapM (λ p => do
pure ⟨p.id, p.effect,
← evaluatePolicy schema p req.asPartialRequest (actionEntities schema.acts)⟩)
pure (batchedAuthorizeLoop residualPolicies req loader (actionEntities schema.acts) iters)
match schema.environment? req.principal.ty req.resource.ty req.action with
| .none => .error .invalidEnvironment
| .some env =>
let residualPolicies ← policies.mapM (λ p => do
pure ⟨p.id, p.effect,
← evaluatePolicy schema p (req.asPartialRequest env.reqty.context)
(actionEntities schema.acts)⟩)
pure (batchedAuthorizeLoop env residualPolicies req loader (actionEntities schema.acts) iters)

/--
Create an entity loader for a given entity store.
Expand Down
Loading
Loading