diff --git a/alembic/versions/20230920-115300_add_from_to_to_event_organization_to__39a62618eacb.py b/alembic/versions/20230920-115300_add_from_to_to_event_organization_to__39a62618eacb.py new file mode 100644 index 0000000..f149856 --- /dev/null +++ b/alembic/versions/20230920-115300_add_from_to_to_event_organization_to__39a62618eacb.py @@ -0,0 +1,50 @@ +"""add from/to to event, organization to event + +Revision ID: 39a62618eacb +Revises: a09857ec9721 +Create Date: 2023-09-20 11:53:00.566736 + +""" +from alembic import op +import sqlalchemy as sa +from sqlalchemy.dialects import mysql + +# revision identifiers, used by Alembic. +revision = '39a62618eacb' +down_revision = 'a09857ec9721' +branch_labels = None +depends_on = None + + +def upgrade() -> None: + # ### commands auto generated by Alembic - please adjust! ### + op.drop_table('aircraft_notes') + op.add_column('event', sa.Column('event_from', sa.DateTime(), nullable=True)) + op.add_column('event', sa.Column('event_to', sa.DateTime(), nullable=True)) + op.add_column('event', sa.Column('organization_id', sa.Integer(), nullable=True)) + op.create_foreign_key(None, 'event', 'organization', ['organization_id'], ['id']) + # ### end Alembic commands ### + + +def downgrade() -> None: + # ### commands auto generated by Alembic - please adjust! ### + op.drop_constraint(None, 'event', type_='foreignkey') + op.drop_column('event', 'organization_id') + op.drop_column('event', 'event_to') + op.drop_column('event', 'event_from') + op.create_table('aircraft_notes', + sa.Column('id', mysql.INTEGER(display_width=11), autoincrement=True, nullable=False), + sa.Column('aircraft_id', mysql.INTEGER(display_width=11), autoincrement=False, nullable=False), + sa.Column('name', mysql.VARCHAR(length=128), nullable=False), + sa.Column('description', mysql.TEXT(), nullable=False), + sa.Column('is_public', mysql.TINYINT(display_width=1), server_default=sa.text('0'), autoincrement=False, nullable=False), + sa.Column('created_by_id', mysql.INTEGER(display_width=11), autoincrement=False, nullable=False), + sa.Column('created_at', mysql.DATETIME(), server_default=sa.text('current_timestamp()'), nullable=False), + sa.ForeignKeyConstraint(['aircraft_id'], ['aircraft.id'], name='aircraft_notes_ibfk_1'), + sa.ForeignKeyConstraint(['created_by_id'], ['user.id'], name='aircraft_notes_ibfk_2'), + sa.PrimaryKeyConstraint('id'), + mysql_collate='utf8mb4_general_ci', + mysql_default_charset='utf8mb4', + mysql_engine='InnoDB' + ) + # ### end Alembic commands ### diff --git a/db/csvToDb.py b/db/csvToDb.py index bf8ffe5..61f94c0 100644 --- a/db/csvToDb.py +++ b/db/csvToDb.py @@ -7,6 +7,7 @@ def deg_to_dec(val: str): deg = int(val) frac = val % 1 + # https://www.pgc.umn.edu/apps/convert/ N47 17.6 E012 47.3 print(frac, frac * 60, frac * 3600) diff --git a/src/database/models.py b/src/database/models.py index 7e2827d..081b717 100644 --- a/src/database/models.py +++ b/src/database/models.py @@ -154,19 +154,19 @@ class Aircraft(BaseModel): notes: Mapped['AircraftNotes'] = relationship() -class AircraftNotes(BaseModel): - __tablename__ = "aircraft_notes" - - id: Mapped[int] = mapped_column(primary_key=True) - aircraft_id: Mapped[int] = mapped_column(Integer, ForeignKey("aircraft.id"), nullable=False) - name: Mapped[str] = mapped_column(String(128), nullable=False) - description: Mapped[str] = mapped_column(Text, nullable=False) - 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()) - - created_by: Mapped['User'] = relationship() - aircraft: Mapped['Aircraft'] = relationship(back_populates="notes") +# class AircraftNotes(BaseModel): +# __tablename__ = "aircraft_notes" +# +# id: Mapped[int] = mapped_column(primary_key=True) +# aircraft_id: Mapped[int] = mapped_column(Integer, ForeignKey("aircraft.id"), nullable=False) +# name: Mapped[str] = mapped_column(String(128), nullable=False) +# description: Mapped[str] = mapped_column(Text, nullable=False) +# 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()) +# +# created_by: Mapped['User'] = relationship() +# aircraft: Mapped['Aircraft'] = relationship(back_populates="notes") class Organization(BaseModel): @@ -215,11 +215,15 @@ class Event(BaseModel): id: Mapped[int] = mapped_column(primary_key=True) name: Mapped[str] = mapped_column(String(128), nullable=False) description: Mapped[str] = mapped_column(Text, nullable=False) + event_from: Mapped[datetime] = mapped_column(DateTime, nullable=True) + event_to: Mapped[datetime] = mapped_column(DateTime, nullable=True) + organization_id: Mapped[int] = mapped_column(Integer, ForeignKey('organization.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()) deleted: Mapped[bool] = mapped_column(Boolean, nullable=False, server_default='0') + organization: Mapped['Organization'] = relationship() created_by: Mapped['User'] = relationship() diff --git a/src/graphql_schema/dataloaders/aircraft.py b/src/graphql_schema/dataloaders/aircraft.py index cf18b72..17bb6e7 100644 --- a/src/graphql_schema/dataloaders/aircraft.py +++ b/src/graphql_schema/dataloaders/aircraft.py @@ -1,8 +1,9 @@ -from typing import List +from collections import defaultdict +from typing import List, Optional from sqlalchemy import select from strawberry.dataloader import DataLoader from database import async_session -from database.models import Aircraft +from database.models import Aircraft, Organization async def load(ids: List[int]): @@ -13,4 +14,37 @@ async def load(ids: List[int]): return [models_by_id.get(id_) for id_ in ids] +class OrganizationLoader: + def __init__(self, relationship_column, extra_join: Optional[list] = None): + if extra_join is None: + extra_join = [] + + self.relationship_column = relationship_column + self.extra_join = extra_join + + async def load(self, ids: List[int]): + async with async_session() as session: + rel_column = self.relationship_column + query = ( + select(Aircraft, rel_column) + .filter(rel_column.in_(ids)) + ) + + for table in self.extra_join: + query = query.join(table) + + data = (await session.execute(query)).all() + + result_data = defaultdict(list) + for item, rel_id in data: + result_data[rel_id].append(item) + + return [result_data[id_] for id_ in ids] + + +aircrafts_from_organization_dataloader = DataLoader( + load_fn=OrganizationLoader(Organization.id, extra_join=[Aircraft.organization]).load, + cache=False +) + aircraft_dataloader = DataLoader(load_fn=load, cache=False) diff --git a/src/graphql_schema/dataloaders/event.py b/src/graphql_schema/dataloaders/event.py new file mode 100644 index 0000000..0d6f4ba --- /dev/null +++ b/src/graphql_schema/dataloaders/event.py @@ -0,0 +1,23 @@ +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 Event + + +async def load(ids: List[int]): + async with async_session() as session: + models = (await session.scalars( + select(Event) + .filter(Event.id.in_(ids)) + )).all() + + events_by_id = {} + for event in models: + events_by_id[event.id] = event + + return [events_by_id.get(id_) for id_ in ids] + + +event_dataloader = DataLoader(load_fn=load, cache=False) diff --git a/src/graphql_schema/dataloaders/flight.py b/src/graphql_schema/dataloaders/flight.py index be15066..fe17fe1 100644 --- a/src/graphql_schema/dataloaders/flight.py +++ b/src/graphql_schema/dataloaders/flight.py @@ -3,7 +3,7 @@ from typing import List, Optional from sqlalchemy import select from strawberry.dataloader import DataLoader from database import async_session -from database.models import Flight, Copilot, PointOfInterest +from database.models import Flight, Copilot, PointOfInterest, Event class FlightsLoader: @@ -43,3 +43,8 @@ flight_by_poi_dataloader = DataLoader( load_fn=FlightsLoader(PointOfInterest.id, extra_join=[Flight.track, PointOfInterest]).load, cache=False ) + +flights_by_event_dataloader = DataLoader( + load_fn=FlightsLoader(Event.id, extra_join=[Flight.event]).load, + cache=False +) diff --git a/src/graphql_schema/dataloaders/organizations.py b/src/graphql_schema/dataloaders/organizations.py new file mode 100644 index 0000000..7beb3c4 --- /dev/null +++ b/src/graphql_schema/dataloaders/organizations.py @@ -0,0 +1,39 @@ +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 Organization, Flight, user_is_in_organization + + +async def load(ids: List[int]): + async with async_session() as session: + models = (await session.scalars( + select(Organization) + .filter(Organization.id.in_(ids)) + )).all() + + organizations_by_id = {} + for organization in models: + organizations_by_id[organization.id] = organization + + return [organizations_by_id.get(id_) for id_ in ids] + + +async def load_organizations(ids: List[int]): + async with async_session() as session: + models = (await session.execute( + select(Organization, user_is_in_organization.c.user_id) + .join(user_is_in_organization) + .filter(user_is_in_organization.c.user_id.in_(ids)) + )).all() + + organizations_by_user_id = defaultdict(list) + for organization, user_id in models: + organizations_by_user_id[user_id].append(organization) + return [organizations_by_user_id.get(id_, []) for id_ in ids] + + +organizations_dataloader = DataLoader(load_fn=load, cache=False) + +user_organizations_dataloader = DataLoader(load_fn=load_organizations, cache=False) \ No newline at end of file diff --git a/src/graphql_schema/dataloaders/users.py b/src/graphql_schema/dataloaders/users.py new file mode 100644 index 0000000..d6d1511 --- /dev/null +++ b/src/graphql_schema/dataloaders/users.py @@ -0,0 +1,40 @@ +from collections import defaultdict +from typing import List, Optional +from sqlalchemy import select +from strawberry.dataloader import DataLoader +from database import async_session +from database.models import Organization, User + + +class UsersLoader: + def __init__(self, relationship_column, extra_join: Optional[list] = None): + if extra_join is None: + extra_join = [] + + self.relationship_column = relationship_column + self.extra_join = extra_join + + async def load(self, ids: List[int]): + async with async_session() as session: + rel_column = self.relationship_column + query = ( + select(User, rel_column) + .filter(rel_column.in_(ids)) + ) + + for table in self.extra_join: + query = query.join(table) + + data = (await session.execute(query)).all() + + result_data = defaultdict(list) + for item, rel_id in data: + result_data[rel_id].append(item) + + return [result_data[id_] for id_ in ids] + + +users_in_organization_dataloader = DataLoader( + load_fn=UsersLoader(Organization.id, extra_join=[Organization.users]).load, + cache=False +) diff --git a/src/graphql_schema/entities/aircraft.py b/src/graphql_schema/entities/aircraft.py index a6e4183..c79a3d1 100644 --- a/src/graphql_schema/entities/aircraft.py +++ b/src/graphql_schema/entities/aircraft.py @@ -1,16 +1,20 @@ -from typing import List, Optional, Annotated, TYPE_CHECKING +from typing import List, Optional, Annotated, TYPE_CHECKING, Set import strawberry from strawberry.file_uploads import Upload -from sqlalchemy import select +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 ..dataloaders.flight import flights_by_aircraft_dataloader +from ..dataloaders.organizations import organizations_dataloader +from ..types import ComboboxInput if TYPE_CHECKING: from .flight import Flight + from .organization import Organization AIRCRAFT_UPLOAD_DEST_PATH = "/app/uploads/aircrafts/" @@ -20,18 +24,29 @@ class Aircraft: async def load_flights(root): return await flights_by_aircraft_dataloader.load(root.id) + async def load_organization(root): + return await organizations_dataloader.load(root.organization_id) + photo_url: Optional[str] = strawberry.field( resolver=lambda root: get_public_url(f"aircrafts/{root.photo_filename}") if root.photo_filename else None ) - flights: List[Annotated["Flight", strawberry.lazy('.flight')]] = strawberry.field(resolver=load_flights) + organization: Optional[Annotated["Organization", strawberry.lazy(".organization")]] = strawberry.field( + resolver=load_organization + ) -def get_base_query(user_id: int): +def get_base_query(user_id: int, organization_ids: Set[int]): return ( select(models.Aircraft) - .filter(models.Aircraft.created_by_id == user_id) + .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()) ) @@ -41,12 +56,10 @@ class AircraftQueries: @strawberry.field() @authenticated_user_only() async def aircrafts(root, info) -> List[Aircraft]: - query = ( - get_base_query(info.context.user_id) - .order_by(models.Aircraft.id.desc()) - ) async with get_session() as db: - aircrafts = (await db.scalars(query)).all() + 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] @@ -54,7 +67,7 @@ class AircraftQueries: @authenticated_user_only() async def aircraft(root, info, id: int) -> Aircraft: query = ( - get_base_query(info.context.user_id) + get_base_query(info.context.user_id, info.context.organization_ids) .filter(models.Aircraft.id == id) ) async with get_session() as db: @@ -67,6 +80,7 @@ class CreateAircraftMutation: @strawberry_sqlalchemy_input(models.Aircraft, exclude_fields=['id', 'photo_filename']) class CreateAircraftInput: photo: Optional[Upload] + organization: Optional[ComboboxInput] = None @strawberry.mutation @authenticated_user_only() @@ -78,6 +92,15 @@ 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, + models.Organization, + input=input.organization, + user_id=info.context.user_id, + ) + aircraft = await models.Aircraft.create( db, data=dict( @@ -94,6 +117,7 @@ class EditAircraftMutation: @strawberry_sqlalchemy_input(models.Aircraft, exclude_fields=['photo_filename']) class EditAircraftInput: photo: Optional[Upload] + organization: Optional[ComboboxInput] = None @strawberry.mutation @authenticated_user_only() @@ -102,17 +126,27 @@ class EditAircraftMutation: update_data = input.to_dict() async with get_session() as db: + if input.organization: + update_data['organization_id'] = await handle_combobox_save( + db, + models.Organization, + input=input.organization, + user_id=info.context.user_id, + ) + aircraft = (await db.scalars( - get_base_query(info.context.user_id) + get_base_query(info.context.user_id, set()) # TODO: bude fungovat prazdny set? .filter(models.Aircraft.id == id) )).one() - if input.photo: - if aircraft.photo_filename: - delete_file(AIRCRAFT_UPLOAD_DEST_PATH + "/" + aircraft.photo_filename, silent=True) - update_data['photo_filename'] = await handle_file_upload(input.photo, AIRCRAFT_UPLOAD_DEST_PATH) + if input.photo: + if aircraft.photo_filename: + delete_file(AIRCRAFT_UPLOAD_DEST_PATH + "/" + aircraft.photo_filename, silent=True) + update_data['photo_filename'] = await handle_file_upload(input.photo, AIRCRAFT_UPLOAD_DEST_PATH) - return await models.Aircraft.update(db, obj=aircraft, data=update_data) + aircraft = await models.Aircraft.update(db, obj=aircraft, data=update_data) + + return Aircraft(**aircraft.as_dict()) @strawberry.type diff --git a/src/graphql_schema/entities/event.py b/src/graphql_schema/entities/event.py new file mode 100644 index 0000000..13d3cd9 --- /dev/null +++ b/src/graphql_schema/entities/event.py @@ -0,0 +1,91 @@ +from typing import List, Annotated, TYPE_CHECKING +import strawberry +from sqlalchemy import select +from database import models +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 + +if TYPE_CHECKING: + from .flight import Flight + + +@strawberry_sqlalchemy_type(models.Event) +class Event: + async def load_flights(root): + return await flights_by_event_dataloader.load(root.id) + + 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.name) + ) + + +@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] + + @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()) + + +@strawberry.type +class CreateEventMutation: + @strawberry_sqlalchemy_input(model=models.Event, exclude_fields=["id"]) + class CreateEventInput: + pass + + @strawberry.mutation + @authenticated_user_only() + async def create_event(root, info, input: CreateEventInput) -> Event: + input_data = input.to_dict() + async with get_session() as db: + event = await models.Event.create( + db, + data=dict( + **input_data, + created_by_id=info.context.user_id, + ) + ) + + return Event(**event.as_dict()) + + +@strawberry.type +class EditEventMutation: + @strawberry_sqlalchemy_input(model=models.Event, exclude_fields=["id"]) + class EditEventInput: + pass + + @strawberry.mutation + @authenticated_user_only() + 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) + )).one() + + updated_event = await models.Event.update(db, obj=event, data=input.to_dict()) + return Event(**updated_event.as_dict()) diff --git a/src/graphql_schema/entities/flight.py b/src/graphql_schema/entities/flight.py index 3ffdc95..a2b6983 100644 --- a/src/graphql_schema/entities/flight.py +++ b/src/graphql_schema/entities/flight.py @@ -26,12 +26,14 @@ from graphql_schema.sqlalchemy_to_strawberry_type import strawberry_sqlalchemy_t from upload_utils import get_public_url from .helpers.flight import ( handle_aircraft_save, handle_track_edit, handle_copilots_edit, handle_weather_info, get_airports, - handle_upload_gpx, handle_airport_changed, add_terrain_elevation + handle_upload_gpx, handle_airport_changed, add_terrain_elevation, handle_combobox_save ) +from ..dataloaders.event import event_dataloader from ..types import ComboboxInput if TYPE_CHECKING: from .copilot import Copilot + from .event import Event @strawberry_sqlalchemy_type(models.FlightTrack) @@ -135,8 +137,13 @@ class Flight: async def load_copilots(root): return await flight_copilots_dataloader.load(root.id) + @authenticated_user_only(raise_when_unauthorized=False, return_value_unauthorized=[]) + async def load_event(root): + return await event_dataloader.load(root.event_id) + duration_min_calculated: int = strawberry.field(resolver=duration_min_calculated) copilots: Optional[List[Annotated["Copilot", strawberry.lazy(".copilot")]]] = strawberry.field(resolver=load_copilots) # noqa + event: Optional[Annotated["Event", strawberry.lazy(".event")]] = strawberry.field(resolver=load_event) aircraft: Aircraft = strawberry.field(resolver=load_aircraft) takeoff_airport: Airport = strawberry.field(resolver=load_takeoff_airport) landing_airport: Airport = strawberry.field(resolver=load_landing_airport) @@ -259,6 +266,7 @@ class EditFlightMutation: aircraft: Optional[ComboboxInput] = None landing_airport: Optional[ComboboxInput] = None takeoff_airport: Optional[ComboboxInput] = None + event: Optional[ComboboxInput] = None @strawberry.mutation @authenticated_user_only() @@ -271,10 +279,6 @@ class EditFlightMutation: get_base_query(user_id=user_id, is_auth=bool(user_id)).filter(models.Flight.id == id) )).one() - takeoff_airport, landing_airport = await get_airports( - db, input.takeoff_airport, input.landing_airport, info.context.user_id, - ) - data = input.to_dict() if input.gpx_track is not None: @@ -283,6 +287,11 @@ class EditFlightMutation: add_terrain_elevation, flight=flight.as_dict(), gpx_filename=data['gpx_track_filename'] ) + if input.takeoff_airport and input.landing_airport: + takeoff_airport, landing_airport = await get_airports( + db, input.takeoff_airport, input.landing_airport, info.context.user_id, + ) + if ( (input.takeoff_airport and input.takeoff_airport.id != flight.takeoff_airport_id) or (data.get('takeoff_datetime') and data.get('takeoff_datetime') != flight.takeoff_datetime) @@ -306,10 +315,20 @@ class EditFlightMutation: type_="landing", input_datetime=data.get('landing_datetime') ) + # ///////////// konec editace s letistem - je to hnusny if input.aircraft is not None: data['aircraft_id'] = await handle_aircraft_save(db, user_id, input.aircraft) + if input.event is not None: + data['event_id'] = await handle_combobox_save( + db, + model=models.Event, + input=input.event, + extra_data={"description": "", "is_public": False}, + user_id=info.context.user_id + ) + if input.track is not None: await handle_track_edit(db=db, flight=flight, track=input.track, user_id=user_id) diff --git a/src/graphql_schema/entities/organization.py b/src/graphql_schema/entities/organization.py new file mode 100644 index 0000000..aeaf468 --- /dev/null +++ b/src/graphql_schema/entities/organization.py @@ -0,0 +1,139 @@ +from typing import List, Annotated, TYPE_CHECKING +import strawberry +from sqlalchemy import select, or_, 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 ..dataloaders.aircraft import aircrafts_from_organization_dataloader +from ..dataloaders.users import users_in_organization_dataloader + +if TYPE_CHECKING: + from .user import User + from .aircraft import Aircraft + + +@strawberry_sqlalchemy_type(models.Organization) +class Organization: + + async def load_users(self): + return await users_in_organization_dataloader.load(self.id) + + async def load_aircrafts(self): + return await aircrafts_from_organization_dataloader.load(self.id) + + users: List[Annotated["User", strawberry.lazy(".user")]] = strawberry.field(resolver=load_users) + 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] + + @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()) + + +@strawberry.type +class CreateOrganizationMutation: + @strawberry_sqlalchemy_input(model=models.Organization, exclude_fields=["id"]) + class CreateOrganizationInput: + pass + + @strawberry.mutation + @authenticated_user_only() + async def create_organization(root, info, input: CreateOrganizationInput) -> Organization: + input_data = input.to_dict() + async with get_session() as db: + organization = await models.Organization.create( + db, + data=dict( + **input_data, + created_by_id=info.context.user_id, + ) + ) + + return Organization(**organization.as_dict()) + + +@strawberry.type +class OrganizationUserMutation: + + @strawberry.mutation + @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() + + try: + await db.execute( + insert(models.user_is_in_organization).values( + user_id=info.context.user_id, + organization_id=organization_id + ) + ) + except IntegrityError: + print("jiz existuje") + pass + + return Organization(**organization.as_dict()) + + @strawberry.mutation + @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() + + await db.execute( + delete(models.user_is_in_organization).filter_by( + user_id=info.context.user_id, + organization_id=organization_id + ) + ) + + return Organization(**organization.as_dict()) + + + +@strawberry.type +class EditOrganizationMutation: + @strawberry_sqlalchemy_input(model=models.Organization, exclude_fields=["id"]) + class EditOrganizationInput: + pass + + @strawberry.mutation + @authenticated_user_only() + 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) + )).one() + + updated_organization = await models.Organization.update(db, obj=organization, data=input.to_dict()) + return Organization(**updated_organization.as_dict()) diff --git a/src/graphql_schema/entities/user.py b/src/graphql_schema/entities/user.py index 8858226..141c176 100644 --- a/src/graphql_schema/entities/user.py +++ b/src/graphql_schema/entities/user.py @@ -1,4 +1,4 @@ -from typing import Optional +from typing import Optional, List, Annotated, TYPE_CHECKING import strawberry from graphql import GraphQLError from passlib.hash import bcrypt @@ -9,9 +9,12 @@ from database import models from decorators.endpoints import authenticated_user_only from decorators.error_logging import error_logging from dependencies.db import get_session +from graphql_schema.dataloaders.organizations import user_organizations_dataloader from graphql_schema.sqlalchemy_to_strawberry_type import strawberry_sqlalchemy_type from upload_utils import handle_file_upload, delete_file, get_public_url, resize_image +if TYPE_CHECKING: + from .organization import Organization @strawberry_sqlalchemy_type(models.User, exclude_fields=['password_hashed']) class User: @@ -27,8 +30,12 @@ class User: return get_public_url(f"profile/{root.id}/{root.title_image_filename}") + async def load_organizations(root): + return await user_organizations_dataloader.load(root.id) + avatar_image_url: Optional[str] = strawberry.field(resolver=load_avatar_image_url) title_image_url: str = strawberry.field(resolver=load_title_image_url) + organizations: List[Annotated['Organization', strawberry.lazy(".organization")]] = strawberry.field(resolver=load_organizations) @strawberry.type diff --git a/src/graphql_schema/mutation.py b/src/graphql_schema/mutation.py index 281e8ac..bbe0801 100644 --- a/src/graphql_schema/mutation.py +++ b/src/graphql_schema/mutation.py @@ -1,7 +1,9 @@ from strawberry.tools import merge_types from graphql_schema.entities.aircraft import CreateAircraftMutation, EditAircraftMutation, DeleteAircraftMutation from graphql_schema.entities.copilot import CreateCopilotMutation, EditCopilotMutation +from graphql_schema.entities.event import CreateEventMutation, EditEventMutation from graphql_schema.entities.flight import CreateFlightMutation, EditFlightMutation, DeleteFlightMutation +from graphql_schema.entities.organization import CreateOrganizationMutation, EditOrganizationMutation, OrganizationUserMutation from graphql_schema.entities.photo import UploadPhotoMutation, DeletePhotoMutation, EditPhotoMutation from graphql_schema.entities.poi import CreatePointOfInterestMutation, EditPointOfInterestMutation from graphql_schema.entities.user import EditUserMutation @@ -21,4 +23,9 @@ Mutation = merge_types("Mutation", ( CreateCopilotMutation, EditCopilotMutation, EditUserMutation, + CreateEventMutation, + EditEventMutation, + CreateOrganizationMutation, + EditOrganizationMutation, + OrganizationUserMutation, )) diff --git a/src/graphql_schema/query.py b/src/graphql_schema/query.py index 27a191c..00f2cfe 100644 --- a/src/graphql_schema/query.py +++ b/src/graphql_schema/query.py @@ -2,7 +2,9 @@ from strawberry.tools import merge_types from .entities.aircraft import AircraftQueries from .entities.airport import AirportQueries from .entities.copilot import CopilotQueries +from .entities.event import EventQueries from .entities.flight import FlightQueries +from .entities.organization import OrganizationQueries from .entities.photo import PhotoQueries from .entities.poi import PointOfInterestQueries from .entities.poi_type import PointOfInterestTypeQueries @@ -20,4 +22,6 @@ Query = merge_types('Query', ( PhotoQueries, PointOfInterestQueries, PointOfInterestTypeQueries, + EventQueries, + OrganizationQueries, )) diff --git a/src/graphql_schema/schema.py b/src/graphql_schema/schema.py index bfaa002..a86d29b 100644 --- a/src/graphql_schema/schema.py +++ b/src/graphql_schema/schema.py @@ -1,4 +1,6 @@ import dataclasses +from typing import Set + import strawberry from fastapi_jwt import JwtAuthorizationCredentials from fastapi_jwt.jwt import JwtAccessBearerCookie @@ -30,6 +32,7 @@ class LoggingExtension(SchemaExtension): @dataclasses.dataclass class GraphQLContext(BaseContext): user_id: int + organization_ids: Set[int] jwt_auth_credentials: JwtAuthorizationCredentials jwt: JwtAccessBearerCookie background_tasks: BackgroundTasks diff --git a/src/main.py b/src/main.py index de9d614..228aeda 100644 --- a/src/main.py +++ b/src/main.py @@ -1,12 +1,15 @@ from datetime import timedelta from fastapi import FastAPI, APIRouter, Depends, Security from fastapi_jwt import JwtAuthorizationCredentials, JwtAccessBearerCookie, JwtRefreshBearerCookie +from sqlalchemy import select from starlette.background import BackgroundTasks from starlette.middleware.cors import CORSMiddleware from starlette.responses import RedirectResponse, Response from starlette.staticfiles import StaticFiles from strawberry.fastapi import GraphQLRouter from config import APP_SECRET_KEY, GRAPHIQL, APP_DEBUG, ALLOW_CORS_ORIGINS +from database import models, async_session +from dependencies.db import get_session from endpoints.login import LoginEndpoint, LoginInput, RefreshEndpoint, LogoutEndpoint from endpoints.registration import RegistrationInput, RegistrationEndpoint from graphql_schema.schema import schema, GraphQLContext @@ -55,9 +58,22 @@ class App: app.mount("/static", StaticFiles(directory="/app/static"), name="static") def setup_graphql_endpoint(self, app: FastAPI): - def setup_graphql_context(credentials: JwtAuthorizationCredentials = Security(self.access_security)): + async def setup_graphql_context(credentials: JwtAuthorizationCredentials = Security(self.access_security)): + user_id = credentials['id'] if credentials else None + organization_ids = set() + + if user_id: + async with async_session() as db: + organization_ids = set((await db.scalars( + select(models.user_is_in_organization.c.organization_id) + .filter(models.user_is_in_organization.c.user_id == user_id) + )).all()) + + print("ORGANIZATION IDS", organization_ids) + return GraphQLContext( - user_id=credentials['id'] if credentials else None, + user_id=user_id, + organization_ids=organization_ids, jwt_auth_credentials=credentials, jwt=self.access_security, background_tasks=Depends(BackgroundTasks)