feat: add extensible operators and golden flows
This commit is contained in:
@@ -0,0 +1,426 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from decimal import Decimal
|
||||
from typing import Any, Literal
|
||||
|
||||
import sqlglot
|
||||
from sqlglot import exp
|
||||
from sqlglot.errors import ParseError
|
||||
|
||||
|
||||
ExpressionDataType = Literal[
|
||||
"unknown",
|
||||
"null",
|
||||
"string",
|
||||
"integer",
|
||||
"number",
|
||||
"boolean",
|
||||
"date",
|
||||
"datetime",
|
||||
]
|
||||
|
||||
|
||||
class ExpressionError(ValueError):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ParsedExpression:
|
||||
source: str
|
||||
expression: exp.Expression
|
||||
columns: tuple[str, ...]
|
||||
|
||||
def sql(self) -> str:
|
||||
return self.expression.sql(dialect="duckdb")
|
||||
|
||||
|
||||
_LEAF_TYPES = (
|
||||
exp.Column,
|
||||
exp.Identifier,
|
||||
exp.Literal,
|
||||
exp.Null,
|
||||
exp.Boolean,
|
||||
exp.DataType,
|
||||
)
|
||||
_BINARY_TYPES = (
|
||||
exp.Add,
|
||||
exp.Sub,
|
||||
exp.Mul,
|
||||
exp.Div,
|
||||
exp.Mod,
|
||||
exp.EQ,
|
||||
exp.NEQ,
|
||||
exp.GT,
|
||||
exp.GTE,
|
||||
exp.LT,
|
||||
exp.LTE,
|
||||
exp.And,
|
||||
exp.Or,
|
||||
exp.Is,
|
||||
)
|
||||
_UNARY_TYPES = (exp.Not, exp.Neg, exp.Paren)
|
||||
_FUNCTION_TYPES = (
|
||||
exp.Lower,
|
||||
exp.Upper,
|
||||
exp.Trim,
|
||||
exp.Coalesce,
|
||||
exp.Concat,
|
||||
exp.Length,
|
||||
exp.Abs,
|
||||
exp.Round,
|
||||
exp.Replace,
|
||||
exp.Substring,
|
||||
exp.Cast,
|
||||
)
|
||||
_CONTROL_TYPES = (exp.Case, exp.If)
|
||||
_ALLOWED_TYPES = (
|
||||
*_LEAF_TYPES,
|
||||
*_BINARY_TYPES,
|
||||
*_UNARY_TYPES,
|
||||
*_FUNCTION_TYPES,
|
||||
*_CONTROL_TYPES,
|
||||
)
|
||||
|
||||
|
||||
def parse_expression(source: str) -> ParsedExpression:
|
||||
normalized = str(source or "").strip()
|
||||
if not normalized:
|
||||
raise ExpressionError("Enter an expression.")
|
||||
if len(normalized) > 10_000:
|
||||
raise ExpressionError("Expressions are limited to 10,000 characters.")
|
||||
try:
|
||||
expression = sqlglot.parse_one(normalized, read="duckdb")
|
||||
except ParseError as exc:
|
||||
raise ExpressionError(f"Expression could not be parsed: {exc}") from exc
|
||||
for item in expression.walk():
|
||||
if not isinstance(item, _ALLOWED_TYPES):
|
||||
raise ExpressionError(
|
||||
f"Expression element {type(item).__name__} is not allowed."
|
||||
)
|
||||
if isinstance(item, exp.Column) and item.table:
|
||||
raise ExpressionError("Expressions use column names without table qualifiers.")
|
||||
if isinstance(item, exp.Cast):
|
||||
_cast_type(item)
|
||||
columns = tuple(
|
||||
dict.fromkeys(
|
||||
column.name
|
||||
for column in expression.find_all(exp.Column)
|
||||
if column.name
|
||||
)
|
||||
)
|
||||
return ParsedExpression(
|
||||
source=normalized,
|
||||
expression=expression,
|
||||
columns=columns,
|
||||
)
|
||||
|
||||
|
||||
def evaluate_expression(
|
||||
parsed: ParsedExpression | str,
|
||||
row: dict[str, Any],
|
||||
) -> Any:
|
||||
expression = (
|
||||
parse_expression(parsed).expression
|
||||
if isinstance(parsed, str)
|
||||
else parsed.expression
|
||||
)
|
||||
return _evaluate(expression, row)
|
||||
|
||||
|
||||
def infer_expression_type(
|
||||
parsed: ParsedExpression | str,
|
||||
schema: dict[str, ExpressionDataType] | None = None,
|
||||
) -> ExpressionDataType:
|
||||
expression = (
|
||||
parse_expression(parsed).expression
|
||||
if isinstance(parsed, str)
|
||||
else parsed.expression
|
||||
)
|
||||
return _infer(expression, schema or {})
|
||||
|
||||
|
||||
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))
|
||||
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."
|
||||
)
|
||||
|
||||
|
||||
def convert_value(
|
||||
value: Any,
|
||||
target_type: str,
|
||||
*,
|
||||
on_error: Literal["fail", "null", "keep"] = "fail",
|
||||
) -> Any:
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
if target_type == "string":
|
||||
return str(value)
|
||||
if target_type == "integer":
|
||||
if isinstance(value, bool):
|
||||
return int(value)
|
||||
return int(str(value).strip())
|
||||
if target_type == "number":
|
||||
return Decimal(str(value).strip())
|
||||
if target_type == "boolean":
|
||||
if isinstance(value, bool):
|
||||
return value
|
||||
normalized = str(value).strip().casefold()
|
||||
if normalized in {"true", "yes", "y", "1", "on"}:
|
||||
return True
|
||||
if normalized in {"false", "no", "n", "0", "off"}:
|
||||
return False
|
||||
raise ValueError(f"{value!r} is not a boolean")
|
||||
if target_type in {"date", "datetime"}:
|
||||
from datetime import datetime
|
||||
|
||||
parsed = datetime.fromisoformat(str(value).strip().replace("Z", "+00:00"))
|
||||
return parsed.date() if target_type == "date" else parsed
|
||||
raise ValueError(f"Unsupported conversion type {target_type!r}")
|
||||
except (ArithmeticError, TypeError, ValueError):
|
||||
if on_error == "null":
|
||||
return None
|
||||
if on_error == "keep":
|
||||
return value
|
||||
raise
|
||||
|
||||
|
||||
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,
|
||||
(
|
||||
exp.EQ,
|
||||
exp.NEQ,
|
||||
exp.GT,
|
||||
exp.GTE,
|
||||
exp.LT,
|
||||
exp.LTE,
|
||||
exp.And,
|
||||
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"
|
||||
|
||||
|
||||
def _literal(expression: exp.Literal) -> Any:
|
||||
if expression.is_string:
|
||||
return str(expression.this)
|
||||
source = str(expression.this)
|
||||
try:
|
||||
return int(source)
|
||||
except ValueError:
|
||||
return Decimal(source)
|
||||
|
||||
|
||||
def _string_call(
|
||||
expression: exp.Expression,
|
||||
row: dict[str, Any],
|
||||
operation,
|
||||
) -> str | None:
|
||||
value = _evaluate(expression, row)
|
||||
return None if value is None else operation(str(value))
|
||||
|
||||
|
||||
def _cast_type(expression: exp.Cast) -> str:
|
||||
data_type = expression.args.get("to")
|
||||
name = str(data_type.this).casefold() if isinstance(data_type, exp.DataType) else ""
|
||||
mapping = {
|
||||
"dtype.int": "integer",
|
||||
"dtype.bigint": "integer",
|
||||
"dtype.smallint": "integer",
|
||||
"dtype.float": "number",
|
||||
"dtype.double": "number",
|
||||
"dtype.decimal": "number",
|
||||
"dtype.boolean": "boolean",
|
||||
"dtype.text": "string",
|
||||
"dtype.varchar": "string",
|
||||
"dtype.date": "date",
|
||||
"dtype.datetime": "datetime",
|
||||
"dtype.timestamp": "datetime",
|
||||
"dtype.timestamptz": "datetime",
|
||||
}
|
||||
target = mapping.get(name)
|
||||
if target is None:
|
||||
raise ExpressionError(f"CAST target {data_type} is not supported.")
|
||||
return target
|
||||
|
||||
|
||||
__all__ = [
|
||||
"ExpressionDataType",
|
||||
"ExpressionError",
|
||||
"ParsedExpression",
|
||||
"convert_value",
|
||||
"evaluate_expression",
|
||||
"infer_expression_type",
|
||||
"parse_expression",
|
||||
]
|
||||
Reference in New Issue
Block a user