"""SQL templating using JinjaSql. Provides Jinja2-based SQL rendering with safe substitution: - {{ value }} - rendered as a bind parameter placeholder - {{ name | identifier }} - rendered as a quoted SQL identifier Supported engines: - 'athena': ? placeholders, double-quote-quoted identifiers - 'snowflake': %(name)s placeholders, double-quote-quoted identifiers """ import re from typing import Iterable import jinja2 from jinjasql import JinjaSql from jinjasql.core import _bind_param, _thread_local from jinjasql.core import markupsafe def _make_env(): return jinja2.Environment( trim_blocks=True, lstrip_blocks=True, undefined=jinja2.StrictUndefined ) class Engine: """Base class for SQL engines.""" param_style: str identifier_quote_character: str def __init__(self): """Init.""" self._jinja = JinjaSql( env=_make_env(), param_style=self.param_style, identifier_quote_character=self.identifier_quote_character, ) def render(self, template, params=None): """Render a JinjaSQL template, returning query and bind params.""" query, bind_params = self._jinja.prepare_query(template, params or {}) return query, bind_params class AthenaEngine(Engine): """Athena engine.""" param_style = 'qmark' identifier_quote_character = '"' def __init__(self): """Init.""" super().__init__() self._jinja.env.filters['bind'] = self._bind_filter @staticmethod def _encode(value): """Encode a Python value as an Athena execution parameter string. Athena requires string values to be enclosed in single quotes. https://docs.aws.amazon.com/athena/latest/ug/querying-with-prepared-statements.html """ if value is None: return 'NULL' if isinstance(value, bool): return 'true' if value else 'false' if isinstance(value, (int, float)): return str(value) if isinstance(value, str): return f"'{value.replace(chr(39), chr(39) * 2)}'" raise ValueError( f'Unsupported type for Athena bind: {type(value).__name__}' ) @classmethod def _bind_filter(cls, value, name): """Override JinjaSql bind. It stores Athena-encoded value, returns ? placeholder. """ if isinstance(value, markupsafe.Markup): return value return _bind_param(_thread_local.bind_params, name, cls._encode(value)) class SnowflakeEngine(Engine): """Snowflake engine.""" param_style = 'pyformat' identifier_quote_character = '"' MAX_IDENTIFIER_LENGTH = 255 UNQUOTED_IDENTIFIER_RE = re.compile(r'^[A-Za-z_][A-Za-z0-9_$]*$') def __init__(self): """Init.""" super().__init__() self._jinja.env.filters['identifier'] = self._identifier_filter @classmethod def _identifier_filter(cls, raw_identifier): if isinstance(raw_identifier, str): raw_identifier = (raw_identifier,) if not isinstance(raw_identifier, Iterable): raise ValueError( 'identifier filter expects a string or an Iterable' ) # Validate Snowflake identifier rules: # https://docs.snowflake.com/en/sql-reference/identifiers-syntax for identifier in raw_identifier: if not isinstance(identifier, str): raise ValueError('identifier must be string') if not identifier: raise ValueError('identifier must not be empty') if identifier.startswith('"') and identifier.endswith('"'): inner = identifier[1:-1].replace('""', '') if not inner: raise ValueError('identifier must not be empty') if inner.find('"') != -1: raise ValueError( 'quoted identifier contains unescaped double quote' ) if len(inner) > cls.MAX_IDENTIFIER_LENGTH: raise ValueError( f'identifier exceeds Snowflake maximum length of ' f'{cls.MAX_IDENTIFIER_LENGTH} characters' ) else: if len(identifier) > cls.MAX_IDENTIFIER_LENGTH: raise ValueError( f'identifier exceeds Snowflake maximum length of ' f'{cls.MAX_IDENTIFIER_LENGTH} characters' ) if not cls.UNQUOTED_IDENTIFIER_RE.match(identifier): raise ValueError( f'invalid Snowflake identifier: {identifier!r}. ' ) return markupsafe.Markup('.'.join(raw_identifier)) _ENGINES = { 'athena': AthenaEngine(), 'snowflake': SnowflakeEngine(), } def render(template, engine, params=None): """Render a Jinja2 SQL template, returning query and bind params. Args: template (str): Jinja2 SQL template string. engine (str): Target engine: 'athena' or 'snowflake'. params (dict): Template context params. Returns: tuple[str, collection]: Rendered SQL string with placeholders and bind parameter values. Example:: query, params = render( 'SELECT * FROM {{ db | identifier }} WHERE date = {{ date }}', engine='athena', params={'db': 'my_db', 'date': '2024-01-15'}, ) # query: 'SELECT * FROM "my_db" WHERE date = ?' # params: ('2024-01-15',) """ if engine not in _ENGINES: raise ValueError( f'Unsupported engine "{engine}". ' f'Valid engines: {list(_ENGINES.keys())}' ) return _ENGINES[engine].render(template, params)