diff --git a/alembic/versions/20230711-074214_make_poi_type_nullable_804e55bbf855.py b/alembic/versions/20230711-074214_make_poi_type_nullable_804e55bbf855.py new file mode 100644 index 0000000..38873cd --- /dev/null +++ b/alembic/versions/20230711-074214_make_poi_type_nullable_804e55bbf855.py @@ -0,0 +1,32 @@ +"""make poi type nullable + +Revision ID: 804e55bbf855 +Revises: f19757a615cc +Create Date: 2023-07-11 07:42:14.437885 + +""" +from alembic import op +import sqlalchemy as sa +from sqlalchemy.dialects import mysql + +# revision identifiers, used by Alembic. +revision = '804e55bbf855' +down_revision = 'f19757a615cc' +branch_labels = None +depends_on = None + + +def upgrade() -> None: + # ### commands auto generated by Alembic - please adjust! ### + op.alter_column('point_of_interest', 'type_id', + existing_type=mysql.INTEGER(display_width=11), + nullable=True) + # ### end Alembic commands ### + + +def downgrade() -> None: + # ### commands auto generated by Alembic - please adjust! ### + op.alter_column('point_of_interest', 'type_id', + existing_type=mysql.INTEGER(display_width=11), + nullable=False) + # ### end Alembic commands ### diff --git a/requirements.txt b/requirements.txt index 26789dc..584d78d 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,8 +1,10 @@ -fastapi -strawberry-graphql[fastapi] -uvicorn +fastapi==0.100.0 +fastapi-jwt==0.1.12 +strawberry-graphql[fastapi]==0.194.4 +uvicorn==0.22.0 sqlalchemy[asyncio] >= 2.0.9 -aiomysql -fastapi-jwt -alembic -passlib +aiomysql==0.2.0 +alembic==1.11.1 +passlib==1.7.4 +pydantic==1.10.11 # vysla uz 2.0, ale nejak mi to nefunguje + diff --git a/src/database/models.py b/src/database/models.py index bea41c0..e8fa010 100644 --- a/src/database/models.py +++ b/src/database/models.py @@ -43,7 +43,7 @@ class BaseModel: if getattr(obj, key) != value: setattr(obj, key, value) - await db_session.commit() + # await db_session.commit() return obj @@ -93,7 +93,7 @@ class PointOfInterest(BaseModel): name: Mapped[str] = mapped_column(String(128), nullable=False) gps_latitude: Mapped[float] = mapped_column(Float, nullable=True) gps_longitude: Mapped[float] = mapped_column(Float, nullable=True) - type_id: Mapped[int] = mapped_column(Integer, ForeignKey("point_of_interest_type.id")) + type_id: Mapped[int] = mapped_column(Integer, ForeignKey("point_of_interest_type.id"), nullable=True) is_public: Mapped[bool] = mapped_column(Boolean, server_default='0') created_by_id: Mapped[int] = mapped_column(Integer, ForeignKey('user.id')) created_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now()) @@ -208,6 +208,7 @@ class Flight(BaseModel): takeoff_airport: Mapped['Airport'] = relationship(foreign_keys=[takeoff_airport_id]) landing_airport: Mapped['Airport'] = relationship(foreign_keys=[landing_airport_id]) + track: Mapped['FlightTrack'] = relationship() copilot: Mapped['Copilot'] = relationship(back_populates="flights") aircraft: Mapped['Aircraft'] = relationship(back_populates="flights") photos: Mapped[List['Photo']] = relationship(foreign_keys=[Photo.flight_id]) diff --git a/src/graphql_schema/dataloaders/aircraft.py b/src/graphql_schema/dataloaders/aircraft.py index 6bb1152..cf18b72 100644 --- a/src/graphql_schema/dataloaders/aircraft.py +++ b/src/graphql_schema/dataloaders/aircraft.py @@ -13,4 +13,4 @@ async def load(ids: List[int]): return [models_by_id.get(id_) for id_ in ids] -aircraft_dataloader = DataLoader(load_fn=load) +aircraft_dataloader = DataLoader(load_fn=load, cache=False) diff --git a/src/graphql_schema/dataloaders/airport.py b/src/graphql_schema/dataloaders/airport.py index 052279e..a98d0da 100644 --- a/src/graphql_schema/dataloaders/airport.py +++ b/src/graphql_schema/dataloaders/airport.py @@ -13,4 +13,4 @@ async def load(ids: List[int]): return [models_by_id.get(id_) for id_ in ids] -airport_dataloader = DataLoader(load_fn=load) +airport_dataloader = DataLoader(load_fn=load, cache=False) diff --git a/src/graphql_schema/dataloaders/copilots.py b/src/graphql_schema/dataloaders/copilots.py index 74121a5..ef4f1f4 100644 --- a/src/graphql_schema/dataloaders/copilots.py +++ b/src/graphql_schema/dataloaders/copilots.py @@ -13,4 +13,4 @@ async def load(ids: List[int]): return [models_by_id.get(id_) for id_ in ids] -copilots_dataloader = DataLoader(load_fn=load) +copilots_dataloader = DataLoader(load_fn=load, cache=False) diff --git a/src/graphql_schema/dataloaders/photos.py b/src/graphql_schema/dataloaders/photos.py index 0e1b316..5aa6560 100644 --- a/src/graphql_schema/dataloaders/photos.py +++ b/src/graphql_schema/dataloaders/photos.py @@ -31,4 +31,4 @@ async def load(ids: List[int]): photos_dataloader = DataLoader(load_fn=load_collection, cache=False) -cover_photo_loader = DataLoader(load_fn=load) +cover_photo_loader = DataLoader(load_fn=load, cache=False) diff --git a/src/graphql_schema/dataloaders/poi.py b/src/graphql_schema/dataloaders/poi.py new file mode 100644 index 0000000..3d8ede1 --- /dev/null +++ b/src/graphql_schema/dataloaders/poi.py @@ -0,0 +1,36 @@ +from collections import defaultdict +from typing import List +from sqlalchemy import select +from strawberry.dataloader import DataLoader +from database import async_session +from database.models import PointOfInterest, FlightTrack + + +async def load_flight_track(flight_ids: List[int]): + async with async_session() as session: + query = ( + select(FlightTrack) + .filter(FlightTrack.flight_id.in_(flight_ids)) + .order_by(FlightTrack.order) + ) + data = (await session.scalars(query)).all() + + pois_by_flight_id = defaultdict(list) + for poi in data: + pois_by_flight_id[poi.flight_id].append(poi) + + return [pois_by_flight_id[id_] for id_ in flight_ids] + + +flight_track_dataloader = DataLoader(load_fn=load_flight_track, cache=False) + + +async def load_poi(ids: List[int]): + async with async_session() as session: + models = (await session.scalars(select(PointOfInterest).filter(PointOfInterest.id.in_(ids)))).all() + + models_by_id = {model.id: model for model in models} + return [models_by_id.get(id_) for id_ in ids] + + +poi_dataloader = DataLoader(load_fn=load_poi, cache=False) diff --git a/src/graphql_schema/entities/flight.py b/src/graphql_schema/entities/flight.py index ea07afc..f42af52 100644 --- a/src/graphql_schema/entities/flight.py +++ b/src/graphql_schema/entities/flight.py @@ -1,27 +1,53 @@ from datetime import timedelta from typing import List, Optional import strawberry -from sqlalchemy import select +from sqlalchemy import select, delete +from sqlalchemy.ext.asyncio import AsyncSession from database import models from graphql_schema.dataloaders import copilots_dataloader from graphql_schema.dataloaders.aircraft import aircraft_dataloader from graphql_schema.dataloaders.airport import airport_dataloader from graphql_schema.dataloaders.photos import photos_dataloader, cover_photo_loader +from graphql_schema.dataloaders.poi import flight_track_dataloader, poi_dataloader from graphql_schema.entities.aircraft import Aircraft from graphql_schema.entities.airport import Airport from graphql_schema.entities.copilot import CopilotType from graphql_schema.entities.photo import Photo +from graphql_schema.entities.poi import PointOfInterest from graphql_schema.sqlalchemy_to_strawberry_type import strawberry_sqlalchemy_type, strawberry_sqlalchemy_input # Bude se hodit: https://strawberry.rocks/docs/types/lazy +@strawberry.input() +class PointOfInterestInput: + id: Optional[int] = None + name: str + + +@strawberry.input() +class CopilotInput: + id: Optional[int] = None + name: str + + +@strawberry_sqlalchemy_type(models.FlightTrack) +class FlightTrack: + async def load_poi(root): + return await poi_dataloader.load(root.point_of_interest_id) + + point_of_interest: PointOfInterest = strawberry.field(resolver=load_poi) + + @strawberry_sqlalchemy_type(models.Flight) class Flight: async def load_takeoff_airport(root): return await airport_dataloader.load(root.takeoff_airport_id) + async def load_track(root): + return await flight_track_dataloader.load(root.id) + async def load_landing_airport(root): return await airport_dataloader.load(root.landing_airport_id) @@ -53,6 +79,7 @@ class Flight: takeoff_airport: Airport = strawberry.field(resolver=load_takeoff_airport) landing_airport: Airport = strawberry.field(resolver=load_landing_airport) cover_photo: Optional[Photo] = strawberry.field(resolver=load_cover_photo) + track: List[FlightTrack] = strawberry.field(resolver=load_track) photos: List[Photo] = strawberry.field(resolver=load_photos) @@ -91,41 +118,79 @@ class FlightQueries: class CreateFlightMutation: @strawberry_sqlalchemy_input(models.Flight, exclude_fields=["id"]) class CreateFlightInput: - # photos: Optional[List[Upload]] pass @strawberry.mutation async def create_flight(self, info, input: CreateFlightInput) -> Flight: - input_data = input.to_dict() - - flight = await models.Flight.create(info.context.db, data={ - **input_data, - # "photos": [], # aby se nedelal select pri vytvareni fotek + return await models.Flight.create(info.context.db, data={ + **input.to_dict(), "created_by_id": info.context.user_id }) - await info.context.db.flush() - # - # if input.photos: - # photos_dest = f"/app/uploads/photos/{flight.id}/" - # for photo in input.photos: - # filename = await handle_file_upload(photo, photos_dest) - # flight.photos.append(models.Photo(**{ - # "name": "", - # "filename": filename, - # "description": "", - # "created_by_id": info.context.user_id, - # })) - return flight +async def handle_track_edit(db: AsyncSession, flight: models.Flight, track: List[PointOfInterestInput], user_id: int): + await db.execute(delete(models.FlightTrack).filter(models.FlightTrack.flight_id == flight.id)) + + existing_poi_ids = [i.id for i in track if i.id] + poi_query = ( + select(models.PointOfInterest) + .filter(models.PointOfInterest.created_by_id == user_id) + .filter(models.PointOfInterest.id.in_(existing_poi_ids)) + ) + pois = (await db.scalars(poi_query)).all() + poi_map = {poi.id: poi for poi in pois} + + order = 0 + for item in track: + poi_object = None + if item.id: + poi_object = poi_map.get(item.id) + + if not poi_object: + poi_object = await models.PointOfInterest.create(db, data=dict(created_by_id=user_id, name=item.name)) + await db.flush() + + await models.FlightTrack.create( + db, + data={ + "flight_id": flight.id, + "point_of_interest_id": poi_object.id, + "order": order + } + ) + order += 1 + + +def handle_copilot_edit(): + pass @strawberry.type class EditFlightMutation: - @strawberry_sqlalchemy_input(models.Flight, exclude_fields=["id"], all_optional=True) + @strawberry_sqlalchemy_input(models.Flight, exclude_fields=["id", "copilot_id"], all_optional=True) class EditFlightInput: - pass + track: Optional[List[PointOfInterestInput]] = None + copilot: Optional[CopilotInput] = None @strawberry.mutation async def edit_flight(self, info, id: int, input: EditFlightInput) -> Flight: - return await models.Flight.update(info.context.db, id=id, data=input.to_dict()) + flight = await models.Flight.update(info.context.db, id=id, data=input.to_dict()) + + if input.track is not None: + await handle_track_edit(db=info.context.db, flight=flight, track=input.track, user_id=info.context.user_id) + + if input.copilot is not None and not flight.solo: + if input.copilot.id: + flight.copilot_id = input.copilot.id + else: + copilot = await models.Copilot.create( + info.context.db, + data={ + "created_by_id": info.context.user_id, + "name": input.copilot.name + } + ) + info.context.db.flush() + flight.copilot_id = copilot.id + + return flight diff --git a/src/graphql_schema/entities/poi.py b/src/graphql_schema/entities/poi.py new file mode 100644 index 0000000..81aea51 --- /dev/null +++ b/src/graphql_schema/entities/poi.py @@ -0,0 +1,85 @@ +from typing import List, Optional +import strawberry +from strawberry.file_uploads import Upload +from sqlalchemy import select +from database import models +from graphql_schema.sqlalchemy_to_strawberry_type import strawberry_sqlalchemy_type, strawberry_sqlalchemy_input + + +@strawberry_sqlalchemy_type(models.PointOfInterest) +class PointOfInterest: + pass + + +def get_base_query(user_id: int): + return ( + select(models.PointOfInterest) + .filter(models.PointOfInterest.created_by_id == user_id) + .filter(models.PointOfInterest.deleted.is_(False)) + ) + + +@strawberry.type +class PointOfInterestQueries: + + @strawberry.field + async def points_of_interest(root, info) -> List[PointOfInterest]: + query = ( + get_base_query(info.context.user_id) + .order_by(models.PointOfInterest.id.desc()) + ) + + return (await info.context.db.scalars(query)).all() + + @strawberry.field + async def point_of_interest(root, info, id: int) -> PointOfInterest: + query = ( + get_base_query(info.context.user_id) + .filter(models.PointOfInterest.id == id) + ) + return (await info.context.db.scalars(query)).one() + + +@strawberry.type +class CreatePointOfInterestMutation: + @strawberry_sqlalchemy_input(models.PointOfInterest, exclude_fields=['id']) + class CreatePointOfInterestInput: + pass + + @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, + data=dict( + **input_data, + created_by_id=info.context.user_id, + ) + ) + + +@strawberry.type +class EditPointOfInterestMutation: + @strawberry_sqlalchemy_input(models.PointOfInterest, exclude_fields=['photo_filename']) + class EditPointOfInterestInput: + photo: Optional[Upload] + + @strawberry.mutation + async def edit_PointOfInterest(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) + return await models.PointOfInterest.update(info.context.db, obj=poi, data=input.to_dict()) + + +@strawberry.type +class DeletePointOfInterestMutation: + + @strawberry.mutation + async def delete_point_of_interest(self, info, id: int) -> PointOfInterest: + # TODO: kontrola opravneni na akci + + return await models.PointOfInterest.update(info.context.db, id=id, data=dict(deleted=True)) diff --git a/src/graphql_schema/query.py b/src/graphql_schema/query.py index 3f0d954..4192722 100644 --- a/src/graphql_schema/query.py +++ b/src/graphql_schema/query.py @@ -3,6 +3,7 @@ from .entities.aircraft import AircraftQueries from .entities.airport import AirportQueries from .entities.copilot import CopilotQueries from .entities.flight import FlightQueries +from .entities.poi import PointOfInterestQueries from .entities.user import UserQueries # https://github.com/strawberry-graphql/examples/blob/main/fastapi-sqlalchemy/api/schema.py @@ -13,5 +14,6 @@ Query = merge_types('Query', ( AirportQueries, FlightQueries, CopilotQueries, - UserQueries + UserQueries, + PointOfInterestQueries, ))