import datetime from unittest import mock import fakeredis import pytest from fansifter_common.utils import timezone import app.handler as handler from app.config import settings from app.enums import CampaignCancelReason, CampaignStatus, MessageChannel from app.handler import ( _dispatch_batch, _plan_batch_sends, _prepare_campaign, _resolve_batch_gate, handle, ) from app.models import AudienceTextFan, BatchRecipient, Campaign, CampaignBatch from app.types import BatchGate, BatchSend, CampaignSendCountry, CampaignSendTimezone from tests.unit.helpers import build_model, create_model class TestPrepareCampaign: @pytest.mark.db def test_skips_when_already_prepared(self) -> None: campaign = create_model(Campaign, prepared_at=timezone.now()) _prepare_campaign(campaign) assert ( CampaignBatch.query.where(CampaignBatch.campaign_id == campaign.id).all() == [] ) @pytest.mark.db def test_does_nothing_when_audience_missing(self) -> None: campaign = create_model(Campaign, audience_id=None, prepared_at=None) _prepare_campaign(campaign) assert campaign.status == CampaignStatus.SCHEDULED assert not campaign.is_prepared @pytest.mark.db def test_refreshes_audience_fans_before_reading_them( self, ows_dmp_client_mock: mock.MagicMock ) -> None: campaign = create_model(Campaign, audience_id="audience-1", prepared_at=None) _prepare_campaign(campaign) ows_dmp_client_mock.upsert_audience_fans.assert_called_once_with( audience_id="audience-1" ) @pytest.mark.db def test_does_nothing_when_audience_refresh_fails( self, ows_dmp_client_mock: mock.MagicMock ) -> None: ows_dmp_client_mock.upsert_audience_fans.side_effect = RuntimeError("boom") campaign = create_model(Campaign, audience_id="audience-1", prepared_at=None) create_model( AudienceTextFan, audience_id="audience-1", fan_country="US", fan_state="NY", channel="SMS", ) _prepare_campaign(campaign) assert campaign.status == CampaignStatus.SCHEDULED assert not campaign.is_prepared assert ( CampaignBatch.query.where(CampaignBatch.campaign_id == campaign.id).all() == [] ) @pytest.mark.db def test_cancels_campaign_when_no_recipients(self) -> None: campaign = create_model(Campaign, audience_id="audience-1", prepared_at=None) _prepare_campaign(campaign) assert campaign.status == CampaignStatus.CANCELLED assert campaign.cancel_reason == CampaignCancelReason.NO_FANS @pytest.mark.db def test_cancels_campaign_when_dmp_reports_no_fans( self, ows_dmp_client_mock: mock.MagicMock ) -> None: ows_dmp_client_mock.upsert_audience_fans.return_value = 0 campaign = create_model(Campaign, audience_id="audience-1", prepared_at=None) # Even if AudienceTextFan rows exist, a 0 fan_count from ows-dmp # should short-circuit before we even query them. create_model( AudienceTextFan, audience_id="audience-1", fan_country="US", fan_state="NY", channel="SMS", ) _prepare_campaign(campaign) assert campaign.status == CampaignStatus.CANCELLED assert campaign.cancel_reason == CampaignCancelReason.NO_FANS assert ( CampaignBatch.query.where(CampaignBatch.campaign_id == campaign.id).all() == [] ) @pytest.mark.db def test_creates_a_batch_per_country_and_state_group(self) -> None: campaign = create_model(Campaign, audience_id="audience-1", prepared_at=None) create_model( AudienceTextFan, audience_id="audience-1", fan_country="US", fan_state="NY", channel="SMS", ) create_model( AudienceTextFan, audience_id="audience-1", fan_country="US", fan_state="NY", channel="SMS", ) create_model( AudienceTextFan, audience_id="audience-1", fan_country="GB", fan_state=None, channel="SMS", ) # different channel -- must be ignored create_model( AudienceTextFan, audience_id="audience-1", fan_country="US", fan_state="NY", channel="WHATSAPP", ) _prepare_campaign(campaign) batches = CampaignBatch.query.where( CampaignBatch.campaign_id == campaign.id ).all() assert { (batch.country_code, batch.state_province, batch.batch_size) for batch in batches } == { ("US", "NY", 2), ("GB", None, 1), } assert campaign.status == CampaignStatus.IN_PROGRESS assert campaign.recipients_count == 3 assert campaign.prepared_at is not None assert len(BatchRecipient.query.all()) == 3 @pytest.mark.db def test_excludes_blocked_channel_country_groups(self) -> None: campaign = create_model( Campaign, audience_id="audience-1", prepared_at=None, channel=MessageChannel.WHATSAPP, ) # WhatsApp is not allowed in the US -- this group must be dropped. create_model( AudienceTextFan, audience_id="audience-1", fan_country="US", fan_state="NY", channel="WHATSAPP", ) create_model( AudienceTextFan, audience_id="audience-1", fan_country="GB", fan_state=None, channel="WHATSAPP", ) _prepare_campaign(campaign) batches = CampaignBatch.query.where( CampaignBatch.campaign_id == campaign.id ).all() assert {(b.country_code, b.state_province) for b in batches} == {("GB", None)} assert campaign.status == CampaignStatus.IN_PROGRESS assert campaign.recipients_count == 1 # No US recipient should ever be materialized into the pipeline. assert {r.fan_country for r in BatchRecipient.query.all()} == {"GB"} @pytest.mark.db def test_cancels_when_all_groups_blocked_for_channel(self) -> None: campaign = create_model( Campaign, audience_id="audience-1", prepared_at=None, channel=MessageChannel.WHATSAPP, ) create_model( AudienceTextFan, audience_id="audience-1", fan_country="US", fan_state="NY", channel="WHATSAPP", ) _prepare_campaign(campaign) assert campaign.status == CampaignStatus.CANCELLED assert campaign.cancel_reason == CampaignCancelReason.CHANNEL_NOT_ALLOWED assert ( CampaignBatch.query.where(CampaignBatch.campaign_id == campaign.id).all() == [] ) class TestPlanBatchSends: @pytest.mark.db def test_returns_empty_when_no_active_batches(self) -> None: assert _plan_batch_sends(dispatch_time=timezone.now()) == [] @pytest.mark.db def test_returns_batch_send_when_country_is_in_safe_hours( self, ows_text_campaigns_client_mock: mock.MagicMock ) -> None: campaign = create_model( Campaign, status=CampaignStatus.IN_PROGRESS, prepared_at=timezone.now(), cancelled_at=None, deleted_at=None, ) batch = create_model( CampaignBatch, campaign_id=campaign.id, country_code="US", state_province="NY", batch_size=100, batch_offset=0, ) ows_text_campaigns_client_mock.get_send_timezones.return_value = [ build_model( CampaignSendCountry, countryCode="US", timezones=[build_model(CampaignSendTimezone, stateProvinces=["NY"])], ) ] batch_sends = _plan_batch_sends(dispatch_time=timezone.now()) assert len(batch_sends) == 1 assert batch_sends[0].batch_id == batch.id assert batch_sends[0].messages_to_send == 100 @pytest.mark.db def test_skips_batch_when_country_not_in_safe_hours( self, ows_text_campaigns_client_mock: mock.MagicMock ) -> None: campaign = create_model( Campaign, status=CampaignStatus.IN_PROGRESS, prepared_at=timezone.now() ) create_model( CampaignBatch, campaign_id=campaign.id, country_code="US", state_province="NY", batch_size=100, ) ows_text_campaigns_client_mock.get_send_timezones.return_value = [ build_model( CampaignSendCountry, countryCode="US", timezones=[ build_model( CampaignSendTimezone, stateProvinces=["NY"], isInSafeHours=False, isAdjusted=False, ) ], ) ] assert _plan_batch_sends(dispatch_time=timezone.now()) == [] @pytest.mark.db def test_fails_closed_when_quiet_hours_lookup_raises( self, ows_text_campaigns_client_mock: mock.MagicMock ) -> None: campaign = create_model( Campaign, status=CampaignStatus.IN_PROGRESS, prepared_at=timezone.now() ) create_model( CampaignBatch, campaign_id=campaign.id, country_code="US", state_province="NY", batch_size=100, ) ows_text_campaigns_client_mock.get_send_timezones.side_effect = Exception( "boom" ) assert _plan_batch_sends(dispatch_time=timezone.now()) == [] @pytest.mark.db def test_caps_messages_to_send_by_configured_max( self, ows_text_campaigns_client_mock: mock.MagicMock, monkeypatch: pytest.MonkeyPatch, ) -> None: monkeypatch.setattr(settings, "max_recipients_per_batch_send", 10) campaign = create_model( Campaign, status=CampaignStatus.IN_PROGRESS, prepared_at=timezone.now() ) create_model( CampaignBatch, campaign_id=campaign.id, country_code="US", state_province="NY", batch_size=100, batch_offset=0, ) ows_text_campaigns_client_mock.get_send_timezones.return_value = [ build_model( CampaignSendCountry, countryCode="US", timezones=[build_model(CampaignSendTimezone, stateProvinces=["NY"])], ) ] batch_sends = _plan_batch_sends(dispatch_time=timezone.now()) assert batch_sends[0].messages_to_send == 10 @pytest.mark.db def test_skips_batch_with_no_remaining_messages( self, ows_text_campaigns_client_mock: mock.MagicMock ) -> None: campaign = create_model( Campaign, status=CampaignStatus.IN_PROGRESS, prepared_at=timezone.now() ) create_model( CampaignBatch, campaign_id=campaign.id, country_code="US", state_province="NY", batch_size=10, batch_offset=10, ) ows_text_campaigns_client_mock.get_send_timezones.return_value = [ build_model( CampaignSendCountry, countryCode="US", timezones=[build_model(CampaignSendTimezone, stateProvinces=["NY"])], ) ] assert _plan_batch_sends(dispatch_time=timezone.now()) == [] @pytest.mark.db def test_one_campaigns_lookup_failure_does_not_block_others( self, ows_text_campaigns_client_mock: mock.MagicMock ) -> None: broken_campaign = create_model( Campaign, status=CampaignStatus.IN_PROGRESS, prepared_at=timezone.now() ) broken_batch = create_model( CampaignBatch, campaign_id=broken_campaign.id, country_code="US", state_province="NY", batch_size=10, ) healthy_campaign = create_model( Campaign, status=CampaignStatus.IN_PROGRESS, prepared_at=timezone.now() ) healthy_batch = create_model( CampaignBatch, campaign_id=healthy_campaign.id, country_code="US", state_province="NY", batch_size=10, ) def _get_send_timezones_side_effect( *, campaign_id: str, **_kwargs: object ) -> list[CampaignSendCountry]: if campaign_id == broken_campaign.id: raise RuntimeError("boom") return [ build_model( CampaignSendCountry, countryCode="US", timezones=[ build_model(CampaignSendTimezone, stateProvinces=["NY"]) ], ) ] ows_text_campaigns_client_mock.get_send_timezones.side_effect = ( _get_send_timezones_side_effect ) batch_sends = _plan_batch_sends(dispatch_time=timezone.now()) assert {batch_send.batch_id for batch_send in batch_sends} == {healthy_batch.id} assert broken_batch.id not in { batch_send.batch_id for batch_send in batch_sends } class TestResolveBatchGate: def test_not_sendable_when_no_timezones_for_country(self) -> None: batch = build_model(CampaignBatch, country_code="US", state_province="NY") gate = _resolve_batch_gate(batch, []) assert gate == BatchGate(is_sendable=False, safe_until=None) def test_sendable_when_matching_state_in_safe_hours(self) -> None: batch = build_model(CampaignBatch, country_code="US", state_province="NY") matching_timezone = build_model( CampaignSendTimezone, stateProvinces=["NY"], isInSafeHours=True ) gate = _resolve_batch_gate(batch, [matching_timezone]) assert gate.is_sendable assert gate.safe_until is not None def test_not_sendable_when_matching_state_outside_safe_hours(self) -> None: batch = build_model(CampaignBatch, country_code="US", state_province="NY") matching_timezone = build_model( CampaignSendTimezone, stateProvinces=["NY"], isInSafeHours=False, isAdjusted=False, ) gate = _resolve_batch_gate(batch, [matching_timezone]) assert not gate.is_sendable def test_sendable_with_no_boundary_when_adjusted_outside_safe_hours(self) -> None: batch = build_model(CampaignBatch, country_code="US", state_province="NY") matching_timezone = build_model( CampaignSendTimezone, stateProvinces=["NY"], isInSafeHours=False, isAdjusted=True, ) gate = _resolve_batch_gate(batch, [matching_timezone]) assert gate.is_sendable assert gate.safe_until is None def test_falls_back_to_requiring_all_country_timezones_safe(self) -> None: batch = build_model(CampaignBatch, country_code="US", state_province="TX") safe_timezone = build_model( CampaignSendTimezone, stateProvinces=["NY"], isInSafeHours=True ) unsafe_timezone = build_model( CampaignSendTimezone, stateProvinces=["CA"], isInSafeHours=False, isAdjusted=False, ) gate = _resolve_batch_gate(batch, [safe_timezone, unsafe_timezone]) assert not gate.is_sendable def test_takes_earliest_close_across_matching_timezones(self) -> None: batch = build_model(CampaignBatch, country_code="US", state_province=None) early_close = build_model( CampaignSendTimezone, ianaTzId="America/New_York", stateProvinces=None, quietHoursMax="18:00:00", ) late_close = build_model( CampaignSendTimezone, ianaTzId="America/Los_Angeles", stateProvinces=None, quietHoursMax="21:00:00", ) gate = _resolve_batch_gate(batch, [early_close, late_close]) assert gate.is_sendable assert gate.safe_until == early_close.safe_until class TestDispatchBatch: def test_skips_when_locked( self, lambda_client_mock: mock.MagicMock, redis_client: fakeredis.FakeRedis, ) -> None: batch_send = BatchSend( batch_id=1, campaign_id="campaign-1", country_code="US", state_province="NY", messages_to_send=10, safe_until=None, ) redis_client.set(settings.redis_sender_lock_key.format(batch_id=1), "1") _dispatch_batch(batch_send) lambda_client_mock.invoke.assert_not_called() def test_invokes_sender_with_expected_payload( self, lambda_client_mock: mock.MagicMock ) -> None: safe_until = datetime.datetime(2026, 7, 8, 20, 0, tzinfo=datetime.UTC) batch_send = BatchSend( batch_id=1, campaign_id="campaign-1", country_code="US", state_province="NY", messages_to_send=10, safe_until=safe_until, ) _dispatch_batch(batch_send) lambda_client_mock.invoke.assert_called_once_with( function_name=settings.sender_lambda_function_name, invocation_type="Event", data={ "batch_id": 1, "campaign_id": "campaign-1", "messages_to_send": 10, "safe_until": safe_until.isoformat(), }, ) def test_invokes_sender_with_none_safe_until( self, lambda_client_mock: mock.MagicMock ) -> None: batch_send = BatchSend( batch_id=2, campaign_id="campaign-1", country_code="US", state_province=None, messages_to_send=5, safe_until=None, ) _dispatch_batch(batch_send) assert lambda_client_mock.invoke.call_args.kwargs["data"]["safe_until"] is None def test_successful_dispatch_acquires_the_lock( self, redis_client: fakeredis.FakeRedis ) -> None: batch_send = BatchSend( batch_id=3, campaign_id="campaign-1", country_code="US", state_province="NY", messages_to_send=10, safe_until=None, ) _dispatch_batch(batch_send) assert redis_client.exists(settings.redis_sender_lock_key.format(batch_id=3)) def test_second_dispatch_for_the_same_batch_is_skipped( self, lambda_client_mock: mock.MagicMock ) -> None: batch_send = BatchSend( batch_id=4, campaign_id="campaign-1", country_code="US", state_province="NY", messages_to_send=10, safe_until=None, ) _dispatch_batch(batch_send) _dispatch_batch(batch_send) lambda_client_mock.invoke.assert_called_once() def test_failed_invoke_releases_the_lock( self, lambda_client_mock: mock.MagicMock, redis_client: fakeredis.FakeRedis, ) -> None: lambda_client_mock.invoke.return_value = False batch_send = BatchSend( batch_id=5, campaign_id="campaign-1", country_code="US", state_province="NY", messages_to_send=10, safe_until=None, ) _dispatch_batch(batch_send) assert not redis_client.exists( settings.redis_sender_lock_key.format(batch_id=5) ) class TestHandle: @pytest.mark.db def test_prepares_and_dispatches_a_ready_campaign( self, lambda_client_mock: mock.MagicMock, ows_text_campaigns_client_mock: mock.MagicMock, ) -> None: campaign = create_model( Campaign, status=CampaignStatus.SCHEDULED, send_at=timezone.now() - datetime.timedelta(minutes=5), audience_id="audience-1", prepared_at=None, ) create_model( AudienceTextFan, audience_id="audience-1", fan_country="US", fan_state="NY", channel="SMS", ) ows_text_campaigns_client_mock.get_send_timezones.return_value = [ build_model( CampaignSendCountry, countryCode="US", timezones=[build_model(CampaignSendTimezone, stateProvinces=["NY"])], ) ] batch_sends = handle() campaign.refresh() assert campaign.status == CampaignStatus.IN_PROGRESS assert len(batch_sends) == 1 assert batch_sends[0].messages_to_send == 1 lambda_client_mock.invoke.assert_called_once() @pytest.mark.db def test_returns_empty_when_nothing_ready(self) -> None: assert handle() == [] @pytest.mark.db def test_one_campaign_prepare_failure_does_not_block_others( self, monkeypatch: pytest.MonkeyPatch, lambda_client_mock: mock.MagicMock, ows_text_campaigns_client_mock: mock.MagicMock, ) -> None: broken_campaign = create_model( Campaign, status=CampaignStatus.SCHEDULED, send_at=timezone.now() - datetime.timedelta(minutes=5), audience_id="audience-broken", prepared_at=None, ) healthy_campaign = create_model( Campaign, status=CampaignStatus.SCHEDULED, send_at=timezone.now() - datetime.timedelta(minutes=5), audience_id="audience-healthy", prepared_at=None, ) create_model( AudienceTextFan, audience_id="audience-healthy", fan_country="US", fan_state="NY", channel="SMS", ) ows_text_campaigns_client_mock.get_send_timezones.return_value = [ build_model( CampaignSendCountry, countryCode="US", timezones=[build_model(CampaignSendTimezone, stateProvinces=["NY"])], ) ] original_prepare_campaign = handler._prepare_campaign def _prepare_campaign_side_effect(campaign: Campaign) -> None: if campaign.id == broken_campaign.id: raise RuntimeError("boom") original_prepare_campaign(campaign) monkeypatch.setattr(handler, "_prepare_campaign", _prepare_campaign_side_effect) batch_sends = handle() broken_campaign.refresh() healthy_campaign.refresh() assert broken_campaign.status == CampaignStatus.SCHEDULED assert healthy_campaign.status == CampaignStatus.IN_PROGRESS assert len(batch_sends) == 1 lambda_client_mock.invoke.assert_called_once() @pytest.mark.db def test_one_dispatch_failure_does_not_block_others( self, monkeypatch: pytest.MonkeyPatch, lambda_client_mock: mock.MagicMock, ows_text_campaigns_client_mock: mock.MagicMock, ) -> None: campaign_1 = create_model( Campaign, status=CampaignStatus.IN_PROGRESS, send_at=timezone.now() - datetime.timedelta(minutes=5), prepared_at=timezone.now(), ) broken_batch = create_model( CampaignBatch, campaign_id=campaign_1.id, country_code="US", state_province="NY", batch_size=10, ) campaign_2 = create_model( Campaign, status=CampaignStatus.IN_PROGRESS, send_at=timezone.now() - datetime.timedelta(minutes=5), prepared_at=timezone.now(), ) healthy_batch = create_model( CampaignBatch, campaign_id=campaign_2.id, country_code="US", state_province="NY", batch_size=10, ) ows_text_campaigns_client_mock.get_send_timezones.return_value = [ build_model( CampaignSendCountry, countryCode="US", timezones=[build_model(CampaignSendTimezone, stateProvinces=["NY"])], ) ] original_dispatch_batch = handler._dispatch_batch def _dispatch_batch_side_effect(batch_send: BatchSend) -> None: if batch_send.batch_id == broken_batch.id: raise RuntimeError("boom") original_dispatch_batch(batch_send) monkeypatch.setattr(handler, "_dispatch_batch", _dispatch_batch_side_effect) batch_sends = handle() assert {batch_send.batch_id for batch_send in batch_sends} == { broken_batch.id, healthy_batch.id, } lambda_client_mock.invoke.assert_called_once() assert ( lambda_client_mock.invoke.call_args.kwargs["data"]["batch_id"] == healthy_batch.id )