122 lines
3.6 KiB
Python
122 lines
3.6 KiB
Python
from collections import defaultdict
|
|
from typing import Type
|
|
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: list | None = 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]):
|
|
ids_set = {id_ for id_ in set(ids) if id_ is not None}
|
|
if not ids_set:
|
|
return [None for _ in ids]
|
|
|
|
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_(ids_set))
|
|
.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: list | None = None,
|
|
extra_select: list | None = None,
|
|
extra_join: list | None = None,
|
|
order_by: list | None = 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: set[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_(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]) -> list:
|
|
ids_set = {id_ for id_ in set(ids) if id_ is not None}
|
|
if not ids_set:
|
|
return [[] for _ in ids]
|
|
|
|
query = self.get_query(ids_set)
|
|
|
|
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
|