Skip to main content

gluesql_core/executor/
validate.rs

1use {
2    crate::{
3        ast::{ColumnDef, ColumnUniqueOption},
4        data::{Key, Value},
5        result::Result,
6        store::Store,
7    },
8    im::HashSet,
9    serde::Serialize,
10    std::fmt::Debug,
11    thiserror::Error as ThisError,
12};
13
14#[derive(ThisError, Debug, PartialEq, Serialize)]
15pub enum ValidateError {
16    #[error("conflict! storage row has no column on index {0}")]
17    ConflictOnStorageColumnIndex(usize),
18
19    #[error("duplicate entry '{}' for unique column '{1}'", String::from(.0))]
20    DuplicateEntryOnUniqueField(Value, String),
21
22    #[error("duplicate entry '{0:?}' for primary_key field")]
23    DuplicateEntryOnPrimaryKeyField(Key),
24}
25
26pub enum ColumnValidation<'column_def> {
27    /// `INSERT`
28    All(&'column_def [ColumnDef]),
29    /// `UPDATE`
30    SpecifiedColumns(&'column_def [ColumnDef], Vec<String>),
31}
32
33#[derive(Debug)]
34struct UniqueConstraint {
35    column_index: usize,
36    column_name: String,
37    keys: HashSet<Key>,
38}
39
40impl UniqueConstraint {
41    fn new(column_index: usize, column_name: String) -> Self {
42        Self {
43            column_index,
44            column_name,
45            keys: HashSet::new(),
46        }
47    }
48
49    fn add(self, value: &Value) -> Result<Self> {
50        let new_key = self.check(value)?;
51
52        if matches!(new_key, Key::None) {
53            return Ok(self);
54        }
55
56        let keys = self.keys.update(new_key);
57
58        Ok(Self {
59            column_index: self.column_index,
60            column_name: self.column_name,
61            keys,
62        })
63    }
64
65    fn check(&self, value: &Value) -> Result<Key> {
66        let key = Key::try_from(value)?;
67
68        if self.keys.contains(&key) {
69            Err(
70                ValidateError::DuplicateEntryOnUniqueField(value.clone(), self.column_name.clone())
71                    .into(),
72            )
73        } else {
74            Ok(key)
75        }
76    }
77}
78
79pub fn validate_unique<'a, T: Store>(
80    storage: &T,
81    table_name: &str,
82    column_validation: &ColumnValidation<'_>,
83    row_iter: impl Iterator<Item = &'a [Value]> + Clone,
84) -> Result<()> {
85    enum Columns {
86        /// key index
87        PrimaryKeyOnly(usize),
88        /// `[(key_index, table_name)]`
89        All(Vec<(usize, String)>),
90    }
91
92    let columns = match &column_validation {
93        ColumnValidation::All(column_defs) => {
94            let primary_key_index = column_defs
95                .iter()
96                .enumerate()
97                .find(|(_, ColumnDef { unique, .. })| {
98                    unique == &Some(ColumnUniqueOption { is_primary: true })
99                })
100                .map(|(i, _)| i);
101            let other_unique_column_def_count = column_defs
102                .iter()
103                .filter(|ColumnDef { unique, .. }| {
104                    unique == &Some(ColumnUniqueOption { is_primary: false })
105                })
106                .count();
107
108            match (primary_key_index, other_unique_column_def_count) {
109                (Some(primary_key_index), 0) => Columns::PrimaryKeyOnly(primary_key_index),
110                _ => Columns::All(fetch_all_unique_columns(column_defs)),
111            }
112        }
113        ColumnValidation::SpecifiedColumns(column_defs, specified_columns) => Columns::All(
114            fetch_specified_unique_columns(column_defs, specified_columns),
115        ),
116    };
117
118    match columns {
119        Columns::PrimaryKeyOnly(primary_key_index) => {
120            for primary_key in
121                row_iter.filter_map(|row| row.get(primary_key_index).map(Key::try_from))
122            {
123                let key = primary_key?;
124
125                if storage.fetch_data(table_name, &key)?.is_some() {
126                    return Err(ValidateError::DuplicateEntryOnPrimaryKeyField(key).into());
127                }
128            }
129
130            Ok(())
131        }
132        Columns::All(columns) => {
133            let unique_constraints = create_unique_constraints(columns, &row_iter)?;
134            if unique_constraints.is_empty() {
135                return Ok(());
136            }
137
138            for row in storage.scan_data(table_name)? {
139                let (_, values) = row?;
140                for constraint in &unique_constraints {
141                    let col_idx = constraint.column_index;
142                    let val = values
143                        .get(col_idx)
144                        .ok_or(ValidateError::ConflictOnStorageColumnIndex(col_idx))?;
145
146                    constraint.check(val)?;
147                }
148            }
149
150            Ok(())
151        }
152    }
153}
154
155fn create_unique_constraints<'a>(
156    unique_columns: Vec<(usize, String)>,
157    row_iter: &(impl Iterator<Item = &'a [Value]> + Clone),
158) -> Result<Vec<UniqueConstraint>> {
159    let mut constraints = Vec::with_capacity(unique_columns.len());
160
161    for (col_idx, col_name) in unique_columns {
162        let new_constraint = UniqueConstraint::new(col_idx, col_name);
163        let new_constraint = row_iter
164            .clone()
165            .try_fold(new_constraint, |constraint, row| {
166                let val = row
167                    .get(col_idx)
168                    .ok_or(ValidateError::ConflictOnStorageColumnIndex(col_idx))?;
169
170                constraint.add(val)
171            })?;
172
173        constraints.push(new_constraint);
174    }
175
176    Ok(constraints)
177}
178
179fn fetch_all_unique_columns(column_defs: &[ColumnDef]) -> Vec<(usize, String)> {
180    column_defs
181        .iter()
182        .enumerate()
183        .filter_map(|(i, table_col)| table_col.unique.map(|_| (i, table_col.name.clone())))
184        .collect()
185}
186
187fn fetch_specified_unique_columns(
188    all_column_defs: &[ColumnDef],
189    specified_columns: &[String],
190) -> Vec<(usize, String)> {
191    all_column_defs
192        .iter()
193        .enumerate()
194        .filter_map(|(i, table_col)| {
195            (table_col.unique.is_some()
196                && specified_columns.iter().any(|col| col == &table_col.name))
197            .then_some((i, table_col.name.clone()))
198        })
199        .collect()
200}