use super::trait_::{Gate, GateDecision, GateRequest};
use async_trait::async_trait;
use cedar_policy::{Authorizer, Context, Entities, EntityUid, PolicySet, Request};
use std::str::FromStr;
use thiserror::Error;
#[derive(Debug, Error)]
#[non_exhaustive]
pub enum PolicyGateError {
#[error("cedar policy parse: {0}")]
Parse(String),
}
pub struct PolicyGate {
policies: PolicySet,
authorizer: Authorizer,
}
impl PolicyGate {
pub fn from_cedar_source(src: &str) -> Result<Self, PolicyGateError> {
let policies =
PolicySet::from_str(src).map_err(|e| PolicyGateError::Parse(e.to_string()))?;
Ok(Self {
policies,
authorizer: Authorizer::new(),
})
}
fn build_uid(prefix: &str, name: &str) -> Result<EntityUid, GateDecision> {
EntityUid::from_str(&format!("{prefix}::\"{name}\"")).map_err(|e| GateDecision::Deny {
code: format!("cedar.{}.invalid", prefix.to_ascii_lowercase()),
reason: format!("could not form {prefix} uid: {e}"),
})
}
}
#[async_trait]
impl Gate for PolicyGate {
async fn evaluate(&self, req: GateRequest) -> GateDecision {
let principal = match Self::build_uid("Tool", &req.tool_name) {
Ok(u) => u,
Err(d) => return d,
};
let action = match EntityUid::from_str("Action::\"invoke\"") {
Ok(u) => u,
Err(e) => {
return GateDecision::Deny {
code: "cedar.action.invalid".into(),
reason: format!("could not form action uid: {e}"),
};
}
};
let resource = match Self::build_uid("Args", &req.tool_name) {
Ok(u) => u,
Err(d) => return d,
};
let ctx_pairs: Vec<(String, cedar_policy::RestrictedExpression)> = match &req.args {
serde_json::Value::Object(map) => {
let mut pairs = Vec::with_capacity(map.len());
for (k, v) in map {
match json_to_cedar_expr(v) {
Some(expr) => pairs.push((k.clone(), expr)),
None => {
return GateDecision::Deny {
code: "cedar.context.unrepresentable".into(),
reason: format!("arg `{k}` is not Cedar-representable"),
};
}
}
}
pairs
}
_ => Vec::new(),
};
let context = match Context::from_pairs(ctx_pairs) {
Ok(c) => c,
Err(e) => {
return GateDecision::Deny {
code: "cedar.context.invalid".into(),
reason: format!("context build: {e}"),
};
}
};
let request = match Request::new(principal, action, resource, context, None) {
Ok(r) => r,
Err(e) => {
return GateDecision::Deny {
code: "cedar.request.invalid".into(),
reason: format!("request build: {e}"),
};
}
};
let entities = Entities::empty();
let response = self
.authorizer
.is_authorized(&request, &self.policies, &entities);
match response.decision() {
cedar_policy::Decision::Allow => GateDecision::Allow,
cedar_policy::Decision::Deny => GateDecision::Deny {
code: "cedar.deny".into(),
reason: "cedar deny".to_string(),
},
}
}
fn name(&self) -> &'static str {
"PolicyGate"
}
}
fn json_to_cedar_expr(v: &serde_json::Value) -> Option<cedar_policy::RestrictedExpression> {
use cedar_policy::RestrictedExpression as RE;
match v {
serde_json::Value::Bool(b) => Some(RE::new_bool(*b)),
serde_json::Value::Number(n) => {
if let Some(i) = n.as_i64() {
Some(RE::new_long(i))
} else if let Some(u) = n.as_u64() {
i64::try_from(u).ok().map(RE::new_long)
} else {
n.as_f64().map(|f| RE::new_decimal(format!("{f:.4}")))
}
}
serde_json::Value::String(s) => Some(RE::new_string(s.clone())),
serde_json::Value::Array(xs) => {
let elems: Option<Vec<RE>> = xs.iter().map(json_to_cedar_expr).collect();
elems.map(RE::new_set)
}
serde_json::Value::Object(map) => {
let pairs: Option<Vec<(String, RE)>> = map
.iter()
.map(|(k, v)| json_to_cedar_expr(v).map(|expr| (k.clone(), expr)))
.collect();
pairs.and_then(|p| RE::new_record(p).ok())
}
serde_json::Value::Null => None,
}
}