Files
api/src/graphql_schema/entities/resolvers/base.py
T

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())