import abc import os import pathlib import sys from contextlib import contextmanager from copy import deepcopy 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 sqlalchemy.orm import Session, sessionmaker import playlist.config as config 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 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, } DATABASE_CONFIG = { "database_name": config.SNOWFLAKE_DATABASE, "schema_name": config.SNOWFLAKE_SCHEMA, } if config.ENVIRONMENT == config.TEST_ENVIRONMENT: # pragma: no cover POOL_ARGS = {"poolclass": sqlalchemy.pool.NullPool} if not config.SNOWFLAKE_USER: # pragma: no cover raise Exception("Need to setup .env creds") 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 # only specify macro paths, as templates are loaded directly from reading the string data, not via a path. SQL_MACRO_LOCATIONS = [ BASE_DIR.joinpath("playlist/queries/placements/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 SnowflakeQuery: @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 {} @property def table_params(self) -> dict: """ Default database and table config related args, not validated by the schema """ return {} def __init__(self, user_params: dict = {}): self.user_params = user_params table_params = deepcopy(self.table_params) validated_params = self.validate_params({**self.default_params, **user_params}) for key in ( "public_placement_table", "public_placement_table_name", "recent_placements_table_name", ): val = table_params.get(key) if val and not val.lower().startswith("v_"): table_params[key] = "v_" + val self._query_params = { **validated_params, **table_params, **DATABASE_CONFIG, } @tracer.wrap(name="snowflake_execute") def execute(self): span = tracer.current_span() query, bind_params = self.prepare_query() span.set_tag("query", query) with snowflake_session() as sess: return sess.execute(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.__class__.__module__].__file__) ) ) def load_query_template(self): class_base_dir = self.base_dir path = class_base_dir.joinpath(f"sql/{self.filename}") with open(path) 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)