protocheck_core/validators/
cel.rs1use 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});