import uuid from unittest import mock import pytest from pytest_mock import MockerFixture from app.dsp.models import DSPClient from app.fandata.enums import FanCredentialsStatus from app.fandata.models import FanCredentials from app.fandata.types import FanBatch from app.pipeline.enums import RunFinishedReason, RunStatus from app.pipeline.fan_fanout import fanout from app.pipeline.models import PipelineRun from tests.unit.helpers import create_model, override_settings class TestFanFanout: @pytest.fixture(autouse=True) def dispatch_fan_collect_mock(self, mocker: MockerFixture) -> mock.MagicMock: return mocker.patch( "app.pipeline.fan_fanout.dispatch_fan_collect", new_callable=mock.MagicMock ) @pytest.mark.db def test_cancels_run_if_no_credentials( self, dispatch_fan_collect_mock: mock.MagicMock ) -> None: client = create_model(DSPClient) run = create_model( PipelineRun, dsp_client_id=client.id, status=RunStatus.queued ) fanout(run.id) assert run.status == RunStatus.cancelled assert run.finished_reason == RunFinishedReason.no_credentials dispatch_fan_collect_mock.assert_not_called() @pytest.mark.db def test_finalizes_run_if_client_not_found( self, dispatch_fan_collect_mock: mock.MagicMock ) -> None: run = create_model(PipelineRun, dsp_client_id=999999, status=RunStatus.queued) fanout(run.id) assert run.status == RunStatus.done assert run.finished_reason == RunFinishedReason.done dispatch_fan_collect_mock.assert_not_called() @pytest.mark.db def test_noop_if_run_not_found( self, dispatch_fan_collect_mock: mock.MagicMock ) -> None: fanout(uuid.uuid4()) dispatch_fan_collect_mock.assert_not_called() @pytest.mark.db def test_activates_run_and_sets_fan_counts(self) -> None: client = create_model(DSPClient) run = create_model( PipelineRun, dsp_client_id=client.id, status=RunStatus.queued ) create_model( FanCredentials, dsp_id=client.dsp_id, client_id=client.client_id, dsp_user_id="user_1", token_status=FanCredentialsStatus.active, ) create_model( FanCredentials, dsp_id=client.dsp_id, client_id=client.client_id, dsp_user_id="user_2", token_status=FanCredentialsStatus.active, ) with override_settings(fan_fanout_batch_size=10): fanout(run.id) assert run.status == RunStatus.running assert run.total_fans == 2 assert run.total_batches == 1 @pytest.mark.db def test_dispatches_one_batch_per_page( self, dispatch_fan_collect_mock: mock.MagicMock ) -> None: client = create_model(DSPClient) run = create_model( PipelineRun, dsp_client_id=client.id, status=RunStatus.queued ) create_model( FanCredentials, dsp_id=client.dsp_id, client_id=client.client_id, dsp_user_id="user_1", token_status=FanCredentialsStatus.active, ) create_model( FanCredentials, dsp_id=client.dsp_id, client_id=client.client_id, dsp_user_id="user_2", token_status=FanCredentialsStatus.active, ) dispatched: list[FanBatch] = [] dispatch_fan_collect_mock.side_effect = lambda b, _: dispatched.append(b) with override_settings(fan_fanout_batch_size=10): fanout(run.id) assert len(dispatched) == 1 assert len(dispatched[0].fans) == 2 assert dispatched[0].run_id == run.id @pytest.mark.db def test_dispatches_multiple_batches( self, dispatch_fan_collect_mock: mock.MagicMock ) -> None: client = create_model(DSPClient) run = create_model( PipelineRun, dsp_client_id=client.id, status=RunStatus.queued ) for i in range(5): create_model( FanCredentials, dsp_id=client.dsp_id, client_id=client.client_id, dsp_user_id=f"user_{i}", token_status=FanCredentialsStatus.active, ) dispatched: list[FanBatch] = [] dispatch_fan_collect_mock.side_effect = lambda b, _: dispatched.append(b) with override_settings(fan_fanout_batch_size=2): fanout(run.id) assert len(dispatched) == 3 assert sum(len(b.fans) for b in dispatched) == 5 assert all(b.run_id == run.id for b in dispatched)