from collections import defaultdict from typing import Type, List, Optional from logger import log from database import models, async_session from database.query_builder import QueryBuilder class BaseDataloader: def __init__( self, model: Type[models.BaseModel], relationship_column, filters: Optional[list] = None ): super().__init__() self.model = model self.query_builder = QueryBuilder(self.model) if relationship_column is None: relationship_column = model.id self.relationship_column = relationship_column if filters is None: filters = [] self.filters = filters class SingleModelByIdDataloader(BaseDataloader): async def load(self, ids: List[int]): async with async_session() as session: query = ( self.query_builder.get_simple_query(extra_select=[self.relationship_column], include_deleted=True) .filter(self.relationship_column.in_(set(ids))) .filter(*self.filters) ) items = (await session.execute(query)).all() items_by_id = {rel_id: item for item, rel_id in items} return [items_by_id.get(id_) for id_ in ids] class MultiModelsDataloader(BaseDataloader): def __init__( self, model: Type[models.BaseModel], relationship_column=None, filters: Optional[list] = None, extra_select: Optional[list] = None, extra_join: Optional[list] = None, order_by: Optional[list] = None, ): super().__init__(model, relationship_column, filters) if extra_join is None: extra_join = [] self.extra_join = extra_join if extra_select is None: extra_select = [] self.extra_select = extra_select if order_by is None: order_by = [model.id.desc()] # defaultne radit od nejnovejsich zaznamu self.order_by = order_by def get_query(self, ids: list[int]): query = ( self.query_builder.get_simple_query( extra_select=[self.relationship_column] + self.extra_select, order_by=self.order_by ) .filter(self.relationship_column.in_(set(ids))) .filter(*self.filters) ) for joined_table in self.extra_join: query = query.join(joined_table) if self.filters: query = query.filter(*self.filters) return query async def load(self, ids: List[int]): query = self.get_query(ids) async with async_session() as db: data = (await db.execute(query)).all() result_data = self.process_data(data) return [result_data[id_] for id_ in ids] def process_data(self, data): result_data = defaultdict(list) for row in data: item, rel_id = row[0:2] extra = row[2:] if extra: log.warning(f"Override function process_data, extra params={extra} are going to be discarded!") result_data[rel_id].append(item) return result_data class FlightCopilotDataloader(MultiModelsDataloader): def process_data(self, data): result_data = defaultdict(list) for row in data: item, rel_id = row[0:2] token = row[2] item.token = token result_data[rel_id].append(item) return result_data