Files
govoplan-dataflow/src/govoplan_dataflow/backend/expressions.py
T

644 lines
17 KiB
Python

from __future__ import annotations
import operator
from dataclasses import dataclass
from decimal import Decimal
from typing import Any, Callable, 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:
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 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,
*,
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:
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,
exp.GTE,
exp.LT,
exp.LTE,
exp.And,
exp.Or,
exp.Is,
exp.Not,
)
},
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:
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",
]