"""Contract Exclusion blueprint.""" from http import HTTPStatus from abacus_common_logic.utils.authorization import permissions_authorize_many_accounts from abacus_common_logic.views.validations import validated_request_body from common_apispec import doc, marshal_with, use_kwargs from flask import Blueprint, g, request from owsrequest import flask_request from owsresponse import response from owsresponse.adaptors.flask import flaskify from abacus_contract.constants import error from abacus_contract.logic import ( contract as contract_logic, contract_exclusion as logic, ) from abacus_contract.schemas.contract_exclusion import ( ContractExclusionListResponseSchema, ContractExclusionSchema, ExclusionsSchema, ) from abacus_contract.utils.authorization import pdp_authorize_many_accounts from abacus_contract.utils.dataloader import account_scoped_dataloader from core.config import ows_client contract_exclusion_api = Blueprint('contract_exclusion_api', __name__) @contract_exclusion_api.route( '/contract//exclusions/', methods=['POST'] ) @doc( summary='Create countries and stores that should be excluded from distribution', params={ 'contract_id': { 'description': 'The ID of the contract', }, }, ) @use_kwargs(ExclusionsSchema, location='json', required=True, apply=False) @marshal_with( ContractExclusionSchema, code=HTTPStatus.CREATED, description=HTTPStatus.CREATED.phrase, ) def create_contract_exclusions(contract_id): """POST exclusions for a specified contract.""" post_schema = ExclusionsSchema() exclusions = validated_request_body(post_schema) access_rule_decision = flask_request.verify_rules_access_standalone(request) if not access_rule_decision: return flaskify( response.create_error_response( code=error.ERROR_CODE_AUTHORIZATION, message='Unauthorized', status=401, ) ) return flaskify(logic.create_contract_exclusions(contract_id, exclusions)) @contract_exclusion_api.route( '/contract//exclusions/', methods=['GET'] ) @doc( summary='Get excluded countries by contract id', params={ 'contract_id': { 'description': 'The ID of the contract', }, }, ) @marshal_with( ContractExclusionListResponseSchema, code=HTTPStatus.OK, description=HTTPStatus.OK.phrase, ) def get_exclusions_by_contract(contract_id): """GET exclusions object for a specified contract.""" account_id = contract_logic.get_account_id_by_contract_id(contract_id) if not account_id: return flaskify( response.create_error_response( code=error.ERROR_CODE_NOT_FOUND, message=error.ERROR_ACCOUNT_NOT_FOUND_FOR_CONTRACT, status=404, ) ) access_rule_decision = flask_request.verify_rules_access_standalone(request) if not access_rule_decision: authorized = pdp_authorize_many_accounts([account_id]) if not authorized: return flaskify( response.create_error_response( code=error.ERROR_CODE_AUTHORIZATION, message=error.ERROR_MESSAGE_FORBIDDEN_USER, status=403, ) ) authorized = permissions_authorize_many_accounts( ows_client, g.request_context.profile_type, g.request_context.profile_id, [account_id], ) if not authorized: return flaskify( response.create_error_response( code=error.ERROR_CODE_FORBIDDEN, message=error.ERROR_MESSAGE_FORBIDDEN_USER, status=403, ) ) return flaskify(logic.get_exclusions_by_contract(contract_id)) @contract_exclusion_api.route( '/contract//exclusions/', methods=['PUT'] ) @doc( summary='Update countries and stores that should be excluded from distribution', params={ 'contract_id': { 'description': 'The ID of the contract', }, }, ) @use_kwargs(ExclusionsSchema, location='json', required=True, apply=False) @marshal_with( ContractExclusionSchema, code=HTTPStatus.OK, description=HTTPStatus.OK.phrase ) def update_contract_exclusions(contract_id): """PUT exclusions for a specified contract.""" put_schema = ExclusionsSchema() exclusions = validated_request_body(put_schema) access_rule_decision = flask_request.verify_rules_access_standalone(request) if not access_rule_decision: return flaskify( response.create_error_response( code=error.ERROR_CODE_AUTHORIZATION, message='Unauthorized', status=401, ) ) return flaskify(logic.update_contract_exclusions(contract_id, exclusions)) @contract_exclusion_api.route('/contract-exclusions/dataloader', methods=['POST']) def get_exclusions_by_contract_ids_dataloader(): """Batch-fetch contract exclusions for the given contract_ids.""" return account_scoped_dataloader( entity_name='Contract', resolve_accounts=contract_logic.get_account_id_map_by_contract_ids, fetch_records=logic.get_exclusion_records_by_contract_ids, key_field='contract_id', as_list=False, )