import asyncio from sqlalchemy import Boolean, Column, DateTime, and_, bindparam, delete, func, insert, update from sqlalchemy.exc import IntegrityError from sqlalchemy.ext.declarative import declarative_base from sqlalchemy.future import select from sqlalchemy.orm import DeclarativeMeta, Query, Session from sqlalchemy.sql.expression import ClauseElement from typing import Any, Dict, Iterable, List, Optional, Union from server.db.session import db_session, metadata from server.utils.logger import logger class TimestampMixin: created_at = Column(DateTime, default=func.now()) updated_at = Column(DateTime, default=func.now(), onupdate=func.now()) class BaseCRUDMixin: """Basic mixin to provide model CRUD interface.""" @classmethod def q_filter( cls, query: Optional[ClauseElement] = None, filters: Optional[List[ClauseElement]] = None, **by_params ) -> ClauseElement: """Basic query filter. all Executables passed in 'filters' is being applied with 'and' operator. """ if query is None: query = select(cls) filters = [(getattr(cls, param) == value) for param, value in by_params.items()] + ( list(filters) if filters is not None else [] ) if filters: query = query.where(and_(*filters)) return query @classmethod async def get(cls, session: Session = None, **params): query = cls.q_filter(**params) async with db_session(session=session) as _session: result = (await _session.execute(query)).fetchone() if result: return result[0] @classmethod async def create(cls, params, session: Session = None): async with db_session(session=session) as _session: inst = cls(**params) _session.add(inst) return inst @classmethod async def get_or_create(cls, session: Session = None, **params): created = False async with db_session(session=session) as _session: item = await cls.get(session=_session, **params) if not item: try: item = await cls.create(params, session=_session) created = True except IntegrityError: await session.rollback() item = await cls.get(session=_session, **params) if not item: raise return item, created @classmethod def _extract_cursor_data(cls, result, return_many: bool = True, return_single_column: bool = True): if return_many: if return_single_column: return [i[0] for i in result.fetchall()] return result.fetchall() return result.fetchone() @classmethod async def _get_data( cls, query, session: Session = None, return_data: bool = True, return_many: bool = True, return_single_column: bool = True, ): async with db_session(session=session) as _session: result = await _session.execute(query) if return_data: return cls._extract_cursor_data(result, return_many, return_single_column) return result @classmethod async def _handle_limit( cls, filters: Optional[List[ClauseElement]], limit: int, session: Session = None, return_data: bool = True, **params, ): """Handle limit in delete or update queries. If we need to return data then execute get ID list query and then perform the original operation for them. If we don't neet to return data then create a subquery and execute the original operation with IN subquery. Args: filters: Original filters. limit: Items limit. return_data: Return data or not. Returns: Updated filters and params. """ query = cls.q_filter(select(cls.id), filters=filters, **params).limit(limit) if return_data: id_list = await cls._get_data( query=query, session=session, return_data=True, return_many=True, return_single_column=True, ) return [cls.id.in_(id_list)], {} return [cls.id.in_(query.subquery())], {} @classmethod async def update( cls, data: dict, filters: Optional[List[ClauseElement]] = None, session: Session = None, return_source: Union[Column, DeclarativeMeta] = None, return_many: bool = False, return_single_column: bool = False, limit: Optional[int] = None, synchronize_session = None, **filter_params, ): if filters is None and not filter_params: raise ValueError(f"At least on of 'filters', 'filter_params' should be passed.") if limit: filters, filter_params = await cls._handle_limit( filters, limit, session, bool(return_source), **filter_params ) query = cls.q_filter(update(cls), filters=filters, **filter_params).values(**data) if return_source: query = query.returning(return_source) if synchronize_session is not None: query = query.execution_options(synchronize_session=synchronize_session) return await cls._get_data( query=query, session=session, return_data=bool(return_source), return_many=return_many, return_single_column=return_single_column, ) @classmethod async def delete( cls, filters: Optional[List[ClauseElement]] = None, session: Session = None, return_source: Union[Column, DeclarativeMeta] = None, return_single_column: bool = False, return_row_count: bool = False, limit: Optional[int] = None, **params, ): query = delete(cls) if limit: filters, params = await cls._handle_limit(filters, limit, session, bool(return_source), **params) query = cls.q_filter(query, filters=filters, **params) if return_source: query = query.returning(return_source).execution_options(synchronize_session=False) else: query = query.execution_options(synchronize_session="fetch") result = await cls._get_data( query=query, session=session, return_data=bool(return_source), return_many=True, return_single_column=return_single_column, ) if return_source: return result return result.rowcount if return_row_count else bool(result.rowcount) class BulkCRUDMixin(BaseCRUDMixin): """Basic mixin to provide model bulk CRUD interface.""" @staticmethod async def get_bulk_unique(query: Query, session: Session = None) -> list: async with db_session(session=session) as _session: return (await _session.execute(query)).scalars().unique().all() @classmethod def q_count(cls, query=None, **params): return select(func.count()).select_from(cls.q_filter(query, **params).alias("count_query")) @classmethod async def count(cls, query=None, session: Session = None, **params): async with db_session(session=session) as _session: return (await _session.execute(cls.q_count(query, **params))).fetchone()[0] @classmethod async def paginate(cls, query, session: Session = None, limit=None, offset=None): page_query = query if limit is not None: page_query = page_query.limit(limit) if offset is not None: page_query = page_query.offset(offset) async with db_session(session=session) as _session: items, count = await asyncio.gather(_session.execute(page_query), cls.count(query)) return items.scalars().all(), count @classmethod async def list( cls, query=None, session: Session = None, filters=None, order_by=None, limit=None, offset=None, **params ): query = cls.q_filter(query, filters=filters, **params) if order_by is None: order_by = cls.created_at.desc() if isinstance(order_by, Iterable): query = query.order_by(*order_by) else: query = query.order_by(order_by) async with db_session(session=session) as _session: if limit is not None or offset is not None: return await cls.paginate(query, session=_session, limit=limit, offset=offset) return (await _session.execute(query)).scalars().all(), None @classmethod async def create_bulk(cls, values: Iterable[Dict[str, Any]], query=None, session: Session = None): if query is None: query = insert(cls) query = query.values(values).returning(cls) async with db_session(session=session) as _session: return (await _session.execute(query)).all() @classmethod async def update_bulk( cls, values: Iterable[Dict[str, Any]], bind_values=None, filters=None, bind_filters=None, query=None, session: Session = None, ): # Pay attention: every dict from 'values' should contain all the keys, mentioned in the 'bind_values' & # 'bind_filters'. So sometimes you firstly have to select updated records for filling absent data in the values. def get_binds(binds_mapping, encode=False): binds = {} if binds_mapping: if isinstance(binds_mapping, dict): binds = {k: bindparam(v) for k, v in binds_mapping.items()} else: # bind names can not be equal to columns names, so we have to encode them binds = {v: bindparam(f"_{v}" if encode else v) for v in binds_mapping} return binds if query is None: query = update(cls) q = cls.q_filter(query, filters=filters, **get_binds(bind_filters, encode=True)) q = q.values(**get_binds(bind_values)) async with db_session(session=session) as _session: return await _session.execute(q, values) class HiddenDeleteMixin(BulkCRUDMixin): get_or_activate: bool = False is_active = Column(Boolean, nullable=False, default=True) @classmethod def q_filter(cls, query=None, filters=None, _active_only: bool = True, **by_params) -> ClauseElement: if _active_only: by_params["is_active"] = True return super().q_filter(query=query, filters=filters, **by_params) @classmethod async def get(cls, session: Session = None, **params): query = cls.q_filter(**params, _active_only=not cls.get_or_activate) async with db_session(session=session) as _session: if cls.get_or_activate: query = query.order_by(cls.is_active.desc(), cls.id).limit(1) result = (await _session.execute(query)).fetchone() if result: result = result[0] if cls.get_or_activate: if not result.is_active: result.is_active = True await cls.update({"is_active": True}, id=result.id) logger.warning(f"{cls.__name__} id={result.id} reactivated") return result @classmethod async def delete( cls, filters=None, session: Session = None, force_delete: bool = False, return_source: Union[Column, DeclarativeMeta] = None, return_single_column: bool = False, limit: Optional[int] = None, **params, ): if force_delete: return await super().delete( filters=filters, session=session, return_source=return_source, return_single_column=return_single_column, limit=limit, _active_only=False, **params, ) async with db_session(session=session) as _session: result = await cls.update( {"is_active": False}, filters=filters, session=_session, return_source=return_source, return_single_column=return_single_column, limit=limit, **params, ) return result if return_source else bool(result.rowcount) class Model(BulkCRUDMixin, TimestampMixin): pass Base = declarative_base(metadata=metadata, cls=Model)