Skip to main content

protocheck_core/validators/
cel.rs

1use std::{sync::LazyLock, vec};
2
3use cel::{Context, Program, Value as CelValue};
4use chrono::Utc;
5use proto_types::cel::CelConversionError;
6
7use crate::{
8  field_data::FieldContext,
9  protovalidate::{FieldPath, FieldPathElement, Violation},
10  validators::static_data::base_violations::create_violation,
11  ProtoType,
12};
13
14pub struct CelRule {
15  pub id: &'static str,
16  pub error_message: &'static str,
17  pub program: &'static Program,
18  pub item_full_name: &'static str,
19}
20
21pub fn validate_cel_field_with_val(
22  field_context: &FieldContext,
23  rule: CelRule,
24  value: CelValue,
25) -> Result<(), Violation>
26where
27{
28  let CelRule {
29    id: rule_id,
30    error_message,
31    program,
32    ..
33  } = rule;
34
35  let error_prefix = format!(
36    "Error during Cel validation for field {}:",
37    field_context.proto_name
38  );
39
40  let mut cel_context = Context::default();
41  cel_context.add_variable_from_value("now", CelValue::Timestamp(Utc::now().into()));
42
43  cel_context.add_variable_from_value("this", value);
44
45  let result = program.execute(&cel_context);
46
47  match result {
48    Ok(value) => {
49      if let CelValue::Bool(bool_value) = value {
50        if bool_value {
51          Ok(())
52        } else {
53          Err(create_violation(
54            field_context,
55            &CEL_VIOLATION,
56            rule_id,
57            error_message,
58          ))
59        }
60      } else {
61        eprintln!(
62          "{} expected boolean result from expression, got `{:?}`",
63          error_prefix,
64          value.type_of()
65        );
66        Err(create_violation(
67          field_context,
68          &CEL_VIOLATION,
69          "internal server error",
70          "internal server error",
71        ))
72      }
73    }
74    Err(e) => {
75      eprintln!("{} {:?}", error_prefix, e);
76      Err(create_violation(
77        field_context,
78        &CEL_VIOLATION,
79        "internal server error",
80        "internal server error",
81      ))
82    }
83  }
84}
85
86pub fn validate_cel_field_try_into<T>(
87  field_context: &FieldContext,
88  rule: CelRule,
89  value: T,
90) -> Result<(), Violation>
91where
92  T: TryInto<CelValue> + Clone,
93  <T as std::convert::TryInto<cel::Value>>::Error: std::fmt::Display,
94{
95  let cel_conversion: Result<CelValue, _> = value.try_into();
96
97  match cel_conversion {
98    Ok(cel_val) => validate_cel_field_with_val(field_context, rule, cel_val),
99    Err(e) => {
100      eprintln!(
101        "Failed to convert field {} to Cel value: {}",
102        rule.item_full_name, e
103      );
104
105      Err(create_violation(
106        field_context,
107        &CEL_VIOLATION,
108        "internal server error",
109        "internal server error",
110      ))
111    }
112  }
113}
114
115pub fn validate_cel_message<T>(
116  parent_elements: &[FieldPathElement],
117  rule: CelRule,
118  value: T,
119) -> Result<(), Violation>
120where
121  T: TryInto<CelValue, Error = CelConversionError>,
122{
123  let CelRule {
124    id: rule_id,
125    error_message,
126    program,
127    item_full_name: message_name,
128  } = rule;
129
130  let error_prefix = format!("Error during Cel validation for message {}:", message_name);
131
132  let mut cel_context = Context::default();
133  cel_context.add_variable_from_value("now", CelValue::Timestamp(Utc::now().into()));
134
135  let cel_conversion: Result<CelValue, CelConversionError> = value.try_into();
136
137  match cel_conversion {
138    Ok(cel_val) => {
139      cel_context.add_variable_from_value("this", cel_val);
140      let result = program.execute(&cel_context);
141
142      match result {
143        Ok(value) => {
144          if let CelValue::Bool(bool_value) = value {
145            if bool_value {
146              Ok(())
147            } else {
148              Err(create_cel_message_violation(
149                rule_id,
150                error_message,
151                parent_elements,
152              ))
153            }
154          } else {
155            eprintln!(
156              "{} expected boolean result from expression, got `{:?}`",
157              error_prefix,
158              value.type_of()
159            );
160            Err(create_cel_message_violation(
161              "internal_server_error",
162              "internal server error",
163              parent_elements,
164            ))
165          }
166        }
167        Err(e) => {
168          eprintln!("{} program failed to compile: {:?}", error_prefix, e);
169          Err(create_cel_message_violation(
170            "internal_server_error",
171            "internal server error",
172            parent_elements,
173          ))
174        }
175      }
176    }
177    Err(e) => {
178      eprintln!(
179        "{} could not convert message to Cel value: {:?}",
180        error_prefix, e
181      );
182      Err(create_cel_message_violation(
183        "internal_server_error",
184        "internal server error",
185        parent_elements,
186      ))
187    }
188  }
189}
190
191fn create_cel_message_violation(
192  rule_id: &str,
193  error_message: &str,
194  parent_elements: &[FieldPathElement],
195) -> Violation {
196  let is_nested = !parent_elements.is_empty();
197  let field_path = is_nested.then(|| FieldPath {
198    elements: parent_elements.to_vec(),
199  });
200
201  Violation {
202    message: Some(error_message.to_string()),
203    rule_id: Some(rule_id.to_string()),
204    rule: Some(FieldPath {
205      elements: CEL_VIOLATION.clone(),
206    }),
207    field: field_path,
208    for_key: None,
209  }
210}
211
212static CEL_VIOLATION: LazyLock<Vec<FieldPathElement>> = LazyLock::new(|| {
213  vec![FieldPathElement {
214    field_name: Some("cel".to_string()),
215    field_number: Some(23),
216    field_type: Some(ProtoType::Message as i32),
217    key_type: None,
218    value_type: None,
219    subscript: None,
220  }]
221});