diff --git a/alembic/versions/20230922-085320_rename_from_to_in_events_8b0c020dc0c4.py b/alembic/versions/20230922-085320_rename_from_to_in_events_8b0c020dc0c4.py new file mode 100644 index 0000000..04a777b --- /dev/null +++ b/alembic/versions/20230922-085320_rename_from_to_in_events_8b0c020dc0c4.py @@ -0,0 +1,34 @@ +"""rename from/to in events + +Revision ID: 8b0c020dc0c4 +Revises: 39a62618eacb +Create Date: 2023-09-22 08:53:20.901142 + +""" +from alembic import op +import sqlalchemy as sa +from sqlalchemy.dialects import mysql + +# revision identifiers, used by Alembic. +revision = '8b0c020dc0c4' +down_revision = '39a62618eacb' +branch_labels = None +depends_on = None + + +def upgrade() -> None: + # ### commands auto generated by Alembic - please adjust! ### + op.add_column('event', sa.Column('date_from', sa.DateTime(), nullable=True)) + op.add_column('event', sa.Column('date_to', sa.DateTime(), nullable=True)) + op.drop_column('event', 'event_from') + op.drop_column('event', 'event_to') + # ### end Alembic commands ### + + +def downgrade() -> None: + # ### commands auto generated by Alembic - please adjust! ### + op.add_column('event', sa.Column('event_to', mysql.DATETIME(), nullable=True)) + op.add_column('event', sa.Column('event_from', mysql.DATETIME(), nullable=True)) + op.drop_column('event', 'date_to') + op.drop_column('event', 'date_from') + # ### end Alembic commands ### diff --git a/src/background_jobs/flight.py b/src/background_jobs/flight.py new file mode 100644 index 0000000..3469397 --- /dev/null +++ b/src/background_jobs/flight.py @@ -0,0 +1,26 @@ +from aiohttp import ClientResponseError +from database import models +from dependencies.db import get_session +from external.elevation import elevation_api +from external.gpx_parser import GPXParser + + +async def add_terrain_elevation(flight: dict, gpx_filename: str): + path = "/app/uploads/tracks" # TODO vytahnout do configu + + gpx_parser = GPXParser(f"{path}/{gpx_filename}") + coordinates = await gpx_parser.get_coordinates() + + try: + elevation = await elevation_api.get_elevation_for_points(coordinates) + tree_with_elevation = gpx_parser.add_terrain_elevation(elevation) + output_name = f"terrain_{gpx_filename}" + gpx_parser.write(tree_with_elevation, f"{path}/{output_name}") + + async with get_session() as db: + await models.Flight.update( + db, {"gpx_track_filename": output_name, "has_terrain_elevation": True}, + id=flight['id']) + + except ClientResponseError as e: + print(e) diff --git a/src/graphql_schema/entities/aircraft.py b/src/graphql_schema/entities/aircraft.py index c79a3d1..a3ce0f7 100644 --- a/src/graphql_schema/entities/aircraft.py +++ b/src/graphql_schema/entities/aircraft.py @@ -1,13 +1,14 @@ from typing import List, Optional, Annotated, TYPE_CHECKING, Set import strawberry from strawberry.file_uploads import Upload -from sqlalchemy import select, or_ from database import models from decorators.endpoints import authenticated_user_only from dependencies.db import get_session from graphql_schema.sqlalchemy_to_strawberry_type import strawberry_sqlalchemy_type, strawberry_sqlalchemy_input from upload_utils import handle_file_upload, delete_file, get_public_url -from .helpers.flight import handle_combobox_save +from graphql_schema.entities.helpers.combobox import handle_combobox_save +from .resolvers.aircraft import get_aircraft_resolver +from .resolvers.base import get_list, get_one from ..dataloaders.flight import flights_by_aircraft_dataloader from ..dataloaders.organizations import organizations_dataloader from ..types import ComboboxInput @@ -36,43 +37,19 @@ class Aircraft: ) -def get_base_query(user_id: int, organization_ids: Set[int]): - return ( - select(models.Aircraft) - .filter( - or_( - models.Aircraft.created_by_id == user_id, - models.Aircraft.organization_id.in_(organization_ids) - ) - ) - .filter(models.Aircraft.deleted.is_(False)) - .order_by(models.Aircraft.id.desc()) - ) - - @strawberry.type class AircraftQueries: - @strawberry.field() @authenticated_user_only() async def aircrafts(root, info) -> List[Aircraft]: - async with get_session() as db: - aircrafts = (await db.scalars( - get_base_query(info.context.user_id, info.context.organization_ids) - )).all() - - return [Aircraft(**a.as_dict()) for a in aircrafts] + query = get_aircraft_resolver(info.context.user_id, info.context.organization_ids) + return await get_list(models.Aircraft, query) @strawberry.field() @authenticated_user_only() async def aircraft(root, info, id: int) -> Aircraft: - query = ( - get_base_query(info.context.user_id, info.context.organization_ids) - .filter(models.Aircraft.id == id) - ) - async with get_session() as db: - aircraft = (await db.scalars(query)).one() - return Aircraft(**aircraft.as_dict()) + query = get_aircraft_resolver(info.context.user_id, info.context.organization_ids, id) + return await get_one(models.Aircraft, query) @strawberry.type @@ -92,7 +69,6 @@ class CreateAircraftMutation: input_data['photo_filename'] = await handle_file_upload(input.photo, AIRCRAFT_UPLOAD_DEST_PATH) async with get_session() as db: - if input.organization: input_data['organization_id'] = await handle_combobox_save( db, @@ -134,10 +110,7 @@ class EditAircraftMutation: user_id=info.context.user_id, ) - aircraft = (await db.scalars( - get_base_query(info.context.user_id, set()) # TODO: bude fungovat prazdny set? - .filter(models.Aircraft.id == id) - )).one() + aircraft = (await db.scalars(get_aircraft_resolver(user_id=info.context.user_id, aircraft_id=id))).one() if input.photo: if aircraft.photo_filename: @@ -156,10 +129,7 @@ class DeleteAircraftMutation: @authenticated_user_only() async def delete_aircraft(self, info, id: int) -> Aircraft: async with get_session() as db: - aircraft = (await db.scalars( - get_base_query(info.context.user_id) - .filter(models.Aircraft.id == id) - )).one() + aircraft = (await db.scalars(get_aircraft_resolver(info.context.user_id, aircraft_id=id))).one() aircraft = await models.Aircraft.update(db, obj=aircraft, data=dict(deleted=True)) return Aircraft(**aircraft.as_dict()) diff --git a/src/graphql_schema/entities/airport.py b/src/graphql_schema/entities/airport.py index e2fec8b..5822adb 100644 --- a/src/graphql_schema/entities/airport.py +++ b/src/graphql_schema/entities/airport.py @@ -1,9 +1,9 @@ from typing import List import strawberry -from sqlalchemy import select, or_ from database import models from decorators.endpoints import authenticated_user_only -from dependencies.db import get_session +from graphql_schema.entities.resolvers.airport import get_airport_resolver +from graphql_schema.entities.resolvers.base import get_list, get_one from graphql_schema.sqlalchemy_to_strawberry_type import strawberry_sqlalchemy_type @@ -12,37 +12,16 @@ class Airport: pass -def get_base_query(user_id: int): - return ( - select(models.Airport) - .filter(models.Airport.deleted.is_(False)) - .filter(or_( - models.Airport.created_by_id == user_id, - models.Airport.created_by_id.is_(None), - )) - .order_by(models.Airport.icao_code) - ) - - @strawberry.type class AirportQueries: @strawberry.field() @authenticated_user_only() async def airports(root, info) -> List[Airport]: - query = get_base_query(info.context.user_id) - - async with get_session() as db: - airports = (await db.scalars(query)).all() - return [Airport(**a.as_dict()) for a in airports] + query = get_airport_resolver(info.context.user_id) + return await get_list(models.Airport, query) @strawberry.field() @authenticated_user_only() async def airport(root, info, id: int) -> Airport: - query = ( - get_base_query(info.context.user_id) - .filter(models.Airport.id == id) - ) - - async with get_session() as db: - airport = (await db.scalars(query)).one() - return Airport(**airport.as_dict()) + query = get_airport_resolver(info.context.user_id, id) + return await get_one(models.Airport, query) diff --git a/src/graphql_schema/entities/copilot.py b/src/graphql_schema/entities/copilot.py index 4915465..6f68b30 100644 --- a/src/graphql_schema/entities/copilot.py +++ b/src/graphql_schema/entities/copilot.py @@ -6,6 +6,7 @@ from decorators.endpoints import authenticated_user_only from dependencies.db import get_session from graphql_schema.dataloaders.flight import flights_by_copilot_dataloader from graphql_schema.sqlalchemy_to_strawberry_type import strawberry_sqlalchemy_type, strawberry_sqlalchemy_input +from .resolvers.base import get_base_resolver, get_list, get_one if TYPE_CHECKING: from .flight import Flight @@ -19,36 +20,19 @@ class Copilot: flights: List[Annotated["Flight", strawberry.lazy('.flight')]] = strawberry.field(resolver=load_flights) -def get_base_query(user_id: int): - return ( - select(models.Copilot) - .filter(models.Copilot.created_by_id == user_id) - .filter(models.Copilot.deleted.is_(False)) - .order_by(models.Copilot.name) - ) - - @strawberry.type class CopilotQueries: @strawberry.field() @authenticated_user_only() async def copilots(root, info) -> List[Copilot]: - async with get_session() as db: - copilots = (await db.scalars( - get_base_query(info.context.user_id) - )).all() - - return [Copilot(**c.as_dict()) for c in copilots] + query = get_base_resolver(models.Copilot, user_id=info.context.user_id, order_by=[models.Copilot.name]) + return await get_list(models.Copilot, query) @strawberry.field() @authenticated_user_only() async def copilot(root, info, id: int) -> Copilot: - async with get_session() as db: - copilot = (await db.scalars( - get_base_query(info.context.user_id) - .filter(models.Copilot.id == id) - )).one() - return Copilot(**copilot.as_dict()) + query = get_base_resolver(models.Copilot, object_id=id, user_id=info.context.user_id) + return await get_one(models.Copilot, query) @strawberry.type @@ -84,7 +68,7 @@ class EditCopilotMutation: async def edit_copilot(root, info, id: int, input: EditCopilotInput) -> Copilot: async with get_session() as db: copilot = (await db.scalars( - get_base_query(info.context.user_id).filter(models.Copilot.id == id) + query=get_base_resolver(models.Copilot, object_id=id, user_id=info.context.user_id) )).one() updated_copilot = await models.Copilot.update(db, obj=copilot, data=input.to_dict()) diff --git a/src/graphql_schema/entities/event.py b/src/graphql_schema/entities/event.py index fa80686..e569efe 100644 --- a/src/graphql_schema/entities/event.py +++ b/src/graphql_schema/entities/event.py @@ -6,6 +6,7 @@ from decorators.endpoints import authenticated_user_only from dependencies.db import get_session from graphql_schema.dataloaders.flight import flights_by_event_dataloader from graphql_schema.sqlalchemy_to_strawberry_type import strawberry_sqlalchemy_type, strawberry_sqlalchemy_input +from .resolvers.base import get_base_resolver, get_list, get_one if TYPE_CHECKING: from .flight import Flight @@ -19,36 +20,22 @@ class Event: flights: List[Annotated["Flight", strawberry.lazy('.flight')]] = strawberry.field(resolver=load_flights) -def get_base_query(user_id: int): - return ( - select(models.Event) - .filter(models.Event.created_by_id == user_id) - .filter(models.Event.deleted.is_(False)) - .order_by(models.Event.date_from.desc(), models.Event.id.desc()) - ) - - @strawberry.type class EventQueries: @strawberry.field() @authenticated_user_only() async def events(root, info) -> List[Event]: - async with get_session() as db: - events = (await db.scalars( - get_base_query(info.context.user_id) - )).all() - - return [Event(**c.as_dict()) for c in events] + query = get_base_resolver( + models.Event, user_id=info.context.user_id, + order_by=[models.Event.date_from.desc(), models.Event.id.desc()] + ) + return await get_list(models.Event, query) @strawberry.field() @authenticated_user_only() async def event(root, info, id: int) -> Event: - async with get_session() as db: - event = (await db.scalars( - get_base_query(info.context.user_id) - .filter(models.Event.id == id) - )).one() - return Event(**event.as_dict()) + query = get_base_resolver(models.Event, user_id=info.context.user_id, object_id=id) + return await get_one(models.Event, query) @strawberry.type @@ -84,7 +71,7 @@ class EditEventMutation: async def edit_event(root, info, id: int, input: EditEventInput) -> Event: async with get_session() as db: event = (await db.scalars( - get_base_query(info.context.user_id).filter(models.Event.id == id) + get_base_resolver(models.Event, 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/helpers/combobox.py b/src/graphql_schema/entities/helpers/combobox.py new file mode 100644 index 0000000..ec3979d --- /dev/null +++ b/src/graphql_schema/entities/helpers/combobox.py @@ -0,0 +1,28 @@ +from typing import Type, Optional +from sqlalchemy.ext.asyncio import AsyncSession +from database import models +from graphql_schema.types import ComboboxInput + + +async def handle_combobox_save( + db: AsyncSession, + model: Type[models.BaseModel], + input: ComboboxInput, + user_id: int, + name_column: str = "name", + extra_data: Optional[dict] = None +) -> int: + if input.id: + return input.id + else: + + if not extra_data: + extra_data = {} + + data = {name_column: input.name, **extra_data} + if hasattr(model, "created_by_id"): + data["created_by_id"] = user_id + + obj = await model.create(db, data) + await db.flush() + return obj.id diff --git a/src/graphql_schema/entities/helpers/flight.py b/src/graphql_schema/entities/helpers/flight.py index d38591f..c67b4f7 100644 --- a/src/graphql_schema/entities/helpers/flight.py +++ b/src/graphql_schema/entities/helpers/flight.py @@ -1,15 +1,12 @@ import asyncio from datetime import datetime from typing import List, Type, Literal, Optional, Tuple -from aiohttp import ClientResponseError from sqlalchemy import select, delete from sqlalchemy.ext.asyncio import AsyncSession from strawberry.file_uploads import Upload from database import models -from dependencies.db import get_session -from external.elevation import elevation_api -from external.gpx_parser import GPXParser from external.weather import Weather +from graphql_schema.entities.helpers.combobox import handle_combobox_save from graphql_schema.types import ComboboxInput from upload_utils import delete_file, handle_file_upload @@ -141,27 +138,6 @@ async def handle_airport_changed( setattr(flight, f"{type_}_datetime", input_datetime) -async def add_terrain_elevation(flight: dict, gpx_filename: str): - path = "/app/uploads/tracks" # TODO vytahnout do configu - - gpx_parser = GPXParser(f"{path}/{gpx_filename}") - coordinates = await gpx_parser.get_coordinates() - - try: - elevation = await elevation_api.get_elevation_for_points(coordinates) - tree_with_elevation = gpx_parser.add_terrain_elevation(elevation) - output_name = f"terrain_{gpx_filename}" - gpx_parser.write(tree_with_elevation, f"{path}/{output_name}") - - async with get_session() as db: - await models.Flight.update( - db, {"gpx_track_filename": output_name, "has_terrain_elevation": True}, - id=flight['id']) - - except ClientResponseError as e: - print(e) - - async def handle_upload_gpx(flight: models.Flight, gpx_track: Upload): path = "/app/uploads/tracks" @@ -176,25 +152,3 @@ async def handle_copilots_edit(db: AsyncSession, copilots: List[ComboboxInput], return await asyncio.gather(*cors) -async def handle_combobox_save( - db: AsyncSession, - model: Type[models.BaseModel], - input: ComboboxInput, - user_id: int, - name_column: str = "name", - extra_data: Optional[dict] = None -) -> int: - if input.id: - return input.id - else: - - if not extra_data: - extra_data = {} - - data = {name_column: input.name, **extra_data} - if hasattr(model, "created_by_id"): - data["created_by_id"] = user_id - - obj = await model.create(db, data) - await db.flush() - return obj.id diff --git a/src/graphql_schema/entities/organization.py b/src/graphql_schema/entities/organization.py index aeaf468..2032107 100644 --- a/src/graphql_schema/entities/organization.py +++ b/src/graphql_schema/entities/organization.py @@ -1,13 +1,13 @@ from typing import List, Annotated, TYPE_CHECKING import strawberry -from sqlalchemy import select, or_, delete +from sqlalchemy import select, delete from sqlalchemy.dialects.mysql import insert from sqlalchemy.exc import IntegrityError - from database import models from decorators.endpoints import authenticated_user_only from dependencies.db import get_session from graphql_schema.sqlalchemy_to_strawberry_type import strawberry_sqlalchemy_type, strawberry_sqlalchemy_input +from .resolvers.base import get_base_resolver, get_list, get_one from ..dataloaders.aircraft import aircrafts_from_organization_dataloader from ..dataloaders.users import users_in_organization_dataloader @@ -29,33 +29,19 @@ class Organization: aircrafts: List[Annotated["Aircraft", strawberry.lazy(".aircraft")]] = strawberry.field(resolver=load_aircrafts) -def get_base_query(): - return ( - select(models.Organization) - .filter(models.Organization.deleted.is_(False)) - .order_by(models.Organization.name) - ) - - @strawberry.type class OrganizationQueries: @strawberry.field() @authenticated_user_only() async def organizations(root, info) -> List[Organization]: - async with get_session() as db: - organizations = (await db.scalars(get_base_query())).all() - - return [Organization(**c.as_dict()) for c in organizations] + query = get_base_resolver(models.Organization, order_by=[models.Organization.name]) + return await get_list(models.Organization, query) @strawberry.field() @authenticated_user_only() async def organization(root, info, id: int) -> Organization: - async with get_session() as db: - organization = (await db.scalars( - get_base_query() - .filter(models.Organization.id == id) - )).one() - return Organization(**organization.as_dict()) + query = get_base_resolver(models.Organization, object_id=id) + return await get_one(models.Organization, query) @strawberry.type @@ -87,7 +73,7 @@ class OrganizationUserMutation: @authenticated_user_only() async def add_to_organization(root, info, organization_id: int) -> Organization: async with get_session() as db: - organization = (await db.scalars(get_base_query().filter(models.Organization.id == organization_id))).one() + organization = (await db.scalars(get_base_resolver(models.Organization, object_id=organization_id))).one() try: await db.execute( @@ -106,7 +92,7 @@ class OrganizationUserMutation: @authenticated_user_only() async def remove_from_organization(root, info, organization_id: int) -> Organization: async with get_session() as db: - organization = (await db.scalars(get_base_query().filter(models.Organization.id == organization_id))).one() + organization = (await db.scalars(get_base_resolver(models.Organization, object_id=organization_id))).one() await db.execute( delete(models.user_is_in_organization).filter_by( @@ -118,7 +104,6 @@ class OrganizationUserMutation: return Organization(**organization.as_dict()) - @strawberry.type class EditOrganizationMutation: @strawberry_sqlalchemy_input(model=models.Organization, exclude_fields=["id"]) @@ -130,9 +115,7 @@ class EditOrganizationMutation: async def edit_organization(root, info, id: int, input: EditOrganizationInput) -> Organization: async with get_session() as db: organization = (await db.scalars( - get_base_query() - .filter(models.Organization.created_by_id == info.context.user_id) - .filter(models.Organization.id == id) + get_base_resolver(models.Organization, user_id=info.context.user_id, object_id=id) )).one() updated_organization = await models.Organization.update(db, obj=organization, data=input.to_dict()) diff --git a/src/graphql_schema/entities/photo.py b/src/graphql_schema/entities/photo.py index 518f902..bd2deca 100644 --- a/src/graphql_schema/entities/photo.py +++ b/src/graphql_schema/entities/photo.py @@ -13,7 +13,8 @@ from graphql_schema.types import ComboboxInput from upload_utils import ( get_public_url, handle_file_upload, delete_file, parse_exif_info, generate_thumbnail, file_exists, resize_image, rotate_image ) -from .helpers.flight import handle_combobox_save +from graphql_schema.entities.helpers.combobox import handle_combobox_save +from .resolvers.base import get_base_resolver, get_list if TYPE_CHECKING: from .poi import PointOfInterest @@ -58,11 +59,8 @@ def get_photo_basepath(flight_id: int) -> str: class PhotoQueries: @strawberry.field() async def photos(root, info) -> List[Photo]: - query = get_base_query(info.context.user_id) - - async with get_session() as db: - photos = (await db.scalars(query)).all() - return [Photo(**photo.as_dict()) for photo in photos] + query = get_base_resolver(models.Photo, user_id=info.context.user_id) + return await get_list(models.Photo, query) @strawberry.type @@ -97,6 +95,7 @@ class UploadPhotoMutation: }, db_session=db) photo = Photo(**photo_model.as_dict()) + # TODO: udelat primo konkretni bg joby na resize a thumbnaily info.context.background_tasks.add_task(resize_image, path=path, filename=filename, new_width=2500, quality=85) info.context.background_tasks.add_task(generate_thumbnail, path=path, filename=filename) @@ -118,15 +117,16 @@ class EditPhotoMutation: @strawberry.mutation() @authenticated_user_only() async def edit_photo(self, info, id: int, input: EditPhotoInput) -> Photo: - query = get_base_query(info.context.user_id) + # TODO: base trida pro inputy s definici to_dict/as_dict? data = { key: getattr(input, key) for key in ('name', 'description', 'is_flight_cover') if getattr(input, key) is not None } + query = get_base_resolver(models.Photo, user_id=info.context.user_id, object_id=id) async with get_session() as db: - photo = (await db.scalars(query.filter(models.Photo.id == id))).one() + photo = (await db.scalars(query)).one() if input.point_of_interest: data['point_of_interest_id'] = await handle_combobox_save( @@ -155,8 +155,8 @@ class EditPhotoMutation: async def rotate_photo(self, info, id: int, angle: int) -> Photo: async with get_session() as db: - query = get_base_query(info.context.user_id) - photo = (await db.scalars(query.filter(models.Photo.id == id))).one() + query = get_base_resolver(models.Photo, user_id=info.context.user_id, object_id=id) + photo = (await db.scalars(query)).one() await asyncio.gather( rotate_image( diff --git a/src/graphql_schema/entities/poi.py b/src/graphql_schema/entities/poi.py index ef96874..8e9625e 100644 --- a/src/graphql_schema/entities/poi.py +++ b/src/graphql_schema/entities/poi.py @@ -7,10 +7,11 @@ from dependencies.db import get_session from graphql_schema.dataloaders.flight import flight_by_poi_dataloader from graphql_schema.dataloaders.photos import poi_photos_dataloader from graphql_schema.dataloaders.poi import poi_type_dataloader -from graphql_schema.entities.helpers.flight import handle_combobox_save +from graphql_schema.entities.helpers.combobox import handle_combobox_save from graphql_schema.entities.poi_type import PointOfInterestType from graphql_schema.sqlalchemy_to_strawberry_type import strawberry_sqlalchemy_type, strawberry_sqlalchemy_input from graphql_schema.types import ComboboxInput +from .resolvers.base import get_base_resolver, get_list, get_one if TYPE_CHECKING: from .flight import Flight @@ -51,28 +52,17 @@ def get_base_query(user_id: int, only_my: bool = False): @strawberry.type class PointOfInterestQueries: - @strawberry.field() @authenticated_user_only() async def points_of_interest(root, info) -> List[PointOfInterest]: - query = ( - get_base_query(info.context.user_id) - .order_by(models.PointOfInterest.id.desc()) - ) - async with get_session() as db: - pois = (await db.scalars(query)).all() - return [PointOfInterest(**poi.as_dict()) for poi in pois] + query = get_base_resolver(models.PointOfInterest, user_id=info.context.user_id) + return await get_list(models.PointOfInterest, query) @strawberry.field() @authenticated_user_only() async def point_of_interest(root, info, id: int) -> PointOfInterest: - query = ( - get_base_query(info.context.user_id) - .filter(models.PointOfInterest.id == id) - ) - async with get_session() as db: - poi = (await db.scalars(query)).one() - return PointOfInterest(**poi.as_dict()) + query = get_base_resolver(models.PointOfInterest, user_id=info.context.user_id, object_id=id) + return await get_one(models.PointOfInterest, query) @strawberry.type @@ -105,20 +95,19 @@ class EditPointOfInterestMutation: @strawberry.mutation @authenticated_user_only() async def edit_point_of_interest(root, info, id: int, input: EditPointOfInterestInput) -> PointOfInterest: - # TODO: kontrola organizace input_data = input.to_dict() + query = get_base_resolver( + models.PointOfInterest, user_id=info.context.user_id, object_id=id, include_public=False + ) + async with get_session() as db: if input.type is not None: input_data['type_id'] = await handle_combobox_save( db, models.PointOfInterestType, input.type, info.context.user_id ) - poi = ( - await db.scalars( - get_base_query(info.context.user_id, only_my=True) - .filter(models.PointOfInterest.id == id) - )).one() + poi = (await db.scalars(query)).one() updated_poi = await models.PointOfInterest.update(db, obj=poi, data=input_data) return PointOfInterest(**updated_poi.as_dict()) @@ -129,11 +118,12 @@ class DeletePointOfInterestMutation: @strawberry.mutation @authenticated_user_only() async def delete_point_of_interest(self, info, id: int) -> PointOfInterest: + query = get_base_resolver( + models.PointOfInterest, user_id=info.context.user_id, object_id=id, include_public=False + ) + async with get_session() as db: - poi = (await db.scalars( - get_base_query(info.context.user_id, only_my=True) - .filter(models.PointOfInterest.id == id) - )).one() + poi = (await db.scalars(query)).one() updated_poi = await models.PointOfInterest.update(db, obj=poi, data=dict(deleted=True)) return PointOfInterest(**updated_poi.as_dict()) diff --git a/src/graphql_schema/entities/poi_type.py b/src/graphql_schema/entities/poi_type.py index 94394cd..6340baf 100644 --- a/src/graphql_schema/entities/poi_type.py +++ b/src/graphql_schema/entities/poi_type.py @@ -3,7 +3,7 @@ import strawberry from sqlalchemy import select, or_ from database import models from decorators.endpoints import authenticated_user_only -from dependencies.db import get_session +from graphql_schema.entities.resolvers.base import get_base_resolver, get_list, get_one from graphql_schema.sqlalchemy_to_strawberry_type import strawberry_sqlalchemy_type @@ -12,48 +12,20 @@ class PointOfInterestType: pass -def get_base_query(user_id: int, only_my: bool = False): - query = ( - select(models.PointOfInterestType) - .filter(models.PointOfInterestType.deleted.is_(False)) - ) - - if only_my: - query = query.filter(models.PointOfInterestType.created_by_id == user_id) - else: - query = query.filter(or_( - models.PointOfInterestType.created_by_id == user_id, - models.PointOfInterestType.is_public.is_(True) - )) - - return query - - @strawberry.type class PointOfInterestTypeQueries: @strawberry.field() @authenticated_user_only() async def point_of_interest_types(root, info) -> List[PointOfInterestType]: - query = ( - get_base_query(info.context.user_id, only_my=False) - .order_by(models.PointOfInterestType.id.desc()) - ) - - async with get_session() as db: - poi_types = (await db.scalars(query)).all() - return [PointOfInterestType(**poi_type.as_dict()) for poi_type in poi_types] + query = get_base_resolver(models.PointOfInterestType, user_id=info.context.user_id) + return await get_list(models.PointOfInterestType, query) @strawberry.field() @authenticated_user_only() async def point_of_interest_type(root, info, id: int) -> PointOfInterestType: - query = ( - get_base_query(info.context.user_id) - .filter(models.PointOfInterestType.id == id) - ) - async with get_session() as db: - poi_type = (await db.scalars(query)).one() - return PointOfInterestType(**poi_type.as_dict()) + query = get_base_resolver(models.PointOfInterestType, user_id=info.context.user_id, object_id=id) + return await get_one(models.PointOfInterestType, query) # # @strawberry.type diff --git a/src/graphql_schema/entities/resolvers/__init__.py b/src/graphql_schema/entities/resolvers/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/graphql_schema/entities/resolvers/aircraft.py b/src/graphql_schema/entities/resolvers/aircraft.py new file mode 100644 index 0000000..4a6334c --- /dev/null +++ b/src/graphql_schema/entities/resolvers/aircraft.py @@ -0,0 +1,24 @@ +from operator import or_ +from typing import Set, Optional +from database import models +from graphql_schema.entities.resolvers.base import get_base_resolver + + +def get_aircraft_resolver(user_id: int, organization_ids: Optional[Set[int]] = None, aircraft_id: Optional[int] = None): + query = get_base_resolver( + model=models.Aircraft, + order_by=[models.Aircraft.id.desc()], + object_id=aircraft_id, + user_id=user_id if not organization_ids else None + ) + if organization_ids: + query = ( + query.filter( + or_( + models.Aircraft.created_by_id == user_id, + models.Aircraft.organization_id.in_(organization_ids) + ) + ) + ) + + return query \ No newline at end of file diff --git a/src/graphql_schema/entities/resolvers/airport.py b/src/graphql_schema/entities/resolvers/airport.py new file mode 100644 index 0000000..4415ea6 --- /dev/null +++ b/src/graphql_schema/entities/resolvers/airport.py @@ -0,0 +1,17 @@ +from typing import Optional +from sqlalchemy import or_ +from database import models +from graphql_schema.entities.resolvers.base import get_base_resolver + + +def get_airport_resolver(user_id: int, airport_id: Optional[int] = None): + if airport_id: + return get_base_resolver(model=models.Airport, object_id=airport_id, user_id=user_id) + + return ( + get_base_resolver(model=models.Airport) + .filter(or_( + models.Airport.created_by_id == user_id, + models.Airport.created_by_id.is_(None), + )) + ) \ No newline at end of file diff --git a/src/graphql_schema/entities/resolvers/base.py b/src/graphql_schema/entities/resolvers/base.py new file mode 100644 index 0000000..e41dd1a --- /dev/null +++ b/src/graphql_schema/entities/resolvers/base.py @@ -0,0 +1,53 @@ +from typing import Optional, Type +from sqlalchemy import select, or_ +from database import models +from dependencies.db import get_session + + +def get_base_resolver( + model: Type[models.BaseModel], + object_id: Optional[int] = None, + user_id: Optional[int] = None, + order_by: Optional[list] = None, + include_public: Optional[bool] = True, +): + query = select(model) + + if object_id: + if hasattr(model, "id"): + query = query.filter(model.id == object_id) + else: + raise AssertionError(f"Model {model} has no ID column! Cannot query by ID!") + + if hasattr(model, "deleted"): + query = query.filter(model.deleted.is_(False)) + + ownership_clause = [] + if hasattr(model, "is_public") and include_public: + ownership_clause.append(model.is_public.is_(True)) + + if hasattr(model, "created_by_id") and user_id: + ownership_clause.append(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) + + return query + + +async def get_list(model, query): + async with get_session() as db: + items = (await db.scalars(query)).all() + + return [model(**m.as_dict()) for m in items] + + +async def get_one(model, query): + async with get_session() as db: + data = (await db.scalars(query)).one() + return model(**data.as_dict())