klieo-ops 0.3.0

Operational layer above klieo-core: supervisor, governor, gates, escalation, worklog, handoff.
Documentation
//! Cedar-backed policy gate.
//!
//! `PolicyGate::from_cedar_source` accepts a Cedar policy source string and
//! evaluates it for every effectful tool invocation. Each `GateRequest` is
//! lifted into a Cedar `Request` of shape:
//!   - principal: `Tool::"<tool_name>"`
//!   - action:    `Action::"invoke"`
//!   - resource:  `Args::"<tool_name>"` with the JSON args as Cedar attrs.
//!
//! JSON args are lifted into the Cedar context as follows: `bool` → Cedar
//! bool; `i64`-compatible numbers → Cedar long; floats → Cedar decimal
//! (4-fractional-digit precision); strings → Cedar string; arrays → Cedar
//! set (recursively lifted); objects → Cedar record (recursively lifted);
//! `null` and out-of-range `u64` → unrepresentable, treated as
//! `GateDecision::Deny { code: "cedar.context.unrepresentable" }`.
//!
//! Phase A scope: source-string load + evaluate. Hot-reload, bundle sha
//! digest in audit, and per-tenant bundle scoping land in Phase B.

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;

/// Errors raised by `PolicyGate::from_cedar_source`.
#[derive(Debug, Error)]
#[non_exhaustive]
pub enum PolicyGateError {
    /// The Cedar source did not parse.
    #[error("cedar policy parse: {0}")]
    Parse(String),
}

/// Cedar policy bundle gate.
pub struct PolicyGate {
    policies: PolicySet,
    authorizer: Authorizer,
}

impl PolicyGate {
    /// Construct from a Cedar source string.
    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,
        };

        // Lift all top-level args into Cedar context (recursive for nested
        // objects and arrays). Any unrepresentable value (null, out-of-range
        // u64) causes an immediate Deny so we fail-CLOSED rather than
        // silently dropping an attribute a forbid rule may reference.
        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() {
                // u64 values >= 2^63 cannot fit in Cedar long; treat as
                // unrepresentable so the caller can deny fail-CLOSED.
                i64::try_from(u).ok().map(RE::new_long)
            } else {
                // Cedar decimal extension stores fixed-point with up to 4
                // fractional digits. Format to exactly that precision so
                // `decimal("0.5000") == context.field` evaluations match.
                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,
    }
}