"""The SQL data base SQLAlchemy connector.""" from sqlalchemy import create_engine from sqlalchemy import exc from sqlalchemy import sql as alchemy_sql from sqlalchemy.pool import NullPool from accounting.flows.reserve_payouts import connectors from accounting.flows.reserve_payouts import setting from accounting.flows.reserve_payouts.constants import sql def _create_engines(): """Create DB engines that are suitable for select queries. Returns: dict: database engines for each database. """ db_engines = {} for db_name in setting.DATABASES: connection_string = setting.DB_CONNECTION_STRINGS[db_name] db_engine = create_engine( connection_string, connect_args=setting.DB_CONNECT_ARGS, poolclass=NullPool) db_engines[db_name] = db_engine return db_engines def get_db_engine(db_name): """Get an SQLAlchemy database engine for db_name. Args: db_name (str): the name of the database, should be present in settings Returns: sqlalchemy.engine.base.Engine: database engine instance """ return _db_engines.get(db_name) def health_check(db_name): """Perform a simple query to do the health check. Args: db_name (srt): name of the database Returns: namedtuple: with bool and message attributes, (True, '') if connection is ok, (False, 'Error message') otherwise. """ engine = get_db_engine(db_name) try: query_result = engine.execute(sql.DB_HEALTH_CHECK_QUERY) query_result.close() except (exc.SQLAlchemyError, exc.ProgrammingError) as e: result = connectors.HealthCheckResult(False, str(e)) else: result = connectors.HealthCheckResult(True, '') return result def _execute_sql(db_name, query, params): """Execute the query against the database. Args: db_name (str): name of the database query (str): SQL query params (dict): query parameters Returns: sqlalchemy.engine.result.ResultProxy: query result """ engine = get_db_engine(db_name) return engine.execute(query, params) def get_physical_transactions_sum(period_id): """Fetch physical transactions data from the database. Args: period_id (int): period identifier Returns: sqlalchemy.engine.result.ResultProxy: query result """ raw_query = sql.PHYSICAL_TRANSACTIONS_SUM_SQL.format( accountingflat=setting.DB_ACC_FLAT_NAME, art_relations=setting.DB_ART_RELATIONS_NAME) query = alchemy_sql.text(raw_query) query_params = {'period_id': period_id} result = _execute_sql(setting.DB_ACC_FLAT_NAME, query, query_params) return result def get_vendor_contracts(period_id, label_ids): """Fetch contract information for reserve payout calculation. Args: period_id (int): period identifier label_ids (list|tuple|dict_keys): list of vendors Returns: sqlalchemy.engine.result.ResultProxy: query result """ query = alchemy_sql.text(sql.VENDOR_CONTRACT_SQL) query_params = { 'label_ids': tuple(label_ids), 'period_id': period_id} result = _execute_sql(setting.DB_ART_RELATIONS_NAME, query, query_params) return result def _make_dict_from_sql_result(keys, data_tuple): """Make a dict from sql execution result item and result.keys(). Args: keys (list): list of query fields data_tuple (tuple): query result row Returns: dict: combined result """ return dict(zip(keys, data_tuple)) def vendor_query_result_to_dict(query_result): """Transform vendor query result into dict. Args: query_result (sqlalchemy.engine.result.ResultProxy): query result Returns: dict: resulting dict, where vendor_id is a key """ keys = query_result.keys() processed = {} for row in query_result: item = _make_dict_from_sql_result(keys, row) vendor_id = item['vendor_id'] if vendor_id in processed: msg = 'Duplicate row in query result! vendor_id: {}'.format( vendor_id) query_result.close() raise ValueError(msg) processed[vendor_id] = item query_result.close() return processed def truncate_reserves_temp_table(): """Truncate the temp table for reserve payout calculation.""" truncate_query = alchemy_sql.text(sql.TRUNCATE_TEMP_TABLE) result = _execute_sql(setting.DB_ACC_FLAT_NAME, truncate_query, {}) result.close() def populate_reserves_temp_table(period_id): """Populate the temp table for reserve payout calculation. Args: period_id (int): period id to filter transactions data Returns: int: Row count (affected rows) """ insert_query = alchemy_sql.text(sql.INSERT_PHYS_TRANSACTIONS_FOR_PERIOD) query_params = {'period_id': period_id} result = _execute_sql( setting.DB_ACC_FLAT_NAME, insert_query, query_params) result.close() return result.rowcount _db_engines = _create_engines()