134 lines
4.8 KiB
Python
134 lines
4.8 KiB
Python
from typing import Optional, Type, TypeVar, Generic, List
|
|
|
|
from sqlalchemy import or_
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
from database import models
|
|
from database.query_builder import QueryBuilder
|
|
from database.transaction import get_session
|
|
from graphql_schema.entities.types.base import BaseGraphqlInputType
|
|
|
|
GQL_TYPE = TypeVar('GQL_TYPE')
|
|
|
|
|
|
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: int | None = None,
|
|
object_id: int | None = None,
|
|
order_by: list | None = None,
|
|
only_public: bool | None = False,
|
|
only_my: bool | None = False,
|
|
include_others_public: bool | None = False,
|
|
url_slug: str | None = None,
|
|
filters: list | None = None,
|
|
**kwargs,
|
|
):
|
|
query = self.query_builder.get_simple_query(
|
|
created_by_id=user_id,
|
|
order_by=order_by,
|
|
only_public=only_public,
|
|
include_others_public=include_others_public,
|
|
only_my=only_my,
|
|
url_slug=url_slug
|
|
)
|
|
|
|
if object_id:
|
|
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)
|
|
|
|
query = self._handle_search(query, kwargs.pop("search", None))
|
|
|
|
if filters:
|
|
query = query.filter(*filters)
|
|
|
|
return query
|
|
|
|
def _handle_search(self, query, search):
|
|
if not search:
|
|
return query
|
|
|
|
search_clauses = []
|
|
if hasattr(self.model, "name"):
|
|
search_clauses.append(self.model.name.contains(search))
|
|
|
|
if hasattr(self.model, "description"):
|
|
search_clauses.append(self.model.description.contains(search))
|
|
|
|
if search_clauses:
|
|
query = query.filter(or_(*search_clauses))
|
|
return query
|
|
|
|
async def get_list(self, user_id: int | None = None, **kwargs) -> List[GQL_TYPE]:
|
|
query = self.get_query(user_id=user_id, **kwargs)
|
|
return await self._get_list(query)
|
|
|
|
async def get_one(self, user_id: int | None = None, **kwargs) -> GQL_TYPE:
|
|
query = self.get_query(user_id=user_id, **kwargs)
|
|
return await self._get_one(query)
|
|
|
|
|
|
class BaseMutationResolver(BaseResolver):
|
|
async def _get_one(self, db: AsyncSession, id: int, created_by_id: int):
|
|
query = self.query_builder.get_simple_query(created_by_id=created_by_id).filter(self.model.id == id)
|
|
return (await db.scalars(query)).one()
|
|
|
|
async def _do_create(self, db: AsyncSession, data: dict) -> GQL_TYPE:
|
|
model = await self.model.create(db, data=data)
|
|
return self.graphql_type(**model.as_dict())
|
|
|
|
async def _do_update(self, db: AsyncSession, obj: models.BaseModel | dict | int, data: dict) -> GQL_TYPE:
|
|
update_where = {}
|
|
if isinstance(obj, models.BaseModel):
|
|
update_where['obj'] = obj
|
|
elif isinstance(obj, dict):
|
|
update_where['id'] = obj['id']
|
|
else:
|
|
update_where['id'] = obj
|
|
|
|
model = await self.model.update(db, data=data, **update_where)
|
|
return self.graphql_type(**model.as_dict())
|
|
|
|
async def create(self, context, data: BaseGraphqlInputType) -> GQL_TYPE:
|
|
input_data = data.to_dict()
|
|
|
|
if hasattr(self.model, "created_by_id"):
|
|
input_data['created_by_id'] = context.user_id
|
|
|
|
async with get_session() as db:
|
|
return await self._do_create(db, input_data)
|
|
|
|
async def update(self, context, id: int, data: BaseGraphqlInputType, user_id: int) -> GQL_TYPE:
|
|
async with get_session() as db:
|
|
item = await self._get_one(db, id, user_id)
|
|
return await self._do_update(db, item, data.to_dict())
|
|
|
|
async def delete(self, context, id: int, **kwargs) -> GQL_TYPE:
|
|
async with get_session() as db:
|
|
model = await self._get_one(db, id, context.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())
|