Source code for sqlspec.builder._parsing_utils

"""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]
def parse_table_expression( table_input: str, explicit_alias: "str | None" = None, dialect: "DialectType | None" = None ) -> exp.Expr: r"""Parses a table string that can be a name, a name with an alias, or a subquery string. The ``dialect`` selects the identifier-quoting rules so dialect-quoted identifiers such as BigQuery's ``\`project.dataset.table\``` are parsed into their qualified parts instead of a single literal name. """ if explicit_alias is None and " " in table_input.strip(): parts = table_input.strip().split(None, 1) if len(parts) == ALIAS_PARTS_EXPECTED_COUNT: base_table, alias = parts return exp.to_table(base_table, alias=alias, dialect=dialect) if _is_simple_identifier(table_input): return exp.to_table(table_input, alias=explicit_alias, dialect=dialect) with contextlib.suppress(Exception): parsed: exp.Expr | None = exp.maybe_parse(f"SELECT * FROM {table_input}", dialect=dialect) if isinstance(parsed, exp.Select): from_clause = parsed.find(exp.From) if from_clause is not None: table_expr = from_clause.this if explicit_alias: return exp.alias_(table_expr, explicit_alias) return table_expr # type: ignore[no-any-return] return exp.to_table(table_input, alias=explicit_alias, dialect=dialect)