"""Contract Exclusion logic.""" from owsresponse import response from abacus_contract import models from abacus_contract.constants.constants import COUNTRIES_POST_LIMIT from abacus_contract.constants.error import ( ERROR_CONTRACT_EXCLUSION_ALREADY_EXISTS, ERROR_CONTRACT_EXCLUSION_DOES_NOT_EXIST, ERROR_COUNTRIES_POST_LIMIT, ERROR_OA_CONTRACT_EXISTS, ) from abacus_contract.schemas.contract_exclusion import ContractExclusionSchema from abacus_contract.utils.validations import validate_country_code def create_contract_exclusions(contract_id, exclusions): """Create contract exclusions. @param contract_id: Contract id @param exclusions: Dict containing excluded countries and/or stores. For example: {"countries": ["USA", "AFG"], "stores": ["123"]} """ contract = models.Contract.get_by_id_or_error(contract_id) if contract.contract_exclusion: return response.create_error_response( 'error', ERROR_CONTRACT_EXCLUSION_ALREADY_EXISTS.format(contract_id=contract_id), status=400, ) countries = exclusions.get('countries', []) if len(countries) > COUNTRIES_POST_LIMIT: return response.create_error_response( 'error', ERROR_COUNTRIES_POST_LIMIT.format(limit=COUNTRIES_POST_LIMIT), status=400, ) if contract.legacy_contract and contract.legacy_contract.oa_contract_id: return response.create_error_response( 'error', ERROR_OA_CONTRACT_EXISTS, status=400 ) try: for country_code in countries: validate_country_code(country_code) exclusion_params = {'contract_id': contract_id, 'exclusions': exclusions} exclusion = models.ContractExclusion.create(**exclusion_params) message = ContractExclusionSchema().dump(exclusion) return response.Response(message=message, status=201) except Exception as e: return response.create_error_response('error', str(e), status=400) def get_exclusions_by_contract(contract_id): """GET exclusions object for a specified contract.""" result = models.Contract.get_by_id_or_error(contract_id) if not result.contract_exclusion: return response.Response(message={}, status=200) message = ContractExclusionSchema().dump(result.contract_exclusion) return response.Response(message=message, status=200) def get_exclusion_records_by_contract_ids(authorized_contract_ids: list) -> list: """Serialize contract exclusions for the authorized contract ids as flat records. Each record carries contract_id; the caller groups and shapes them (one per id). """ if not authorized_contract_ids: return [] exclusions = models.ContractExclusion.get_by_contract_ids(authorized_contract_ids) return ContractExclusionSchema(many=True).dump(exclusions) def update_contract_exclusions(contract_id, exclusions): """Update contract exclusions. @param contract_id: Contract id @param exclusions: Dict containing excluded countries and/or stores. For example: {"countries": ["USA", "AFG"], "stores": ["123"]} """ contract = models.Contract.get_by_id_or_error(contract_id) if not contract.contract_exclusion: return response.create_error_response( 'error', ERROR_CONTRACT_EXCLUSION_DOES_NOT_EXIST.format(contract_id=contract_id), status=404, ) countries = exclusions.get('countries', []) try: for country_code in countries: validate_country_code(country_code) exclusion_params = {'exclusions': exclusions} exclusion = contract.contract_exclusion.update_attributes(**exclusion_params) models.ContractExclusion.commit_changes() message = ContractExclusionSchema().dump(exclusion) return response.Response(message=message, status=200) except Exception as e: return response.create_error_response('error', str(e), status=400)