Refactor dataflow operators around runtime registry
This commit is contained in:
@@ -1,8 +1,9 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import operator
|
||||
from dataclasses import dataclass
|
||||
from decimal import Decimal
|
||||
from typing import Any, Literal
|
||||
from typing import Any, Callable, Literal
|
||||
|
||||
import sqlglot
|
||||
from sqlglot import exp
|
||||
@@ -141,133 +142,255 @@ def infer_expression_type(
|
||||
|
||||
|
||||
def _evaluate(expression: exp.Expression, row: dict[str, Any]) -> Any:
|
||||
if isinstance(expression, exp.Paren):
|
||||
return _evaluate(expression.this, row)
|
||||
if isinstance(expression, exp.Column):
|
||||
return row.get(expression.name)
|
||||
if isinstance(expression, exp.Null):
|
||||
return None
|
||||
if isinstance(expression, exp.Boolean):
|
||||
return bool(expression.this)
|
||||
if isinstance(expression, exp.Literal):
|
||||
return _literal(expression)
|
||||
if isinstance(expression, exp.Neg):
|
||||
value = _evaluate(expression.this, row)
|
||||
return None if value is None else -value
|
||||
if isinstance(expression, exp.Not):
|
||||
return not bool(_evaluate(expression.this, row))
|
||||
evaluator = _EVALUATORS.get(type(expression))
|
||||
if evaluator is None:
|
||||
raise ExpressionError(
|
||||
f"Expression element {type(expression).__name__} cannot be evaluated."
|
||||
)
|
||||
return evaluator(expression, row)
|
||||
|
||||
|
||||
def _evaluate_child(expression: exp.Expression, row: dict[str, Any]) -> Any:
|
||||
return _evaluate(expression.this, row)
|
||||
|
||||
|
||||
def _evaluate_negation(expression: exp.Expression, row: dict[str, Any]) -> Any:
|
||||
value = _evaluate_child(expression, row)
|
||||
return None if value is None else -value
|
||||
|
||||
|
||||
def _evaluate_boolean_not(
|
||||
expression: exp.Expression,
|
||||
row: dict[str, Any],
|
||||
) -> bool:
|
||||
return not bool(_evaluate_child(expression, row))
|
||||
|
||||
|
||||
def _evaluate_boolean_binary(
|
||||
expression: exp.Expression,
|
||||
row: dict[str, Any],
|
||||
) -> bool:
|
||||
left = bool(_evaluate(expression.this, row))
|
||||
if isinstance(expression, exp.And):
|
||||
return bool(_evaluate(expression.this, row)) and bool(
|
||||
_evaluate(expression.expression, row)
|
||||
)
|
||||
if isinstance(expression, exp.Or):
|
||||
return bool(_evaluate(expression.this, row)) or bool(
|
||||
_evaluate(expression.expression, row)
|
||||
)
|
||||
if isinstance(expression, exp.Is):
|
||||
left = _evaluate(expression.this, row)
|
||||
right = _evaluate(expression.expression, row)
|
||||
return left is right if right is None else left == right
|
||||
if isinstance(expression, (exp.Add, exp.Sub, exp.Mul, exp.Div, exp.Mod)):
|
||||
left = _evaluate(expression.this, row)
|
||||
right = _evaluate(expression.expression, row)
|
||||
if left is None or right is None:
|
||||
return None
|
||||
if isinstance(expression, exp.Add):
|
||||
return left + right
|
||||
if isinstance(expression, exp.Sub):
|
||||
return left - right
|
||||
if isinstance(expression, exp.Mul):
|
||||
return left * right
|
||||
if isinstance(expression, exp.Div):
|
||||
return left / right
|
||||
return left % right
|
||||
if isinstance(expression, (exp.EQ, exp.NEQ, exp.GT, exp.GTE, exp.LT, exp.LTE)):
|
||||
left = _evaluate(expression.this, row)
|
||||
right = _evaluate(expression.expression, row)
|
||||
if isinstance(expression, exp.EQ):
|
||||
return left == right
|
||||
if isinstance(expression, exp.NEQ):
|
||||
return left != right
|
||||
if left is None or right is None:
|
||||
return False
|
||||
if isinstance(expression, exp.GT):
|
||||
return left > right
|
||||
if isinstance(expression, exp.GTE):
|
||||
return left >= right
|
||||
if isinstance(expression, exp.LT):
|
||||
return left < right
|
||||
return left <= right
|
||||
if isinstance(expression, exp.Lower):
|
||||
return _string_call(expression.this, row, str.lower)
|
||||
if isinstance(expression, exp.Upper):
|
||||
return _string_call(expression.this, row, str.upper)
|
||||
if isinstance(expression, exp.Trim):
|
||||
return _string_call(expression.this, row, str.strip)
|
||||
if isinstance(expression, exp.Length):
|
||||
value = _evaluate(expression.this, row)
|
||||
return None if value is None else len(value)
|
||||
if isinstance(expression, exp.Abs):
|
||||
value = _evaluate(expression.this, row)
|
||||
return None if value is None else abs(value)
|
||||
if isinstance(expression, exp.Round):
|
||||
value = _evaluate(expression.this, row)
|
||||
decimals = expression.args.get("decimals")
|
||||
places = int(_evaluate(decimals, row)) if decimals is not None else 0
|
||||
return None if value is None else round(value, places)
|
||||
if isinstance(expression, exp.Coalesce):
|
||||
arguments = [expression.this, *expression.expressions]
|
||||
return next(
|
||||
(
|
||||
value
|
||||
for argument in arguments
|
||||
if (value := _evaluate(argument, row)) is not None
|
||||
),
|
||||
None,
|
||||
)
|
||||
if isinstance(expression, exp.Concat):
|
||||
return "".join(
|
||||
"" if (value := _evaluate(item, row)) is None else str(value)
|
||||
for item in expression.expressions
|
||||
)
|
||||
if isinstance(expression, exp.Replace):
|
||||
value = _evaluate(expression.this, row)
|
||||
old = _evaluate(expression.expression, row)
|
||||
new = _evaluate(expression.args["replacement"], row)
|
||||
return None if value is None else str(value).replace(str(old), str(new))
|
||||
if isinstance(expression, exp.Substring):
|
||||
value = _evaluate(expression.this, row)
|
||||
if value is None:
|
||||
return None
|
||||
start = int(_evaluate(expression.args["start"], row)) - 1
|
||||
length = expression.args.get("length")
|
||||
return (
|
||||
str(value)[start:]
|
||||
if length is None
|
||||
else str(value)[start : start + int(_evaluate(length, row))]
|
||||
)
|
||||
if isinstance(expression, exp.Cast):
|
||||
return convert_value(
|
||||
_evaluate(expression.this, row),
|
||||
_cast_type(expression),
|
||||
on_error="fail",
|
||||
)
|
||||
if isinstance(expression, exp.Case):
|
||||
base = _evaluate(expression.this, row) if expression.this is not None else None
|
||||
for condition in expression.args.get("ifs") or []:
|
||||
if not isinstance(condition, exp.If):
|
||||
continue
|
||||
candidate = _evaluate(condition.this, row)
|
||||
matched = candidate == base if expression.this is not None else bool(candidate)
|
||||
if matched:
|
||||
return _evaluate(condition.args["true"], row)
|
||||
default = expression.args.get("default")
|
||||
return _evaluate(default, row) if default is not None else None
|
||||
raise ExpressionError(
|
||||
f"Expression element {type(expression).__name__} cannot be evaluated."
|
||||
return left and bool(_evaluate(expression.expression, row))
|
||||
return left or bool(_evaluate(expression.expression, row))
|
||||
|
||||
|
||||
def _evaluate_is(expression: exp.Expression, row: dict[str, Any]) -> bool:
|
||||
left = _evaluate(expression.this, row)
|
||||
right = _evaluate(expression.expression, row)
|
||||
return left is right if right is None else left == right
|
||||
|
||||
|
||||
def _evaluate_binary(
|
||||
expression: exp.Expression,
|
||||
row: dict[str, Any],
|
||||
operations: dict[type[exp.Expression], Callable[[Any, Any], Any]],
|
||||
*,
|
||||
null_result: Any,
|
||||
) -> Any:
|
||||
left = _evaluate(expression.this, row)
|
||||
right = _evaluate(expression.expression, row)
|
||||
if left is None or right is None:
|
||||
return null_result
|
||||
return operations[type(expression)](left, right)
|
||||
|
||||
|
||||
def _evaluate_arithmetic(expression: exp.Expression, row: dict[str, Any]) -> Any:
|
||||
return _evaluate_binary(
|
||||
expression,
|
||||
row,
|
||||
_ARITHMETIC_OPERATIONS,
|
||||
null_result=None,
|
||||
)
|
||||
|
||||
|
||||
def _evaluate_comparison(
|
||||
expression: exp.Expression,
|
||||
row: dict[str, Any],
|
||||
) -> bool:
|
||||
if isinstance(expression, (exp.EQ, exp.NEQ)):
|
||||
left = _evaluate(expression.this, row)
|
||||
right = _evaluate(expression.expression, row)
|
||||
return _COMPARISON_OPERATIONS[type(expression)](left, right)
|
||||
return _evaluate_binary(
|
||||
expression,
|
||||
row,
|
||||
_COMPARISON_OPERATIONS,
|
||||
null_result=False,
|
||||
)
|
||||
|
||||
|
||||
def _evaluate_string_function(
|
||||
expression: exp.Expression,
|
||||
row: dict[str, Any],
|
||||
) -> str | None:
|
||||
return _string_call(
|
||||
expression.this,
|
||||
row,
|
||||
_STRING_OPERATIONS[type(expression)],
|
||||
)
|
||||
|
||||
|
||||
def _evaluate_length(expression: exp.Expression, row: dict[str, Any]) -> int | None:
|
||||
value = _evaluate_child(expression, row)
|
||||
return None if value is None else len(value)
|
||||
|
||||
|
||||
def _evaluate_abs(expression: exp.Expression, row: dict[str, Any]) -> Any:
|
||||
value = _evaluate_child(expression, row)
|
||||
return None if value is None else abs(value)
|
||||
|
||||
|
||||
def _evaluate_round(expression: exp.Expression, row: dict[str, Any]) -> Any:
|
||||
value = _evaluate_child(expression, row)
|
||||
decimals = expression.args.get("decimals")
|
||||
places = int(_evaluate(decimals, row)) if decimals is not None else 0
|
||||
return None if value is None else round(value, places)
|
||||
|
||||
|
||||
def _evaluate_coalesce(expression: exp.Expression, row: dict[str, Any]) -> Any:
|
||||
arguments = [expression.this, *expression.expressions]
|
||||
return next(
|
||||
(
|
||||
value
|
||||
for argument in arguments
|
||||
if (value := _evaluate(argument, row)) is not None
|
||||
),
|
||||
None,
|
||||
)
|
||||
|
||||
|
||||
def _evaluate_concat(expression: exp.Expression, row: dict[str, Any]) -> str:
|
||||
return "".join(
|
||||
"" if (value := _evaluate(item, row)) is None else str(value)
|
||||
for item in expression.expressions
|
||||
)
|
||||
|
||||
|
||||
def _evaluate_replace(expression: exp.Expression, row: dict[str, Any]) -> Any:
|
||||
value = _evaluate_child(expression, row)
|
||||
old = _evaluate(expression.expression, row)
|
||||
new = _evaluate(expression.args["replacement"], row)
|
||||
return None if value is None else str(value).replace(str(old), str(new))
|
||||
|
||||
|
||||
def _evaluate_substring(expression: exp.Expression, row: dict[str, Any]) -> Any:
|
||||
value = _evaluate_child(expression, row)
|
||||
if value is None:
|
||||
return None
|
||||
start = int(_evaluate(expression.args["start"], row)) - 1
|
||||
length = expression.args.get("length")
|
||||
if length is None:
|
||||
return str(value)[start:]
|
||||
return str(value)[start : start + int(_evaluate(length, row))]
|
||||
|
||||
|
||||
def _evaluate_cast(expression: exp.Expression, row: dict[str, Any]) -> Any:
|
||||
return convert_value(
|
||||
_evaluate_child(expression, row),
|
||||
_cast_type(expression), # type: ignore[arg-type]
|
||||
on_error="fail",
|
||||
)
|
||||
|
||||
|
||||
def _evaluate_case(expression: exp.Expression, row: dict[str, Any]) -> Any:
|
||||
base_expression = expression.this
|
||||
base = _evaluate(base_expression, row) if base_expression is not None else None
|
||||
for condition in expression.args.get("ifs") or []:
|
||||
if isinstance(condition, exp.If) and _case_matches(
|
||||
condition,
|
||||
row,
|
||||
base=base,
|
||||
has_base=base_expression is not None,
|
||||
):
|
||||
return _evaluate(condition.args["true"], row)
|
||||
default = expression.args.get("default")
|
||||
return _evaluate(default, row) if default is not None else None
|
||||
|
||||
|
||||
def _case_matches(
|
||||
condition: exp.If,
|
||||
row: dict[str, Any],
|
||||
*,
|
||||
base: Any,
|
||||
has_base: bool,
|
||||
) -> bool:
|
||||
candidate = _evaluate(condition.this, row)
|
||||
return candidate == base if has_base else bool(candidate)
|
||||
|
||||
|
||||
def _evaluate_if(expression: exp.Expression, row: dict[str, Any]) -> Any:
|
||||
branch = "true" if bool(_evaluate(expression.this, row)) else "false"
|
||||
selected = expression.args.get(branch)
|
||||
return _evaluate(selected, row) if selected is not None else None
|
||||
|
||||
|
||||
_ARITHMETIC_OPERATIONS: dict[
|
||||
type[exp.Expression],
|
||||
Callable[[Any, Any], Any],
|
||||
] = {
|
||||
exp.Add: operator.add,
|
||||
exp.Sub: operator.sub,
|
||||
exp.Mul: operator.mul,
|
||||
exp.Div: operator.truediv,
|
||||
exp.Mod: operator.mod,
|
||||
}
|
||||
_COMPARISON_OPERATIONS: dict[
|
||||
type[exp.Expression],
|
||||
Callable[[Any, Any], bool],
|
||||
] = {
|
||||
exp.EQ: operator.eq,
|
||||
exp.NEQ: operator.ne,
|
||||
exp.GT: operator.gt,
|
||||
exp.GTE: operator.ge,
|
||||
exp.LT: operator.lt,
|
||||
exp.LTE: operator.le,
|
||||
}
|
||||
_STRING_OPERATIONS: dict[type[exp.Expression], Callable[[str], str]] = {
|
||||
exp.Lower: str.lower,
|
||||
exp.Upper: str.upper,
|
||||
exp.Trim: str.strip,
|
||||
}
|
||||
_EVALUATORS: dict[
|
||||
type[exp.Expression],
|
||||
Callable[[exp.Expression, dict[str, Any]], Any],
|
||||
] = {
|
||||
exp.Paren: _evaluate_child,
|
||||
exp.Column: lambda expression, row: row.get(expression.name),
|
||||
exp.Null: lambda _expression, _row: None,
|
||||
exp.Boolean: lambda expression, _row: bool(expression.this),
|
||||
exp.Literal: lambda expression, _row: _literal(expression), # type: ignore[arg-type]
|
||||
exp.Neg: _evaluate_negation,
|
||||
exp.Not: _evaluate_boolean_not,
|
||||
exp.And: _evaluate_boolean_binary,
|
||||
exp.Or: _evaluate_boolean_binary,
|
||||
exp.Is: _evaluate_is,
|
||||
**{
|
||||
expression_type: _evaluate_arithmetic
|
||||
for expression_type in _ARITHMETIC_OPERATIONS
|
||||
},
|
||||
**{
|
||||
expression_type: _evaluate_comparison
|
||||
for expression_type in _COMPARISON_OPERATIONS
|
||||
},
|
||||
**{
|
||||
expression_type: _evaluate_string_function
|
||||
for expression_type in _STRING_OPERATIONS
|
||||
},
|
||||
exp.Length: _evaluate_length,
|
||||
exp.Abs: _evaluate_abs,
|
||||
exp.Round: _evaluate_round,
|
||||
exp.Coalesce: _evaluate_coalesce,
|
||||
exp.Concat: _evaluate_concat,
|
||||
exp.Replace: _evaluate_replace,
|
||||
exp.Substring: _evaluate_substring,
|
||||
exp.Cast: _evaluate_cast,
|
||||
exp.Case: _evaluate_case,
|
||||
exp.If: _evaluate_if,
|
||||
}
|
||||
|
||||
|
||||
def convert_value(
|
||||
value: Any,
|
||||
target_type: str,
|
||||
@@ -312,19 +435,115 @@ def _infer(
|
||||
expression: exp.Expression,
|
||||
schema: dict[str, ExpressionDataType],
|
||||
) -> ExpressionDataType:
|
||||
if isinstance(expression, exp.Column):
|
||||
return schema.get(expression.name, "unknown")
|
||||
if isinstance(expression, exp.Null):
|
||||
return "null"
|
||||
if isinstance(expression, exp.Boolean):
|
||||
return "boolean"
|
||||
if isinstance(expression, exp.Literal):
|
||||
if expression.is_string:
|
||||
return "string"
|
||||
return "number" if "." in str(expression.this) else "integer"
|
||||
if isinstance(
|
||||
expression,
|
||||
(
|
||||
inferer = _TYPE_INFERERS.get(type(expression))
|
||||
return inferer(expression, schema) if inferer is not None else "unknown"
|
||||
|
||||
|
||||
def _infer_literal(
|
||||
expression: exp.Expression,
|
||||
_schema: dict[str, ExpressionDataType],
|
||||
) -> ExpressionDataType:
|
||||
literal = expression
|
||||
if not isinstance(literal, exp.Literal):
|
||||
return "unknown"
|
||||
if literal.is_string:
|
||||
return "string"
|
||||
return "number" if "." in str(literal.this) else "integer"
|
||||
|
||||
|
||||
def _infer_cast(
|
||||
expression: exp.Expression,
|
||||
_schema: dict[str, ExpressionDataType],
|
||||
) -> ExpressionDataType:
|
||||
return _cast_type(expression) # type: ignore[arg-type,return-value]
|
||||
|
||||
|
||||
def _infer_numeric(
|
||||
expression: exp.Expression,
|
||||
schema: dict[str, ExpressionDataType],
|
||||
) -> ExpressionDataType:
|
||||
child_types = {
|
||||
_infer(item, schema)
|
||||
for item in expression.iter_expressions()
|
||||
}
|
||||
if "number" in child_types or isinstance(expression, exp.Div):
|
||||
return "number"
|
||||
return "integer"
|
||||
|
||||
|
||||
def _infer_coalesce(
|
||||
expression: exp.Expression,
|
||||
schema: dict[str, ExpressionDataType],
|
||||
) -> ExpressionDataType:
|
||||
candidates = (
|
||||
_infer(item, schema)
|
||||
for item in (expression.this, *expression.expressions)
|
||||
)
|
||||
return next(
|
||||
(item for item in candidates if item not in {"null", "unknown"}),
|
||||
"unknown",
|
||||
)
|
||||
|
||||
|
||||
def _infer_case(
|
||||
expression: exp.Expression,
|
||||
schema: dict[str, ExpressionDataType],
|
||||
) -> ExpressionDataType:
|
||||
types = [
|
||||
_infer(item.args["true"], schema)
|
||||
for item in expression.args.get("ifs") or []
|
||||
if isinstance(item, exp.If)
|
||||
]
|
||||
default = expression.args.get("default")
|
||||
if default is not None:
|
||||
types.append(_infer(default, schema))
|
||||
return _common_expression_type(types)
|
||||
|
||||
|
||||
def _infer_if(
|
||||
expression: exp.Expression,
|
||||
schema: dict[str, ExpressionDataType],
|
||||
) -> ExpressionDataType:
|
||||
return _common_expression_type(
|
||||
[
|
||||
_infer(branch, schema)
|
||||
for key in ("true", "false")
|
||||
if (branch := expression.args.get(key)) is not None
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
def _common_expression_type(
|
||||
types: list[ExpressionDataType],
|
||||
) -> ExpressionDataType:
|
||||
concrete = {item for item in types if item not in {"null", "unknown"}}
|
||||
return concrete.pop() if len(concrete) == 1 else "unknown"
|
||||
|
||||
|
||||
def _infer_child(
|
||||
expression: exp.Expression,
|
||||
schema: dict[str, ExpressionDataType],
|
||||
) -> ExpressionDataType:
|
||||
return _infer(expression.this, schema)
|
||||
|
||||
|
||||
_TYPE_INFERERS: dict[
|
||||
type[exp.Expression],
|
||||
Callable[
|
||||
[exp.Expression, dict[str, ExpressionDataType]],
|
||||
ExpressionDataType,
|
||||
],
|
||||
] = {
|
||||
exp.Column: lambda expression, schema: schema.get(
|
||||
expression.name,
|
||||
"unknown",
|
||||
),
|
||||
exp.Null: lambda _expression, _schema: "null",
|
||||
exp.Boolean: lambda _expression, _schema: "boolean",
|
||||
exp.Literal: _infer_literal,
|
||||
**{
|
||||
expression_type: lambda _expression, _schema: "boolean"
|
||||
for expression_type in (
|
||||
exp.EQ,
|
||||
exp.NEQ,
|
||||
exp.GT,
|
||||
@@ -335,41 +554,39 @@ def _infer(
|
||||
exp.Or,
|
||||
exp.Is,
|
||||
exp.Not,
|
||||
),
|
||||
):
|
||||
return "boolean"
|
||||
if isinstance(expression, (exp.Length,)):
|
||||
return "integer"
|
||||
if isinstance(expression, (exp.Lower, exp.Upper, exp.Trim, exp.Concat, exp.Replace, exp.Substring)):
|
||||
return "string"
|
||||
if isinstance(expression, exp.Cast):
|
||||
return _cast_type(expression) # type: ignore[return-value]
|
||||
if isinstance(expression, (exp.Add, exp.Sub, exp.Mul, exp.Div, exp.Mod, exp.Abs, exp.Round)):
|
||||
child_types = {
|
||||
_infer(item, schema)
|
||||
for item in expression.iter_expressions()
|
||||
}
|
||||
return "number" if "number" in child_types or isinstance(expression, exp.Div) else "integer"
|
||||
if isinstance(expression, exp.Coalesce):
|
||||
types = [
|
||||
_infer(item, schema)
|
||||
for item in (expression.this, *expression.expressions)
|
||||
]
|
||||
return next((item for item in types if item not in {"null", "unknown"}), "unknown")
|
||||
if isinstance(expression, exp.Case):
|
||||
types = [
|
||||
_infer(item.args["true"], schema)
|
||||
for item in expression.args.get("ifs") or []
|
||||
if isinstance(item, exp.If)
|
||||
]
|
||||
default = expression.args.get("default")
|
||||
if default is not None:
|
||||
types.append(_infer(default, schema))
|
||||
concrete = {item for item in types if item not in {"null", "unknown"}}
|
||||
return concrete.pop() if len(concrete) == 1 else "unknown"
|
||||
if isinstance(expression, (exp.Paren, exp.Neg)):
|
||||
return _infer(expression.this, schema)
|
||||
return "unknown"
|
||||
)
|
||||
},
|
||||
exp.Length: lambda _expression, _schema: "integer",
|
||||
**{
|
||||
expression_type: lambda _expression, _schema: "string"
|
||||
for expression_type in (
|
||||
exp.Lower,
|
||||
exp.Upper,
|
||||
exp.Trim,
|
||||
exp.Concat,
|
||||
exp.Replace,
|
||||
exp.Substring,
|
||||
)
|
||||
},
|
||||
exp.Cast: _infer_cast,
|
||||
**{
|
||||
expression_type: _infer_numeric
|
||||
for expression_type in (
|
||||
exp.Add,
|
||||
exp.Sub,
|
||||
exp.Mul,
|
||||
exp.Div,
|
||||
exp.Mod,
|
||||
exp.Abs,
|
||||
exp.Round,
|
||||
)
|
||||
},
|
||||
exp.Coalesce: _infer_coalesce,
|
||||
exp.Case: _infer_case,
|
||||
exp.If: _infer_if,
|
||||
exp.Paren: _infer_child,
|
||||
exp.Neg: _infer_child,
|
||||
}
|
||||
|
||||
|
||||
def _literal(expression: exp.Literal) -> Any:
|
||||
|
||||
Reference in New Issue
Block a user