Refaktoring a bugfixing

This commit is contained in:
Michal Kváček
2023-10-13 23:24:53 +02:00
parent ce2c023eaf
commit 7f7f79dad6
36 changed files with 698 additions and 722 deletions
+101 -46
View File
@@ -1,53 +1,108 @@
from typing import Optional, Type
from typing import Optional, Type, T
from sqlalchemy import select, or_
from sqlalchemy.ext.asyncio import AsyncSession
from database import models
from dependencies.db import get_session
from database.transaction import get_session
def get_base_resolver(
model: Type[models.BaseModel],
object_id: Optional[int] = None,
user_id: Optional[int] = None,
order_by: Optional[list] = None,
include_public: Optional[bool] = True,
):
query = select(model)
class BaseQueryResolver:
if object_id:
if hasattr(model, "id"):
query = query.filter(model.id == object_id)
def __init__(self, graphql_type, model):
self.graphql_type = graphql_type
self.model = model
def get_query(
self,
user_id: Optional[int] = None,
object_id: Optional[int] = None,
order_by: Optional[list] = None,
include_public: Optional[bool] = True,
*args,
**kwargs,
):
query = select(self.model)
if object_id:
if hasattr(self.model, "id"):
query = query.filter(self.model.id == object_id)
else:
raise AssertionError(f"Model {self.model} has no ID column! Cannot query by 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)
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:
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):
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
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 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)
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:
update_where = {}
if isinstance(obj, models.BaseModel):
update_where['obj'] = obj
else:
raise AssertionError(f"Model {model} has no ID column! Cannot query by ID!")
update_where['id'] = obj['id']
if hasattr(model, "deleted"):
query = query.filter(model.deleted.is_(False))
ownership_clause = []
if hasattr(model, "is_public") and include_public:
ownership_clause.append(model.is_public.is_(True))
if hasattr(model, "created_by_id") and user_id:
ownership_clause.append(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)
return query
async def get_list(model, query):
async with get_session() as db:
items = (await db.scalars(query)).all()
return [model(**m.as_dict()) for m in items]
async def get_one(model, query):
async with get_session() as db:
data = (await db.scalars(query)).one()
return model(**data.as_dict())
model = await self.model.update(db, data=data, **update_where)
return self.graphql_type(**model.as_dict())