"""The MySQL-related helpers.""" from contextlib import contextmanager from urllib import parse import pymysql from sqlalchemy import create_engine from sqlalchemy.orm import sessionmaker from sqlalchemy.pool import NullPool from sqlalchemy.sql import text from feed_sender.conf import config from feed_sender.util import sentry FETCHMANY_SIZE = 100 # number of rows you fetch at a time. def _get_db_connection_url(): """Construct db connection url from credentials from config. Returns: str: db connection url. """ mysql_config = config.ART_RELATIONS return ( 'mysql+pymysql://{MYSQL_USER}:{MYSQL_PASSWORD}' '@{MYSQL_HOST}/{MYSQL_DATABASE}').format( MYSQL_USER=mysql_config.get('user'), MYSQL_PASSWORD=parse.quote(mysql_config.get('password').encode('utf-8')), MYSQL_HOST=mysql_config.get('host'), MYSQL_DATABASE=mysql_config.get('database')) def _get_pool_class(): """Get pool class depends on environment. Returns: mixed: NullPool or None for dev environment. """ env = config.ENV POOLCLASS = NullPool if env == 'dev': POOLCLASS = None # sqlite does not play nice with NullPool return POOLCLASS def _session_factory(): """Produce session object. Returns: session object. """ db_engine = create_engine( _get_db_connection_url(), poolclass=_get_pool_class()) return sessionmaker(bind=db_engine)() @contextmanager def db_session(): """Provide a transactional scope around a series of operations. Taken from http://docs.sqlalchemy.org/en/latest/orm/session_basics.html. This handles rollback and closing of session, so there is no need to do that throughout the code. Usage: with db_session() as session: session.execute(query) """ session = _session_factory() try: yield session session.commit() except Exception: session.rollback() sentry_client = sentry.get_client() sentry_client.captureException() raise finally: session.close() def execute(sql, **args): """Execute MySQL query. Args: sql (str): SQL query. args (dict): dictionary of parameters. Example: dict(vendor_id=vendor_id) """ with db_session() as session: session.execute(text(sql), args) def fetchall(sql, **args): """Fetch all MySQL query. Args: sql (str): SQL query. args (dict): dictionary of parameters. Example: dict(vendor_id=vendor_id) Yields: tuple: Fetched values. """ with db_session() as session: results = session.execute(text(sql), args) while True: rows = results.mappings().fetchmany(FETCHMANY_SIZE) if not rows: break for row in rows: yield row def get_art_db_connection_pymysql(autocommit=True): """Get the DB connection for PyMSQL. Return: obj: PyMYSQL connection object. """ mysql_config = config.ART_RELATIONS return pymysql.connect( host=mysql_config.get('host'), user=mysql_config.get('user'), password=mysql_config.get('password'), database='art_relations', autocommit=autocommit, charset='utf8mb4', cursorclass=pymysql.cursors.DictCursor) def get_dd_db_connection_pymysql(autocommit=True): """Get the DB connection for PyMSQL. Return: Connection: PyMYSQL connection object. """ mysql_config = config.DIRECT_DELIVERY return pymysql.connect( host=mysql_config.get('host'), user=mysql_config.get('user'), password=mysql_config.get('password'), database='direct_delivery', autocommit=autocommit, charset='utf8', cursorclass=pymysql.cursors.DictCursor) def execute_query(connection, sql): """Execute READ query using the PyMySQL connection object. Args: connection (Connection): PyMySQL connection object. sql (str): fully formed SQL query. Return: dict: DictCursor object of the resulting rows. """ with connection.cursor() as cursor: cursor.execute(sql) result = cursor.fetchall() return result def execute_write_query(connection, sql, params=None): """Execute insert/update query using the PyMySQL connection object. Args: connection (Connection): PyMySQL connection object. sql (str): fully formed SQL query. params (dict|tuple): parameters passed in to query. Return: int: Number of rows affected. """ with connection.cursor() as cursor: cursor.execute(sql, params) return cursor.rowcount