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
+52 -65
View File
@@ -1,103 +1,72 @@
from typing import Optional, Type, T
from sqlalchemy import select, or_
from typing import Optional, Type, TypeVar, Generic, List
from sqlalchemy.ext.asyncio import AsyncSession
from database import models
from database.query_builder import QueryBuilder
from database.transaction import get_session
GQL_TYPE = TypeVar('GQL_TYPE')
class BaseQueryResolver:
def __init__(self, graphql_type, model):
class BaseResolver(Generic[GQL_TYPE]):
def __init__(self, graphql_type: GQL_TYPE, model: Type[models.BaseModel]):
self.graphql_type = graphql_type
self.model = model
self.query_builder = QueryBuilder(self.model)
class BaseQueryResolver(BaseResolver):
async def _get_list(self, query) -> List[GQL_TYPE]:
async with get_session() as db:
items = (await db.scalars(query)).all()
return [self.graphql_type(**m.as_dict()) for m in items]
async def _get_one(self, query) -> GQL_TYPE:
async with get_session() as db:
data = (await db.scalars(query)).one()
return self.graphql_type(**data.as_dict())
def get_query(
self,
user_id: Optional[int] = None,
object_id: Optional[int] = None,
order_by: Optional[list] = None,
include_public: Optional[bool] = True,
only_public: Optional[bool] = False,
*args,
**kwargs,
):
query = select(self.model)
query = self.query_builder.get_simple_query(created_by_id=user_id, order_by=order_by)
if object_id:
if hasattr(self.model, "id"):
query = query.filter(self.model.id == object_id)
else:
if not hasattr(self.model, "id"):
raise AssertionError(f"Model {self.model} has no ID column! Cannot query by ID!")
query = query.filter(self.model.id == object_id)
if hasattr(self.model, "deleted"):
query = query.filter(self.model.deleted.is_(False))
ownership_clause = []
if hasattr(self.model, "is_public") and include_public:
ownership_clause.append(self.model.is_public.is_(True))
if hasattr(self.model, "created_by_id") and user_id:
ownership_clause.append(self.model.created_by_id == user_id)
if len(ownership_clause) > 1:
query = query.filter(or_(*ownership_clause))
elif len(ownership_clause) == 1:
query = query.filter(*ownership_clause)
if order_by:
query = query.order_by(*order_by)
if only_public and hasattr(self.model, "is_public"):
query = query.filter(self.model.is_public.is_(True))
return query
async def _get_list(self, query):
async with get_session() as db:
items = (await db.scalars(query)).all()
return [self.model(**m.as_dict()) for m in items]
async def _get_one(self, query):
async with get_session() as db:
data = (await db.scalars(query)).one()
return self.model(**data.as_dict())
async def get_list(self, user_id: Optional[int] = None, **kwargs) -> list:
async def get_list(self, user_id: Optional[int] = None, **kwargs) -> List[GQL_TYPE]:
query = self.get_query(user_id, **kwargs)
return await self._get_list(query)
async def get_one(self, id: int, user_id: Optional[int] = None, **kwargs):
async def get_one(self, id: int, user_id: Optional[int] = None, **kwargs) -> GQL_TYPE:
query = self.get_query(user_id, object_id=id, **kwargs)
return await self._get_one(query)
class BaseMutationResolver:
model: Type[models.BaseModel]
graphql_type: Type[T] = None
class BaseMutationResolver(BaseResolver):
async def _get_one(self, db: AsyncSession, id: int, created_by_id: int) -> models.BaseModel:
query = self.query_builder.get_simple_query(created_by_id=created_by_id).filter(self.model.id == id)
return (await db.scalars(query)).one()
def __init__(self, graphql_type: Type[T], model: Type[models.BaseModel]):
self.graphql_type = graphql_type
self.model = model
async def delete(self, user_id: int, id: int) -> T:
async def _do_create(self, data) -> GQL_TYPE:
async with get_session() as db:
query = BaseQueryResolver(self.graphql_type, self.model).get_query(user_id, object_id=id)
model = (await db.scalars(query)).one()
if hasattr(self.model, "deleted"):
model = await self.model.update(db, obj=model, data=dict(deleted=True))
else:
await db.delete(model)
model = await self.model.create(db, data=data)
return self.graphql_type(**model.as_dict())
async def create(self, data: dict, user_id: Optional[int] = None) -> T:
input_data = {**data}
if hasattr(self.model, "created_by_id"):
input_data['created_by_id'] = user_id
async with get_session() as db:
model = await self.model.create(db, data=input_data)
return self.graphql_type(**model.as_dict())
async def _do_update(self, db: AsyncSession, obj: models.BaseModel | dict, data: dict) -> T:
async def _do_update(self, db: AsyncSession, obj: models.BaseModel | dict, data: dict) -> GQL_TYPE:
update_where = {}
if isinstance(obj, models.BaseModel):
update_where['obj'] = obj
@@ -106,3 +75,21 @@ class BaseMutationResolver:
model = await self.model.update(db, data=data, **update_where)
return self.graphql_type(**model.as_dict())
async def delete(self, user_id: int, id: int) -> GQL_TYPE:
async with get_session() as db:
model = await self._get_one(db, id, user_id)
if hasattr(self.model, "deleted"):
model = await self.model.update(db, obj=model, data=dict(deleted=True))
else:
await db.delete(model)
return self.graphql_type(**model.as_dict())
async def create(self, data: dict, user_id: Optional[int] = None) -> GQL_TYPE:
input_data = {**data}
if hasattr(self.model, "created_by_id"):
input_data['created_by_id'] = user_id
return await self._do_create(input_data)