import enum import os from collections.abc import Iterable, Sequence from contextvars import ContextVar from typing import Any, cast import humps import jinja2 from jinja2 import Environment, Template, nodes from jinja2.ext import Extension from jinja2.lexer import Token, TokenStream from jinja2.parser import Parser from markupsafe import Markup from sqlalchemy_utils import escape_like BindParams = dict[str, Any] DEFAULT_IDENTIFIER_QUOTE_CHAR = "" class ParamStyle(enum.StrEnum): NAMED = "named" _bind_params_var: ContextVar[BindParams] = ContextVar("_bind_params_var", default={}) _param_style_var: ContextVar[ParamStyle] = ContextVar( "_param_style_var", default=ParamStyle.NAMED ) _param_index_var: ContextVar[int] = ContextVar("_param_index_var", default=0) _identifier_quote_char_var: ContextVar[str] = ContextVar( "_identifier_quote_char_var", default=DEFAULT_IDENTIFIER_QUOTE_CHAR ) class JinjaSQLExtension(Extension): def parse(self, parser: Parser) -> nodes.Node | list[nodes.Node]: return [] def filter_stream(self, stream: TokenStream) -> Iterable[Token]: while not stream.eos: token = next(stream) if token.test("variable_begin"): var_expr = [] while not token.test("variable_end"): var_expr.append(token) token = next(stream) variable_end = token last_token = var_expr[-1] lineno = last_token.lineno if not last_token.test("name") or last_token.value not in ( "bind", "inclause", "sqlsafe", "identifier", "orderby", ): param_name = extract_param_name(var_expr) var_expr.insert(1, Token(lineno, "lparen", "(")) var_expr.append(Token(lineno, "rparen", ")")) var_expr.append(Token(lineno, "pipe", "|")) var_expr.append(Token(lineno, "name", "bind")) var_expr.append(Token(lineno, "lparen", "(")) var_expr.append(Token(lineno, "string", param_name)) var_expr.append(Token(lineno, "rparen", ")")) var_expr.append(variable_end) for token in var_expr: yield token else: yield token def extract_param_name(tokens: list[Token]) -> str: name = "" for token in tokens: if token.test("variable_begin"): continue elif token.test("name"): name += token.value elif token.test("dot"): name += token.value else: break if not name: name = "bind#0" return name def sql_safe(value: Any) -> Markup: return Markup(value) def bind(value: Any, name: str) -> Markup | str: if isinstance(value, Markup): return value else: return _bind_param(_bind_params_var.get(), name, value) def bind_in_clause(value: Any) -> str: values = list(value) results = [] for v in values: results.append(_bind_param(_bind_params_var.get(), "inclause", v)) clause = ",".join(results) clause = "(" + clause + ")" return clause def identifier_filter(value: Any) -> Markup: if isinstance(value, str): identifier = (value,) else: identifier = value if not isinstance(value, Iterable): raise ValueError("identifier filter expects a string or an Iterable") identifier_quote_char = _identifier_quote_char_var.get() return Markup( ".".join( "".join( [ identifier_quote_char, s.replace(identifier_quote_char, identifier_quote_char * 2), identifier_quote_char, ] ) for s in identifier ) ) def orderby_filter(value: str, decamelize: bool = True) -> str: column, direction = value.rsplit(".", maxsplit=1) orderby = f"{column} {direction.upper()}" if decamelize: return humps.decamelize(orderby) return orderby def _bind_param(bound: dict[str, Any], key: str, value: Any) -> str: param_index = _param_index_var.get() param_index += 1 _param_index_var.set(param_index) new_key = f"{key.replace('.', '__')}_{param_index}" bound[new_key] = value param_style = _param_style_var.get() if param_style == ParamStyle.NAMED: return f":{new_key}" else: raise ValueError(f"Invalid param_style - {param_style}") class JinjaSQL: def __init__( self, searchpath: str | os.PathLike[str] | Sequence[str | os.PathLike[str]], param_style: ParamStyle = ParamStyle.NAMED, identifier_quote_char: str = DEFAULT_IDENTIFIER_QUOTE_CHAR, ) -> None: self.default_param_style = param_style self.default_identifier_quote_char = identifier_quote_char self.env = Environment(loader=jinja2.FileSystemLoader(searchpath)) self.env.autoescape = True self.env.add_extension(JinjaSQLExtension) self.env.filters["bind"] = bind self.env.filters["inclause"] = bind_in_clause self.env.filters["sqlsafe"] = sql_safe self.env.filters["identifier"] = identifier_filter self.env.filters["orderby"] = orderby_filter self.env.filters["escape_like"] = escape_like def prepare_query( self, template: str | Template, *, context: dict[str, Any] | None = None, param_style: ParamStyle | None = None, identifier_quote_char: str | None = None, ) -> tuple[str, BindParams]: if isinstance(template, str): try: template = self.env.get_template(template) except jinja2.TemplateNotFound: template = self.env.from_string(cast(str, template)) return self._prepare_query( template, context=context, param_style=param_style, identifier_quote_char=identifier_quote_char, ) def _prepare_query( self, template: Template, *, context: dict[str, Any] | None = None, param_style: ParamStyle | None = None, identifier_quote_char: str | None = None, ) -> tuple[str, BindParams]: context = context or {} param_style = param_style or self.default_param_style identifier_quote_char = ( identifier_quote_char or self.default_identifier_quote_char ) bind_params_token = _bind_params_var.set({}) param_style_token = _param_style_var.set(param_style) param_index_token = _param_index_var.set(0) identifier_quote_char_token = _identifier_quote_char_var.set( identifier_quote_char ) try: query = template.render(context) bind_params = _bind_params_var.get() return query, bind_params finally: _bind_params_var.reset(bind_params_token) _param_style_var.reset(param_style_token) _param_index_var.reset(param_index_token) _identifier_quote_char_var.reset(identifier_quote_char_token)