Skip to main content

proofframe/contract/
v2.rs

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}