1use std::collections::{BTreeMap, BTreeSet};
2
3use serde::Deserialize;
4use serde_json::{Map, Value};
5
6use super::{BoundAst, ContractVersion, NaNPolicyAst};
7use crate::{ErrorCode, ProofFrameError};
8
9const DEFAULT_MAX_FINDINGS: usize = 100;
10
11#[derive(Debug, Clone, Copy, Eq, PartialEq, Deserialize)]
12#[serde(rename_all = "snake_case")]
13pub enum PrimitiveTypeAst {
14 Boolean,
15 Int8,
16 Int16,
17 Int32,
18 Int64,
19 Uint8,
20 Uint16,
21 Uint32,
22 Uint64,
23 Float32,
24 Float64,
25 Date32,
26 Date64,
27 Utf8,
28 LargeUtf8,
29 Utf8View,
30 Binary,
31 LargeBinary,
32 BinaryView,
33}
34
35#[derive(Debug, Clone, Eq, PartialEq, Deserialize)]
36#[serde(tag = "name", rename_all = "snake_case")]
37pub enum ParameterizedTypeAst {
38 Decimal128 {
39 precision: u8,
40 scale: i8,
41 },
42 Timestamp {
43 unit: TimeUnitAst,
44 #[serde(default)]
45 timezone: Option<String>,
46 },
47}
48
49#[derive(Debug, Clone, Eq, PartialEq, Deserialize)]
50#[serde(untagged)]
51pub enum TypeAst {
52 Primitive(PrimitiveTypeAst),
53 Parameterized(ParameterizedTypeAst),
54}
55
56#[derive(Debug, Clone, Copy, Eq, PartialEq, Deserialize)]
57#[serde(rename_all = "snake_case")]
58pub enum TimeUnitAst {
59 S,
60 Ms,
61 Us,
62 Ns,
63}
64
65#[derive(Debug, Clone, PartialEq, Deserialize)]
66#[serde(deny_unknown_fields)]
67pub struct RuleAstV2 {
68 #[serde(default)]
69 pub required: bool,
70 #[serde(default)]
71 pub not_null: bool,
72 #[serde(default)]
73 pub unique: bool,
74 #[serde(rename = "type")]
75 pub expected_type: Option<TypeAst>,
76 pub min: Option<BoundAst>,
77 pub max: Option<BoundAst>,
78 pub nan: Option<NaNPolicyAst>,
79 pub pattern: Option<String>,
80 pub allowed: Option<BTreeSet<String>>,
81}
82
83#[derive(Debug, Clone, PartialEq, Deserialize)]
84#[serde(untagged)]
85pub enum OperandAst {
86 Column(ColumnOperandAst),
87 Literal(LiteralOperandAst),
88}
89
90#[derive(Debug, Clone, PartialEq, Deserialize)]
91#[serde(deny_unknown_fields)]
92pub struct ColumnOperandAst {
93 pub column: String,
94}
95
96#[derive(Debug, Clone, PartialEq, Deserialize)]
97#[serde(deny_unknown_fields)]
98pub struct LiteralOperandAst {
99 pub literal: Value,
100}
101
102#[derive(Debug, Clone, Copy, Eq, PartialEq, Deserialize)]
103#[serde(rename_all = "snake_case")]
104pub enum CompareOpAst {
105 Eq,
106 Ne,
107 Lt,
108 Lte,
109 Gt,
110 Gte,
111}
112
113#[derive(Debug, Clone, Copy, Default, Eq, PartialEq, Deserialize)]
114#[serde(rename_all = "snake_case")]
115pub enum NullPolicyAst {
116 #[default]
117 Skip,
118 Fail,
119 Equal,
120}
121
122#[derive(Debug, Clone, PartialEq, Deserialize)]
123#[serde(deny_unknown_fields)]
124pub struct CompareAst {
125 pub left: OperandAst,
126 pub op: CompareOpAst,
127 pub right: OperandAst,
128 #[serde(default)]
129 pub nulls: NullPolicyAst,
130}
131
132#[derive(Debug, Clone, PartialEq, Deserialize)]
133#[serde(deny_unknown_fields)]
134pub struct AssertionAst {
135 pub column: String,
136 #[serde(default)]
137 pub not_null: bool,
138 pub min: Option<BoundAst>,
139 pub max: Option<BoundAst>,
140 pub nan: Option<NaNPolicyAst>,
141 pub pattern: Option<String>,
142 pub allowed: Option<BTreeSet<String>>,
143}
144
145#[derive(Debug, Clone, PartialEq, Deserialize)]
146#[serde(deny_unknown_fields)]
147pub struct RowRuleAst {
148 pub name: String,
149 pub compare: Option<CompareAst>,
150 pub when: Option<CompareAst>,
151 #[serde(rename = "assert")]
152 pub assertion: Option<AssertionAst>,
153}
154
155#[derive(Debug, Clone, Default, PartialEq, Deserialize)]
156#[serde(deny_unknown_fields)]
157pub struct CountRangeAst {
158 pub exact: Option<u64>,
159 pub min: Option<u64>,
160 pub max: Option<u64>,
161}
162
163#[derive(Debug, Clone, Default, PartialEq, Deserialize)]
164#[serde(deny_unknown_fields)]
165pub struct RatioRangeAst {
166 pub min: Option<f64>,
167 pub max: Option<f64>,
168}
169
170#[derive(Debug, Clone, Copy, Default, Eq, PartialEq, Deserialize)]
171#[serde(rename_all = "snake_case")]
172pub enum CompositeNullPolicyAst {
173 #[default]
174 Equal,
175 Reject,
176}
177
178#[derive(Debug, Clone, PartialEq, Deserialize)]
179#[serde(deny_unknown_fields)]
180pub struct CompositeUniqueAst {
181 pub name: String,
182 pub columns: Vec<String>,
183 #[serde(default)]
184 pub nulls: CompositeNullPolicyAst,
185}
186
187#[derive(Debug, Clone, Default, PartialEq, Deserialize)]
188#[serde(deny_unknown_fields)]
189pub struct DatasetRulesAst {
190 pub row_count: Option<CountRangeAst>,
191 #[serde(default)]
192 pub composite_unique: Vec<CompositeUniqueAst>,
193 #[serde(default)]
194 pub null_ratio: BTreeMap<String, RatioRangeAst>,
195 #[serde(default)]
196 pub distinct_count: BTreeMap<String, CountRangeAst>,
197 #[serde(default)]
198 pub distinct_ratio: BTreeMap<String, RatioRangeAst>,
199}
200
201#[derive(Debug, Clone, PartialEq, Deserialize)]
202#[serde(deny_unknown_fields)]
203pub struct ContractAstV2 {
204 pub version: ContractVersion,
205 #[serde(default)]
206 pub columns: BTreeMap<String, RuleAstV2>,
207 #[serde(default)]
208 pub row_rules: Vec<RowRuleAst>,
209 #[serde(default)]
210 pub dataset_rules: DatasetRulesAst,
211 #[serde(default = "default_max_findings")]
212 pub max_findings: usize,
213}
214
215pub(crate) fn parse(value: Value) -> Result<ContractAstV2, ProofFrameError> {
216 validate_known_fields(&value)?;
217 let contract: ContractAstV2 = serde_json::from_value(value).map_err(|error| {
218 ProofFrameError::contract(
219 ErrorCode::ContractInvalidJson,
220 format!("Invalid V2 contract value: {error}"),
221 None,
222 )
223 })?;
224 if contract.version != ContractVersion::V2 {
225 return Err(invalid(
226 "V2 contract has an incompatible version",
227 "$.version",
228 ));
229 }
230 validate_semantics(&contract)?;
231 Ok(contract)
232}
233
234fn default_max_findings() -> usize {
235 DEFAULT_MAX_FINDINGS
236}
237
238fn validate_semantics(contract: &ContractAstV2) -> Result<(), ProofFrameError> {
239 let mut names = BTreeSet::new();
240 for (index, rule) in contract.row_rules.iter().enumerate() {
241 if rule.name.is_empty() || !names.insert(rule.name.as_str()) {
242 return Err(ProofFrameError::contract(
243 ErrorCode::ContractDuplicateRule,
244 format!("Row rule name `{}` is empty or duplicated", rule.name),
245 Some(format!("$.row_rules[{index}].name")),
246 ));
247 }
248 let direct = rule.compare.is_some();
249 let conditional = rule.when.is_some() && rule.assertion.is_some();
250 if direct == conditional {
251 return Err(invalid(
252 "A row rule must contain either compare or both when and assert",
253 &format!("$.row_rules[{index}]"),
254 ));
255 }
256 }
257 for (column, range) in &contract.dataset_rules.null_ratio {
258 validate_ratio(range, &format!("$.dataset_rules.null_ratio.{column}"))?;
259 }
260 for (column, range) in &contract.dataset_rules.distinct_ratio {
261 validate_ratio(range, &format!("$.dataset_rules.distinct_ratio.{column}"))?;
262 }
263 for (index, rule) in contract.dataset_rules.composite_unique.iter().enumerate() {
264 if rule.columns.len() < 2 {
265 return Err(invalid(
266 "Composite uniqueness requires at least two columns",
267 &format!("$.dataset_rules.composite_unique[{index}].columns"),
268 ));
269 }
270 }
271 Ok(())
272}
273
274fn validate_ratio(range: &RatioRangeAst, path: &str) -> Result<(), ProofFrameError> {
275 for (name, value) in [("min", range.min), ("max", range.max)] {
276 if value.is_some_and(|value| !value.is_finite() || !(0.0..=1.0).contains(&value)) {
277 return Err(ProofFrameError::contract(
278 ErrorCode::ContractInvalidRatio,
279 "Ratio bounds must be finite values between zero and one",
280 Some(format!("{path}.{name}")),
281 ));
282 }
283 }
284 if range.min.zip(range.max).is_some_and(|(min, max)| min > max) {
285 return Err(ProofFrameError::contract(
286 ErrorCode::ContractInvalidRatio,
287 "Ratio minimum exceeds maximum",
288 Some(path.to_string()),
289 ));
290 }
291 Ok(())
292}
293
294fn invalid(message: &str, path: &str) -> ProofFrameError {
295 ProofFrameError::contract(
296 ErrorCode::ContractInvalidJson,
297 message,
298 Some(path.to_string()),
299 )
300}
301
302fn validate_known_fields(value: &Value) -> Result<(), ProofFrameError> {
303 let root = object(value, "$")?;
304 reject_unknown(
305 root,
306 &[
307 "columns",
308 "dataset_rules",
309 "max_findings",
310 "row_rules",
311 "version",
312 ],
313 "$",
314 )?;
315 if let Some(columns) = root.get("columns") {
316 for (name, rules) in object(columns, "$.columns")? {
317 reject_unknown(
318 object(rules, &format!("$.columns.{name}"))?,
319 &[
320 "allowed", "max", "min", "nan", "not_null", "pattern", "required", "type",
321 "unique",
322 ],
323 &format!("$.columns.{name}"),
324 )?;
325 }
326 }
327 if let Some(row_rules) = root.get("row_rules") {
328 let rules = row_rules
329 .as_array()
330 .ok_or_else(|| invalid("row_rules must be an array", "$.row_rules"))?;
331 for (index, rule) in rules.iter().enumerate() {
332 let path = format!("$.row_rules[{index}]");
333 let rule = object(rule, &path)?;
334 reject_unknown(rule, &["assert", "compare", "name", "when"], &path)?;
335 for field in ["compare", "when"] {
336 if let Some(compare) = rule.get(field) {
337 reject_unknown(
338 object(compare, &format!("{path}.{field}"))?,
339 &["left", "nulls", "op", "right"],
340 &format!("{path}.{field}"),
341 )?;
342 }
343 }
344 if let Some(assertion) = rule.get("assert") {
345 reject_unknown(
346 object(assertion, &format!("{path}.assert"))?,
347 &[
348 "allowed", "column", "max", "min", "nan", "not_null", "pattern",
349 ],
350 &format!("{path}.assert"),
351 )?;
352 }
353 }
354 }
355 Ok(())
356}
357
358fn object<'a>(value: &'a Value, path: &str) -> Result<&'a Map<String, Value>, ProofFrameError> {
359 value
360 .as_object()
361 .ok_or_else(|| invalid("Expected a JSON object", path))
362}
363
364fn reject_unknown(
365 object: &Map<String, Value>,
366 allowed: &[&str],
367 path: &str,
368) -> Result<(), ProofFrameError> {
369 if let Some(field) = object
370 .keys()
371 .find(|field| !allowed.contains(&field.as_str()))
372 {
373 return Err(ProofFrameError::contract(
374 ErrorCode::ContractUnknownField,
375 format!("Unknown contract field `{field}`"),
376 Some(format!("{path}.{field}")),
377 ));
378 }
379 Ok(())
380}