"""Contract Terms logic.""" import json from typing import Type from abacus_common_logic.connectors.database import db from marshmallow import ValidationError from owsresponse import response from sqlalchemy.exc import SQLAlchemyError from sqlalchemy.orm import joinedload from sqlalchemy.sql import null from werkzeug.exceptions import HTTPException from abacus_contract import models from abacus_contract.constants import error from abacus_contract.constants.constants import ( CONTRACT_TERM_TYPES, CONTRACT_TERMS_SNAPSHOT_HEADER, ) from abacus_contract.logic.contract_term_schedule import ( create_or_update_contract_term_schedules, ) from abacus_contract.schemas.contract_term import ( ContractTermBulkPostSchema, ContractTermSchema, ) from abacus_contract.schemas.contract_term_condition import ( ContractTermConditionPostSchema, ) from abacus_contract.utils.features import ( is_abacus_terms_with_no_tracks_and_products_enabled, ) from abacus_contract.utils.format_error import validation_error valid_nr_term_types = [ CONTRACT_TERM_TYPES.CONTRIBUTION_SCHEDULE, CONTRACT_TERM_TYPES.CONTRIBUTOR_SCHEDULE, ] def bulk_create_contract_terms_and_conditions(create_params: list): """Bulk create multiple contract terms and conditions.""" if not isinstance(create_params, list) or len(create_params) == 0: return validation_error( error.ERROR_INVALID_INPUT_TYPE.format( type_received=type(create_params).__name__, type_expected='List[dict]' ) ) new_contract_terms = list() invalid_contract_terms = list() for params in create_params: try: contract_term_params = ContractTermBulkPostSchema( exclude=('contract_term_id',) ).load(params) contract_term_conditions = ContractTermConditionPostSchema(many=True).load( params.get('contract_term_conditions') ) _validate_contract_term(contract_term_params) attachments_relations = ( contract_term_params.get('attachments_relations') if contract_term_params.get('attachments_relations') else null() ) new_contract_term = models.ContractTerm.build( attachments=contract_term_params.get('attachments'), attachments_relations=attachments_relations, contract_id=contract_term_params.get('contract_id'), is_base_term=contract_term_params.get('is_base_term'), term_type=contract_term_params.get('term_type'), ) db.session.flush() for term_condition in contract_term_conditions: contract_term_condition_name = term_condition.get( 'contract_term_condition_name', None ) term_rate = term_condition.get('term_rate', None) commission = term_condition.get('commission', None) if term_rate is not None and commission is None: commission = 100 - term_rate if commission is not None and term_rate is None: term_rate = 100 - commission models.ContractTermCondition.build( contract_term_id=new_contract_term.contract_term_id, contract_term_condition_name=contract_term_condition_name, conditions=term_condition.get('conditions'), priority=term_condition.get('priority'), term_rate=term_rate, commission=commission, ) new_contract_terms.append(new_contract_term) except (ValidationError, HTTPException) as e: invalid_contract_terms.append({'data': params, 'error': str(e)}) if new_contract_terms: models.ContractTermCondition.commit_changes() message = { 'contract_terms': ContractTermSchema(many=True).dump(new_contract_terms), 'invalid_contract_terms': invalid_contract_terms, } return response.Response(message=message, status=201) def create_contract_term(**params: dict): """Create contract terms logic. The function will also create the contract_term_schedules if the term_type is contributor_schedule/contribution_schedule. """ schedule_ids = params.get('schedule_ids', None) term_type = params.get('term_type', None) attachments = params.get('attachments', None) attachments_relations = params.get('attachments_relations', None) try: # TODO: use FF condition when FF is removed if is_abacus_terms_with_no_tracks_and_products_enabled(): _validate_contract_term(params) if attachments is None or term_type in valid_nr_term_types: params.update({'attachments': null()}) if attachments_relations is None or term_type in valid_nr_term_types: params.update({'attachments_relations': null()}) if 'schedule_ids' in params: del params['schedule_ids'] else: if ( 'attachments_relations' in params and params.get('attachments_relations') is None ): del params['attachments_relations'] _validate_contract_term(params) if 'schedule_ids' in params: del params['schedule_ids'] if term_type in valid_nr_term_types: params.update({'attachments': None, 'attachments_relations': None}) contract_term = models.ContractTerm.build(**params) if term_type in valid_nr_term_types: db.session.flush() create_or_update_contract_term_schedules( schedule_ids, contract_term.contract_term_id ) except ValidationError as e: return validation_error(e.messages[0]) except SQLAlchemyError as e: db.session.rollback() raise e db.session.commit() message = ContractTermSchema().dump(contract_term) return response.Response(message=message, status=201) def get_contract_terms_by_contract(contract_id): """Get contract terms belonging to the specified contract.""" result = models.Contract.get_by_id_or_error(contract_id) message = ContractTermSchema(many=True).dump(result.contract_terms) return response.Response(message=message, status=200) def update_contract_term(contract_term, **kwargs): """Update contract terms logic.""" contract_terms = contract_term.contract.contract_terms if contract_term.term_type in valid_nr_term_types: try: term_type = ( kwargs.get('term_type') if 'term_type' in kwargs else contract_term.term_type ) schedule_ids = kwargs.get('schedule_ids') if 'schedule_ids' not in kwargs and contract_term.term_type != term_type: return validation_error(error.ERROR_MISSING_SCHEDULE_IDS) if 'schedule_ids' in kwargs: del kwargs['schedule_ids'] if kwargs: kwargs.update({'attachments': None, 'attachments_relations': None}) contract_term.update_attributes(**kwargs) if schedule_ids: db.session.flush() create_or_update_contract_term_schedules( schedule_ids, contract_term.contract_term_id ) except ValidationError as e: return validation_error(e.messages[0]) except SQLAlchemyError as e: db.session.rollback() raise e models.ContractTerm.commit_changes() else: try: attachments = kwargs['attachments'] # TODO: Remove this condition when FF is removed if ( attachments is None and is_abacus_terms_with_no_tracks_and_products_enabled() is False ): raise ValidationError(error.ERROR_ATTACHMENTS_CAN_NOT_BE_NULL) # TODO: Remove only FF check if ( attachments is None and contract_term.term_type == CONTRACT_TERM_TYPES.LABEL and is_abacus_terms_with_no_tracks_and_products_enabled() is True ): raise ValidationError(error.ERROR_ATTACHMENTS_CAN_NOT_BE_NULL) _validate_conflicting_attachments( contract_terms=contract_terms, term_type=contract_term.term_type, attachments=attachments, skip_term_id=contract_term.contract_term_id, ) except ValidationError as err: return response.create_error_response(code='error', message=err.messages[0]) if 'attachments_relations' in kwargs: attachments_relations = ( kwargs['attachments_relations'] if kwargs['attachments_relations'] else null() ) contract_term.update_attributes(attachments_relations=attachments_relations) # TODO: Use FF condition when FF is removed if is_abacus_terms_with_no_tracks_and_products_enabled(): contract_term.update_attributes( attachments=attachments if attachments else null() ) else: contract_term.update_attributes(attachments=attachments) if 'contract_term_name' in kwargs: contract_term.update_attributes( contract_term_name=kwargs['contract_term_name'] ) models.ContractTerm.commit_changes() return response.Response( message=ContractTermSchema().dump(contract_term), status=200 ) def contract_term_export(contract_ids=None): """Export contract terms as csv.""" yield '\t'.join(CONTRACT_TERMS_SNAPSHOT_HEADER) + '\n' contract_terms = models.ContractTerm.stream_all(contract_ids=contract_ids) for row in contract_terms: contract_term = row[1] term_condition = row[2] attachments = ( _format_attachments(contract_term) if contract_term.attachments else '' ) conditions = _format_term_conditions(term_condition) if term_condition else '' yield _build_contract_term_row( contract_term, term_condition, attachments, conditions ) def _build_contract_term_row(contract_term, term_condition, attachments, conditions): """Format a contract term row.""" return ( '\t'.join( [ str(contract_term.contract_id), str(contract_term.contract_term_id), contract_term.term_type, attachments, json.dumps(contract_term.attachments_relations or dict()), 'true' if contract_term.is_base_term else 'false', str(term_condition.contract_term_condition_id) if term_condition else '', conditions, str(term_condition.term_rate) if term_condition else '', str(term_condition.priority) if term_condition else '', ] ) + '\n' ) def _format_attachments(contract_term): """Format attachments as a list.""" return ','.join([str(attachment) for attachment in contract_term.attachments]) def _format_term_conditions(term_condition): """Format contract_term_conditions as stringified json.""" conditions = '' conditions_with_values = { key: value for (key, value) in term_condition.conditions.items() if len(value) } if len(conditions_with_values): conditions = json.dumps(conditions_with_values) return conditions def _validate_contract_term(params: dict): """Validate contract term parameters.""" attachments = params.get('attachments') contract_id = params.get('contract_id') term_type = params.get('term_type') contract = models.Contract.get_by_id_or_error(contract_id) if term_type in valid_nr_term_types: return True # TODO: remove this condition when FF is removed if ( attachments is None and is_abacus_terms_with_no_tracks_and_products_enabled() is False and term_type in [ CONTRACT_TERM_TYPES.LABEL, CONTRACT_TERM_TYPES.PRODUCT, CONTRACT_TERM_TYPES.TRACK, ] ): raise ValidationError(error.ERROR_ATTACHMENTS_CAN_NOT_BE_NULL) # TODO: remove this condition when FF is removed if ( not attachments and is_abacus_terms_with_no_tracks_and_products_enabled() is False ): raise ValidationError(error.ERROR_MISSING_ATTACHMENT) # TODO: remove FF check if is_abacus_terms_with_no_tracks_and_products_enabled() is True: if attachments is not None and len(attachments) == 0: raise ValidationError(error.ERROR_MISSING_ATTACHMENT) if attachments is None: if term_type == CONTRACT_TERM_TYPES.LABEL: raise ValidationError(error.ERROR_ATTACHMENTS_CAN_NOT_BE_NULL) else: return True _validate_conflicting_attachments(contract.contract_terms, term_type, attachments) if term_type in [CONTRACT_TERM_TYPES.PRODUCT, CONTRACT_TERM_TYPES.TRACK]: _validate_attachment_uniqueness(contract, attachments, term_type) def _validate_conflicting_attachments( contract_terms, term_type, attachments, skip_term_id=None ): """Validate attachments to avoid duplications. Args: contract_terms (list): existing contract terms of a contract term_type (str): new/updated type of contract_term attachments (list): new/updated list of attachments skip_term_id (int): skip this term_id for edit purposes """ same_type_terms = [term for term in contract_terms if term.term_type == term_type] for term in same_type_terms: if skip_term_id and term.contract_term_id == skip_term_id: continue if term.attachments and attachments: conflict_attachments = set(term.attachments) & set(attachments) if conflict_attachments: attachments_msg = ', '.join(map(str, sorted(conflict_attachments))) raise ValidationError( error.ERROR_ATTACHMENT_ALREADY_EXISTS.format( attachment=attachments_msg ) ) def _validate_attachment_uniqueness(contract, attachments, term_type): """Ensure that the attachments don't exist on other contracts. Args: contract (dict): contract attachments are for. attachments (list): attachments. term_type (string): type of contract term """ account_id = contract.account_contract.account_id params = dict(attachments=attachments, term_type=term_type) contract_terms = models.ContractTerm.get_contract_terms_by_account_and_term_type( account_id, **params ) contract_attachments = [ attachment for term in contract_terms for attachment in term.attachments ] assigned_attachments = [ attachment for attachment in attachments if str(attachment) in contract_attachments ] if assigned_attachments: raise ValidationError( error.ERROR_ATTACHMENTS_ALREADY_EXISTS_OTHER_CONTRACT.format( attachments=list(set(assigned_attachments)), term_type=term_type.capitalize(), ) ) def get_contract_terms_for_account_and_term_type(account_id, params): """Get contract terms for a specific account and term_type.""" result = models.ContractTerm.get_contract_terms_by_account_and_term_type( account_id, **params ) if not result: return response.create_not_found_response() return response.Response( message=ContractTermSchema(many=True).dump(result), status=200 ) def soft_delete_contract_term_and_conditions( contract_term: Type[models.ContractTerm], ) -> Type[response.Response]: """Soft delete specified contract term and attached term conditions, schedules. Args: contract_term (class): ContactTerm instance """ contract_term_conditions = contract_term.conditions contract_term_schedules = contract_term.contract_term_schedule try: for condition in contract_term_conditions: if condition.deleted_by is None or condition.deleted_at is None: condition._soft_delete() for contract_term_schedule in contract_term_schedules: if ( contract_term_schedule.deleted_by is None or contract_term_schedule.deleted_at is None ): contract_term_schedule._soft_delete() contract_term._soft_delete() except SQLAlchemyError as e: db.session.rollback() return e db.session.commit() return response.Response(status=204) def get_account_id_by_contract_term_id(contract_term_id: int) -> int | None: """Get account id by contract term id.""" contract_term = ( db.session.query(models.ContractTerm) .options( joinedload(models.ContractTerm.contract).joinedload( models.Contract.account_contract ) ) .filter(models.ContractTerm.contract_term_id == contract_term_id) .one_or_none() ) if ( contract_term and contract_term.contract and contract_term.contract.account_contract ): return contract_term.contract.account_contract.account_id def get_account_ids_by_contract_term_ids( contract_term_ids: list[int], ) -> dict[int, int]: """Get a map of contract_term_id to account_id for a batch of contract term ids.""" contract_terms = ( db.session.query(models.ContractTerm) .options( joinedload(models.ContractTerm.contract).joinedload( models.Contract.account_contract ) ) .filter(models.ContractTerm.contract_term_id.in_(contract_term_ids)) .all() ) return { contract_term.contract_term_id: contract_term.contract.account_contract.account_id for contract_term in contract_terms if contract_term.contract and contract_term.contract.account_contract } def get_contract_term_records_by_contract_ids(authorized_contract_ids: list) -> list: """Serialize non-deleted contract terms for the authorized contract ids as flat records. Each record carries contract_id; the caller groups and shapes them into the dataload response. """ if not authorized_contract_ids: return [] contract_terms = models.ContractTerm.get_non_deleted_by_contract_ids( authorized_contract_ids ) return ContractTermSchema(many=True).dump(contract_terms)