cedar_policy_symcc/symcc/
ext.rs1use std::str::FromStr;
20
21use cedar_policy::EvaluationError;
22use cedar_policy_core::ast::{Extension, Name, PartialValue, RestrictedExpr, Value, ValueKind};
23use miette::Diagnostic;
24use thiserror::Error;
25
26use crate::err::IPError;
27use crate::extension_types::datetime::DatetimeError;
28use crate::extension_types::decimal::DecimalError;
29
30use super::extension_types::datetime::{Datetime, Duration};
31use super::extension_types::decimal::Decimal;
32use super::extension_types::ipaddr::IPNet;
33
34#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord)]
36#[expect(missing_docs, reason = "self-explanatory")]
37pub enum Ext {
38 Decimal { d: Decimal },
39 Ipaddr { ip: IPNet },
40 Datetime { dt: Datetime },
41 Duration { d: Duration },
42}
43
44#[derive(Debug, Diagnostic, Error)]
46pub enum ExtError {
47 #[error("fail to convert value to an extension term: {0}")]
49 FromValue(Value),
50 #[error("fail to convert expression to an extension term: {0}")]
52 FromRestrictedExpr(RestrictedExpr),
53 #[error("evaluation error when converting to value")]
55 EvaluationError(#[from] EvaluationError),
56 #[error("extension function `{0}` not found")]
58 ExtensionFunctionNotFound(String),
59 #[error("extension function returned a partial value")]
61 UnsupportedPartialValue,
62 #[error("failed to parse extension function name")]
64 ExtensionFunctionParseError,
65 #[error("datetime error")]
67 DatetimeError(#[from] DatetimeError),
68 #[error("decimal error")]
70 DecimalError(#[from] DecimalError),
71 #[error("IP error")]
73 IPError(#[from] IPError),
74}
75
76impl Ext {
77 pub fn parse_decimal(s: &str) -> Result<Ext, ExtError> {
79 Ok(Decimal::from_str(s).map(|d| Ext::Decimal { d })?)
80 }
81
82 pub fn parse_datetime(s: &str) -> Result<Ext, ExtError> {
84 Ok(Datetime::from_str(s).map(|dt| Ext::Datetime { dt })?)
85 }
86
87 pub fn parse_duration(s: &str) -> Result<Ext, ExtError> {
89 Ok(Duration::from_str(s).map(|d| Ext::Duration { d })?)
90 }
91
92 pub fn parse_ip(s: &str) -> Result<Ext, ExtError> {
94 Ok(IPNet::from_str(s).map(|ip| Ext::Ipaddr { ip })?)
95 }
96}
97
98impl Ext {
99 fn from_ext_value(rexp: &RestrictedExpr) -> Option<Self> {
101 let (name, args) = rexp.as_extn_fn_call()?;
102 let args = args.collect::<Vec<_>>();
103
104 match (name.as_ref().to_string().as_str(), args.as_slice()) {
107 ("decimal", &[arg]) => Self::parse_decimal(arg.as_string()?.as_str()).ok(),
108 ("duration", &[arg]) => Self::parse_duration(arg.as_string()?.as_str()).ok(),
109 ("datetime", &[arg]) => Self::parse_datetime(arg.as_string()?.as_str()).ok(),
110 ("offset", &[arg1, arg2]) => {
112 let (arg1_name, arg1_args) = arg1.as_extn_fn_call()?;
113 let (arg2_name, arg2_args) = arg2.as_extn_fn_call()?;
114 let arg1_args = arg1_args.collect::<Vec<_>>();
115 let arg2_args = arg2_args.collect::<Vec<_>>();
116 if arg1_name.as_ref().to_string() != "datetime"
117 || arg1_args.len() != 1
118 || arg2_name.as_ref().to_string() != "duration"
119 || arg2_args.len() != 1
120 {
121 return None;
122 }
123
124 #[expect(
125 clippy::indexing_slicing,
126 reason = "arg1_args.len() == 1 thus indexing by 0 should not panic"
127 )]
128 let dt = Datetime::from_str(arg1_args[0].as_string()?.as_str()).ok()?;
129 #[expect(
130 clippy::indexing_slicing,
131 reason = "arg2_args.len() == 1 thus indexing by 0 should not panic"
132 )]
133 let d = Duration::from_str(arg2_args[0].as_string()?.as_str()).ok()?;
134 Some(Ext::Datetime { dt: dt.offset(&d)? })
135 }
136 ("ip", &[arg]) => Self::parse_ip(arg.as_string()?.as_str()).ok(),
137 _ => None,
138 }
139 }
140}
141
142impl TryFrom<&RestrictedExpr> for Ext {
146 type Error = ExtError;
147
148 fn try_from(rexp: &RestrictedExpr) -> Result<Self, Self::Error> {
149 Self::from_ext_value(rexp).ok_or_else(|| ExtError::FromRestrictedExpr(rexp.clone()))
150 }
151}
152
153impl TryFrom<&Value> for Ext {
154 type Error = ExtError;
155
156 fn try_from(v: &Value) -> Result<Self, Self::Error> {
157 let ValueKind::ExtensionValue(ext) = v.value_kind() else {
158 return Err(ExtError::FromValue(v.clone()));
159 };
160 let rexp = RestrictedExpr::from(ext.as_ref().clone());
161 Self::from_ext_value(&rexp).ok_or_else(|| ExtError::FromValue(v.clone()))
162 }
163}
164
165fn call_extension_func(ext: &Extension, name: &str, args: &[Value]) -> Result<Value, ExtError> {
167 let name = Name::parse_unqualified_name(name).or(Err(ExtError::ExtensionFunctionParseError))?;
168 match ext
169 .get_func(&name)
170 .ok_or_else(|| ExtError::ExtensionFunctionNotFound(name.to_string()))?
171 .call(args)?
172 {
173 PartialValue::Value(v) => Ok(v),
174 _ => Err(ExtError::UnsupportedPartialValue),
175 }
176}
177
178impl TryFrom<&Ext> for Value {
179 type Error = ExtError;
180
181 fn try_from(ext: &Ext) -> Result<Self, Self::Error> {
182 use cedar_policy_core::extensions::{datetime, decimal, ipaddr};
183 match ext {
184 Ext::Decimal { d } => {
185 call_extension_func(&decimal::extension(), "decimal", &[format!("{}", d).into()])
186 }
187 Ext::Datetime { dt } => {
188 let epoch = call_extension_func(
190 &datetime::extension(),
191 "datetime",
192 &["1970-01-01".into()],
193 )?;
194 let offset: i64 = dt.into();
196 let offset = call_extension_func(
197 &datetime::extension(),
198 "duration",
199 &[format!("{}ms", offset).into()],
200 )?;
201 call_extension_func(&datetime::extension(), "offset", &[epoch, offset])
203 }
204 Ext::Duration { d } => {
205 let offset: i64 = d.into();
206 call_extension_func(
207 &datetime::extension(),
208 "duration",
209 &[format!("{}ms", offset).into()],
210 )
211 }
212 Ext::Ipaddr { ip } => {
213 call_extension_func(&ipaddr::extension(), "ip", &[format!("{}", ip).into()])
214 }
215 }
216 }
217}