"""Abacus state model. Model for managing abacus states """ from typing import cast from abacus_common_logic.connectors.database import db from abacus_common_logic.models.base import BaseModel from sqlalchemy import Enum from abacus_state.constants.constants import ACTION_STATUSES class AbacusState(BaseModel): """Abacus state model.""" __tablename__ = 'abacus_state' abacus_state_id = db.Column(db.Integer, primary_key=True) parent_table_id = db.Column(db.Integer, nullable=False) parent_table_name = db.Column(db.String(64), nullable=False) action_name = db.Column(db.String(32), nullable=False) action_status = db.Column( Enum(*ACTION_STATUSES, name='action_status', create_type=False), default=ACTION_STATUSES.INIT, nullable=False, ) message = db.Column(db.String(180), nullable=True) @classmethod def get_action_status_list(cls, parent_table_name, parent_table_id): """Get action statuses by specified parent table.""" return ( cls.query.filter( cls.parent_table_name == parent_table_name, cls.parent_table_id == parent_table_id, ) .order_by(cls.abacus_state_id.asc()) .all() ) @classmethod def get_formatted_action_statuses( cls, parent_table_name: str, parent_table_id: int ) -> dict[str, str]: """Get formatted action statuses by specified parent table.""" actions_list = cls.get_action_status_list(parent_table_name, parent_table_id) actions = dict() for action in actions_list: actions[action.action_name] = action.action_status return actions @classmethod def get_filtered_query( cls, state_ids=None, parent_table_name=None, parent_table_ids=None, action_name=None, limit=None, offset=None, ): """Get states by field values. Args: state_ids (list): List of state IDs to filter by parent_table_name (str): Parent table name to filter by parent_table_ids (list): List of parent table IDs to filter by (requires parent_table_name) action_name (str): Action name to filter by limit (int): Maximum number of results to return offset (int): Number of results to skip Returns: Query object with applied filters """ query = cls.query if state_ids is not None: query = query.filter(cls.abacus_state_id.in_(state_ids)) if parent_table_ids is not None: query = query.filter(cls.parent_table_id.in_(parent_table_ids)) if parent_table_name is not None: query = query.filter(cls.parent_table_name == parent_table_name) if action_name is not None: query = query.filter(cls.action_name == action_name) query = query.order_by(cls.abacus_state_id.asc()) if limit is not None: query = query.limit(limit) if offset is not None: query = query.offset(offset) return query @classmethod def reset_abacus_state_actions( cls, parent_table_name: str, parent_table_id: int ) -> list['AbacusState']: """Reset all action states for specified parent table name and id. Args: parent_table_name (str): name of the parent table parent_table_id (int): id of the parent table Returns: list of updated states """ query = db.session.query(AbacusState).filter( AbacusState.parent_table_id == parent_table_id, AbacusState.parent_table_name == parent_table_name, ) query.update( { AbacusState.action_status: ACTION_STATUSES.INIT, }, synchronize_session=False, ) db.session.commit() result = query.all() return cast('AbacusState', result)