""" Snowflake Adapter Tests Description: Snowflake Adapter tests. """ import os import pytest import numpy as np import pandas as pd from absl import logging from forecasting_toolkit.datastore.adapters.snowflake import ( SnowflakeDatasetAdapter, AlchemyDatasetAdapter ) from forecasting_toolkit.datastore.connectors.snowflake import ( snowflake_connector_factory, alchemy_connector_factory, set_snowflake_environment ) from forecasting_toolkit.datastore.adapters.helpers import ( create_snowflake_table ) SNOWFLAKE_WAREHOUSE = os.environ.get("SNOWFLAKE_WAREHOUSE", "DEV_OWS_WAREHOUSE") SNOWFLAKE_DB = os.environ.get("SNWOFLAKE_DB", "DEV_ENGINEERING") SNOWFLAKE_SCHEMA = os.environ.get("SNOWFLAKE_SCHEMA", "AADAMU_DEBUT_FORECASTING_DBT") SNOWFLAKE_ROLE = os.environ.get("SNOWFLAKE_ROLE", "DEV_ENGINEERING") @pytest.fixture(autouse=True, scope='session') def snowflake_dataset_adapter(): with snowflake_connector_factory() as conn: # set snowflake environment logging.debug("setting up DB Env") # TODO: at a later point - makes this ENV vars that get injected set_snowflake_environment(conn_cursor=conn, warehouse=SNOWFLAKE_WAREHOUSE, database=SNOWFLAKE_DB, schema=SNOWFLAKE_SCHEMA) # create snowflake dataset adapter snowflake_dataset_adapter = SnowflakeDatasetAdapter(conn=conn) yield snowflake_dataset_adapter @pytest.fixture(autouse=True, scope='session') def alchemy_dataset_adapter(): kwargs = { "database": "DEV_ENGINEERING", "schema": "AADAMU_DEBUT_FORECASTING_DBT", "warehouse": "DEV_OWS_WAREHOUSE", "role": "DEV_ENGINEERING" } with alchemy_connector_factory(**kwargs) as conn: # create snowflake dataset adapter alchemy_dataset_adapter = AlchemyDatasetAdapter(conn=conn) yield alchemy_dataset_adapter def test_snowflake_adapter(snowflake_dataset_adapter): """tests snowflake adapter """ dataset_table = "DATASET_STREAMS_DAILY_2022" limit = 10 filters = {} # fetch dataset logging.debug("Fetching dataset from snowflake") dataset_df = snowflake_dataset_adapter.fetch_dataset(snowflake_table=dataset_table, filters=filters, limit=limit) # ensure it has data less than or equal to the limit assert len(dataset_df) <= 10 # ensure its a data frame assert isinstance(dataset_df, pd.DataFrame) def test_store_filters_with_limit(snowflake_dataset_adapter): """Test filters with limit""" dataset_table = "DATASET_STREAMS_DAILY_2022" limit = 10 filters = {"STORE_ID": 286} # fetch dataset logging.debug("Fetching dataset from snowflake") dataset_df = snowflake_dataset_adapter.fetch_dataset(snowflake_table=dataset_table, filters=filters, limit=limit) # ensure its a dataframe contains data from same store id assert len(dataset_df[dataset_df["STORE_ID"] == 286]) == len(dataset_df) def test_country_with_limit(snowflake_dataset_adapter): """test filters only""" dataset_table = "DATASET_STREAMS_DAILY_2022" limit = 10 filters = {"COUNTRY_CODE": "US"} # fetch dataset logging.debug("Fetching dataset from snowflake") dataset_df = snowflake_dataset_adapter.fetch_dataset(snowflake_table=dataset_table, filters=filters, limit=limit) # ensure its a dataframe contains data from same store id assert len(dataset_df[dataset_df["COUNTRY_CODE"] == "US"]) == len(dataset_df) def test_alchemy_with_country_code_and_limit(alchemy_dataset_adapter): dataset_table = "DATASET_STREAMS_DAILY_2022" limit = 10 filters = {"COUNTRY_CODE": "US"} # fetch dataset logging.debug("Fetching dataset from snowflake") dataset_df = alchemy_dataset_adapter.fetch_dataset(snowflake_table=dataset_table, filters=filters, limit=limit) logging.debug(dataset_df) # ensure its a dataframe contains data from same store id assert len(dataset_df[dataset_df["country_code"] == "US"]) == len(dataset_df) def test_alchemy_with_snowflake_write(alchemy_dataset_adapter): """test filters only""" dataset_table = "DATASET_STREAMS_DAILY_2022" limit = 10 filters = {} # fetch dataset logging.debug("Fetching dataset from snowflake") dataset_df = alchemy_dataset_adapter.fetch_dataset(snowflake_table=dataset_table, filters=filters, limit=limit) dataset_df.columns = [col.upper() for col in dataset_df.columns] alchemy_dataset_adapter.to_snowflake( snowflake_table="TEST_SNOWFLAKE_TABLE", data_df=dataset_df[['SNAPSHOT_DATE', 'ISRC', 'STREAMS', 'UPC']], if_exists='replace') dataset_new_df = alchemy_dataset_adapter.fetch_dataset(snowflake_table="TEST_SNOWFLAKE_TABLE", filters={}, limit=None) assert (len(dataset_df), 4) == dataset_new_df.shape def test_create_snowflake_table(): """tests creation of snowflake tables using sql alchecmy""" snowflake_table = "AADAMU_TEST_SNOWFLAKE_TBL" create_snowflake_table(snowflake_table=snowflake_table)