"""Payee model.""" from typing import Protocol, cast from abacus_common_logic.models.base import db from sqlalchemy import select from sqlalchemy.sql import literal_column, table from abacus_state.constants.constants import PAYEE_TYPES class GetTypedPayeeException(Exception): """Get typed payee exception.""" def __init__(self, message): """Init GetTypedPayeeException.""" super().__init__(f'Getting payee subtype failed: {message}') class PayeeCollaborator(Protocol): """Payee Collaborator protocol.""" collaborator_id: int class Payee(Protocol): """Payee protocol.""" payee_id: int payee_type: str @classmethod def get_by_id(cls, payee_id) -> 'Payee': """Get payee data by payee_id.""" query = ( select( [ literal_column('p.payee_id').label('payee_id'), literal_column('p.payee_type').label('payee_type'), ] ) .where(literal_column('p.payee_id') == payee_id) .select_from(table('payee').alias('p')) ) result = db.session.execute(query).first() return cast('Payee', result) @classmethod def get_typed_payee(cls, payee_id) -> 'PayeeCollaborator': """Get the payee's concrete subtype entity.""" payee = cls.get_by_id(payee_id) payee_type = payee.payee_type if payee else None if payee_type != PAYEE_TYPES.COLLABORATOR: raise GetTypedPayeeException(payee_type) query = ( select( [ literal_column('pc.collaborator_id').label('collaborator_id'), ] ) .where(literal_column('pc.payee_id') == payee_id) .select_from( table('payee') .alias('p') .join( table('payee_collaborator').alias('pc'), literal_column('p.payee_id') == literal_column('pc.payee_id'), ) ) ) result = db.session.execute(query).first() return cast('PayeeCollaborator', result)