"""Parsing utilities for SQL builders.
Provides common parsing functions to handle SQL expressions
passed as strings to builder methods.
"""
import contextlib
import re
from collections.abc import Callable
from typing import TYPE_CHECKING, Any, Final
from sqlglot import exp, maybe_parse
from sqlspec.builder._column import Column
from sqlspec.builder._expression_wrappers import ExpressionWrapper
from sqlspec.core import ParameterStyle, ParameterValidator
from sqlspec.utils.type_guards import (
has_expression_and_parameters,
has_expression_and_sql,
has_expression_attr,
has_parameter_builder,
)
if TYPE_CHECKING:
from sqlglot.dialects.dialect import DialectType
__all__ = (
"extract_expression",
"extract_sql_object_expression",
"parse_column_expression",
"parse_condition_expression",
"parse_order_expression",
"parse_table_expression",
"to_expression",
)
ALIAS_PARTS_EXPECTED_COUNT = 2
QUALIFIED_IDENTIFIER_PARTS = 2
_SIMPLE_IDENTIFIER_RE: Final["re.Pattern[str]"] = re.compile(
r"^[A-Za-z_][A-Za-z0-9_$]*(?:\.[A-Za-z_][A-Za-z0-9_$]*){0,2}$"
)
_BARE_KEYWORDS: Final[frozenset[str]] = frozenset({
"all",
"and",
"any",
"asc",
"between",
"case",
"current_date",
"current_time",
"current_timestamp",
"current_user",
"default",
"delete",
"desc",
"distinct",
"end",
"exists",
"false",
"from",
"in",
"insert",
"interval",
"is",
"like",
"localtime",
"localtimestamp",
"not",
"null",
"or",
"select",
"session_user",
"some",
"true",
"update",
"user",
"where",
})
_PARAMETER_VALIDATOR = ParameterValidator()
def extract_column_name(column: str | exp.Column) -> str:
"""Extract column name from column expression for parameter naming.
Args:
column: Column expression (string or SQLGlot Column)
Returns:
Column name as string for use as parameter name
"""
if isinstance(column, str):
col_expr: exp.Expr | None = exp.maybe_parse(column)
if isinstance(col_expr, exp.Column):
return col_expr.name
return column.split(".")[-1] if "." in column else column
if isinstance(column, exp.Column):
return column.name
return "column"
def _merge_sql_parameters(sql_obj: Any, builder: Any) -> None:
"""Merge parameters from SQL object into builder.
Args:
sql_obj: SQL object with parameters attribute
builder: Builder instance with add_parameter method
"""
if not (builder and has_expression_and_parameters(sql_obj) and has_parameter_builder(builder)):
return
for param_name, param_value in sql_obj.parameters.items():
builder.add_parameter(param_value, name=param_name)
def _is_simple_identifier(value: str) -> bool:
stripped = value.strip()
if not _SIMPLE_IDENTIFIER_RE.fullmatch(stripped):
return False
return "." in stripped or stripped.lower() not in _BARE_KEYWORDS
def _simple_column_expression(value: str) -> exp.Column:
parts = value.strip().split(".")
identifiers = [exp.Identifier(this=part, quoted=False) for part in parts]
if len(parts) == 1:
return exp.Column(this=identifiers[0])
if len(parts) == QUALIFIED_IDENTIFIER_PARTS:
return exp.Column(this=identifiers[1], table=identifiers[0])
return exp.Column(this=identifiers[2], table=identifiers[1], db=identifiers[0])
def parse_column_expression(column_input: str | exp.Expr | Any, builder: Any | None = None) -> exp.Expr:
"""Parse a column input that might be a complex expression.
Handles cases like:
- Simple column names: "name" -> Column(this=name)
- Qualified names: "users.name" -> Column(table=users, this=name)
- Aliased columns: "name AS user_name" -> Alias(this=Column(name), alias=user_name)
- Function calls: "MAX(price)" -> Max(this=Column(price))
- Complex expressions: "CASE WHEN ... END" -> Case(...)
- Custom Column objects from our builder
- SQL objects with raw SQL expressions
Args:
column_input: String, SQLGlot expression, SQL object, or Column object
builder: Optional builder instance for parameter merging
Returns:
exp.Expr: Parsed SQLGlot expression
"""
if isinstance(column_input, exp.Expr):
return column_input
if isinstance(column_input, str):
if _is_simple_identifier(column_input):
return _simple_column_expression(column_input)
return exp.maybe_parse(column_input) or exp.column(column_input)
if has_expression_and_sql(column_input):
if column_input.expression is not None and isinstance(column_input.expression, exp.Expr):
_merge_sql_parameters(column_input, builder)
return column_input.expression
_merge_sql_parameters(column_input, builder)
sql_str = getattr(column_input, "raw_sql", None)
if sql_str is None:
sql_str = column_input.sql
return exp.maybe_parse(sql_str) or exp.column(sql_str)
if has_expression_attr(column_input) and isinstance(column_input._expression, exp.Expr): # pyright: ignore[reportPrivateUsage]
return column_input._expression # pyright: ignore[reportPrivateUsage]
return exp.maybe_parse(column_input) or exp.column(str(column_input)) # pyright: ignore[reportArgumentType]