"""Base Mixin.""" from typing import Any, Self, Sequence from sqlalchemy import ColumnElement, Select, UnaryExpression, func, select from sqlalchemy.orm import Session from ..errors import ERROR_ENTITY_DOES_NOT_EXIST from ..types.executor import Executor from ..utils import get_pk class BaseMixin: """Base mixin for common functionality.""" @classmethod def build(cls, session: Session, **kwargs): """Construct a new object without persisting it.""" new_object = cls(**kwargs) session.add(new_object) return new_object @classmethod def create(cls, session: Session, **kwargs): """Persist a new object.""" new_object = cls.build(session, **kwargs) session.commit() return new_object @classmethod def default_filter(cls, stmt: Select[Any] | None = None) -> Select[Any]: """Get select statement with default filter applied. Models should override `_get_default_filter()` to define their filtering logic. If no filter is defined, the statement is returned unchanged. Args: stmt: Optional base select statement. If not provided, uses select(cls) Returns: Select statement with default filter applied (or unchanged if no filter defined) Example: ```python # Create new statement with default filter stmt = AbacusEvent.default_filter() # Apply default filter to existing statement stmt = select(AbacusEvent).join(SomeTable) stmt = AbacusEvent.default_filter(stmt) ``` """ base_stmt: Select[Any] = select(cls) if stmt is None else stmt filter_condition = cls._get_default_filter() if filter_condition is not None: return base_stmt.where(filter_condition) return base_stmt @classmethod def default_order(cls, stmt: Select[Any] | None = None) -> Select[Any]: """Get select statement with default ordering applied. Models should override `_get_default_order()` to define their ordering logic. If no ordering is defined, the statement is returned unchanged. Args: stmt: Optional base select statement. If not provided, uses select(cls) Returns: Select statement with default ordering applied (or unchanged if no ordering defined) Example: ```python # Create new statement with default ordering stmt = AbacusEvent.default_order() # Apply default ordering to existing statement stmt = select(AbacusEvent).where(AbacusEvent.id > 100) stmt = AbacusEvent.default_order(stmt) # Combine with default_filter stmt = AbacusEvent.default_filter() stmt = AbacusEvent.default_order(stmt) ``` """ base_stmt: Select[Any] = select(cls) if stmt is None else stmt order_expr = cls._get_default_order() if order_expr is not None: if isinstance(order_expr, Sequence) and not isinstance(order_expr, str): return base_stmt.order_by(*order_expr) return base_stmt.order_by(order_expr) return base_stmt @classmethod def get_by_id(cls, executor: Executor, obj_id): """Get object from DB by ID property.""" if obj_id is None: return None return executor.get(cls, obj_id) @classmethod def get_by_id_or_error(cls, executor: Executor, obj_id): """Find object by ID. Abort request if not found.""" obj = cls.get_by_id(executor, obj_id) if not obj: raise ValueError( ERROR_ENTITY_DOES_NOT_EXIST.format( object_type=cls.get_class_name(), object_id=obj_id ) ) return obj @classmethod def get_class_name(cls) -> str: """Return name of the subclass.""" return cls.__name__ @classmethod def _get_default_filter(cls) -> ColumnElement[bool] | None: """Define the default filter condition for this model. Override this method in subclasses to provide model-specific filtering. Returns: SQLAlchemy column expression for the WHERE clause, or None for no filter """ return None @classmethod def _get_default_order( cls, ) -> UnaryExpression[Any] | Sequence[UnaryExpression[Any]] | None: """Define the default ordering for this model. Override this method in subclasses to provide model-specific ordering. Returns: SQLAlchemy ordering expression(s), or None for no ordering """ return None @classmethod def count(cls, executor: Executor, *where_clauses: ColumnElement[bool]) -> int: """Count rows, optionally filtered. Args: executor: SQLAlchemy Session or Connection *where_clauses: Optional filter conditions Returns: Number of matching rows """ stmt = select(func.count()).select_from(cls) for clause in where_clauses: stmt = stmt.where(clause) result = executor.execute(stmt).scalar() return result or 0 def update_attributes(self, session: Session, **attrs) -> Self: """Update instance attributes. UpdateMixin handles last_modified/last_modified_by automatically via before_update event listener when the session flushes. Args: session: SQLAlchemy session (for API consistency with create/build) **attrs: Attribute key-value pairs to update Raises: AttributeError: If an attribute does not exist on the model """ for key, value in attrs.items(): if not hasattr(self, key): raise AttributeError( f'{self.get_class_name()} has no attribute {key!r}' ) setattr(self, key, value) return self def __repr__(self) -> str: """Return string representation.""" pk = get_pk(self) pk_str = ', '.join(f'{key!s}: {val!r}' for key, val in pk.items()) return f'<{self.get_class_name()}({pk_str})>'