from unittest.mock import patch import pandas as pd import snowflake.snowpark.session as ses import streamlit as st from pandas._testing import assert_frame_equal from snowflake.snowpark.types import StructType, StructField, StringType, IntegerType, DateType from common.data_utilities import extract_fan_data, prepare_streamlit_app @patch("common.data_utilities.DATABASE_NAME", "MOCK_DATABASE") @patch("common.data_utilities.SCHEMA_NAME", "MOCK_SCHEMA") @patch("common.data_utilities.FANS", "FANS") def test_extract_fan_data(session: ses.Session, fan_ids_for_export: pd.DataFrame, fans: pd.DataFrame)-> None: """ Test the basket input creation :param session: Snowpark session object :param fan_ids_for_export: Dataframe containing potential fan ids used for finding related data :param fans: Dataframe referencing/mocking FANS_C table in Snowflake :return: True if test passes, otherwise False """ st.session_state.email_column_key_2 = True st.session_state.mobile_phone_column_key_2 = True tbl = session.create_dataframe(fans) tbl.write.mode('overwrite').save_as_table(["MOCK_DATABASE", "MOCK_SCHEMA", 'FANS'], mode='overwrite') result_df = fans.merge(fan_ids_for_export, left_on='ID', right_on='FAN_ID')[['FAN_ID', 'ID', 'EMAIL_C', 'MOBILE_PHONE_C']] result_df['EXPORT_DATA'] = result_df.apply(lambda row: row.to_dict(), axis=1) result_df.columns = pd.Index(['FAN_ID', 'ID', 'EMAIL', 'MOBILE_PHONE', 'EXPORT_DATA']) returned_df, returned_query, email_col_exclusion = extract_fan_data(fan_ids_for_export, session, st.session_state) assert returned_query == "SELECT MOCK_TEST_FAKE_QUERY()" assert_frame_equal(result_df, returned_df) @patch("common.data_utilities.DATABASE_NAME", "MOCK_DATABASE") @patch("common.data_utilities.SCHEMA_NAME", "MOCK_SCHEMA") @patch("common.data_utilities.FORMS", "FORMS") @patch("common.data_utilities.MAILING_LISTS", "MAILING_LISTS") @patch("common.data_utilities.TLAS", "TLAS") @patch("common.data_utilities.TLS", "TLS") @patch("common.data_utilities.TERRITORIES", "TERRITORIES") @patch("common.data_utilities.LABELS", "LABELS") @patch("common.data_utilities.ARTISTS", "ARTISTS") def test_prepare_streamlit_app(session: ses.Session, forms: pd.DataFrame, mailing_lists: pd.DataFrame, tlas: pd.DataFrame, tls: pd.DataFrame, territories: pd.DataFrame, labels: pd.DataFrame, artists: pd.DataFrame)-> None: """ Test function that populates data for main export drop down selections :param forms: Pandas dataframe replicating data from FORM_C table in Snowflake :param mailing_lists: Pandas dataframe replicating data from MAILING_LIST_C table in Snowflake :param tlas: Pandas dataframe replicating data from TLA_C table in Snowflake :param tls: Pandas dataframe replicating data from TL_C table in Snowflake :param territories: Pandas dataframe replicating data from TERRITORY_C table in Snowflake :param labels: Pandas dataframe replicating data from LABEL_C table in Snowflake :param artists: Pandas dataframe replicating data from ARTIST_C table in Snowflake """ forms_tbl = session.create_dataframe(forms) forms_tbl.write.mode('overwrite').save_as_table(["MOCK_DATABASE", "MOCK_SCHEMA", 'FORMS'], mode='overwrite') mailing_lists_tbl = session.create_dataframe(mailing_lists) mailing_lists_tbl.write.mode('overwrite').save_as_table(["MOCK_DATABASE", "MOCK_SCHEMA", 'MAILING_LISTS'], mode='overwrite') tlas_tbl = session.create_dataframe(tlas) tlas_tbl.write.mode('overwrite').save_as_table(["MOCK_DATABASE", "MOCK_SCHEMA", 'TLAS'], mode='overwrite') tls_tbl = session.create_dataframe(tls) tls_tbl.write.mode('overwrite').save_as_table(["MOCK_DATABASE", "MOCK_SCHEMA", 'TLS'], mode='overwrite') territories_tbl = session.create_dataframe(territories) territories_tbl.write.mode('overwrite').save_as_table(["MOCK_DATABASE", "MOCK_SCHEMA", 'TERRITORIES'], mode='overwrite') labels_tbl = session.create_dataframe(labels) labels_tbl.write.mode('overwrite').save_as_table(["MOCK_DATABASE", "MOCK_SCHEMA", 'LABELS'], mode='overwrite') artists_tbl = session.create_dataframe(artists) artists_tbl.write.mode('overwrite').save_as_table(["MOCK_DATABASE", "MOCK_SCHEMA", 'ARTISTS'], mode='overwrite') forms_df, mailing_lists_df = prepare_streamlit_app(session) session.file.put("streamlit_app/tests/expected_forms.csv", "@mystage", auto_compress=False) session.file.put("streamlit_app/tests/expected_mailing_lists.csv", "@mystage", auto_compress=False) schema = StructType( [ StructField("MAILING_LIST_ID", StringType()), StructField("TLA_ID", StringType()), StructField("TERRITORY_C", StringType()), StructField("LABEL_NAME", StringType()), StructField("ARTIST_C", StringType()), StructField("MAILING_LIST_NAME_C", StringType()) ] ) expected_mailing_lists = session.read.schema(schema).option("SKIP_HEADER", 1).csv("@mystage/expected_mailing_lists.csv") expected_mailing_lists_df = expected_mailing_lists.to_pandas() schema = StructType( [ StructField("FORM_ID", StringType()), StructField("TLA_ID", StringType()), StructField("TERRITORY_C", StringType()), StructField("LABEL_NAME", StringType()), StructField("ARTIST_C", StringType()), StructField("FORM_NAME_C", StringType()), StructField("FORM_MIGRATION_ID", IntegerType()), StructField("FORM_LAST_MODIFIED_DATE", DateType()), StructField("MIGRATION_ID_FORM_NAME", StringType()), ] ) expected_forms = session.read.schema(schema).option("SKIP_HEADER", 1).csv( "@mystage/expected_forms.csv") expected_forms_df = expected_forms.to_pandas() assert_frame_equal(expected_forms_df, forms_df) assert_frame_equal(expected_mailing_lists_df, mailing_lists_df)