use std::str::FromStr;
use cedar_policy::EvaluationError;
use cedar_policy_core::ast::{Extension, Name, PartialValue, RestrictedExpr, Value, ValueKind};
use miette::Diagnostic;
use thiserror::Error;
use crate::err::IPError;
use crate::extension_types::datetime::DatetimeError;
use crate::extension_types::decimal::DecimalError;
use super::extension_types::datetime::{Datetime, Duration};
use super::extension_types::decimal::Decimal;
use super::extension_types::ipaddr::IPNet;
#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord)]
#[expect(missing_docs, reason = "self-explanatory")]
pub enum Ext {
Decimal { d: Decimal },
Ipaddr { ip: IPNet },
Datetime { dt: Datetime },
Duration { d: Duration },
}
#[derive(Debug, Diagnostic, Error)]
pub enum ExtError {
#[error("fail to convert value to an extension term: {0}")]
FromValue(Value),
#[error("fail to convert expression to an extension term: {0}")]
FromRestrictedExpr(RestrictedExpr),
#[error("evaluation error when converting to value")]
EvaluationError(#[from] EvaluationError),
#[error("extension function `{0}` not found")]
ExtensionFunctionNotFound(String),
#[error("extension function returned a partial value")]
UnsupportedPartialValue,
#[error("failed to parse extension function name")]
ExtensionFunctionParseError,
#[error("datetime error")]
DatetimeError(#[from] DatetimeError),
#[error("decimal error")]
DecimalError(#[from] DecimalError),
#[error("IP error")]
IPError(#[from] IPError),
}
impl Ext {
pub fn parse_decimal(s: &str) -> Result<Ext, ExtError> {
Ok(Decimal::from_str(s).map(|d| Ext::Decimal { d })?)
}
pub fn parse_datetime(s: &str) -> Result<Ext, ExtError> {
Ok(Datetime::from_str(s).map(|dt| Ext::Datetime { dt })?)
}
pub fn parse_duration(s: &str) -> Result<Ext, ExtError> {
Ok(Duration::from_str(s).map(|d| Ext::Duration { d })?)
}
pub fn parse_ip(s: &str) -> Result<Ext, ExtError> {
Ok(IPNet::from_str(s).map(|ip| Ext::Ipaddr { ip })?)
}
}
impl Ext {
fn from_ext_value(rexp: &RestrictedExpr) -> Option<Self> {
let (name, args) = rexp.as_extn_fn_call()?;
let args = args.collect::<Vec<_>>();
match (name.as_ref().to_string().as_str(), args.as_slice()) {
("decimal", &[arg]) => Self::parse_decimal(arg.as_string()?.as_str()).ok(),
("duration", &[arg]) => Self::parse_duration(arg.as_string()?.as_str()).ok(),
("datetime", &[arg]) => Self::parse_datetime(arg.as_string()?.as_str()).ok(),
("offset", &[arg1, arg2]) => {
let (arg1_name, arg1_args) = arg1.as_extn_fn_call()?;
let (arg2_name, arg2_args) = arg2.as_extn_fn_call()?;
let arg1_args = arg1_args.collect::<Vec<_>>();
let arg2_args = arg2_args.collect::<Vec<_>>();
if arg1_name.as_ref().to_string() != "datetime"
|| arg1_args.len() != 1
|| arg2_name.as_ref().to_string() != "duration"
|| arg2_args.len() != 1
{
return None;
}
#[expect(
clippy::indexing_slicing,
reason = "arg1_args.len() == 1 thus indexing by 0 should not panic"
)]
let dt = Datetime::from_str(arg1_args[0].as_string()?.as_str()).ok()?;
#[expect(
clippy::indexing_slicing,
reason = "arg2_args.len() == 1 thus indexing by 0 should not panic"
)]
let d = Duration::from_str(arg2_args[0].as_string()?.as_str()).ok()?;
Some(Ext::Datetime { dt: dt.offset(&d)? })
}
("ip", &[arg]) => Self::parse_ip(arg.as_string()?.as_str()).ok(),
_ => None,
}
}
}
impl TryFrom<&RestrictedExpr> for Ext {
type Error = ExtError;
fn try_from(rexp: &RestrictedExpr) -> Result<Self, Self::Error> {
Self::from_ext_value(rexp).ok_or_else(|| ExtError::FromRestrictedExpr(rexp.clone()))
}
}
impl TryFrom<&Value> for Ext {
type Error = ExtError;
fn try_from(v: &Value) -> Result<Self, Self::Error> {
let ValueKind::ExtensionValue(ext) = v.value_kind() else {
return Err(ExtError::FromValue(v.clone()));
};
let rexp = RestrictedExpr::from(ext.as_ref().clone());
Self::from_ext_value(&rexp).ok_or_else(|| ExtError::FromValue(v.clone()))
}
}
fn call_extension_func(ext: &Extension, name: &str, args: &[Value]) -> Result<Value, ExtError> {
let name = Name::parse_unqualified_name(name).or(Err(ExtError::ExtensionFunctionParseError))?;
match ext
.get_func(&name)
.ok_or_else(|| ExtError::ExtensionFunctionNotFound(name.to_string()))?
.call(args)?
{
PartialValue::Value(v) => Ok(v),
_ => Err(ExtError::UnsupportedPartialValue),
}
}
impl TryFrom<&Ext> for Value {
type Error = ExtError;
fn try_from(ext: &Ext) -> Result<Self, Self::Error> {
use cedar_policy_core::extensions::{datetime, decimal, ipaddr};
match ext {
Ext::Decimal { d } => {
call_extension_func(&decimal::extension(), "decimal", &[format!("{}", d).into()])
}
Ext::Datetime { dt } => {
let epoch = call_extension_func(
&datetime::extension(),
"datetime",
&["1970-01-01".into()],
)?;
let offset: i64 = dt.into();
let offset = call_extension_func(
&datetime::extension(),
"duration",
&[format!("{}ms", offset).into()],
)?;
call_extension_func(&datetime::extension(), "offset", &[epoch, offset])
}
Ext::Duration { d } => {
let offset: i64 = d.into();
call_extension_func(
&datetime::extension(),
"duration",
&[format!("{}ms", offset).into()],
)
}
Ext::Ipaddr { ip } => {
call_extension_func(&ipaddr::extension(), "ip", &[format!("{}", ip).into()])
}
}
}
}