import json import uuid from unittest import mock import pytest from fansifter_common.utils import timezone from pytest_mock import MockerFixture from app.core.encrypter import FernetEncrypter from app.dsp.models import DSPClient from app.fandata.enums import FanCredentialsStatus from app.fandata.models import FanCredentials, FanCredentialsFilter from app.pipeline.enums import RunStatus from app.pipeline.models import PipelineRun from app.runtime.handlers import ( fan_collect_handler, fan_fanout_handler, prepare_rotation_handler, rotate_keys_handler, ) from tests.unit.helpers import create_model, override_settings _FAN_RECORD = {"dsp_user_id": "u1", "refresh_token_encrypted": "tok"} class TestFanFanoutHandler: @pytest.fixture(autouse=True) def mock_fanout(self, mocker: MockerFixture) -> mock.MagicMock: return mocker.patch("app.runtime.handlers.fanout") def test_raises_on_invalid_payload(self) -> None: with pytest.raises(ValueError, match="Invalid event payload"): fan_fanout_handler({}, None) @pytest.mark.db def test_raises_if_client_not_found(self) -> None: with pytest.raises(ValueError, match="not found"): fan_fanout_handler({"dsp_client": "spotify_smf"}, None) @pytest.mark.db def test_skips_if_pipeline_already_running(self) -> None: client = create_model(DSPClient) create_model( PipelineRun, dsp_client_id=client.id, status=RunStatus.running, started_at=timezone.now(), ) result = fan_fanout_handler({"dsp_client": client.name.value}, None) assert result == {"run_id": None, "skipped": True, "reason": mock.ANY} @pytest.mark.db def test_creates_run_and_calls_fanout(self, mock_fanout: mock.MagicMock) -> None: client = create_model(DSPClient) result = fan_fanout_handler({"dsp_client": client.name.value}, None) assert result == {"run_id": mock.ANY, "skipped": False} mock_fanout.assert_called_once_with(uuid.UUID(result["run_id"])) @pytest.mark.db def test_passes_optional_params_to_run(self) -> None: client = create_model(DSPClient) result = fan_fanout_handler( { "dsp_client": client.name.value, "token_status": FanCredentialsStatus.revoked.value, "limit": 50, "max_consecutive_failures": 3, }, None, ) run = PipelineRun.query.where( PipelineRun.id == uuid.UUID(result["run_id"]) ).one() assert run.filters == FanCredentialsFilter( dsp_client_id=client.id, token_status=FanCredentialsStatus.revoked, not_collected_since=None, max_consecutive_failures=3, limit=50, ) class TestFanCollectHandler: @pytest.fixture(autouse=True) def mock_collect(self, mocker: MockerFixture) -> mock.MagicMock: return mocker.patch("app.runtime.handlers.collect") def test_empty_records_returns_no_failures(self) -> None: result = fan_collect_handler({"Records": []}, None) assert result == {"batchItemFailures": []} def test_missing_records_key_returns_no_failures(self) -> None: result = fan_collect_handler({}, None) assert result == {"batchItemFailures": []} def test_invalid_json_body_marks_record_failed(self) -> None: result = fan_collect_handler( {"Records": [{"messageId": "bad-1", "body": "not-json"}]}, None ) assert result == {"batchItemFailures": [{"itemIdentifier": "bad-1"}]} def test_missing_body_key_marks_record_failed(self) -> None: result = fan_collect_handler({"Records": [{"messageId": "bad-2"}]}, None) assert result == {"batchItemFailures": [{"itemIdentifier": "bad-2"}]} def test_collect_exception_marks_record_failed( self, mock_collect: mock.MagicMock ) -> None: mock_collect.side_effect = RuntimeError("boom") body = json.dumps({"run_id": str(uuid.uuid4()), "fans": [_FAN_RECORD]}) result = fan_collect_handler( {"Records": [{"messageId": "msg-3", "body": body}]}, None ) assert result == {"batchItemFailures": [{"itemIdentifier": "msg-3"}]} def test_successful_records_return_no_failures( self, mock_collect: mock.MagicMock ) -> None: records = [ { "messageId": "msg-1", "body": json.dumps( {"run_id": str(uuid.uuid4()), "fans": [_FAN_RECORD]} ), }, { "messageId": "msg-2", "body": json.dumps( {"run_id": str(uuid.uuid4()), "fans": [_FAN_RECORD]} ), }, ] result = fan_collect_handler({"Records": records}, None) assert result == {"batchItemFailures": []} assert mock_collect.call_count == 2 def test_partial_failures_returns_only_failed( self, mock_collect: mock.MagicMock ) -> None: mock_collect.side_effect = [None, RuntimeError("boom"), None] records = [ { "messageId": "ok-1", "body": json.dumps( {"run_id": str(uuid.uuid4()), "fans": [_FAN_RECORD]} ), }, { "messageId": "bad-2", "body": json.dumps( {"run_id": str(uuid.uuid4()), "fans": [_FAN_RECORD]} ), }, { "messageId": "ok-3", "body": json.dumps( {"run_id": str(uuid.uuid4()), "fans": [_FAN_RECORD]} ), }, ] result = fan_collect_handler({"Records": records}, None) assert result == {"batchItemFailures": [{"itemIdentifier": "bad-2"}]} class TestPrepareRotationHandler: @pytest.fixture def secrets_manager_mock(self, mocker: MockerFixture) -> mock.MagicMock: return mocker.patch("app.runtime.handlers.aws_secrets_manager") def test_raises_if_not_fernet(self) -> None: with override_settings(encrypter_backend="kms"): with pytest.raises(ValueError, match="fernet backend"): prepare_rotation_handler({}, None) def test_prepends_new_key_and_updates_sm( self, secrets_manager_mock: mock.MagicMock ) -> None: old_key = FernetEncrypter.generate_key() secrets_manager_mock.get_fernet_keys.return_value = [old_key] with override_settings(encrypter_backend="fernet", fernet_keys=[old_key]): result = prepare_rotation_handler({}, None) assert result["total_keys"] == 2 assert len(result["new_key_id"]) == 8 keys_written: list[str] = secrets_manager_mock.put_fernet_keys.call_args[0][0] assert keys_written[1] == old_key assert keys_written[0].partition(":")[0] == result["new_key_id"] class TestRotateKeysHandler: @pytest.fixture def secrets_manager_mock(self, mocker: MockerFixture) -> mock.MagicMock: return mocker.patch("app.runtime.handlers.aws_secrets_manager") def test_skips_if_not_fernet(self) -> None: with override_settings(encrypter_backend="kms"): result = rotate_keys_handler({}, None) assert result == { "rotated": 0, "errors": 0, "next_after_id": None, "done": True, "finalized": False, } def test_raises_if_sm_has_single_key( self, secrets_manager_mock: mock.MagicMock ) -> None: new_key = FernetEncrypter.generate_key() old_key = FernetEncrypter.generate_key() secrets_manager_mock.get_fernet_keys.return_value = [new_key] with ( override_settings( encrypter_backend="fernet", fernet_keys=[new_key, old_key] ), pytest.raises(ValueError, match="prepare_rotation first"), ): rotate_keys_handler({}, None) def test_raises_if_env_out_of_sync( self, secrets_manager_mock: mock.MagicMock ) -> None: new_key = FernetEncrypter.generate_key() old_key = FernetEncrypter.generate_key() secrets_manager_mock.get_fernet_keys.return_value = [new_key, old_key] with override_settings( encrypter_backend="fernet", fernet_keys=[old_key, new_key] ): with pytest.raises(ValueError, match="out of sync"): rotate_keys_handler({}, None) @pytest.mark.db def test_finalizes_when_no_rows_to_rotate( self, secrets_manager_mock: mock.MagicMock ) -> None: new_key = FernetEncrypter.generate_key() old_key = FernetEncrypter.generate_key() secrets_manager_mock.get_fernet_keys.return_value = [new_key, old_key] with override_settings( encrypter_backend="fernet", fernet_keys=[new_key, old_key] ): result = rotate_keys_handler({}, None) assert result == { "rotated": 0, "errors": 0, "next_after_id": None, "done": True, "finalized": True, } secrets_manager_mock.put_fernet_keys.assert_called_once_with([new_key]) @pytest.mark.db def test_rotates_rows_with_old_key( self, secrets_manager_mock: mock.MagicMock ) -> None: new_key = FernetEncrypter.generate_key() old_key = FernetEncrypter.generate_key() secrets_manager_mock.get_fernet_keys.return_value = [new_key, old_key] old_encrypter = FernetEncrypter([old_key]) create_model( FanCredentials, refresh_token_encrypted=old_encrypter.encrypt("token-a"), ) create_model( FanCredentials, refresh_token_encrypted=old_encrypter.encrypt("token-b"), ) with ( override_settings( encrypter_backend="fernet", fernet_keys=[new_key, old_key] ), mock.patch( "app.core.encrypter.encrypter.__wrapped__", FernetEncrypter([new_key, old_key]), ), ): result = rotate_keys_handler({}, None) assert result == { "rotated": 2, "errors": 0, "next_after_id": None, "done": True, "finalized": True, } secrets_manager_mock.put_fernet_keys.assert_called_once_with([new_key]) new_prefix = new_key.partition(":")[0] + ":" assert all( c.refresh_token_encrypted.startswith(new_prefix) for c in FanCredentials.query.all() ) @pytest.mark.db def test_skips_rows_already_on_new_key( self, secrets_manager_mock: mock.MagicMock ) -> None: new_key = FernetEncrypter.generate_key() old_key = FernetEncrypter.generate_key() secrets_manager_mock.get_fernet_keys.return_value = [new_key, old_key] create_model( FanCredentials, refresh_token_encrypted=FernetEncrypter([new_key]).encrypt("token-a"), ) with override_settings( encrypter_backend="fernet", fernet_keys=[new_key, old_key] ): result = rotate_keys_handler({}, None) assert result == { "rotated": 0, "errors": 0, "next_after_id": None, "done": True, "finalized": True, } @pytest.mark.db def test_counts_errors_and_does_not_finalize( self, secrets_manager_mock: mock.MagicMock, mocker: MockerFixture, ) -> None: new_key = FernetEncrypter.generate_key() old_key = FernetEncrypter.generate_key() secrets_manager_mock.get_fernet_keys.return_value = [new_key, old_key] old_encrypter = FernetEncrypter([old_key]) new_encrypter = FernetEncrypter([new_key]) create_model( FanCredentials, refresh_token_encrypted=old_encrypter.encrypt("token-a"), ) create_model( FanCredentials, refresh_token_encrypted=old_encrypter.encrypt("token-b"), ) mocker.patch( "app.runtime.handlers.decrypt", side_effect=[RuntimeError("bad"), "plaintext"], ) mocker.patch( "app.runtime.handlers.encrypt", return_value=new_encrypter.encrypt("rotated"), ) with override_settings( encrypter_backend="fernet", fernet_keys=[new_key, old_key] ): result = rotate_keys_handler({}, None) assert result == { "rotated": 1, "errors": 1, "next_after_id": None, "done": True, "finalized": False, } secrets_manager_mock.put_fernet_keys.assert_not_called() @pytest.mark.db def test_pagination_returns_next_after_id( self, secrets_manager_mock: mock.MagicMock ) -> None: new_key = FernetEncrypter.generate_key() old_key = FernetEncrypter.generate_key() secrets_manager_mock.get_fernet_keys.return_value = [new_key, old_key] old_encrypter = FernetEncrypter([old_key]) create_model( FanCredentials, refresh_token_encrypted=old_encrypter.encrypt("t1"), ) create_model( FanCredentials, refresh_token_encrypted=old_encrypter.encrypt("t2"), ) create_model( FanCredentials, refresh_token_encrypted=old_encrypter.encrypt("t3"), ) with ( override_settings( encrypter_backend="fernet", fernet_keys=[new_key, old_key] ), mock.patch( "app.core.encrypter.encrypter.__wrapped__", FernetEncrypter([new_key, old_key]), ), ): result = rotate_keys_handler({"batch_size": 2}, None) assert result == { "rotated": 2, "errors": 0, "next_after_id": mock.ANY, "done": False, "finalized": False, }