Refaktoring
This commit is contained in:
@@ -0,0 +1,31 @@
|
||||
from typing import Optional, Type
|
||||
from sqlalchemy import select
|
||||
from database import models
|
||||
|
||||
|
||||
class QueryBuilder:
|
||||
def __init__(self, model: Type[models.BaseModel]):
|
||||
self.model = model
|
||||
|
||||
def get_simple_query(
|
||||
self,
|
||||
extra_select: Optional[list] = None,
|
||||
created_by_id: Optional[int] = None,
|
||||
order_by: Optional[list] = None,
|
||||
include_deleted: bool = False
|
||||
):
|
||||
if not extra_select:
|
||||
extra_select = []
|
||||
|
||||
query = select(self.model, *extra_select)
|
||||
|
||||
if not include_deleted and hasattr(self.model, "deleted"):
|
||||
query = query.filter(self.model.deleted.is_(False))
|
||||
|
||||
if hasattr(self.model, "created_by_id") and created_by_id:
|
||||
query = query.filter(self.model.created_by_id == created_by_id)
|
||||
|
||||
if order_by:
|
||||
query = query.order_by(*order_by)
|
||||
|
||||
return query
|
||||
@@ -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:
|
||||
|
||||
@@ -48,6 +48,15 @@ flights_by_event_dataloader = DataLoader(
|
||||
cache=False
|
||||
)
|
||||
|
||||
public_flights_by_event_dataloader = DataLoader(
|
||||
load_fn=MultiModelsDataloader(
|
||||
models.Flight,
|
||||
relationship_column=models.Event.id,
|
||||
filters=[models.Flight.is_public.is_(True)],
|
||||
extra_join=[models.Flight.event]).load,
|
||||
cache=False
|
||||
)
|
||||
|
||||
user_organizations_dataloader = DataLoader(load_fn=MultiModelsDataloader(
|
||||
models.Organization,
|
||||
relationship_column=models.user_is_in_organization.c.user_id,
|
||||
|
||||
@@ -1,9 +1,13 @@
|
||||
from typing import List
|
||||
from typing import List, Optional
|
||||
import strawberry
|
||||
from fastapi import HTTPException
|
||||
from starlette.status import HTTP_401_UNAUTHORIZED
|
||||
|
||||
from database import models
|
||||
from decorators.endpoints import authenticated_user_only
|
||||
from database.transaction import get_session
|
||||
from graphql_schema.entities.resolvers.base import BaseQueryResolver, BaseMutationResolver
|
||||
from graphql_schema.entities.resolvers.base import BaseMutationResolver
|
||||
from graphql_schema.entities.resolvers.event import EventQueryResolver
|
||||
from graphql_schema.entities.types.mutation_input import CreateEventInput, EditEventInput
|
||||
from graphql_schema.entities.types.types import Event
|
||||
|
||||
@@ -11,17 +15,26 @@ from graphql_schema.entities.types.types import Event
|
||||
@strawberry.type
|
||||
class EventQueries:
|
||||
@strawberry.field()
|
||||
@authenticated_user_only()
|
||||
async def events(root, info) -> List[Event]:
|
||||
return await BaseQueryResolver(Event, models.Event).get_list(
|
||||
async def events(root, info, username: Optional[str] = None) -> List[Event]:
|
||||
if not info.context.user_id and not username:
|
||||
raise HTTPException(HTTP_401_UNAUTHORIZED)
|
||||
|
||||
return await EventQueryResolver().get_list(
|
||||
info.context.user_id,
|
||||
username=username,
|
||||
order_by=[models.Event.date_from.desc(), models.Event.id.desc()]
|
||||
)
|
||||
|
||||
@strawberry.field()
|
||||
@authenticated_user_only()
|
||||
async def event(root, info, id: int) -> Event:
|
||||
return await BaseQueryResolver(Event, models.Event).get_one(id, user_id=info.context.user_id)
|
||||
async def event(root, info, id: int, username: Optional[str] = None) -> Event:
|
||||
if not info.context.user_id and not username:
|
||||
raise HTTPException(HTTP_401_UNAUTHORIZED)
|
||||
|
||||
return await EventQueryResolver().get_one(
|
||||
id,
|
||||
username=username,
|
||||
user_id=info.context.user_id
|
||||
)
|
||||
|
||||
|
||||
@strawberry.type
|
||||
@@ -36,7 +49,7 @@ class EventMutation:
|
||||
async def edit_event(root, info, id: int, input: EditEventInput) -> Event:
|
||||
async with get_session() as db:
|
||||
event = (await db.scalars(
|
||||
BaseQueryResolver(Event, models.Event).get_query(user_id=info.context.user_id, object_id=id)
|
||||
EventQueryResolver().get_query(user_id=info.context.user_id, object_id=id)
|
||||
)).one()
|
||||
|
||||
updated_event = await models.Event.update(db, obj=event, data=input.to_dict())
|
||||
|
||||
@@ -63,7 +63,7 @@ class FlightMutation:
|
||||
"name": "",
|
||||
"description": ""
|
||||
})
|
||||
flight = FlightMutationResolver().create(data, info.context.user_id)
|
||||
flight = await FlightMutationResolver().create(data, info.context.user_id)
|
||||
|
||||
info.context.background_tasks.add_task(
|
||||
download_weather,
|
||||
|
||||
@@ -1,21 +1,19 @@
|
||||
import asyncio
|
||||
from typing import List, Optional
|
||||
from typing import List
|
||||
import strawberry
|
||||
from sqlalchemy import update
|
||||
from strawberry.file_uploads import Upload
|
||||
from background_jobs.elevation import add_terrain_elevation_to_photo
|
||||
from background_jobs.photo import generate_thumbnail, resize_photo
|
||||
from database import models
|
||||
from decorators.endpoints import authenticated_user_only
|
||||
from database.transaction import get_session
|
||||
from graphql_schema.entities.resolvers.base import BaseQueryResolver, BaseMutationResolver
|
||||
from graphql_schema.entities.resolvers.base import BaseQueryResolver
|
||||
from graphql_schema.entities.resolvers.photo import PhotoMutationResolver
|
||||
from graphql_schema.entities.types.types import Photo
|
||||
from paths import get_photo_basepath
|
||||
from graphql_schema.entities.helpers.combobox import handle_combobox_save
|
||||
from utils.file import delete_file
|
||||
from utils.image import parse_exif_info, rotate_image
|
||||
from utils.upload import handle_file_upload
|
||||
from graphql_schema.entities.types.mutation_input import ComboboxInput
|
||||
from graphql_schema.entities.types.mutation_input import EditPhotoInput, UploadPhotoInput
|
||||
|
||||
|
||||
@strawberry.type
|
||||
@@ -26,15 +24,7 @@ class PhotoQueries:
|
||||
|
||||
|
||||
@strawberry.type
|
||||
class UploadPhotoMutation:
|
||||
@strawberry.input
|
||||
class UploadPhotoInput:
|
||||
photo: Upload
|
||||
flight_id: int
|
||||
name: Optional[str] = None
|
||||
description: Optional[str] = None
|
||||
point_of_interest: Optional[ComboboxInput] = None
|
||||
|
||||
class PhotoMutation:
|
||||
@strawberry.mutation
|
||||
@authenticated_user_only()
|
||||
async def upload_photo(self, info, input: UploadPhotoInput) -> Photo:
|
||||
@@ -42,8 +32,9 @@ class UploadPhotoMutation:
|
||||
filename = await handle_file_upload(input.photo, path)
|
||||
exif_info = await parse_exif_info(path, filename)
|
||||
|
||||
async with get_session() as db:
|
||||
photo_model = await models.Photo.create(data={
|
||||
photo = await PhotoMutationResolver().create(
|
||||
user_id=info.context.user_id,
|
||||
data={
|
||||
"flight_id": input.flight_id,
|
||||
"name": input.name,
|
||||
"filename": filename,
|
||||
@@ -53,9 +44,8 @@ class UploadPhotoMutation:
|
||||
"gps_longitude": exif_info.get("gps_longitude"),
|
||||
"gps_altitude": exif_info.get("gps_altitude"),
|
||||
"is_flight_cover": False,
|
||||
"created_by_id": info.context.user_id,
|
||||
}, db_session=db)
|
||||
photo = Photo(**photo_model.as_dict())
|
||||
},
|
||||
)
|
||||
|
||||
info.context.background_tasks.add_task(resize_photo, path=path, filename=filename)
|
||||
info.context.background_tasks.add_task(generate_thumbnail, path=path, filename=filename)
|
||||
@@ -65,52 +55,10 @@ class UploadPhotoMutation:
|
||||
|
||||
return photo
|
||||
|
||||
|
||||
@strawberry.type
|
||||
class EditPhotoMutation:
|
||||
@strawberry.input
|
||||
class EditPhotoInput:
|
||||
name: Optional[str] = None
|
||||
description: Optional[str] = None
|
||||
point_of_interest: Optional[ComboboxInput] = None
|
||||
is_flight_cover: Optional[bool] = None
|
||||
|
||||
def to_dict(self):
|
||||
return {
|
||||
key: getattr(self, key) for key in ('name', 'description', 'is_flight_cover')
|
||||
if getattr(self, key) is not None
|
||||
}
|
||||
|
||||
@strawberry.mutation()
|
||||
@authenticated_user_only()
|
||||
async def edit_photo(self, info, id: int, input: EditPhotoInput) -> Photo:
|
||||
data = input.to_dict()
|
||||
async with get_session() as db:
|
||||
photo = (await db.scalars(
|
||||
BaseQueryResolver(Photo, models.Photo).get_query(user_id=info.context.user_id, object_id=id)
|
||||
)).one()
|
||||
|
||||
if input.point_of_interest:
|
||||
data['point_of_interest_id'] = await handle_combobox_save(
|
||||
db,
|
||||
models.PointOfInterest,
|
||||
input.point_of_interest,
|
||||
info.context.user_id,
|
||||
extra_data={
|
||||
"description": ""
|
||||
}
|
||||
)
|
||||
|
||||
if input.is_flight_cover:
|
||||
# reset other covers
|
||||
(await db.execute(
|
||||
update(models.Photo)
|
||||
.filter(models.Photo.flight_id == photo.flight_id)
|
||||
.filter(models.Photo.id != id).values(is_flight_cover=False))
|
||||
)
|
||||
|
||||
updated_model = await models.Photo.update(db, obj=photo, data=data)
|
||||
return Photo(**updated_model.as_dict())
|
||||
return await PhotoMutationResolver().update(id, input, info.context.user_id)
|
||||
|
||||
@strawberry.mutation()
|
||||
@authenticated_user_only()
|
||||
@@ -137,13 +85,10 @@ class EditPhotoMutation:
|
||||
|
||||
return Photo(**photo_as_dict)
|
||||
|
||||
|
||||
@strawberry.type
|
||||
class DeletePhotoMutation:
|
||||
@strawberry.mutation()
|
||||
@authenticated_user_only()
|
||||
async def delete_photo(self, info, id: int) -> Photo:
|
||||
photo = await BaseMutationResolver(Photo, models.Photo).delete(user_id=info.context.user_id, id=id)
|
||||
photo = await PhotoMutationResolver().delete(user_id=info.context.user_id, id=id)
|
||||
|
||||
base_path = get_photo_basepath(photo.flight_id)
|
||||
delete_file(f"{base_path}/{photo.filename}", silent=True)
|
||||
|
||||
@@ -45,7 +45,7 @@ class PointOfInterestMutation:
|
||||
input_data = input.to_dict()
|
||||
|
||||
query = BaseQueryResolver(PointOfInterest, models.PointOfInterest).get_query(
|
||||
user_id=info.context.user_id, object_id=id, include_public=False
|
||||
user_id=info.context.user_id, object_id=id, only_public=False
|
||||
)
|
||||
|
||||
async with get_session() as db:
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -0,0 +1,33 @@
|
||||
from typing import Optional
|
||||
from database import models
|
||||
from graphql_schema.entities.resolvers.base import BaseQueryResolver
|
||||
from graphql_schema.entities.types.types import Event
|
||||
|
||||
|
||||
class EventQueryResolver(BaseQueryResolver):
|
||||
def __init__(self):
|
||||
super().__init__(graphql_type=Event, model=models.Event)
|
||||
|
||||
def get_query(
|
||||
self,
|
||||
user_id: Optional[int] = None,
|
||||
object_id: Optional[int] = None,
|
||||
order_by: Optional[list] = None,
|
||||
only_public: Optional[bool] = True,
|
||||
*args,
|
||||
**kwargs,
|
||||
):
|
||||
query = super().get_query(
|
||||
user_id, object_id,
|
||||
order_by=[models.Event.date_from.desc()],
|
||||
only_public=not bool(user_id)
|
||||
)
|
||||
|
||||
if kwargs.get('username'):
|
||||
query = (
|
||||
query
|
||||
.join(models.Event.created_by)
|
||||
.filter(models.User.public_username == kwargs['username'])
|
||||
)
|
||||
|
||||
return query
|
||||
@@ -25,7 +25,7 @@ class FlightQueryResolver(BaseQueryResolver):
|
||||
query = super().get_query(
|
||||
user_id, object_id,
|
||||
order_by=[models.Flight.takeoff_datetime.desc()],
|
||||
include_public=bool(user_id)
|
||||
only_public=not bool(user_id)
|
||||
)
|
||||
|
||||
if kwargs.get('username'):
|
||||
|
||||
@@ -0,0 +1,40 @@
|
||||
from sqlalchemy import update
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from database import models
|
||||
from database.transaction import get_session
|
||||
from graphql_schema.entities.helpers.combobox import handle_combobox_save
|
||||
from graphql_schema.entities.resolvers.base import BaseMutationResolver
|
||||
from graphql_schema.entities.types.mutation_input import EditPhotoInput
|
||||
from graphql_schema.entities.types.types import Photo
|
||||
|
||||
|
||||
class PhotoMutationResolver(BaseMutationResolver):
|
||||
def __init__(self):
|
||||
super().__init__(Photo, models.Photo)
|
||||
|
||||
async def reset_flight_cover(self, db: AsyncSession, flight_id: int):
|
||||
(await db.execute(
|
||||
update(models.Photo)
|
||||
.filter(models.Photo.flight_id == flight_id)
|
||||
.filter(models.Photo.id != id).values(is_flight_cover=False))
|
||||
)
|
||||
|
||||
async def update(self, id: int, input: EditPhotoInput, user_id: int) -> Photo:
|
||||
data = input.to_dict()
|
||||
|
||||
async with get_session() as db:
|
||||
photo = await self._get_one(db, id, created_by_id=user_id)
|
||||
if input.point_of_interest:
|
||||
data['point_of_interest_id'] = await handle_combobox_save(
|
||||
db,
|
||||
models.PointOfInterest,
|
||||
input.point_of_interest,
|
||||
user_id,
|
||||
extra_data={"description": ""}
|
||||
)
|
||||
|
||||
if input.is_flight_cover:
|
||||
# reset other covers
|
||||
await self.reset_flight_cover(db, photo.flight_id)
|
||||
|
||||
return await self._do_update(db, obj=photo, data=data)
|
||||
@@ -54,6 +54,29 @@ class EditEventInput(BaseGraphqlInputType):
|
||||
pass
|
||||
|
||||
|
||||
@strawberry.input
|
||||
class UploadPhotoInput:
|
||||
photo: Upload
|
||||
flight_id: int
|
||||
name: Optional[str] = None
|
||||
description: Optional[str] = None
|
||||
point_of_interest: Optional[ComboboxInput] = None
|
||||
|
||||
|
||||
@strawberry.input
|
||||
class EditPhotoInput:
|
||||
name: Optional[str] = None
|
||||
description: Optional[str] = None
|
||||
point_of_interest: Optional[ComboboxInput] = None
|
||||
is_flight_cover: Optional[bool] = None
|
||||
|
||||
def to_dict(self):
|
||||
return {
|
||||
key: getattr(self, key) for key in ('name', 'description', 'is_flight_cover')
|
||||
if getattr(self, key) is not None
|
||||
}
|
||||
|
||||
|
||||
@strawberry_sqlalchemy_input(models.Flight, exclude_fields=[
|
||||
"id", "aircraft_id", "deleted", "landing_airport_id", "takeoff_airport_id",
|
||||
"takeoff_weather_info_id", "landing_weather_info_id", "gpx_track_filename", "event_id"
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from __future__ import annotations
|
||||
from datetime import datetime
|
||||
from typing import Optional, Annotated, List
|
||||
from typing import Optional, List
|
||||
import strawberry
|
||||
from database import models
|
||||
from decorators.endpoints import authenticated_user_only
|
||||
@@ -10,7 +10,7 @@ from graphql_schema.dataloaders.multi_models import (
|
||||
poi_photos_dataloader, flight_by_poi_dataloader, flight_copilots_dataloader, flight_track_dataloader,
|
||||
photos_dataloader, flights_by_aircraft_dataloader, users_in_organization_dataloader,
|
||||
aircrafts_from_organization_dataloader, user_organizations_dataloader, flights_by_event_dataloader,
|
||||
flights_by_copilot_dataloader
|
||||
flights_by_copilot_dataloader, public_flights_by_event_dataloader
|
||||
)
|
||||
from graphql_schema.dataloaders.single_model import (
|
||||
poi_dataloader, poi_type_dataloader, event_dataloader, aircraft_dataloader, airport_dataloader, cover_photo_loader,
|
||||
@@ -161,13 +161,16 @@ class Organization:
|
||||
class User:
|
||||
avatar_image_url: Optional[str] = strawberry.field(resolver=lambda root: get_avatar_url(root))
|
||||
title_image_url: str = strawberry.field(resolver=lambda root: get_title_image_url(root))
|
||||
organizations: List[Annotated['Organization', strawberry.lazy(".organization")]] = strawberry.field(
|
||||
organizations: List[Organization] = strawberry.field(
|
||||
resolver=lambda root: user_organizations_dataloader.load(root.id)
|
||||
)
|
||||
|
||||
|
||||
@strawberry_sqlalchemy_type(models.Event)
|
||||
class Event:
|
||||
flights: List[Annotated["Flight", strawberry.lazy('.flight')]] = strawberry.field(
|
||||
resolver=lambda root: flights_by_event_dataloader.load(root.id)
|
||||
)
|
||||
async def load_flights(root, info):
|
||||
is_user_logged_in = bool(info.context.user_id)
|
||||
dataloader = flights_by_event_dataloader if is_user_logged_in else public_flights_by_event_dataloader
|
||||
return await dataloader.load(root.id)
|
||||
|
||||
flights: List[Flight] = strawberry.field(resolver=load_flights)
|
||||
|
||||
@@ -4,16 +4,14 @@ from graphql_schema.entities.copilot import CreateCopilotMutation, EditCopilotMu
|
||||
from graphql_schema.entities.event import EventMutation
|
||||
from graphql_schema.entities.flight import FlightMutation
|
||||
from graphql_schema.entities.organization import OrganizationUserMutation, OrganizationMutation
|
||||
from graphql_schema.entities.photo import UploadPhotoMutation, DeletePhotoMutation, EditPhotoMutation
|
||||
from graphql_schema.entities.photo import PhotoMutation
|
||||
from graphql_schema.entities.poi import PointOfInterestMutation
|
||||
from graphql_schema.entities.user import EditUserMutation
|
||||
|
||||
Mutation = merge_types("Mutation", (
|
||||
AircraftMutation,
|
||||
FlightMutation,
|
||||
UploadPhotoMutation,
|
||||
EditPhotoMutation,
|
||||
DeletePhotoMutation,
|
||||
PhotoMutation,
|
||||
PointOfInterestMutation,
|
||||
CreateCopilotMutation,
|
||||
EditCopilotMutation,
|
||||
|
||||
@@ -3,6 +3,7 @@ from typing import List, Optional
|
||||
import strawberry
|
||||
import sqlalchemy
|
||||
from sqlalchemy import Column
|
||||
from logger import log
|
||||
from database.models import BaseModel
|
||||
from graphql_schema.entities.types.base import BaseGraphqlInputType
|
||||
|
||||
@@ -22,7 +23,7 @@ def get_annotations_for_scalars(model: BaseModel, exclude_fields=None, force_opt
|
||||
type_ = typing.Optional[column.type.python_type] if is_optional else column.type.python_type
|
||||
annotations_[name] = type_
|
||||
except NotImplementedError as e:
|
||||
print(f"Neimplementovano: {e}, {name=}")
|
||||
log.warning(f"Cannot annotate {name} in {model} for GQL type. Exception: {e}")
|
||||
|
||||
return annotations_
|
||||
|
||||
|
||||
+1
-1
@@ -90,7 +90,7 @@ class App:
|
||||
if APP_DEBUG:
|
||||
@self.api_router.get("/graphql/autologin")
|
||||
async def autologin():
|
||||
access_token = self.access_security.create_access_token(subject={"id": 1, "name": "Franta Vomacka"})
|
||||
access_token = self.access_security.create_access_token(subject={"id": 6, "name": "Franta Vomacka"})
|
||||
|
||||
response = RedirectResponse(url="/graphql")
|
||||
self.access_security.set_access_cookie(response, access_token, expires_delta=timedelta(days=14))
|
||||
|
||||
Reference in New Issue
Block a user