"""Unit tests for historical_contract_advance logic.""" from unittest.mock import patch from abacus_contract.logic import historical_contract_advance as logic from abacus_contract.tests.utils.factories import ( AccountContractFactory, ContractFactory, HistoricalContractAdvanceFactory, ) @patch('abacus_contract.logic.historical_contract_advance.HistoricalContractAdvance') @patch('abacus_contract.logic.historical_contract_advance.Contract') def test_get_historical_contract_advances_empty_list(mock_contract_model, mock_model): """Test getting historical contract advances with empty contract_ids list.""" res = logic.get_historical_contract_advances_by_contract_ids([]) assert res.status == 200 assert res.message == [] mock_contract_model.get_by_ids.assert_not_called() mock_model.get_by_account_ids.assert_not_called() @patch('abacus_contract.logic.historical_contract_advance.HistoricalContractAdvance') @patch('abacus_contract.logic.historical_contract_advance.Contract') def test_get_historical_contract_advances_by_single_contract_id( mock_contract_model, mock_model ): """Test getting historical contract advances for a single contract.""" contract = ContractFactory.build(contract_id=1) account_contract = AccountContractFactory.build(contract=contract, account_id=100) contract.account_contract = account_contract advances = HistoricalContractAdvanceFactory.build_batch(2, account_id=100) mock_contract_model.get_by_ids.return_value = [contract] mock_model.get_by_account_ids.return_value = advances res = logic.get_historical_contract_advances_by_contract_ids([contract.contract_id]) assert res.status == 200 assert len(res.message) == 1 assert res.message[0]['data'] is not None assert len(res.message[0]['data']) == 2 mock_contract_model.get_by_ids.assert_called_once_with([contract.contract_id]) mock_model.get_by_account_ids.assert_called_once_with([100]) @patch('abacus_contract.logic.historical_contract_advance.HistoricalContractAdvance') @patch('abacus_contract.logic.historical_contract_advance.Contract') def test_get_historical_contract_advances_by_multiple_contract_ids( mock_contract_model, mock_model ): """Test getting historical contract advances for multiple contracts.""" contract1 = ContractFactory.build(contract_id=1) account_contract1 = AccountContractFactory.build(contract=contract1, account_id=100) contract1.account_contract = account_contract1 contract2 = ContractFactory.build(contract_id=2) account_contract2 = AccountContractFactory.build(contract=contract2, account_id=200) contract2.account_contract = account_contract2 advances_account_100 = HistoricalContractAdvanceFactory.build_batch( 2, account_id=100 ) advances_account_200 = HistoricalContractAdvanceFactory.build_batch( 1, account_id=200 ) all_advances = advances_account_100 + advances_account_200 mock_contract_model.get_by_ids.return_value = [contract1, contract2] mock_model.get_by_account_ids.return_value = all_advances contract_ids = [contract1.contract_id, contract2.contract_id] res = logic.get_historical_contract_advances_by_contract_ids(contract_ids) assert res.status == 200 assert len(res.message) == 2 assert res.message[0]['data'] is not None assert len(res.message[0]['data']) == 2 assert res.message[1]['data'] is not None assert len(res.message[1]['data']) == 1 mock_contract_model.get_by_ids.assert_called_once_with(contract_ids) mock_model.get_by_account_ids.assert_called_once() @patch('abacus_contract.logic.historical_contract_advance.HistoricalContractAdvance') @patch('abacus_contract.logic.historical_contract_advance.Contract') def test_get_historical_contract_advances_no_advances_found( mock_contract_model, mock_model ): """Test getting historical contract advances when no advances exist.""" contract = ContractFactory.build(contract_id=1) account_contract = AccountContractFactory.build(contract=contract, account_id=100) contract.account_contract = account_contract mock_contract_model.get_by_ids.return_value = [contract] mock_model.get_by_account_ids.return_value = [] res = logic.get_historical_contract_advances_by_contract_ids([contract.contract_id]) assert res.status == 200 assert len(res.message) == 1 assert res.message[0]['data'] is None mock_contract_model.get_by_ids.assert_called_once_with([contract.contract_id]) mock_model.get_by_account_ids.assert_called_once_with([100]) @patch('abacus_contract.logic.historical_contract_advance.HistoricalContractAdvance') @patch('abacus_contract.logic.historical_contract_advance.Contract') def test_get_historical_contract_advances_mixed_results( mock_contract_model, mock_model ): """Test getting historical contract advances with mixed results.""" contract1 = ContractFactory.build(contract_id=1) account_contract1 = AccountContractFactory.build(contract=contract1, account_id=100) contract1.account_contract = account_contract1 contract2 = ContractFactory.build(contract_id=2) account_contract2 = AccountContractFactory.build(contract=contract2, account_id=200) contract2.account_contract = account_contract2 contract3 = ContractFactory.build(contract_id=3) account_contract3 = AccountContractFactory.build(contract=contract3, account_id=300) contract3.account_contract = account_contract3 advances = HistoricalContractAdvanceFactory.build_batch( 2, account_id=100 ) + HistoricalContractAdvanceFactory.build_batch(1, account_id=300) mock_contract_model.get_by_ids.return_value = [contract1, contract2, contract3] mock_model.get_by_account_ids.return_value = advances contract_ids = [contract1.contract_id, contract2.contract_id, contract3.contract_id] res = logic.get_historical_contract_advances_by_contract_ids(contract_ids) assert res.status == 200 assert len(res.message) == 3 assert res.message[0]['data'] is not None assert len(res.message[0]['data']) == 2 assert res.message[1]['data'] is None assert res.message[2]['data'] is not None assert len(res.message[2]['data']) == 1 mock_contract_model.get_by_ids.assert_called_once_with(contract_ids) mock_model.get_by_account_ids.assert_called_once()