diff --git a/src/database/models.py b/src/database/models.py index e8fa010..395496d 100644 --- a/src/database/models.py +++ b/src/database/models.py @@ -99,6 +99,7 @@ class PointOfInterest(BaseModel): created_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now()) deleted: Mapped[bool] = mapped_column(Boolean, nullable=False, server_default='0') + photos: Mapped[List[Photo]] = relationship() type: Mapped[PointOfInterestType] = relationship() created_by: Mapped['User'] = relationship() diff --git a/src/graphql_schema/dataloaders/photos.py b/src/graphql_schema/dataloaders/photos.py index 5aa6560..62c5752 100644 --- a/src/graphql_schema/dataloaders/photos.py +++ b/src/graphql_schema/dataloaders/photos.py @@ -1,23 +1,32 @@ from collections import defaultdict -from typing import List +from typing import List, Literal from sqlalchemy import select from strawberry.dataloader import DataLoader from database import async_session from database.models import Photo -async def load_collection(ids: List[int]): - async with async_session() as session: - models = (await session.scalars(select(Photo).filter(Photo.flight_id.in_(ids)))).all() +class PhotoDataloader: - photos_by_flight_id = defaultdict(list) - for photo in models: - photos_by_flight_id[photo.flight_id].append(photo) + def __init__(self, relationship_column: Literal['flight_id', 'point_of_interest_id']) -> None: + super().__init__() + self.relationship_column = relationship_column - return [photos_by_flight_id[id_] for id_ in ids] + async def load_collection(self, ids: List[int]): + async with async_session() as session: + models = (await session.scalars( + select(Photo) + .filter(getattr(Photo, self.relationship_column).in_(ids)) + )).all() + + photos_by_relationship_id = defaultdict(list) + for photo in models: + photos_by_relationship_id[getattr(photo, self.relationship_column)].append(photo) + + return [photos_by_relationship_id[id_] for id_ in ids] -async def load(ids: List[int]): +async def flight_cover_photo_load(ids: List[int]): async with async_session() as session: models = ( await session.scalars( @@ -25,10 +34,10 @@ async def load(ids: List[int]): .filter(Photo.is_flight_cover.is_(True)) .filter(Photo.flight_id.in_(ids))) ).all() - photos = {p.id: p for p in models} + photos = {p.flight_id: p for p in models} return [photos.get(id_) for id_ in ids] -photos_dataloader = DataLoader(load_fn=load_collection, cache=False) - -cover_photo_loader = DataLoader(load_fn=load, cache=False) +cover_photo_loader = DataLoader(load_fn=flight_cover_photo_load, cache=False) +photos_dataloader = DataLoader(load_fn=PhotoDataloader("flight_id").load_collection, cache=False) +poi_photos_dataloader = DataLoader(load_fn=PhotoDataloader("point_of_interest_id").load_collection, cache=False) diff --git a/src/graphql_schema/entities/aircraft.py b/src/graphql_schema/entities/aircraft.py index b43458d..369be15 100644 --- a/src/graphql_schema/entities/aircraft.py +++ b/src/graphql_schema/entities/aircraft.py @@ -11,7 +11,9 @@ AIRCRAFT_UPLOAD_DEST_PATH = "/app/uploads/aircrafts/" @strawberry_sqlalchemy_type(models.Aircraft) class Aircraft: - photo_url: Optional[str] = strawberry.field(resolver=lambda root: get_public_url(root.photo_filename)) + photo_url: Optional[str] = strawberry.field( + resolver=lambda root: get_public_url(f"aircrafts/{root.photo_filename}") if root.photo_filename else None + ) def get_base_query(user_id: int): diff --git a/src/graphql_schema/entities/flight.py b/src/graphql_schema/entities/flight.py index f055793..9f709ce 100644 --- a/src/graphql_schema/entities/flight.py +++ b/src/graphql_schema/entities/flight.py @@ -88,6 +88,7 @@ def get_base_query(user_id: int): return ( select(models.Flight) .filter(models.Flight.created_by_id == user_id) + .filter(models.Flight.deleted.is_(False)) .order_by(models.Flight.id.desc()) ) @@ -178,7 +179,7 @@ async def handle_copilot_edit(db: AsyncSession, copilot: CopilotInput, user_id: @strawberry.type class EditFlightMutation: - @strawberry_sqlalchemy_input(models.Flight, exclude_fields=["id", "copilot_id"], all_optional=True) + @strawberry_sqlalchemy_input(models.Flight, exclude_fields=["id", "copilot_id", "deleted"], all_optional=True) class EditFlightInput: track: Optional[List[PointOfInterestInput]] = None copilot: Optional[CopilotInput] = None @@ -196,3 +197,19 @@ class EditFlightMutation: flight.copilot_id = await handle_copilot_edit(info.context.db, input.copilot, info.context.user_id) return flight + +@strawberry.type +class DeleteFlightMutation: + + @strawberry.mutation + async def delete_flight(self, info, id: int) -> Flight: + flight = ( + (await info.context.db.scalars( + get_base_query(info.context.user_id) + .filter(models.Flight.id == id)) + ) + .one() + ) + flight.deleted = True + + return flight \ No newline at end of file diff --git a/src/graphql_schema/entities/photo.py b/src/graphql_schema/entities/photo.py index 9cedb16..38e852c 100644 --- a/src/graphql_schema/entities/photo.py +++ b/src/graphql_schema/entities/photo.py @@ -1,6 +1,6 @@ from typing import List import strawberry -from sqlalchemy import select +from sqlalchemy import select, update from strawberry.file_uploads import Upload from database import models from graphql_schema.sqlalchemy_to_strawberry_type import strawberry_sqlalchemy_type, strawberry_sqlalchemy_input @@ -49,11 +49,13 @@ class UploadPhotoMutation: # todo: udelat nahled do thumbs slozky + is_flight_cover = False # TODO: pokud k letu neexistuje zadna fotka, vybrat nahodne jednu a tu nastavit jako cover created_photo = await models.Photo.create(data={ "flight_id": input.flight_id, "name": input.name, "filename": filename, "description": input.description, + "is_flight_cover": is_flight_cover, "created_by_id": info.context.user_id, }, db_session=info.context.db) @@ -71,7 +73,18 @@ class EditPhotoMutation: query = get_base_query(info.context.user_id) photo = (await info.context.db.scalars(query.filter(models.Photo.id == id))).one() - return await models.Photo.update(info.context.db, obj=photo, data=input.to_dict()) + updated_model = await models.Photo.update(info.context.db, obj=photo, data=input.to_dict()) + + if input.is_flight_cover: + # reset other covers + (await info.context.db.execute( + update(models.Photo) + .filter(models.Photo.flight_id == photo.flight_id) + .filter(models.Photo.id != id).values(is_flight_cover=False)) + + ) + + return updated_model @strawberry.type diff --git a/src/graphql_schema/entities/poi.py b/src/graphql_schema/entities/poi.py index 81aea51..eb25205 100644 --- a/src/graphql_schema/entities/poi.py +++ b/src/graphql_schema/entities/poi.py @@ -1,23 +1,36 @@ -from typing import List, Optional +from typing import List import strawberry -from strawberry.file_uploads import Upload -from sqlalchemy import select +from sqlalchemy import select, or_ from database import models +from graphql_schema.dataloaders.photos import poi_photos_dataloader +from graphql_schema.entities.photo import Photo from graphql_schema.sqlalchemy_to_strawberry_type import strawberry_sqlalchemy_type, strawberry_sqlalchemy_input @strawberry_sqlalchemy_type(models.PointOfInterest) class PointOfInterest: - pass + async def load_photos(root): + return await poi_photos_dataloader.load(root.id) + + photos: List[Photo] = strawberry.field(resolver=load_photos) -def get_base_query(user_id: int): - return ( +def get_base_query(user_id: int, only_my: bool = False): + query = ( select(models.PointOfInterest) - .filter(models.PointOfInterest.created_by_id == user_id) .filter(models.PointOfInterest.deleted.is_(False)) ) + if only_my: + query = query.filter(models.PointOfInterest.created_by_id == user_id) + else: + query = query.filter(or_( + models.PointOfInterest.created_by_id == user_id, + models.PointOfInterest.is_public.is_(True) + )) + + return query + @strawberry.type class PointOfInterestQueries: @@ -48,8 +61,6 @@ class CreatePointOfInterestMutation: @strawberry.mutation async def create_point_of_interest(root, info, input: CreatePointOfInterestInput) -> PointOfInterest: - # TODO: kontrola organizace - input_data = input.to_dict() return await models.PointOfInterest.create( info.context.db, @@ -62,16 +73,16 @@ class CreatePointOfInterestMutation: @strawberry.type class EditPointOfInterestMutation: - @strawberry_sqlalchemy_input(models.PointOfInterest, exclude_fields=['photo_filename']) + @strawberry_sqlalchemy_input(models.PointOfInterest, exclude_fields=['id']) class EditPointOfInterestInput: - photo: Optional[Upload] + pass @strawberry.mutation - async def edit_PointOfInterest(root, info, id: int, input: EditPointOfInterestInput) -> PointOfInterest: + async def edit_point_of_interest(root, info, id: int, input: EditPointOfInterestInput) -> PointOfInterest: # TODO: kontrola organizace # TODO: kontrola opravneni na akci - poi = await models.PointOfInterest.get_one(info.context.db, id) + poi = (await get_base_query(info.context.user_id, only_my=True).filter(models.Photo.id == id)).one() return await models.PointOfInterest.update(info.context.db, obj=poi, data=input.to_dict()) @@ -80,6 +91,6 @@ class DeletePointOfInterestMutation: @strawberry.mutation async def delete_point_of_interest(self, info, id: int) -> PointOfInterest: - # TODO: kontrola opravneni na akci + poi = (await get_base_query(info.context.user_id, only_my=True).filter(models.Photo.id == id)).one() - return await models.PointOfInterest.update(info.context.db, id=id, data=dict(deleted=True)) + return await models.PointOfInterest.update(info.context.db, obj=poi, data=dict(deleted=True)) diff --git a/src/graphql_schema/mutation.py b/src/graphql_schema/mutation.py index d786941..884e934 100644 --- a/src/graphql_schema/mutation.py +++ b/src/graphql_schema/mutation.py @@ -1,15 +1,19 @@ from strawberry.tools import merge_types from graphql_schema.entities.aircraft import CreateAircraftMutation, EditAircraftMutation, DeleteAircraftMutation -from graphql_schema.entities.flight import CreateFlightMutation, EditFlightMutation +from graphql_schema.entities.flight import CreateFlightMutation, EditFlightMutation, DeleteFlightMutation from graphql_schema.entities.photo import UploadPhotoMutation, DeletePhotoMutation, EditPhotoMutation +from graphql_schema.entities.poi import CreatePointOfInterestMutation, EditPointOfInterestMutation Mutation = merge_types("Mutation", ( CreateAircraftMutation, EditAircraftMutation, DeleteAircraftMutation, EditFlightMutation, + DeleteFlightMutation, CreateFlightMutation, UploadPhotoMutation, EditPhotoMutation, - DeletePhotoMutation + DeletePhotoMutation, + CreatePointOfInterestMutation, + EditPointOfInterestMutation, ))