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