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", ]