"""Bulk Ledger Endpoint Logic.""" import copy from decimal import Decimal from typing import Dict, Iterable, List, Tuple, Type, TypedDict, Union from abacus_common_logic.connectors.database import db from owsresponse import response from owsresponse.adaptors.flask import flaskify from werkzeug.exceptions import abort from ledger.constants.error import ( ERROR_INVALID_MODEL_TYPE, ERROR_INVALID_SYNTAX, ERROR_MIN_RECORDS, ) from ledger.models.ledger_account import LedgerAccount from ledger.models.ledger_account_contract import LedgerAccountContract from ledger.models.ledger_deposit import LedgerDeposit from ledger.schemas.ledger_account import LedgerAccountCurrentBalanceSchema from ledger.schemas.ledger_account_contract import ( LedgerAccountContractCurrentBalanceSchema, ) from ledger.utils.currency import currency_exists from ledger.utils.retry import retry_on_exception class BalanceEntry(TypedDict): """Represents a balance entry.""" previous_balance: Decimal current_balance: Decimal @retry_on_exception(retries=5, delay_secs=0.333) def handle_bulk_insert(request: List[Dict]): """Handle bulk creation of ledger_account, ledger_account_contract, and ledger_deposit. Retries a maximum of 5 times, with exponential backoff, for a total of 5 seconds. Args: request (list): a list of data for account and deposit ledger entries Returns: An ows Response """ validate_request(request) account_entries, deposit_entries = validate_and_group_ledger_entries(request) try: if len(account_entries) > 0: l_account_entries, l_contract_entries = build_account_ledger_entries( account_entries ) db.session.add_all(l_account_entries) db.session.add_all(l_contract_entries) if len(deposit_entries) > 0: db.session.add_all(build_deposit_ledger_entries(deposit_entries)) db.session.commit() except Exception as e: db.session.rollback() raise e finally: db.session.close() def validate_request(request): """Validate request.""" if not request: abort(status=400, description=ERROR_MIN_RECORDS) if not isinstance(request, list): abort(status=400, description=ERROR_INVALID_SYNTAX) def validate_and_group_ledger_entries( records: List[Dict], ) -> Tuple[List[Dict], List[LedgerDeposit]]: """Validate and sort ledger entries by type.""" account_records = [] deposit_records = [] for record in records: validate_record(record) record = copy.copy(record) model_type = record.pop('model_type') if model_type == 'account': account_records.append(record) elif model_type == 'deposit': deposit_records.append(record) else: abort( status=400, description=ERROR_INVALID_MODEL_TYPE.format(model_type='known'), ) return account_records, deposit_records def create_ledger_entry(record): """Save a ledger entry.""" validate_record(record) model_type = record.pop('model_type') if model_type == 'account': build_account_ledger(record) build_account_contract_ledger(record) elif model_type == 'deposit': record['remaining_amount'] = Decimal(record['remaining_amount']) LedgerDeposit.build(**record) else: abort( status=400, description=ERROR_INVALID_MODEL_TYPE.format(model_type='known') ) def build_account_ledger_entries( records: List[dict], ) -> Tuple[List[LedgerAccount], List[LedgerAccountContract]]: """Build ledger_account or ledger_account_contract entries.""" account_ids, contract_ids = _get_account_and_contract_ids(records) balances_by_account_ids = get_ledger_current_balances_by_accounts(account_ids, True) balances_by_contract_ids = get_ledger_current_balances_by_contracts( contract_ids, True ) ledger_account_entries = [] ledger_account_contract_entries = [] zero = Decimal(0) for record in records: account_id = record['account_id'] account_balance = balances_by_account_ids.get(account_id, zero) ledger_acc_entry = build_ledger_entry(record, account_balance, LedgerAccount) ledger_account_entries.append(ledger_acc_entry) # this is to manage multiple ledger entries for a single account # this keeps track of what the account balance will be # based on the records we are adding balances_by_account_ids[account_id] = ledger_acc_entry.current_balance contract_balance = balances_by_contract_ids.get(record['contract_id'], zero) ledger_account_contract_entries.append( build_ledger_entry(record, contract_balance, LedgerAccountContract) ) return ledger_account_entries, ledger_account_contract_entries def build_ledger_entry( record: Dict, current_balance: Decimal, ledger_class: Type[Union[LedgerAccount, LedgerAccountContract]], ) -> Union[LedgerAccount, LedgerAccountContract]: """Build ledger_account entry.""" decimal_amount = Decimal(record['currency_amount']) record.update( { 'previous_balance': current_balance, 'current_balance': current_balance + decimal_amount, } ) return ledger_class(**record) def _get_account_and_contract_ids(records: List[Dict]) -> Tuple[set[int], set[int]]: """Split records into sets of unique account and contract ids.""" account_ids = set() contract_ids = set() for record in records: account_ids.add(record['account_id']) contract_ids.add(record['contract_id']) return account_ids, contract_ids def build_deposit_ledger_entries(records: List[LedgerDeposit]): """Build ledger_deposit entries.""" formatted_record = [] for record in records: record['remaining_amount'] = Decimal(record['remaining_amount']) formatted_record.append(LedgerDeposit(**record)) return formatted_record def build_account_ledger(record): """Build a account ledger.""" balances_dict = calculate_account_balances( record['currency_amount'], record['account_id'] ) record.update(balances_dict) LedgerAccount.build(**record) def build_account_contract_ledger(record): """Build a account_contract ledger.""" balances_dict = calculate_account_contract_balances( record['currency_amount'], record['contract_id'], ) record.update(balances_dict) LedgerAccountContract.build(**record) def validate_record(record: dict): """Validate record. Args: record (dict): data for a ledger_account or ledger_deposit entry """ if 'model_type' not in record: abort(status=400, description=ERROR_INVALID_SYNTAX) if not currency_exists(record.get('currency_code')): abort(status=400, description='unrecognized currency code') def get_ledger_current_balances_by_accounts( account_ids: Iterable[int], for_update: bool = False ) -> Dict[int, Decimal]: """Get ledger_account_contract current balance for given account ids.""" query = LedgerAccount.get_ledger_account_balances(account_ids) if for_update: query = query.with_for_update() entries = query.all() balances_dict = {} for entry in entries: balances_dict[entry.account_id] = Decimal(entry.current_balance) return balances_dict def get_ledger_current_balances_by_contracts( contract_ids: Iterable[int], for_update: bool = False ) -> Dict[int, Decimal]: """Get ledger_account_contract current balance for given contract ids.""" query = LedgerAccountContract.get_ledger_account_contract_balance_by_contracts( contract_ids ) if for_update: query = query.with_for_update() entries = query.all() balances_dict = {} for entry in entries: balances_dict[entry.contract_id] = Decimal(entry.current_balance) return balances_dict def calculate_account_balances(amount: str, account_id: int) -> Dict: """Determine previous and current balance.""" latest_entry = LedgerAccount.get_ledger_account_balance(account_id).first() current_balance = Decimal('0') if latest_entry and latest_entry.current_balance: current_balance = Decimal(latest_entry.current_balance) decimal_amount = Decimal(amount) return { 'previous_balance': current_balance, 'current_balance': current_balance + decimal_amount, } def calculate_account_contract_balances(amount: str, contract_id: int) -> BalanceEntry: """Determine previous and current balance.""" latest_entry = ( LedgerAccountContract.get_ledger_account_contract_balance_by_contract( contract_id ).first() ) current_balance = Decimal('0') if latest_entry and latest_entry.current_balance: current_balance = Decimal(latest_entry.current_balance) decimal_amount = Decimal(amount) return { 'previous_balance': current_balance, 'current_balance': current_balance + decimal_amount, } def get_accounts_balance( account_ids: List[int], balance_min: str | Decimal | None, balance_max: str | Decimal | None, ): """Retrieve current accounts balances. Arguments: account_ids (list): list of account ids balance_min (String): minimum value for current balance to filter balance_max (String): maximum value for current balance to filter """ schema = LedgerAccountCurrentBalanceSchema(many=True) if balance_min: balance_min = Decimal(balance_min) if balance_max: balance_max = Decimal(balance_max) rows = LedgerAccount.get_by_custom_filters( account_ids=account_ids, balance_min=balance_min, balance_max=balance_max ).all() return flaskify(response.Response(message=schema.dump(rows), status=200)) def get_contracts_balance(contract_ids: List[int], account_ids: List[int]): """Retrieve current contracts balances. Arguments: contract_ids(list): list of contract ids to filter by account_ids(list): list of account ids to filter by """ rows = LedgerAccountContract.get_ledger_account_contract_balances( contract_ids, account_ids ).all() schema = LedgerAccountContractCurrentBalanceSchema(many=True) return flaskify(response.Response(message=schema.dump(rows), status=200))