Skip to main content

cedar_policy_symcc/symcc/
ext.rs

1/*
2 * Copyright Cedar Contributors
3 *
4 * Licensed under the Apache License, Version 2.0 (the "License");
5 * you may not use this file except in compliance with the License.
6 * You may obtain a copy of the License at
7 *
8 *      https://www.apache.org/licenses/LICENSE-2.0
9 *
10 * Unless required by applicable law or agreed to in writing, software
11 * distributed under the License is distributed on an "AS IS" BASIS,
12 * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13 * See the License for the specific language governing permissions and
14 * limitations under the License.
15 */
16
17//! Extension values in SymCC.
18
19use 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/// Internal representation of extension values.
35#[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/// Errors in [`Ext`] operations.
45#[derive(Debug, Diagnostic, Error)]
46pub enum ExtError {
47    /// Failed to convert from the given [`Value`].
48    #[error("fail to convert value to an extension term: {0}")]
49    FromValue(Value),
50    /// Failed to convert from the given [`RestrictedExpr`].
51    #[error("fail to convert expression to an extension term: {0}")]
52    FromRestrictedExpr(RestrictedExpr),
53    /// Evaluation error when converting to a value.
54    #[error("evaluation error when converting to value")]
55    EvaluationError(#[from] EvaluationError),
56    /// Extension function not found.
57    #[error("extension function `{0}` not found")]
58    ExtensionFunctionNotFound(String),
59    /// Extension value evaluates to a partial value.
60    #[error("extension function returned a partial value")]
61    UnsupportedPartialValue,
62    /// Failed to parse extension function name.
63    #[error("failed to parse extension function name")]
64    ExtensionFunctionParseError,
65    /// Datetime error.
66    #[error("datetime error")]
67    DatetimeError(#[from] DatetimeError),
68    /// Decimal error.
69    #[error("decimal error")]
70    DecimalError(#[from] DecimalError),
71    /// IP error.
72    #[error("IP error")]
73    IPError(#[from] IPError),
74}
75
76impl Ext {
77    /// Parses a `decimal` extension value from a string.
78    pub fn parse_decimal(s: &str) -> Result<Ext, ExtError> {
79        Ok(Decimal::from_str(s).map(|d| Ext::Decimal { d })?)
80    }
81
82    /// Parses a `datetime` extension value from a string.
83    pub fn parse_datetime(s: &str) -> Result<Ext, ExtError> {
84        Ok(Datetime::from_str(s).map(|dt| Ext::Datetime { dt })?)
85    }
86
87    /// Parses a `duration` extension value from a string.
88    pub fn parse_duration(s: &str) -> Result<Ext, ExtError> {
89        Ok(Duration::from_str(s).map(|d| Ext::Duration { d })?)
90    }
91
92    /// Parses an `ipaddr` extension value from a string.
93    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    /// Helper function to convert a [`RestrictedExpr`] to an [`Ext`].
100    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        // Recover the string representation of supported extension values
105        // and then convert them to corresponding `Term`s.
106        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            // A `datetime` value is sometimes represented as `datetime(<epoch>).offset(<...>)`
111            ("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
142/// Rust SymCC and Rust Cedar use different representations
143/// of extension values (whereas the Lean model uses the same),
144/// so we need these utility functions to convert between them.
145impl 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
165/// A utility function to call an extension function
166fn 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                // First construct `datetime("1970-01-01")`
189                let epoch = call_extension_func(
190                    &datetime::extension(),
191                    "datetime",
192                    &["1970-01-01".into()],
193                )?;
194                // Then construct the actual datetime as an offset duration
195                let offset: i64 = dt.into();
196                let offset = call_extension_func(
197                    &datetime::extension(),
198                    "duration",
199                    &[format!("{}ms", offset).into()],
200                )?;
201                // Finally call the offset function to construct the right datetime value
202                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}