From bc258efa3ab0b7cb64ef5f640253825458e1c18d Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Michal=20Kv=C3=A1=C4=8Dek?= Date: Sat, 14 Oct 2023 22:42:51 +0200 Subject: [PATCH] Refaktoring --- src/database/query_builder.py | 31 +++++ src/graphql_schema/dataloaders/base.py | 63 +++++----- .../dataloaders/multi_models.py | 9 ++ src/graphql_schema/entities/event.py | 31 +++-- src/graphql_schema/entities/flight.py | 2 +- src/graphql_schema/entities/photo.py | 79 ++---------- src/graphql_schema/entities/poi.py | 2 +- src/graphql_schema/entities/resolvers/base.py | 117 ++++++++---------- .../entities/resolvers/event.py | 33 +++++ .../entities/resolvers/flight.py | 2 +- .../entities/resolvers/photo.py | 40 ++++++ .../entities/types/mutation_input.py | 23 ++++ src/graphql_schema/entities/types/types.py | 15 ++- src/graphql_schema/mutation.py | 6 +- .../sqlalchemy_to_strawberry_type.py | 3 +- src/main.py | 2 +- 16 files changed, 274 insertions(+), 184 deletions(-) create mode 100644 src/database/query_builder.py create mode 100644 src/graphql_schema/entities/resolvers/event.py create mode 100644 src/graphql_schema/entities/resolvers/photo.py diff --git a/src/database/query_builder.py b/src/database/query_builder.py new file mode 100644 index 0000000..cf39aec --- /dev/null +++ b/src/database/query_builder.py @@ -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 diff --git a/src/graphql_schema/dataloaders/base.py b/src/graphql_schema/dataloaders/base.py index ac2a3e1..7360e1f 100644 --- a/src/graphql_schema/dataloaders/base.py +++ b/src/graphql_schema/dataloaders/base.py @@ -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: diff --git a/src/graphql_schema/dataloaders/multi_models.py b/src/graphql_schema/dataloaders/multi_models.py index 234a211..bd290b8 100644 --- a/src/graphql_schema/dataloaders/multi_models.py +++ b/src/graphql_schema/dataloaders/multi_models.py @@ -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, diff --git a/src/graphql_schema/entities/event.py b/src/graphql_schema/entities/event.py index f643df6..9caf4cb 100644 --- a/src/graphql_schema/entities/event.py +++ b/src/graphql_schema/entities/event.py @@ -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()) diff --git a/src/graphql_schema/entities/flight.py b/src/graphql_schema/entities/flight.py index 04cbbff..11bee09 100644 --- a/src/graphql_schema/entities/flight.py +++ b/src/graphql_schema/entities/flight.py @@ -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, diff --git a/src/graphql_schema/entities/photo.py b/src/graphql_schema/entities/photo.py index 0a0d2d2..15070cb 100644 --- a/src/graphql_schema/entities/photo.py +++ b/src/graphql_schema/entities/photo.py @@ -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) diff --git a/src/graphql_schema/entities/poi.py b/src/graphql_schema/entities/poi.py index a10cc5f..3201a6c 100644 --- a/src/graphql_schema/entities/poi.py +++ b/src/graphql_schema/entities/poi.py @@ -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: diff --git a/src/graphql_schema/entities/resolvers/base.py b/src/graphql_schema/entities/resolvers/base.py index e10c883..6e8fa71 100644 --- a/src/graphql_schema/entities/resolvers/base.py +++ b/src/graphql_schema/entities/resolvers/base.py @@ -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) diff --git a/src/graphql_schema/entities/resolvers/event.py b/src/graphql_schema/entities/resolvers/event.py new file mode 100644 index 0000000..2c195ad --- /dev/null +++ b/src/graphql_schema/entities/resolvers/event.py @@ -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 diff --git a/src/graphql_schema/entities/resolvers/flight.py b/src/graphql_schema/entities/resolvers/flight.py index 9fd430d..99fde48 100644 --- a/src/graphql_schema/entities/resolvers/flight.py +++ b/src/graphql_schema/entities/resolvers/flight.py @@ -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'): diff --git a/src/graphql_schema/entities/resolvers/photo.py b/src/graphql_schema/entities/resolvers/photo.py new file mode 100644 index 0000000..4089ab8 --- /dev/null +++ b/src/graphql_schema/entities/resolvers/photo.py @@ -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) diff --git a/src/graphql_schema/entities/types/mutation_input.py b/src/graphql_schema/entities/types/mutation_input.py index 2bee532..75ba01b 100644 --- a/src/graphql_schema/entities/types/mutation_input.py +++ b/src/graphql_schema/entities/types/mutation_input.py @@ -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" diff --git a/src/graphql_schema/entities/types/types.py b/src/graphql_schema/entities/types/types.py index 3ad3179..8b44e6c 100644 --- a/src/graphql_schema/entities/types/types.py +++ b/src/graphql_schema/entities/types/types.py @@ -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) diff --git a/src/graphql_schema/mutation.py b/src/graphql_schema/mutation.py index 2afe2cb..9ea061f 100644 --- a/src/graphql_schema/mutation.py +++ b/src/graphql_schema/mutation.py @@ -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, diff --git a/src/graphql_schema/sqlalchemy_to_strawberry_type.py b/src/graphql_schema/sqlalchemy_to_strawberry_type.py index 727c043..73aed31 100644 --- a/src/graphql_schema/sqlalchemy_to_strawberry_type.py +++ b/src/graphql_schema/sqlalchemy_to_strawberry_type.py @@ -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_ diff --git a/src/main.py b/src/main.py index 3f26017..aa33410 100644 --- a/src/main.py +++ b/src/main.py @@ -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))