Refaktoring

This commit is contained in:
Michal Kváček
2023-10-14 22:42:51 +02:00
parent 7f7f79dad6
commit bc258efa3a
16 changed files with 274 additions and 184 deletions
+35 -28
View File
@@ -1,72 +1,79 @@
from collections import defaultdict
from sqlalchemy import select
from typing import Type, List, Optional
from database import models, async_session
from database.query_builder import QueryBuilder
class SingleModelByIdDataloader:
class BaseDataloader:
def __init__(
self, model: Type[models.BaseModel], relationship_column=None, filters: Optional[list] = None
self,
model: Type[models.BaseModel],
relationship_column, filters: Optional[list] = None
) -> None:
super().__init__()
self.model = model
self.relationship_column = relationship_column if relationship_column else model.id
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 = select(self.model, self.relationship_column).filter(self.relationship_column.in_(ids))
for filter_ in self.filters:
query = query.filter(filter_)
query = (
self.query_builder.get_simple_query(extra_select=[self.relationship_column], include_deleted=True)
.filter(self.relationship_column.in_(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:
class MultiModelsDataloader(BaseDataloader):
def __init__(
self,
model: Type[models.BaseModel],
relationship_column,
extra_join: Optional[list] = None,
relationship_column=None,
filters: 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 = []
if filters is None:
filters = []
self.extra_join = extra_join
if order_by is None:
order_by = [model.id.desc()] # defaultne radit od nejnovejsich zaznamu
self.model = model
self.relationship_column = relationship_column
self.extra_join = extra_join
self.order_by = order_by
self.filters = filters
async def load(self, ids: List[int]):
async with async_session() as session:
rel_column = self.relationship_column
async with async_session() as db:
query = (
select(self.model, rel_column)
.filter(rel_column.in_(ids))
.order_by(*self.order_by)
self.query_builder.get_simple_query(
extra_select=[self.relationship_column],
order_by=self.order_by
)
.filter(self.relationship_column.in_(ids))
.filter(*self.filters)
)
for table in self.extra_join:
query = query.join(table)
for joined_table in self.extra_join:
query = query.join(joined_table)
for filter_ in self.filters:
query = query.filter(filter_)
if self.filters:
query = query.filter(*self.filters)
data = (await session.execute(query)).all()
data = (await db.execute(query)).all()
result_data = defaultdict(list)
for item, rel_id in data: