# pylint: disable=unused-argument,redefined-outer-name import enum import os from collections import defaultdict from dataclasses import dataclass from datetime import datetime, timedelta, timezone from typing import Any, Dict, List, Tuple, Union import boto3 import pytest from _pytest.fixtures import SubRequest from db_schema.postgres.connection import Connection from db_schema.schemas import slz from freezegun import freeze_time from freezegun.api import FakeDatetime from moto import mock_s3 from mypy_boto3_s3 import S3Client @dataclass class CommonValues: """Stores values used by tests if value is referenced more than once.""" # pylint: disable=too-many-instance-attributes file_prefix: str = 'spotify/charts_daily_regional/v1/report_date=2021-12-10/report_licensor=sme' corrupted_bucket: str = 'corrupted_test_bucket' decompressed_bucket: str = 'decompressed_test_bucket' control_sample_bucket: str = 'control_bucket' current_execution_arn: str = 'current_stateMachine_arn' step_function_arn: str = 'test_step_function_arn' job_1: str = 'test_job_id_1' job_2: str = 'test_job_id_2' job_3: str = 'test_job_id_3' def get_buckets_list(self) -> List[str]: return [self.corrupted_bucket, self.decompressed_bucket, self.control_sample_bucket] def get_job_ids_list(self) -> List[str]: return [self.job_1, self.job_2, self.job_3] @pytest.fixture def common_values() -> CommonValues: return CommonValues() @pytest.fixture(scope='function', autouse=True) def aws_credentials() -> None: os.environ['AWS_ACCESS_KEY_ID'] = 'testing' os.environ['AWS_SECRET_ACCESS_KEY'] = 'testing' os.environ['AWS_SECURITY_TOKEN'] = 'testing' os.environ['AWS_SESSION_TOKEN'] = 'testing' os.environ['AWS_DEFAULT_REGION'] = 'us-east-1' boto3.setup_default_session( aws_access_key_id='testing', aws_secret_access_key='testing', aws_session_token='testing', region_name='us-east-1', ) @pytest.fixture def s3_client(aws_credentials: None, common_values: CommonValues) -> S3Client: with mock_s3(): client: S3Client = boto3.client('s3') for bucket_name in common_values.get_buckets_list(): client.create_bucket(Bucket=bucket_name) yield client s3_resource = boto3.resource('s3') for bucket_name in common_values.get_buckets_list(): bucket = s3_resource.Bucket(bucket_name) bucket.objects.all().delete() bucket.delete() @pytest.fixture(scope='session') def db() -> Connection: def _credentials_loader() -> Dict[str, Union[str, int]]: return { 'active_endpoint': os.environ.get('PG_HOST', '0.0.0.0'), 'port': os.environ.get('PG_PORT', 5432), 'database': os.environ.get('PG_DB', 'slz'), 'username': os.environ.get('PG_USER', 'admin'), 'password': os.environ.get('PG_PASSWORD', 'admin'), } connection = Connection( credentials_loader=_credentials_loader, name='slz_spotify_charts_clean_up_testing', version='v1', engine_params={ 'pool_size': 2, 'max_overflow': 1 }, ) yield connection connection.session.rollback() for model in (slz.ContentFailureLog, slz.Snapshot, slz.ContentStatus, slz.UnitOfWork): connection.session.query(model).delete() connection.session.commit() connection.disconnect() @pytest.fixture def db_licensors(db: Connection) -> Dict[str, slz.Licensor]: return {item.licensor_name: item for item in db.session.query(slz.Licensor).all()} @pytest.fixture def db_reports(db: Connection) -> Dict[str, slz.Report]: return {item.report_name: item for item in db.session.query(slz.Report).all()} @pytest.fixture def now_frozen() -> FakeDatetime: with freeze_time(datetime(2021, 11, 24, 3, 21, 34, tzinfo=timezone.utc)): yield datetime.now().astimezone(timezone.utc) @enum.unique class UoWSet(enum.Enum): DEFAULT = 1 REPROCESSING = 2 ANOTHER_REPROCESSING = 3 @pytest.fixture def unit_of_work_stubs( db_licensors: Dict[str, slz.Licensor], now_frozen: FakeDatetime, db_reports: Dict[str, slz.Report], ) -> Dict[UoWSet, Dict[str, Any]]: return { UoWSet.DEFAULT: dict( unit_of_work_code='spotify-20211201-sme-charts_daily_viral-v1', reprocess_id='', licensor=db_licensors['sme'], report=db_reports['charts_daily_viral'], report_date='2021-12-01', version='v1', activity_status=slz.ActivityStatusEnum.NOT_IN_PROGRESS.value, completeness_status=slz.CompletenessStatusEnum.ACTIVE.value, is_force_complete=False, priority=slz.UnitOfWorkPriorityEnum.DEFAULT.value, next_run_at=now_frozen + timedelta(minutes=15), created_at=now_frozen, last_updated_at=now_frozen, ), UoWSet.REPROCESSING: dict( unit_of_work_code='spotify-20211201-sme-charts_daily_viral-v1', reprocess_id='123', licensor=db_licensors['sme'], report=db_reports['charts_daily_viral'], report_date='2021-12-01', version='v1', activity_status=slz.ActivityStatusEnum.NOT_IN_PROGRESS.value, completeness_status=slz.CompletenessStatusEnum.ACTIVE.value, is_force_complete=False, priority=slz.UnitOfWorkPriorityEnum.DEFAULT.value, next_run_at=now_frozen + timedelta(minutes=15), created_at=now_frozen, last_updated_at=now_frozen, ), UoWSet.ANOTHER_REPROCESSING: dict( unit_of_work_code='spotify-20211201-sme-charts_daily_viral-v1', reprocess_id='234', licensor=db_licensors['sme'], report=db_reports['charts_daily_viral'], report_date='2021-12-01', version='v1', activity_status=slz.ActivityStatusEnum.IN_PROGRESS.value, completeness_status=slz.CompletenessStatusEnum.ACTIVE.value, is_force_complete=False, priority=slz.UnitOfWorkPriorityEnum.DEFAULT.value, next_run_at=now_frozen + timedelta(minutes=15), created_at=now_frozen, last_updated_at=now_frozen, ), } @enum.unique class CSSet(enum.Enum): COMPLETE_JOB_1 = 1 MISSING_JOB_1 = 2 COMPLETE_JOB_2 = 3 COMPLETE_JOB_3 = 4 @pytest.fixture def content_status_stubs( common_values: CommonValues, now_frozen: FakeDatetime, ) -> Dict[CSSet, Dict[str, Any]]: return { CSSet.COMPLETE_JOB_1: dict( context='US', content_name='us.txt', content_status=slz.ContentStatusEnum.COMPLETE.value, created_at=now_frozen, failure_count=0, sub_content='{}', latest_job_id=common_values.job_1, metadata_process_status=slz.ContentMetadataStatusEnum.NOT_QUEUED.value, ), CSSet.MISSING_JOB_1: dict( context='MX', content_name='mx.txt', content_status=slz.ContentStatusEnum.MISSING.value, created_at=now_frozen, failure_count=0, sub_content='{}', latest_job_id=common_values.job_1, metadata_process_status=slz.ContentMetadataStatusEnum.NOT_QUEUED.value, ), CSSet.COMPLETE_JOB_2: dict( context='IT', content_name='it.txt', content_status=slz.ContentStatusEnum.COMPLETE.value, failure_count=0, created_at=now_frozen, sub_content='{}', latest_job_id=common_values.job_2, metadata_process_status=slz.ContentMetadataStatusEnum.NOT_QUEUED.value, ), CSSet.COMPLETE_JOB_3: dict( context='GB', content_name='GB.txt', content_status=slz.ContentStatusEnum.COMPLETE.value, failure_count=0, created_at=now_frozen, sub_content='{}', latest_job_id=common_values.job_3, metadata_process_status=slz.ContentMetadataStatusEnum.NOT_QUEUED.value, ), } @pytest.fixture def unit_of_work__indirect( db: Connection, unit_of_work_stubs: Dict[UoWSet, dict], request: SubRequest, ) -> Dict[UoWSet, slz.UnitOfWork]: ids_list: List[UoWSet] = request.param uows = {uow_set_id: slz.UnitOfWork(**unit_of_work_stubs[uow_set_id]) for uow_set_id in ids_list} db.session.add_all(uows.values()) db.session.commit() yield uows db.session.query(slz.UnitOfWork) \ .filter(slz.UnitOfWork.unit_of_work_id.in_(uow.unit_of_work_id for uow in uows.values())) \ .delete(synchronize_session=False) @pytest.fixture def content_status__indirect( db: Connection, content_status_stubs: Dict[CSSet, dict], unit_of_work__indirect: Dict[UoWSet, slz.UnitOfWork], request: SubRequest, ) -> Dict[UoWSet, Dict[CSSet, slz.ContentStatus]]: """Fixture creating indirect ContentStatuses. Fixture accepts via SubRequest a list of tuples, each must contain two items: first - UoWSet value, with which ContentStatus must be associated. second - CSSet value which determines what data to use from "content_status_stubs". This fixture can only be used in combination with another fixture: "unit_of_work__indirect". This another fixture must be listed first in arguments of a test to let it create UoWs first. """ content_status_data: List[Tuple[UoWSet, CSSet]] = request.param content_statuses = defaultdict(dict) for uow_key, cs_set_id in content_status_data: content_statuses[uow_key][cs_set_id] = slz.ContentStatus( unit_of_work=unit_of_work__indirect[uow_key], **content_status_stubs[cs_set_id], ) cs_instances = [cs for uow_cs in content_statuses.values() for cs in uow_cs.values()] db.session.add_all(cs_instances) db.session.commit() yield content_statuses cs_ids = (cs.content_status_id for cs in cs_instances) db.session.query(slz.ContentStatus) \ .filter(slz.ContentStatus.content_status_id.in_(cs_ids)) \ .delete(synchronize_session=False)