use super::compatibility::are_join_compatible;
use super::core::SchemaValidator;
use crate::distributed::execution::{AggregateExpr, ExecutionPlan, Operation};
use crate::distributed::expr::{
ColumnMeta, ColumnProjection, ExprDataType, ExprSchema, ExprValidator,
};
use crate::error::{Error, Result};
fn schema_of_columns(schema: &ExprSchema, columns: &[String]) -> ExprSchema {
let mut out = ExprSchema::new();
for col in columns {
if let Some(meta) = schema.column(col) {
out.add_column(meta.clone());
}
}
out
}
fn schema_of_aggregate(
schema: &ExprSchema,
keys: &[String],
aggregates: &[AggregateExpr],
) -> ExprSchema {
let mut out = ExprSchema::new();
for key in keys {
if let Some(meta) = schema.column(key) {
out.add_column(meta.clone());
}
}
for agg in aggregates {
let out_name = if agg.alias.trim().is_empty() {
format!("{}_{}", agg.function.trim().to_lowercase(), agg.column)
} else {
agg.alias.clone()
};
let data_type = match agg.function.trim().to_lowercase().as_str() {
"count" => ExprDataType::Integer,
"avg" | "mean" | "stddev" | "std" | "variance" | "var" | "median" => {
ExprDataType::Float
}
_ => schema
.column(&agg.column)
.map(|m| m.data_type.clone())
.unwrap_or(ExprDataType::Float),
};
out.add_column(ColumnMeta::new(out_name, data_type, true, None));
}
out
}
impl SchemaValidator {
pub fn validate_plan(&self, plan: &ExecutionPlan) -> Result<()> {
if plan.operations().is_empty() {
return Ok(());
}
let mut current = match self.schema(plan.input()) {
Some(s) => s.clone(),
None => return Ok(()),
};
for operation in plan.operations() {
match operation {
Operation::Select(columns) => {
self.validate_select(¤t, columns)?;
current = schema_of_columns(¤t, columns);
}
Operation::Filter(predicate) => {
self.validate_filter(¤t, predicate)?;
}
Operation::Aggregate(keys, aggregates) => {
self.validate_groupby(¤t, keys, aggregates)?;
current = schema_of_aggregate(¤t, keys, aggregates);
}
Operation::GroupBy { keys, aggregates } => {
self.validate_groupby(¤t, keys, aggregates)?;
current = schema_of_aggregate(¤t, keys, aggregates);
}
Operation::OrderBy(sort_exprs) => {
self.validate_orderby(¤t, sort_exprs)?;
}
Operation::Limit(_) | Operation::Distinct => {
}
Operation::Join {
left_keys,
right_keys,
right,
..
} => {
match self.schema(right) {
Some(right_schema) => {
self.validate_join(¤t, right_schema, left_keys, right_keys)?;
}
None => return Ok(()), }
return Ok(());
}
Operation::Window(window_functions) => {
self.validate_window(¤t, window_functions)?;
return Ok(());
}
Operation::Custom { name, params } => {
self.validate_custom(¤t, name, params)?;
return Ok(());
}
Operation::Project(_)
| Operation::Union(_)
| Operation::Intersect(_)
| Operation::Except(_) => {
return Ok(());
}
}
}
Ok(())
}
fn validate_custom(
&self,
schema: &ExprSchema,
name: &str,
params: &std::collections::HashMap<String, String>,
) -> Result<()> {
match name {
"select_expr" => {
let projections_json = params.get("projections").ok_or_else(|| {
Error::InvalidOperation(
"select_expr operation requires projections parameter".to_string(),
)
})?;
let projections: Vec<ColumnProjection> = serde_json::from_str(projections_json)
.map_err(|e| {
Error::DistributedProcessing(format!("Failed to parse projections: {}", e))
})?;
self.validate_select_expr(schema, &projections)
}
"with_column" => {
let column_name = params.get("column_name").ok_or_else(|| {
Error::InvalidOperation(
"with_column operation requires column_name parameter".to_string(),
)
})?;
let projection_json = params.get("projection").ok_or_else(|| {
Error::InvalidOperation(
"with_column operation requires projection parameter".to_string(),
)
})?;
let projection: ColumnProjection =
serde_json::from_str(projection_json).map_err(|e| {
Error::DistributedProcessing(format!("Failed to parse projection: {}", e))
})?;
self.validate_with_column(schema, column_name, &projection)
}
"create_udf" => Ok(()),
_ => Err(Error::NotImplemented(format!(
"Schema validation for custom operation '{}' is not implemented",
name
))),
}
}
fn validate_select(&self, schema: &ExprSchema, columns: &[String]) -> Result<()> {
for column in columns {
if !schema.has_column(column) {
return Err(Error::InvalidOperation(format!(
"Column not found in schema: {}",
column
)));
}
}
Ok(())
}
fn validate_select_expr(
&self,
schema: &ExprSchema,
projections: &[ColumnProjection],
) -> Result<()> {
let validator = ExprValidator::new(schema);
validator.validate_projections(projections)?;
Ok(())
}
fn validate_with_column(
&self,
schema: &ExprSchema,
_column_name: &str,
projection: &ColumnProjection,
) -> Result<()> {
let validator = ExprValidator::new(schema);
validator.validate_expr(&projection.expr)?;
Ok(())
}
fn validate_filter(&self, _schema: &ExprSchema, predicate: &str) -> Result<()> {
if predicate.is_empty() {
return Err(Error::InvalidOperation(
"Empty predicate in filter operation".to_string(),
));
}
let mut paren_count = 0;
for c in predicate.chars() {
if c == '(' {
paren_count += 1;
} else if c == ')' {
paren_count -= 1;
if paren_count < 0 {
return Err(Error::InvalidOperation(format!(
"Unbalanced parentheses in predicate: {}",
predicate
)));
}
}
}
if paren_count != 0 {
return Err(Error::InvalidOperation(format!(
"Unbalanced parentheses in predicate: {}",
predicate
)));
}
Ok(())
}
fn validate_join(
&self,
left_schema: &ExprSchema,
right_schema: &ExprSchema,
left_keys: &[String],
right_keys: &[String],
) -> Result<()> {
if left_keys.len() != right_keys.len() {
return Err(Error::InvalidOperation(format!(
"Number of left keys ({}) does not match number of right keys ({})",
left_keys.len(),
right_keys.len()
)));
}
for (left_key, right_key) in left_keys.iter().zip(right_keys.iter()) {
let left_col = left_schema.column(left_key).ok_or_else(|| {
Error::InvalidOperation(format!("Left join key not found in schema: {}", left_key))
})?;
let right_col = right_schema.column(right_key).ok_or_else(|| {
Error::InvalidOperation(format!(
"Right join key not found in schema: {}",
right_key
))
})?;
if !are_join_compatible(&left_col.data_type, &right_col.data_type) {
return Err(Error::InvalidOperation(format!(
"Incompatible join key types: {:?} and {:?}",
left_col.data_type, right_col.data_type
)));
}
}
Ok(())
}
fn validate_groupby(
&self,
schema: &ExprSchema,
keys: &[String],
aggregates: &[crate::distributed::execution::AggregateExpr],
) -> Result<()> {
for key in keys {
if !schema.has_column(key) {
return Err(Error::InvalidOperation(format!(
"Grouping key not found in schema: {}",
key
)));
}
}
for agg in aggregates {
let func = agg.function.trim().to_lowercase();
let is_count_star = func == "count" && agg.column.trim() == "*";
if !is_count_star && !schema.has_column(&agg.column) {
return Err(Error::InvalidOperation(format!(
"Aggregated column not found in schema: {}",
agg.column
)));
}
match func.as_str() {
"count" | "min" | "max" => {}
"sum" | "avg" | "mean" | "stddev" | "std" | "variance" | "var" | "median" => {
if is_count_star {
} else if let Some(col) = schema.column(&agg.column) {
match col.data_type {
ExprDataType::Integer | ExprDataType::Float => {}
_ => {
return Err(Error::InvalidOperation(format!(
"Aggregation function '{}' requires a numeric column but '{}' has type {:?}",
agg.function, agg.column, col.data_type
)));
}
}
}
}
_ => {
return Err(Error::InvalidOperation(format!(
"Unknown aggregation function: {}",
agg.function
)));
}
}
}
Ok(())
}
fn validate_orderby(
&self,
schema: &ExprSchema,
sort_exprs: &[crate::distributed::execution::SortExpr],
) -> Result<()> {
for sort_expr in sort_exprs {
if !schema.has_column(&sort_expr.column) {
return Err(Error::InvalidOperation(format!(
"Sort column not found in schema: {}",
sort_expr.column
)));
}
}
Ok(())
}
pub(crate) fn validate_window(
&self,
schema: &ExprSchema,
window_functions: &[String],
) -> Result<()> {
use crate::distributed::expr::ExprDataType;
if window_functions.is_empty() {
return Err(Error::InvalidOperation(
"Window operation requires at least one window function".to_string(),
));
}
for sql in window_functions {
let spec = parse_window_sql(sql)?;
for col in &spec.input_columns {
if !schema.has_column(col) {
return Err(Error::InvalidOperation(format!(
"Window function input column '{}' not found in schema",
col
)));
}
}
for col in &spec.partition_by {
if !schema.has_column(col) {
return Err(Error::InvalidOperation(format!(
"Window function PARTITION BY column '{}' not found in schema",
col
)));
}
}
for col in &spec.order_by {
if !schema.has_column(col) {
return Err(Error::InvalidOperation(format!(
"Window function ORDER BY column '{}' not found in schema",
col
)));
}
}
match spec.func_name.as_str() {
"SUM" | "AVG" | "STDDEV" | "VARIANCE" => {
for col in &spec.input_columns {
let col_meta = schema.column(col).ok_or_else(|| {
Error::InvalidOperation(format!(
"Window function input column '{}' not found in schema",
col
))
})?;
match col_meta.data_type {
ExprDataType::Integer | ExprDataType::Float => {
}
_ => {
return Err(Error::InvalidOperation(format!(
"Window function '{}' requires a numeric column but '{}' has type {:?}",
spec.func_name, col, col_meta.data_type
)));
}
}
}
}
_ => {}
}
for col in &spec.order_by {
let col_meta = schema.column(col).ok_or_else(|| {
Error::InvalidOperation(format!(
"Window function ORDER BY column '{}' not found in schema",
col
))
})?;
if col_meta.data_type == ExprDataType::Boolean {
return Err(Error::InvalidOperation(format!(
"Window function ORDER BY column '{}' has unsortable type Boolean",
col
)));
}
}
}
Ok(())
}
}
struct ParsedWindowSpec {
func_name: String,
input_columns: Vec<String>,
partition_by: Vec<String>,
order_by: Vec<String>,
}
fn find_ci(haystack: &str, needle: &str) -> Option<usize> {
let hb = haystack.as_bytes();
let nb = needle.as_bytes();
let (hl, nl) = (hb.len(), nb.len());
if nl == 0 || nl > hl {
return None;
}
for i in 0..=(hl - nl) {
if !haystack.is_char_boundary(i) {
continue;
}
if (0..nl).all(|j| hb[i + j].eq_ignore_ascii_case(&nb[j])) {
return Some(i);
}
}
None
}
fn parse_window_sql(sql: &str) -> Result<ParsedWindowSpec> {
let first_paren = sql
.find('(')
.ok_or_else(|| Error::InvalidOperation(format!("Invalid window function SQL: {}", sql)))?;
let func_name = sql[..first_paren].trim().to_uppercase();
const KNOWN_WINDOW_FUNCTIONS: &[&str] = &[
"ROW_NUMBER",
"RANK",
"DENSE_RANK",
"LAG",
"LEAD",
"SUM",
"AVG",
"MIN",
"MAX",
"COUNT",
"NTILE",
"PERCENT_RANK",
"CUME_DIST",
"FIRST_VALUE",
"LAST_VALUE",
"NTH_VALUE",
"STDDEV",
"VARIANCE",
];
if !KNOWN_WINDOW_FUNCTIONS.contains(&func_name.as_str()) {
return Err(Error::InvalidOperation(format!(
"Unknown window function '{}'; known functions are: {}",
func_name,
KNOWN_WINDOW_FUNCTIONS.join(", ")
)));
}
let over_pos = find_ci(sql, " OVER ").ok_or_else(|| {
Error::InvalidOperation(format!("Missing OVER clause in window function: {}", sql))
})?;
let func_args_region = &sql[first_paren + 1..over_pos];
let close_paren = func_args_region.rfind(')').ok_or_else(|| {
Error::InvalidOperation(format!("Malformed window function SQL: {}", sql))
})?;
let func_args = &func_args_region[..close_paren];
let input_columns: Vec<String> = func_args
.split(',')
.map(|s| s.trim().to_string())
.filter(|s| !s.is_empty() && s != "*")
.collect();
let after_over = &sql[over_pos + 6..]; let over_open = after_over
.find('(')
.ok_or_else(|| Error::InvalidOperation(format!("Missing '(' after OVER in: {}", sql)))?;
let over_close = after_over.rfind(')').ok_or_else(|| {
Error::InvalidOperation(format!("Missing ')' to close OVER clause in: {}", sql))
})?;
let over_content = &after_over[over_open + 1..over_close];
let (partition_by, order_by) = if let Some(pb_pos) = find_ci(over_content, "PARTITION BY") {
let after_pb = &over_content[pb_pos + 12..];
let (pb_part, ob_part) = if let Some(ob_pos) = find_ci(after_pb, "ORDER BY") {
(&after_pb[..ob_pos], &after_pb[ob_pos + 8..])
} else {
(after_pb, "")
};
let pb_cols: Vec<String> = pb_part
.split(',')
.map(|s| s.trim().to_string())
.filter(|s| !s.is_empty())
.collect();
let ob_cols: Vec<String> = ob_part
.split(',')
.map(|s| {
let trimmed = s.trim();
let upper = trimmed.to_uppercase();
if upper.ends_with(" ASC") {
trimmed[..trimmed.len() - 4].trim().to_string()
} else if upper.ends_with(" DESC") {
trimmed[..trimmed.len() - 5].trim().to_string()
} else {
trimmed.to_string()
}
})
.filter(|s| !s.is_empty())
.collect();
(pb_cols, ob_cols)
} else if let Some(ob_pos) = find_ci(over_content, "ORDER BY") {
let after_ob = &over_content[ob_pos + 8..];
let ob_cols: Vec<String> = after_ob
.split(',')
.map(|s| {
let trimmed = s.trim();
let upper = trimmed.to_uppercase();
if upper.ends_with(" ASC") {
trimmed[..trimmed.len() - 4].trim().to_string()
} else if upper.ends_with(" DESC") {
trimmed[..trimmed.len() - 5].trim().to_string()
} else {
trimmed.to_string()
}
})
.filter(|s| !s.is_empty())
.collect();
(vec![], ob_cols)
} else {
(vec![], vec![])
};
Ok(ParsedWindowSpec {
func_name,
input_columns,
partition_by,
order_by,
})
}
#[cfg(test)]
mod tests {
use crate::distributed::expr::{ColumnMeta, ExprDataType, ExprSchema};
use crate::distributed::schema_validator::core::SchemaValidator;
fn make_schema() -> ExprSchema {
let mut schema = ExprSchema::new();
schema.add_column(ColumnMeta::new("amount", ExprDataType::Float, false, None));
schema.add_column(ColumnMeta::new("dept", ExprDataType::String, false, None));
schema.add_column(ColumnMeta::new("date", ExprDataType::Date, false, None));
schema.add_column(ColumnMeta::new("region", ExprDataType::String, false, None));
schema.add_column(ColumnMeta::new(
"active",
ExprDataType::Boolean,
false,
None,
));
schema
}
fn make_validator(schema: ExprSchema) -> SchemaValidator {
let mut v = SchemaValidator::new();
v.register_schema("test", schema);
v
}
#[test]
fn test_validate_window_valid() {
let schema = make_schema();
let validator = make_validator(schema.clone());
let wf =
vec!["SUM(amount) OVER (PARTITION BY dept ORDER BY date ASC) AS total".to_string()];
let result = validator.validate_window(&schema, &wf);
assert!(result.is_ok(), "Expected Ok, got: {:?}", result);
}
#[test]
fn test_validate_window_missing_column() {
let schema = make_schema();
let validator = make_validator(schema.clone());
let wf =
vec!["SUM(salary) OVER (PARTITION BY dept ORDER BY date ASC) AS total".to_string()];
let result = validator.validate_window(&schema, &wf);
assert!(result.is_err());
let msg = result.unwrap_err().to_string();
assert!(
msg.contains("salary"),
"Error should mention 'salary': {}",
msg
);
}
#[test]
fn test_validate_window_nonnumeric_sum() {
let schema = make_schema();
let validator = make_validator(schema.clone());
let wf = vec!["SUM(dept) OVER (PARTITION BY region ORDER BY date ASC) AS bad".to_string()];
let result = validator.validate_window(&schema, &wf);
assert!(result.is_err());
let msg = result.unwrap_err().to_string();
assert!(
msg.contains("numeric") || msg.contains("dept"),
"Error should mention numeric requirement or column: {}",
msg
);
}
#[test]
fn test_validate_window_missing_partition_col() {
let schema = make_schema();
let validator = make_validator(schema.clone());
let wf = vec![
"ROW_NUMBER(*) OVER (PARTITION BY nonexistent ORDER BY date ASC) AS rn".to_string(),
];
let result = validator.validate_window(&schema, &wf);
assert!(result.is_err());
let msg = result.unwrap_err().to_string();
assert!(
msg.contains("nonexistent"),
"Error should mention 'nonexistent': {}",
msg
);
}
#[test]
fn test_validate_window_boolean_order_by() {
let schema = make_schema();
let validator = make_validator(schema.clone());
let wf =
vec!["ROW_NUMBER(*) OVER (PARTITION BY dept ORDER BY active ASC) AS rn".to_string()];
let result = validator.validate_window(&schema, &wf);
assert!(result.is_err());
let msg = result.unwrap_err().to_string();
assert!(
msg.contains("Boolean") || msg.contains("active"),
"Error should mention Boolean or active: {}",
msg
);
}
#[test]
fn test_validate_window_empty_over() {
let schema = make_schema();
let validator = make_validator(schema.clone());
let wf = vec!["AVG(amount) OVER () AS overall_avg".to_string()];
let result = validator.validate_window(&schema, &wf);
assert!(
result.is_ok(),
"Expected Ok for empty OVER clause, got: {:?}",
result
);
}
#[test]
fn test_validate_window_empty_list() {
let schema = make_schema();
let validator = make_validator(schema.clone());
let wf: Vec<String> = vec![];
let result = validator.validate_window(&schema, &wf);
assert!(result.is_err());
let msg = result.unwrap_err().to_string();
assert!(
msg.contains("least one") || msg.contains("empty"),
"Error should mention empty list: {}",
msg
);
}
#[test]
fn test_validate_window_unknown_function() {
let schema = make_schema();
let validator = make_validator(schema.clone());
let wf = vec!["FOOBAR(amount) OVER ()".to_string()];
let result = validator.validate_window(&schema, &wf);
assert!(result.is_err());
let msg = result.unwrap_err().to_string();
assert!(
msg.contains("FOOBAR") || msg.contains("unknown") || msg.contains("Unknown"),
"Error should mention unknown function: {}",
msg
);
}
}