"""Sample Snowflake Modeling / Raw SQL usage.""" from oto import response from snowflake_connector.snowflake_conn import ( SQLLoader, fetchone, execute, fetchall, set_default_sessionmaker, ) from snowflake_proxy.snowflake_proxy import ( fetchproxy, fetchproxy_stream, fetchproxy_chunked, fetchproxy_cursor, ) from src.logger import get_logger from src.utils.collections import ( dedupe_list_preserve_order as dedupe_list, zip_rows_to_dicts, ) from src.config import snowflake as config logger = get_logger() sql_loader = SQLLoader(__file__) def initialize_snowflake_session() -> None: """Lazy initialize Snowflake sessionmaker after env & key are ready. Safe to call multiple times (underlying implementation may reuse). """ sf_config = { "role": config.SNOWFLAKE_ROLE, "account": config.SNOWFLAKE_ACCOUNT, "user": config.SNOWFLAKE_USER, "password": config.SNOWFLAKE_PASSWORD, "database": config.SNOWFLAKE_DATABASE, "schema": config.SNOWFLAKE_SCHEMA, "warehouse": config.SNOWFLAKE_WAREHOUSE, } set_default_sessionmaker( connect_args=config.SNOWFLAKE_CONNECT_ARGS, sf_config=sf_config, ) class SMEFeedFileHistoryPersister: """Handles high level operations against SME.""" @classmethod def get_stores_for_period(cls, period_id, table_name): """Get the list of stores for a period id. Args: period_id (str): Period ID by which to limit table_name (str): Table name to query Returns: dict """ params = {"table_name": table_name, "period_id": period_id} sql_template = sql_loader.load_query("select_stores_for_period") try: res = fetchall(sql_template, params=params) # Convert results to list. results = [r[0] for r in res] # Wrap result with oto.Response before return return response.Response(message=results) except Exception as e: logger.error( "Get stores for period_id {} and store from '{}' has failed: {}".format( period_id, table_name, str(e) ) ) raise @classmethod def get_group_names_for_period(cls, period_id, table_name): """Get the list of group names for a period id. Args: period_id (str): Period ID by which to limit table_name (str): Table name to query Returns: dict """ params = {"table_name": table_name, "period_id": period_id} sql_template = sql_loader.load_query("select_group_names_for_period") try: res = fetchall(sql_template, params=params) # Convert results to list. results = [r[0] for r in res] # Wrap result with oto.Response before return return response.Response(message=results) except Exception as e: logger.error( "Get group_names for period_id {} from '{}' has failed: {}".format( period_id, table_name, str(e) ) ) raise @classmethod def get_stores_for_booking_affiliate(cls, affiliate_id, table_name): """Get the list of stores for a booking affiliate. Args: period_id (str): Period ID by which to limit table_name (str): Table name to query Returns: dict """ params = {"table_name": table_name, "affiliate_id": affiliate_id} sql_template = sql_loader.load_query("select_stores_for_affiliate") try: res = fetchall(sql_template, params=params) # Convert results to list. results = [r[0] for r in res] # Wrap result with oto.Response before return return response.Response(message=results) except Exception as e: logger.error( "Get stores for booking affiliate {} and store from " "'{}' has failed: {}".format(affiliate_id, table_name, str(e)) ) raise @classmethod def get_by_period_store(cls, table_name, period_id, store_id): """Get all result rows filtered by period ID and store ID (as a dict). Args: period_id (str): Period ID by which to filter. table_name (str): Table name to query store_id (str): Store ID by which to filter. Returns: dict """ params = { "period_id": period_id, "table_name": table_name, "store_id": store_id, } sql_template = sql_loader.load_query("select_by_period_store") try: res = fetchproxy(sql_template, params=params) # Convert results to dict with field names as keys. results = zip_rows_to_dicts(res) # Wrap result with oto.Response before return return response.Response(message=results) except Exception as e: logger.error( "Get rows by period and store from '{}' has failed: {}".format( table_name, str(e) ) ) raise @classmethod def get_by_period_group_name(cls, table_name, period_id, group_name): """Get all result rows filtered by period ID and group_name (as a dict). Args: period_id (str): Period ID by which to filter. table_name (str): Table name to query store_id (str): Store ID by which to filter. Returns: dict """ params = { "period_id": period_id, "table_name": table_name, "group_name": group_name, } sql_template = sql_loader.load_query("select_by_period_group_name") try: res = fetchproxy(sql_template, params=params) # Convert results to dict with field names as keys. results = zip_rows_to_dicts(res) # Wrap result with oto.Response before return return response.Response(message=results) except Exception as e: logger.error( "Get rows by period and group_name from '{}' has failed: {}".format( table_name, str(e) ) ) raise @classmethod def get_by_label_booking_affiliate_store( cls, table_name, affiliate_id, store_id, label_id ): """Get all result rows filtered by affiliate ID and store ID (as a dict). Args: affiliate_id (str): Booking Affiliate ID by which to filter. table_name (str): Table name to query store_id (str): Store ID by which to filter. label_id (str): label_id by which to filter. Returns: dict """ params = { "affiliate_id": affiliate_id, "table_name": table_name, "store_id": store_id, "label_id": label_id, } sql_template = sql_loader.load_query("select_by_lable_affiliate_store") try: res = fetchproxy(sql_template, params=params) # Convert results to dict with field names as keys. results = zip_rows_to_dicts(res) # Wrap result with oto.Response before return return response.Response(message=results) except Exception as e: logger.error( "Get rows py booking affiliate and store from '{}' has " "failed: {}".format(table_name, str(e)) ) raise @classmethod def get_by_booking_affiliate_store(cls, table_name, affiliate_id, store_id): """Get all result rows filtered by affiliate ID and store ID (as a dict). Args: affiliate_id (str): Booking Affiliate ID by which to filter. table_name (str): Table name to query store_id (str): Store ID by which to filter. label_id (str): label_id by which to filter. Returns: dict """ params = { "affiliate_id": affiliate_id, "table_name": table_name, "store_id": store_id, } if affiliate_id is None: sql_template = sql_loader.load_query("select_by_affiliate_store_null") del params["affiliate_id"] else: sql_template = sql_loader.load_query("select_by_affiliate_store") try: res = fetchproxy(sql_template, params=params) # Convert results to dict with field names as keys. results = zip_rows_to_dicts(res) # Wrap result with oto.Response before return return response.Response(message=results) except Exception as e: logger.error( "Get rows py booking affiliate and store from '{}' has " "failed: {}".format(table_name, str(e)) ) raise @classmethod def get_count_by_booking_affiliate_store(cls, table_name, affiliate_id, store_id): """Get all result rows filtered by affiliate ID and store ID (as a dict). Args: affiliate_id (str): Booking Affiliate ID by which to filter. table_name (str): Table name to query store_id (str): Store ID by which to filter. Returns: dict """ params = { "affiliate_id": affiliate_id, "table_name": table_name, "store_id": store_id, } if affiliate_id is None: sql_template = sql_loader.load_query("select_count_by_affiliate_store_null") del params["affiliate_id"] else: sql_template = sql_loader.load_query("select_count_by_affiliate_store") try: res = fetchone(sql_template, params=params) # Convert results to list. results = res[0] # Wrap result with oto.Response before return return response.Response(message=results) except Exception as e: logger.error( "Get rows py booking affiliate and store from '{}' has " "failed: {}".format(table_name, str(e)) ) raise @classmethod def get_count_by_group_name(cls, table_name, group_name): """Get all result rows filtered by group_name (as a dict). Args: table_name (str): Table name to query group_name (str): Group name to query Returns: dict """ params = { "group_name": group_name, "table_name": table_name, } if group_name is None: sql_template = sql_loader.load_query("select_count_by_group_name_null") del params["group_name"] else: sql_template = sql_loader.load_query("select_count_by_group_name") try: res = fetchone(sql_template, params=params) # Convert results to list. results = res[0] # Wrap result with oto.Response before return return response.Response(message=results) except Exception as e: logger.error( "Get rows by group_name from '{}' has failed: {}".format( table_name, str(e) ) ) raise @classmethod def get_by_label_booking_affiliate(cls, table_name, affiliate_id, label_id): """Get all result rows filtered by label_id, affiliate ID (as a dict). Args: affiliate_id (str): Booking Affiliate ID by which to filter. table_name (str): Table name to query label_id (str): label_id by which to filter. Returns: dict """ params = { "affiliate_id": affiliate_id, "table_name": table_name, "label_id": label_id, } sql_template = sql_loader.load_query("select_by_lable_affiliate") try: # res = fetchall(sql_template, params=params) # Debug res = fetchproxy(sql_template, params=params) # Convert results to dict with field names as keys. results = zip_rows_to_dicts(res) # Wrap result with oto.Response before return return response.Response(message=results) except Exception as e: logger.error( "Get rows py booking affiliate and label from '{}' has " "failed: {}".format(table_name, str(e)) ) raise @classmethod def get_by_period_id(cls, table_name, period_id, label_id): """Get all result rows as a dict. Args: period_id (str): Period ID by which to filter. table_name (str): The name of the table to query label_id (str): Label ID by which to filter. Returns: dict """ params = { "period_id": period_id, "table_name": table_name, "label_id": label_id, } sql_template = sql_loader.load_query("select_by_period_label_id") try: res = fetchproxy(sql_template, params=params) # Convert results to dict with field names as keys. results = zip_rows_to_dicts(res) # Wrap result with oto.Response before return return response.Response(message=results) except Exception as e: logger.error( "Get rows py period from '{}' has failed: {}".format(table_name, str(e)) ) raise @classmethod def get_by_group_name(cls, table_name, group_name): """Get all result rows grouped by `GROUP_NAME` field as a dict. Args: group_name (str): Group Name by which to filter. table_name (str): The name of the table to query Returns: dict """ params = { "group_name": group_name, "table_name": table_name, } sql_template = sql_loader.load_query("select_by_group_name") try: res = fetchproxy(sql_template, params=params) # Convert results to dict with field names as keys. results = zip_rows_to_dicts(res) # Wrap result with oto.Response before return return response.Response(message=results) except Exception as e: logger.error( "Get rows py period from '{}' has failed: {}".format(table_name, str(e)) ) raise @classmethod def get_by_group_name_stream(cls, table_name, group_name): """Get all result rows via streaming iterator. Args: group_name (str): Group Name by which to filter. table_name (str): The name of the table to query Returns: Context manager yielding result set """ params = { "group_name": group_name, "table_name": table_name, } sql_template = sql_loader.load_query("select_by_group_name") return fetchproxy_stream(sql_template, params=params) @classmethod def get_by_group_name_cursor(cls, table_name, group_name, chunk_size=50000): """Get all result rows via server-side cursor chunks. Args: group_name (str): Group Name by which to filter. table_name (str): The name of the table to query chunk_size (int): Size of chunks to yield Returns: Generator yielding chunks of rows """ params = { "group_name": group_name, "table_name": table_name, } sql_template = sql_loader.load_query("select_by_group_name") return fetchproxy_cursor(sql_template, params=params, chunk_size=chunk_size) @classmethod def get_by_group_name_chunked(cls, table_name, group_name, chunk_size=50000): """Get all result rows via LIMIT/OFFSET pagination. Args: group_name (str): Group Name by which to filter. table_name (str): The name of the table to query chunk_size (int): Size of chunks to yield Returns: Generator yielding chunks of rows """ params = { "group_name": group_name, "table_name": table_name, } sql_template = sql_loader.load_query("select_by_group_name") return fetchproxy_chunked(sql_template, params=params, chunk_size=chunk_size) @classmethod def get_all_rows(cls, table_name=None): """Get all result rows as a dict. Args: table_name (str): The name of the table to query Returns: dict """ if not table_name: return response.Response(message={}) sql_template = sql_loader.load_query("select_all_from_table") try: res = fetchproxy(sql_template, params={"table_name": table_name}) # Convert results to dict with field names as keys. results = zip_rows_to_dicts(res) # Wrap result with oto.Response before return return response.Response(message=results) except Exception as e: logger.error( "Get all rows in '{}' has failed: {}".format(table_name, str(e)) ) raise @classmethod def get_all_booking_affiliates(cls, table_name=None): """Get all distinct booking affiliates as a dict. Args: table_name (str): The name of the table to query Returns: dict """ if not table_name: return response.Response(message={}) params = { "table_name": table_name, } sql_template = sql_loader.load_query("select_all_booking_affiliates_from_table") try: res = fetchall(sql_template, params=params) # Convert results to list. results = [r[0] for r in res] # Wrap result with oto.Response before return return response.Response(message=results) except Exception as e: logger.error( "Get all booking affiliates from '{}' has failed: {}".format( table_name, str(e) ) ) raise @classmethod def get_all_periods(cls, table_name=None): """Get all distinct booking affiliates as a dict. Args: table_name (str): The name of the table to query Returns: dict """ if not table_name: return response.Response(message={}) sql_template = sql_loader.load_query("select_all_periods_from_table") try: res = fetchall(sql_template, params={"table_name": table_name}) # Convert results to list. results = [r[0] for r in res] # Wrap result with oto.Response before return return response.Response(message=results) except Exception as e: logger.error( "Get all periods from '{}' has failed: {}".format(table_name, str(e)) ) raise @classmethod def get_all_stores(cls, table_name=None): """Get all distinct booking affiliates as a dict. Args: table_name (str): The name of the table to query Returns: dict """ if not table_name: return response.Response(message={}) sql_template = sql_loader.load_query("select_all_stores_from_table") try: res = fetchall(sql_template, params={"table_name": table_name}) # Convert results to list. results = [r[0] for r in res] # Wrap result with oto.Response before return return response.Response(message=results) except Exception as e: logger.error( "Get all stores from '{}' has failed: {}".format(table_name, str(e)) ) raise @classmethod def get_all_territories(cls, table_name): """Get the list of distinct territories from the group indicator. Args: table_name (str): Table name to query Returns: dict """ params = { "table_name": table_name, } sql_template = sql_loader.load_query("select_all_group_indicators_from_table") try: res = fetchall(sql_template, params=params) # Convert results to list. results = [r[0].split(" ")[0] for r in res] # Dedupe results results = dedupe_list(results) # Wrap result with oto.Response before return return response.Response(message=results) except Exception as e: logger.error( "Get all territories from '{}' has failed: {}".format( table_name, str(e) ) ) raise @classmethod def get_all_group_indicators(cls, table_name): """Get the list of distinct group indicators. Args: table_name (str): Table name to query Returns: dict """ params = { "table_name": table_name, } sql_template = sql_loader.load_query("select_all_group_indicators_from_table") try: res = fetchall(sql_template, params=params) # Parse results and convert to list. results = [r[0].split(" ")[0] for r in res] # Wrap result with oto.Response before return return response.Response(message=results) except Exception as e: logger.error( "Get all territories from '{}' has failed: {}".format( table_name, str(e) ) ) raise @classmethod def get_all_group_names(cls, table_name): """Get the list of distinct group indicators. Args: table_name (str): Table name to query Returns: dict """ params = { "table_name": table_name, } sql_template = sql_loader.load_query("select_all_group_names_from_table") try: res = fetchall(sql_template, params=params) # Parse results and convert to list. results = [r[0] for r in res] # Wrap result with oto.Response before return return response.Response(message=results) except Exception as e: logger.error( "Get all territories from '{}' has failed: {}".format( table_name, str(e) ) ) raise @classmethod def get_all_labels(cls, table_name): """Get the list of booking affiliates. Args: table_name (str): Table name to query Returns: dict """ params = { "table_name": table_name, } sql_template = sql_loader.load_query("select_all_labels_from_table") try: res = fetchall(sql_template, params=params) # Convert results to list. results = [r[0] for r in res] # Wrap result with oto.Response before return return response.Response(message=results) except Exception as e: logger.error( "Get all territories from '{}' has failed: {}".format( table_name, str(e) ) ) raise @classmethod def get_count_of_rows(cls, table_name=None): """Fetch a count of rows. An example of raw SQL usage (which is preferable, because we want to avoid generated by ORM non-optimal SQL queries.) """ if not table_name: return response.Response(message={}) sql_template = sql_loader.load_query("select_count_rows") try: res = fetchone(sql_template, params={"table_name": table_name}) # Convert results to scalar val. results = res[0] # Wrap result with oto.Response before return return response.Response(message=results) except Exception as e: logger.error( "Get count of rows in '{}' has failed: {}".format(table_name, str(e)) ) raise @classmethod def create_table_as_select_by_period(cls, source_table, target_table, period_id): """Create a table with the content from another table/view.""" params = { "source_table_name": source_table, "target_table_name": target_table, "period_id": period_id, } sql_template = sql_loader.load_query("create_table_as_select_by_period") try: # Execute SQL without result execute(sql_template, params=params) except Exception as e: # Minimal enhancement: surface server response details the user # cares about. # Snowflake errors often expose .msg server_msg = getattr(e, 'msg', str(e)) # Query ID for support / troubleshooting sfqid = getattr(e, 'sfqid', None) error_code = getattr(e, 'error_code', getattr(e, 'errno', None)) details = [server_msg] if error_code: details.append(f'code={error_code}') if sfqid: details.append(f'sfqid={sfqid}') logger.error( f"Create table-as-select failed for target '{target_table}' " f"from source '{source_table}' period '{period_id}'. Server " f"response: {details}" ) raise e @classmethod def drop_table(cls, table_name): """Drop Source Table.""" params = { "table_name": table_name, } sql_template = sql_loader.load_query("drop_table") try: # Execute SQL without result execute(sql_template, params=params) except Exception as e: logger.error("Drop table '{}' has failed: {}".format(table_name, str(e))) raise e