"""Proxy for snowflake requests in order to log them in New Relic.""" import abc import os import pathlib import sys from contextlib import contextmanager from pathlib import Path from typing import List import jinjasql import snowflake.sqlalchemy import sqlalchemy.orm import sqlalchemy.pool from ddtrace import tracer from jinja2 import Environment, FileSystemLoader from jinjasql import JinjaSql from markupsafe import Markup from marshmallow import Schema from snowflake.connector.connection import SnowflakeConnection from snowflake.sqlalchemy.snowdialect import SnowflakeDialect from snowflake_connector import snowflake_conn from sqlalchemy import text from sqlalchemy.orm import Session, sessionmaker def pass_rollback(self) -> None: """We're only doing SELECT queries. Rollback is not necessary. Original implementation in the connector: def rollback(self) -> None: self.cursor().execute("ROLLBACK") """ pass SnowflakeConnection.rollback = pass_rollback import analytics.config as config SnowflakeDialect.supports_statement_cache = False DEFAULT_CONFIG = { "pool_pre_ping": False, "pool_reset_on_return": None, "commit_before_close": False, } @tracer.wrap(name="snowflake_fetchone") def fetchone(*args, **kwargs): """Passthrough to Snowflake to fetchone.""" kwargs.update(DEFAULT_CONFIG) return snowflake_conn.fetchone(*args, **kwargs) @tracer.wrap(name="snowflake_fetchall") def fetchall(*args, **kwargs): """Passthrough to Snowflake to fetchall.""" kwargs.update(DEFAULT_CONFIG) return snowflake_conn.fetchall(*args, **kwargs) def SQLLoader(path): """Create a SQL loader in given path.""" return snowflake_conn.SQLLoader(path) POOL_ARGS = { "poolclass": sqlalchemy.pool.QueuePool, "pool_size": config.SNOWFLAKE_POOL_SIZE, "pool_recycle": config.SNOWFLAKE_POOL_RECYCLE, "pool_pre_ping": config.SNOWFLAKE_POOL_PRE_PING, "pool_reset_on_return": config.SNOWFLAKE_POOL_RESET_ON_RETURN, "max_overflow": config.SNOWFLAKE_POOL_MAX_OVERFLOW, } if config.ENVIRONMENT == config.TEST_ENVIRONMENT: # pragma: no cover POOL_ARGS = {"poolclass": sqlalchemy.pool.NullPool} engine = sqlalchemy.create_engine( snowflake.sqlalchemy.URL( account=config.SNOWFLAKE_ACCOUNT, user=config.SNOWFLAKE_USER, password="dummy_password", # we're using key pair auth, pk is passed in config.SNOWFLAKE_CONNECT_ARGS # noqa database=config.SNOWFLAKE_DATABASE, schema=config.SNOWFLAKE_SCHEMA, warehouse=config.SNOWFLAKE_WAREHOUSE, role=config.SNOWFLAKE_ROLE, client_session_keep_alive=True, ), connect_args=config.SNOWFLAKE_CONNECT_ARGS, **POOL_ARGS, ) DEFAULT_SESSION_FACTORY: sessionmaker = sqlalchemy.orm.sessionmaker(bind=engine) BASE_DIR = pathlib.Path(__file__).resolve().parent.parent.parent SQL_MACRO_LOCATIONS = [ BASE_DIR.joinpath("analytics/queries/sql/macros"), ] def get_template_engine() -> JinjaSql: return jinjasql.JinjaSql( env=Environment( loader=FileSystemLoader(SQL_MACRO_LOCATIONS), trim_blocks=True, lstrip_blocks=True, ), param_style="named", ) @contextmanager def snowflake_session() -> Session: session = DEFAULT_SESSION_FACTORY() try: yield session finally: session.close() class AbstractSnowflakeQuery: @property @abc.abstractmethod def filename(self) -> str: """name of this query's sql file""" raise NotImplementedError @property def query_params(self) -> dict: """combining params containing validated data and database config""" return self._query_params @property @abc.abstractmethod def query_schema(self) -> Schema: """marshmallow schema for this endpoint""" raise NotImplementedError @property def default_params(self) -> dict: """Default parameters, these must exist in the schema""" return {} def __init__(self, user_params: dict = {}): self.user_params = user_params validated_params = self.validate_params({**self.default_params, **user_params}) self._query_params = { **validated_params, } def execute(self): query, bind_params = self.prepare_query() with snowflake_session() as sess: return sess.execute(text(query), bind_params) @property def base_dir(self) -> Path: # evaluates to the directory that a subclass that inherits this class resides in return Path( os.path.dirname( os.path.abspath(sys.modules[self.query_schema.__module__].__file__) ) ) def load_query_template(self): class_base_dir = self.base_dir path = class_base_dir.joinpath(f"sql/{self.filename}") # Explicit utf-8: the prod Docker image (python:3.14-slim-bookworm) # has no locale set, so open()'s default encoding resolves to ASCII # and any non-ASCII byte in a SQL template (e.g. an em-dash in a # comment) raises UnicodeDecodeError at request time. with open(path, encoding="utf-8") as fd: return fd.read() def prepare_query(self): """load query template, and JinjaSql.prepare_query() them""" query_template_str = self.load_query_template() template_engine = get_template_engine() query, bind_params = template_engine.prepare_query( query_template_str, self.query_params ) # In jinjasql, macros come through as markupsafe.Markup # Here we convert Markup to str so they're supported as a database params if bind_params: items = list(bind_params.items()) for key, value in items: if type(value) is Markup: bind_params[key] = value.striptags() return query, bind_params def validate_params(self, kwargs: dict) -> dict: """Return validated user input, Raises if invalid""" schema = self.query_schema() schema_fields: List[str] = schema.fields.keys() for key in list(kwargs): if key not in schema_fields: kwargs.pop(key, None) # Validate schema, will raise exception if incorrect types. return schema.load(kwargs) class SnowflakeQuery(AbstractSnowflakeQuery): """composition over inheritance version of AbstractSnowflakeQuery""" def __init__( self, filename, schema, user_params: dict = {}, default_params: dict = {} ): self._filename = filename self._query_schema = schema self._default_params = default_params super(SnowflakeQuery, self).__init__(user_params) @property def query_schema(self) -> Schema: return self._query_schema @property def filename(self) -> str: return self._filename @property def default_params(self) -> dict: return self._default_params