"""Persister tests.""" from decimal import Decimal from unittest.mock import MagicMock import pytest from snowflake_connector import snowflake_conn from sqlalchemy.exc import SQLAlchemyError from sqlalchemy.orm.session import Session from reporting.models import persister from reporting.utils import error_handling from reporting.utils import single_product_overview as spo def test_get_physical_product_monthly_trend_success( local_product_cd, supply_chain_id, single_product_monthly ): """Test persister get_physical_product_monthly_trend success state.""" response = persister.get_physical_product_monthly_trend( local_product_cd, supply_chain_id ) expected_table = spo.extract_sales_by_month_table(single_product_monthly) assert len(expected_table) == len(response.message) for idx, expected_entry in enumerate(expected_table.values()): for item in expected_entry: assert item in list(response.message.values())[idx] def test_get_physical_product_monthly_trend_empty( monkeypatch, local_product_cd, supply_chain_id ): """Test get_physical_product_monthly_trend not found.""" monkeypatch.setattr(Session, 'execute', value=MagicMock(return_value=[])) persister_response = persister.get_physical_product_monthly_trend( local_product_cd, supply_chain_id ) assert persister_response.status == 200 assert persister_response.message == spo._empty_sales_by_month_table() def test_get_physical_product_monthly_trend_failure( mocker, local_product_cd, supply_chain_id ): """Test get_physical_product_monthly_trend failure.""" mocker.patch.object(snowflake_conn, 'get_session') mocker.patch.object(Session, 'execute', side_effect=SQLAlchemyError) mocker.patch.object(error_handling, 'log_db_exception') response = persister.get_physical_product_monthly_trend( local_product_cd, supply_chain_id ) assert response.status == 500 error_handling.log_db_exception.called def test_get_product_overview_success( monkeypatch, local_product_cd, supply_chain_id, single_product_overview ): """Test persister get_physical_product_monthly_trend success state.""" result_message = MagicMock(fetchone=lambda: single_product_overview) monkeypatch.setattr( Session, 'execute', value=MagicMock(return_value=result_message) ) response = persister.get_product_overview( local_product_cd, supply_chain_id ) expected_table = spo.extract_overview(single_product_overview) assert len(expected_table) == len(response.message) for idx, expected_entry in enumerate(expected_table.values()): for item in expected_entry: assert item in list(response.message.values())[idx] def test_get_product_overview_empty( monkeypatch, local_product_cd, supply_chain_id ): """Test get_product_overview not found.""" result_message = MagicMock(fetchone=lambda: None) monkeypatch.setattr( Session, 'execute', value=MagicMock(return_value=result_message) ) persister_response = persister.get_product_overview( local_product_cd, supply_chain_id ) assert persister_response.status == 200 assert persister_response.message == spo._empty_overview_tables() def test_get_product_overview_failure( mocker, local_product_cd, supply_chain_id ): """Test get_product_overview failure.""" mocker.patch.object(snowflake_conn, 'get_session') mocker.patch.object(Session, 'execute', side_effect=SQLAlchemyError) mocker.patch.object(error_handling, 'log_db_exception') response = persister.get_product_overview( local_product_cd, supply_chain_id ) assert response.status == 500 error_handling.log_db_exception.called def test_account_by_product_success( monkeypatch, local_product_cd, supply_chain_id, account_by_product_response ): """Test persister get_physical_product_monthly_trend success state.""" account_resp = account_by_product_response result_message = MagicMock(fetchone=lambda: account_resp) monkeypatch.setattr( Session, 'execute', value=MagicMock(return_value=result_message) ) response = persister.get_account_by_product( local_product_cd, supply_chain_id ) assert response.message.get('vendor_id') == account_resp.get('vendor_id') assert response.message.get('subaccount_id') == account_resp.get( 'subacct_id' ) assert response.status == 200 def test_get_account_by_product_failure( mocker, local_product_cd, supply_chain_id ): """Test get_account_by_product failure.""" mocker.patch.object(snowflake_conn, 'get_session') mocker.patch.object(Session, 'execute', side_effect=SQLAlchemyError) mocker.patch.object(error_handling, 'log_db_exception') response = persister.get_account_by_product( local_product_cd, supply_chain_id ) assert response.status == 500 error_handling.log_db_exception.called def test_build_vendor_scoped_subaccount_filter_some_vendors_scoped( monkeypatch, ): """Builds scoped and unscoped vendor filters when some vendors match.""" monkeypatch.setattr( persister, '_get_vendor_subaccount_pairs', lambda vendor_ids, supply_chain_id, subacct_ids: [ {'vendor_id': 8869, 'subacct_id': 95969}, ], ) sql, params, _ = persister._build_vendor_scoped_subaccount_filter( [22221, 8869], 738, [95969] ) assert 'unmatched_vendor_ids' in params assert params['unmatched_vendor_ids'] == [22221] assert params['scoped_vendor_id_0'] == 8869 assert params['scoped_subaccount_ids_0'] == [95969] assert 'p.vendor_id IN :unmatched_vendor_ids' in sql assert 'p.vendor_id = :scoped_vendor_id_0' in sql def test_build_vendor_scoped_subaccount_filter_no_matches(monkeypatch): """Returns no extra SQL or params when no vendor/subaccount pairs match.""" monkeypatch.setattr( persister, '_get_vendor_subaccount_pairs', lambda vendor_ids, supply_chain_id, subacct_ids: [], ) sql, params, bindparams = persister._build_vendor_scoped_subaccount_filter( [22221, 8869], 738, [95969] ) assert sql == '' assert params == {} assert bindparams == [] @pytest.mark.parametrize( 'vendor_ids, expected_len, expected_result', [ pytest.param( [22221, 24583], 1, [ { 'local_product_cd': '709422', 'open_order_qt': 500, 'artist_product': 'good riddance/peace in our time', 'vendor_id': 22221, 'subacct_id': 33034, } ], id='Valid vendor_id with open orders', ), pytest.param([999999999], 0, [], id='Unknown vendor_id'), ], ) def test_get_top_open_orders( vendor_ids, expected_len, expected_result, supply_chain_id ): """Test get_top_open_orders against Snowflake test DB.""" result = persister.get_top_open_orders(vendor_ids, supply_chain_id) assert len(result) == expected_len assert result == expected_result def test_get_top_open_orders_failure(mocker, supply_chain_id): """Test get_top_open_orders returns 500 response on DB error.""" mocker.patch.object(snowflake_conn, 'get_session') mocker.patch.object(Session, 'execute', side_effect=SQLAlchemyError) mocker.patch.object(error_handling, 'log_db_exception') result = persister.get_top_open_orders([22221], supply_chain_id) assert result.status == 500 error_handling.log_db_exception.called def test_get_top_open_orders_with_subacct_ids(supply_chain_id): """Test get_top_open_orders filters by subacct_ids when provided.""" result = persister.get_top_open_orders( [22221], supply_chain_id, subacct_ids=[33034] ) assert len(result) == 1 assert result[0]['subacct_id'] == 33034 assert result[0]['vendor_id'] == 22221 def test_get_top_open_orders_with_subacct_ids_no_match(supply_chain_id): """Unknown subacct_ids should leave requested vendors unfiltered.""" unfiltered = persister.get_top_open_orders([22221], supply_chain_id) filtered = persister.get_top_open_orders( [22221], supply_chain_id, subacct_ids=[999999999] ) assert filtered == unfiltered def test_get_top_open_orders_with_subacct_ids_failure(mocker, supply_chain_id): """Test get_top_open_orders with subacct_ids returns 500 on DB error.""" mocker.patch.object(snowflake_conn, 'get_session') mocker.patch.object(Session, 'execute', side_effect=SQLAlchemyError) mocker.patch.object(error_handling, 'log_db_exception') result = persister.get_top_open_orders( [22221], supply_chain_id, subacct_ids=[33034] ) assert result.status == 500 error_handling.log_db_exception.called @pytest.mark.parametrize( 'vendor_ids, expected_len, expected_result', [ pytest.param( [22221, 24583], 1, [ { 'local_product_cd': '709422', 'yesterday_ship_qt': 150, 'artist_product': 'good riddance/peace in our time', 'vendor_id': 22221, 'subacct_id': 33034, } ], id='Valid vendor_id with yesterday shipments', ), pytest.param([999999999], 0, [], id='Unknown vendor_id'), ], ) def test_get_top_yesterday_shipments( vendor_ids, expected_len, expected_result, supply_chain_id ): """Test get_top_yesterday_shipments against Snowflake test DB.""" result = persister.get_top_yesterday_shipments( vendor_ids, supply_chain_id ) assert len(result) == expected_len assert result == expected_result def test_get_top_yesterday_shipments_failure(mocker, supply_chain_id): """Test get_top_yesterday_shipments returns 500 response on DB error.""" mocker.patch.object(snowflake_conn, 'get_session') mocker.patch.object(Session, 'execute', side_effect=SQLAlchemyError) mocker.patch.object(error_handling, 'log_db_exception') result = persister.get_top_yesterday_shipments([22221], supply_chain_id) assert result.status == 500 error_handling.log_db_exception.called def test_get_top_yesterday_shipments_with_subacct_ids(supply_chain_id): """Test get_top_yesterday_shipments filters by subacct_ids.""" result = persister.get_top_yesterday_shipments( [22221], supply_chain_id, subacct_ids=[33034] ) assert len(result) == 1 assert result[0]['subacct_id'] == 33034 assert result[0]['vendor_id'] == 22221 def test_get_top_yesterday_shipments_with_subacct_ids_no_match( supply_chain_id ): """Unknown subacct_ids should leave requested vendors unfiltered.""" unfiltered = persister.get_top_yesterday_shipments( [22221], supply_chain_id ) filtered = persister.get_top_yesterday_shipments( [22221], supply_chain_id, subacct_ids=[999999999] ) assert filtered == unfiltered def test_get_top_yesterday_shipments_with_subacct_ids_failure( mocker, supply_chain_id ): """Test get_top_yesterday_shipments with subacct_ids returns 500.""" mocker.patch.object(snowflake_conn, 'get_session') mocker.patch.object(Session, 'execute', side_effect=SQLAlchemyError) mocker.patch.object(error_handling, 'log_db_exception') result = persister.get_top_yesterday_shipments( [22221], supply_chain_id, subacct_ids=[33034] ) assert result.status == 500 error_handling.log_db_exception.called @pytest.mark.parametrize( 'vendor_ids, expected_len, expected_result', [ pytest.param( [22221, 24583], 1, [ { 'local_product_cd': '709422', 'mtd_ship_qt': 206, 'artist_product': 'good riddance/peace in our time', 'vendor_id': 22221, 'subacct_id': 33034, } ], id='Valid vendor_id with MTD shipments', ), pytest.param([999999999], 0, [], id='Unknown vendor_id'), ], ) def test_get_top_mtd_shipments( vendor_ids, expected_len, expected_result, supply_chain_id ): """Test get_top_mtd_shipments against Snowflake test DB.""" result = persister.get_top_mtd_shipments(vendor_ids, supply_chain_id) assert len(result) == expected_len assert result == expected_result def test_get_top_mtd_shipments_failure(mocker, supply_chain_id): """Test get_top_mtd_shipments returns 500 response on DB error.""" mocker.patch.object(snowflake_conn, 'get_session') mocker.patch.object(Session, 'execute', side_effect=SQLAlchemyError) mocker.patch.object(error_handling, 'log_db_exception') result = persister.get_top_mtd_shipments([22221], supply_chain_id) assert result.status == 500 error_handling.log_db_exception.called def test_get_top_mtd_shipments_with_subacct_ids(supply_chain_id): """Test get_top_mtd_shipments filters by subacct_ids when provided.""" result = persister.get_top_mtd_shipments( [22221], supply_chain_id, subacct_ids=[33034] ) assert len(result) == 1 assert result[0]['subacct_id'] == 33034 assert result[0]['vendor_id'] == 22221 def test_get_top_mtd_shipments_with_subacct_ids_no_match(supply_chain_id): """Unknown subacct_ids should leave requested vendors unfiltered.""" unfiltered = persister.get_top_mtd_shipments([22221], supply_chain_id) filtered = persister.get_top_mtd_shipments( [22221], supply_chain_id, subacct_ids=[999999999] ) assert filtered == unfiltered def test_get_top_mtd_shipments_with_subacct_ids_failure( mocker, supply_chain_id ): """Test get_top_mtd_shipments with subacct_ids returns 500 on DB error.""" mocker.patch.object(snowflake_conn, 'get_session') mocker.patch.object(Session, 'execute', side_effect=SQLAlchemyError) mocker.patch.object(error_handling, 'log_db_exception') result = persister.get_top_mtd_shipments( [22221], supply_chain_id, subacct_ids=[33034] ) assert result.status == 500 error_handling.log_db_exception.called @pytest.mark.parametrize( 'vendor_ids, expected_len, expected_result', [ pytest.param( [22221, 24583], 1, [ { 'label name': 'fat wreck chords', 'sub label': 'fat wreck chords', 'open orders': 500, 'backorders': 0, 'day1s#': 150, 'mtds#': 206, 'cytds#': 50, 'rtds#': 8495, 'day1r#': 0, 'mtdr#': 0, 'cytdr#': -5, 'rtdr#': -1395, 'day1net#': 150, 'mtdnet#': 206, 'cytdnet#': 45, 'rtdnet#': 7100, 'open orders$': 4500.0, 'backorders$': 0.0, 'day1s$': 1200.0, 'mtds$': 1844.82, 'cytds$': 396.45, 'rtds$': 72375.3, 'day1r$': 0.0, 'mtdr$': 0.0, 'cytdr$': -45.0, 'rtdr$': -11635.2, 'day1net$': 1200.0, 'mtdnet$': 1844.82, 'cytdnet$': 351.45, 'rtdnet$': 60740.1, 'vendor_id': 22221, 'subacct_id': 33034, } ], id='Valid vendor_id with label subaccount data', ), pytest.param([999999999], 0, [], id='Unknown vendor_id'), ], ) def test_get_label_subaccount_summary( vendor_ids, expected_len, expected_result, supply_chain_id ): """Test get_label_subaccount_summary against Snowflake test DB.""" result = persister.get_label_subaccount_summary( vendor_ids, supply_chain_id ) assert len(result) == expected_len assert result == expected_result def test_get_label_subaccount_summary_failure(mocker, supply_chain_id): """Test get_label_subaccount_summary returns 500 response on DB error.""" mocker.patch.object(snowflake_conn, 'get_session') mocker.patch.object(Session, 'execute', side_effect=SQLAlchemyError) mocker.patch.object(error_handling, 'log_db_exception') result = persister.get_label_subaccount_summary([22221], supply_chain_id) assert result.status == 500 error_handling.log_db_exception.called def test_get_label_subaccount_summary_with_subacct_ids(supply_chain_id): """Test get_label_subaccount_summary filters by subacct_ids.""" result = persister.get_label_subaccount_summary( [22221], supply_chain_id, subacct_ids=[33034] ) assert len(result) == 1 assert result[0]['subacct_id'] == 33034 assert result[0]['vendor_id'] == 22221 def test_get_label_subaccount_summary_with_subacct_ids_no_match( supply_chain_id ): """Unknown subacct_ids should leave requested vendors unfiltered.""" unfiltered = persister.get_label_subaccount_summary( [22221], supply_chain_id ) filtered = persister.get_label_subaccount_summary( [22221], supply_chain_id, subacct_ids=[999999999] ) assert filtered == unfiltered def test_get_label_subaccount_summary_with_subacct_ids_failure( mocker, supply_chain_id ): """Test get_label_subaccount_summary with subacct_ids returns 500.""" mocker.patch.object(snowflake_conn, 'get_session') mocker.patch.object(Session, 'execute', side_effect=SQLAlchemyError) mocker.patch.object(error_handling, 'log_db_exception') result = persister.get_label_subaccount_summary( [22221], supply_chain_id, subacct_ids=[33034] ) assert result.status == 500 error_handling.log_db_exception.called @pytest.mark.parametrize( 'vendor_ids, expected_len, expected_result', [ pytest.param( [21945], 1, [ { 'label_nm': 'fat wreck chords', 'subacct_nm': 'fat wreck chords', 'artist_nm': 'good riddance', 'product_nm': 'peace in our time', 'product_cd': '709422', 'local_product_cd': '709422', 'upc_cd': '751097094228', 'display_configuration': '1 x cd album', 'status_nm': 'active catalog', 'release_dt': '2015-04-21', 'exclusive_for': '', 'price_cd': '9', 'series_cd': Decimal('13.98'), 'units_per_set': 1, 'genre_nm': 'punk', 'subgenre_nm': 'punk', 'order_qt': 500, 'back_order_qt': 0, 'yest_ship_qt': 150, 'mtd_ship_qt': 206, 'ytd_ship_qt': 50, 'rtd_ship_qt': 8495, 'yest_return_qt': 0, 'mtd_return_qt': 0, 'ytd_return_qt': -5, 'rtd_return_qt': -1395, 'yest_net_qt': 150, 'mtd_net_qt': 206, 'ytd_net_qt': 45, 'rtd_net_qt': 7100, 'order_am': 4500.0, 'back_order_am': 0.0, 'yest_ship_am': 1200.0, 'mtd_ship_am': 1844.82, 'ytd_ship_am': 396.45, 'rtd_ship_am': 72375.3, 'yest_return_am': 0.0, 'mtd_return_am': 0.0, 'ytd_return_am': -45.0, 'rtd_return_am': -11635.2, 'yest_net_am': 1200.0, 'mtd_net_am': 1844.82, 'ytd_net_am': 351.45, 'rtd_net_am': 60740.1, 'vendor_id': 21945, 'subacct_id': 7483, } ], id='Valid vendor_id with product detail data', ), pytest.param([999999999], 0, [], id='Unknown vendor_id'), ], ) def test_get_product_detail( vendor_ids, expected_len, expected_result, supply_chain_id ): """Test get_product_detail against Snowflake test DB.""" result = persister.get_product_detail(vendor_ids, supply_chain_id) assert len(result) == expected_len assert result == expected_result def test_get_product_detail_failure(mocker, supply_chain_id): """Test get_product_detail returns 500 response on DB error.""" mocker.patch.object(snowflake_conn, 'get_session') mocker.patch.object(Session, 'execute', side_effect=SQLAlchemyError) mocker.patch.object(error_handling, 'log_db_exception') result = persister.get_product_detail([21945], supply_chain_id) assert result.status == 500 error_handling.log_db_exception.called def test_get_product_detail_with_subacct_ids(supply_chain_id): """Test get_product_detail filters by subacct_ids when provided.""" result = persister.get_product_detail( [21945], supply_chain_id, subacct_ids=[7483] ) assert len(result) == 1 assert result[0]['subacct_id'] == 7483 assert result[0]['vendor_id'] == 21945 def test_get_product_detail_with_subacct_ids_no_match(supply_chain_id): """Unknown subacct_ids should leave requested vendors unfiltered.""" unfiltered = persister.get_product_detail([21945], supply_chain_id) filtered = persister.get_product_detail( [21945], supply_chain_id, subacct_ids=[999999999] ) assert filtered == unfiltered def test_get_product_detail_with_subacct_ids_failure(mocker, supply_chain_id): """Test get_product_detail with subacct_ids returns 500 on DB error.""" mocker.patch.object(snowflake_conn, 'get_session') mocker.patch.object(Session, 'execute', side_effect=SQLAlchemyError) mocker.patch.object(error_handling, 'log_db_exception') result = persister.get_product_detail( [21945], supply_chain_id, subacct_ids=[7483] ) assert result.status == 500 error_handling.log_db_exception.called