Sprava eventu a organizaci, bugfixing a drobny refaktoring

This commit is contained in:
Michal Kváček
2023-09-20 13:26:34 +02:00
parent 1760dba99d
commit 037c7daef6
17 changed files with 557 additions and 41 deletions
@@ -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 ###
+1
View File
@@ -7,6 +7,7 @@ def deg_to_dec(val: str):
deg = int(val) deg = int(val)
frac = val % 1 frac = val % 1
# https://www.pgc.umn.edu/apps/convert/ N47 17.6 E012 47.3
print(frac, frac * 60, frac * 3600) print(frac, frac * 60, frac * 3600)
+17 -13
View File
@@ -154,19 +154,19 @@ class Aircraft(BaseModel):
notes: Mapped['AircraftNotes'] = relationship() notes: Mapped['AircraftNotes'] = relationship()
class AircraftNotes(BaseModel): # class AircraftNotes(BaseModel):
__tablename__ = "aircraft_notes" # __tablename__ = "aircraft_notes"
#
id: Mapped[int] = mapped_column(primary_key=True) # id: Mapped[int] = mapped_column(primary_key=True)
aircraft_id: Mapped[int] = mapped_column(Integer, ForeignKey("aircraft.id"), nullable=False) # aircraft_id: Mapped[int] = mapped_column(Integer, ForeignKey("aircraft.id"), nullable=False)
name: Mapped[str] = mapped_column(String(128), nullable=False) # name: Mapped[str] = mapped_column(String(128), nullable=False)
description: Mapped[str] = mapped_column(Text, nullable=False) # description: Mapped[str] = mapped_column(Text, nullable=False)
is_public: Mapped[bool] = mapped_column(Boolean, server_default='0') # is_public: Mapped[bool] = mapped_column(Boolean, server_default='0')
created_by_id: Mapped[int] = mapped_column(Integer, ForeignKey('user.id')) # created_by_id: Mapped[int] = mapped_column(Integer, ForeignKey('user.id'))
created_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now()) # created_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now())
#
created_by: Mapped['User'] = relationship() # created_by: Mapped['User'] = relationship()
aircraft: Mapped['Aircraft'] = relationship(back_populates="notes") # aircraft: Mapped['Aircraft'] = relationship(back_populates="notes")
class Organization(BaseModel): class Organization(BaseModel):
@@ -215,11 +215,15 @@ class Event(BaseModel):
id: Mapped[int] = mapped_column(primary_key=True) id: Mapped[int] = mapped_column(primary_key=True)
name: Mapped[str] = mapped_column(String(128), nullable=False) name: Mapped[str] = mapped_column(String(128), nullable=False)
description: Mapped[str] = mapped_column(Text, 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') is_public: Mapped[bool] = mapped_column(Boolean, server_default='0')
created_by_id: Mapped[int] = mapped_column(Integer, ForeignKey('user.id')) created_by_id: Mapped[int] = mapped_column(Integer, ForeignKey('user.id'))
created_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now()) created_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now())
deleted: Mapped[bool] = mapped_column(Boolean, nullable=False, server_default='0') deleted: Mapped[bool] = mapped_column(Boolean, nullable=False, server_default='0')
organization: Mapped['Organization'] = relationship()
created_by: Mapped['User'] = relationship() created_by: Mapped['User'] = relationship()
+36 -2
View File
@@ -1,8 +1,9 @@
from typing import List from collections import defaultdict
from typing import List, Optional
from sqlalchemy import select from sqlalchemy import select
from strawberry.dataloader import DataLoader from strawberry.dataloader import DataLoader
from database import async_session from database import async_session
from database.models import Aircraft from database.models import Aircraft, Organization
async def load(ids: List[int]): 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] 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) aircraft_dataloader = DataLoader(load_fn=load, cache=False)
+23
View File
@@ -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)
+6 -1
View File
@@ -3,7 +3,7 @@ from typing import List, Optional
from sqlalchemy import select from sqlalchemy import select
from strawberry.dataloader import DataLoader from strawberry.dataloader import DataLoader
from database import async_session from database import async_session
from database.models import Flight, Copilot, PointOfInterest from database.models import Flight, Copilot, PointOfInterest, Event
class FlightsLoader: class FlightsLoader:
@@ -43,3 +43,8 @@ flight_by_poi_dataloader = DataLoader(
load_fn=FlightsLoader(PointOfInterest.id, extra_join=[Flight.track, PointOfInterest]).load, load_fn=FlightsLoader(PointOfInterest.id, extra_join=[Flight.track, PointOfInterest]).load,
cache=False cache=False
) )
flights_by_event_dataloader = DataLoader(
load_fn=FlightsLoader(Event.id, extra_join=[Flight.event]).load,
cache=False
)
@@ -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)
+40
View File
@@ -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
)
+51 -17
View File
@@ -1,16 +1,20 @@
from typing import List, Optional, Annotated, TYPE_CHECKING from typing import List, Optional, Annotated, TYPE_CHECKING, Set
import strawberry import strawberry
from strawberry.file_uploads import Upload from strawberry.file_uploads import Upload
from sqlalchemy import select from sqlalchemy import select, or_
from database import models from database import models
from decorators.endpoints import authenticated_user_only from decorators.endpoints import authenticated_user_only
from dependencies.db import get_session from dependencies.db import get_session
from graphql_schema.sqlalchemy_to_strawberry_type import strawberry_sqlalchemy_type, strawberry_sqlalchemy_input 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 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.flight import flights_by_aircraft_dataloader
from ..dataloaders.organizations import organizations_dataloader
from ..types import ComboboxInput
if TYPE_CHECKING: if TYPE_CHECKING:
from .flight import Flight from .flight import Flight
from .organization import Organization
AIRCRAFT_UPLOAD_DEST_PATH = "/app/uploads/aircrafts/" AIRCRAFT_UPLOAD_DEST_PATH = "/app/uploads/aircrafts/"
@@ -20,18 +24,29 @@ class Aircraft:
async def load_flights(root): async def load_flights(root):
return await flights_by_aircraft_dataloader.load(root.id) 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( photo_url: Optional[str] = strawberry.field(
resolver=lambda root: get_public_url(f"aircrafts/{root.photo_filename}") if root.photo_filename else None 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) 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 ( return (
select(models.Aircraft) 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)) .filter(models.Aircraft.deleted.is_(False))
.order_by(models.Aircraft.id.desc())
) )
@@ -41,12 +56,10 @@ class AircraftQueries:
@strawberry.field() @strawberry.field()
@authenticated_user_only() @authenticated_user_only()
async def aircrafts(root, info) -> List[Aircraft]: 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: 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] return [Aircraft(**a.as_dict()) for a in aircrafts]
@@ -54,7 +67,7 @@ class AircraftQueries:
@authenticated_user_only() @authenticated_user_only()
async def aircraft(root, info, id: int) -> Aircraft: async def aircraft(root, info, id: int) -> Aircraft:
query = ( query = (
get_base_query(info.context.user_id) get_base_query(info.context.user_id, info.context.organization_ids)
.filter(models.Aircraft.id == id) .filter(models.Aircraft.id == id)
) )
async with get_session() as db: async with get_session() as db:
@@ -67,6 +80,7 @@ class CreateAircraftMutation:
@strawberry_sqlalchemy_input(models.Aircraft, exclude_fields=['id', 'photo_filename']) @strawberry_sqlalchemy_input(models.Aircraft, exclude_fields=['id', 'photo_filename'])
class CreateAircraftInput: class CreateAircraftInput:
photo: Optional[Upload] photo: Optional[Upload]
organization: Optional[ComboboxInput] = None
@strawberry.mutation @strawberry.mutation
@authenticated_user_only() @authenticated_user_only()
@@ -78,6 +92,15 @@ class CreateAircraftMutation:
input_data['photo_filename'] = await handle_file_upload(input.photo, AIRCRAFT_UPLOAD_DEST_PATH) input_data['photo_filename'] = await handle_file_upload(input.photo, AIRCRAFT_UPLOAD_DEST_PATH)
async with get_session() as db: 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( aircraft = await models.Aircraft.create(
db, db,
data=dict( data=dict(
@@ -94,6 +117,7 @@ class EditAircraftMutation:
@strawberry_sqlalchemy_input(models.Aircraft, exclude_fields=['photo_filename']) @strawberry_sqlalchemy_input(models.Aircraft, exclude_fields=['photo_filename'])
class EditAircraftInput: class EditAircraftInput:
photo: Optional[Upload] photo: Optional[Upload]
organization: Optional[ComboboxInput] = None
@strawberry.mutation @strawberry.mutation
@authenticated_user_only() @authenticated_user_only()
@@ -102,17 +126,27 @@ class EditAircraftMutation:
update_data = input.to_dict() update_data = input.to_dict()
async with get_session() as db: 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( 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) .filter(models.Aircraft.id == id)
)).one() )).one()
if input.photo: if input.photo:
if aircraft.photo_filename: if aircraft.photo_filename:
delete_file(AIRCRAFT_UPLOAD_DEST_PATH + "/" + aircraft.photo_filename, silent=True) 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) 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 @strawberry.type
+91
View File
@@ -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())
+24 -5
View File
@@ -26,12 +26,14 @@ from graphql_schema.sqlalchemy_to_strawberry_type import strawberry_sqlalchemy_t
from upload_utils import get_public_url from upload_utils import get_public_url
from .helpers.flight import ( from .helpers.flight import (
handle_aircraft_save, handle_track_edit, handle_copilots_edit, handle_weather_info, get_airports, 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 from ..types import ComboboxInput
if TYPE_CHECKING: if TYPE_CHECKING:
from .copilot import Copilot from .copilot import Copilot
from .event import Event
@strawberry_sqlalchemy_type(models.FlightTrack) @strawberry_sqlalchemy_type(models.FlightTrack)
@@ -135,8 +137,13 @@ class Flight:
async def load_copilots(root): async def load_copilots(root):
return await flight_copilots_dataloader.load(root.id) 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) duration_min_calculated: int = strawberry.field(resolver=duration_min_calculated)
copilots: Optional[List[Annotated["Copilot", strawberry.lazy(".copilot")]]] = strawberry.field(resolver=load_copilots) # noqa 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) aircraft: Aircraft = strawberry.field(resolver=load_aircraft)
takeoff_airport: Airport = strawberry.field(resolver=load_takeoff_airport) takeoff_airport: Airport = strawberry.field(resolver=load_takeoff_airport)
landing_airport: Airport = strawberry.field(resolver=load_landing_airport) landing_airport: Airport = strawberry.field(resolver=load_landing_airport)
@@ -259,6 +266,7 @@ class EditFlightMutation:
aircraft: Optional[ComboboxInput] = None aircraft: Optional[ComboboxInput] = None
landing_airport: Optional[ComboboxInput] = None landing_airport: Optional[ComboboxInput] = None
takeoff_airport: Optional[ComboboxInput] = None takeoff_airport: Optional[ComboboxInput] = None
event: Optional[ComboboxInput] = None
@strawberry.mutation @strawberry.mutation
@authenticated_user_only() @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) get_base_query(user_id=user_id, is_auth=bool(user_id)).filter(models.Flight.id == id)
)).one() )).one()
takeoff_airport, landing_airport = await get_airports(
db, input.takeoff_airport, input.landing_airport, info.context.user_id,
)
data = input.to_dict() data = input.to_dict()
if input.gpx_track is not None: 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'] 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 ( if (
(input.takeoff_airport and input.takeoff_airport.id != flight.takeoff_airport_id) or (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) (data.get('takeoff_datetime') and data.get('takeoff_datetime') != flight.takeoff_datetime)
@@ -306,10 +315,20 @@ class EditFlightMutation:
type_="landing", type_="landing",
input_datetime=data.get('landing_datetime') input_datetime=data.get('landing_datetime')
) )
# ///////////// konec editace s letistem - je to hnusny
if input.aircraft is not None: if input.aircraft is not None:
data['aircraft_id'] = await handle_aircraft_save(db, user_id, input.aircraft) 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: if input.track is not None:
await handle_track_edit(db=db, flight=flight, track=input.track, user_id=user_id) await handle_track_edit(db=db, flight=flight, track=input.track, user_id=user_id)
+139
View File
@@ -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())
+8 -1
View File
@@ -1,4 +1,4 @@
from typing import Optional from typing import Optional, List, Annotated, TYPE_CHECKING
import strawberry import strawberry
from graphql import GraphQLError from graphql import GraphQLError
from passlib.hash import bcrypt from passlib.hash import bcrypt
@@ -9,9 +9,12 @@ from database import models
from decorators.endpoints import authenticated_user_only from decorators.endpoints import authenticated_user_only
from decorators.error_logging import error_logging from decorators.error_logging import error_logging
from dependencies.db import get_session 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 graphql_schema.sqlalchemy_to_strawberry_type import strawberry_sqlalchemy_type
from upload_utils import handle_file_upload, delete_file, get_public_url, resize_image 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']) @strawberry_sqlalchemy_type(models.User, exclude_fields=['password_hashed'])
class User: class User:
@@ -27,8 +30,12 @@ class User:
return get_public_url(f"profile/{root.id}/{root.title_image_filename}") 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) avatar_image_url: Optional[str] = strawberry.field(resolver=load_avatar_image_url)
title_image_url: str = strawberry.field(resolver=load_title_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 @strawberry.type
+7
View File
@@ -1,7 +1,9 @@
from strawberry.tools import merge_types from strawberry.tools import merge_types
from graphql_schema.entities.aircraft import CreateAircraftMutation, EditAircraftMutation, DeleteAircraftMutation from graphql_schema.entities.aircraft import CreateAircraftMutation, EditAircraftMutation, DeleteAircraftMutation
from graphql_schema.entities.copilot import CreateCopilotMutation, EditCopilotMutation 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.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.photo import UploadPhotoMutation, DeletePhotoMutation, EditPhotoMutation
from graphql_schema.entities.poi import CreatePointOfInterestMutation, EditPointOfInterestMutation from graphql_schema.entities.poi import CreatePointOfInterestMutation, EditPointOfInterestMutation
from graphql_schema.entities.user import EditUserMutation from graphql_schema.entities.user import EditUserMutation
@@ -21,4 +23,9 @@ Mutation = merge_types("Mutation", (
CreateCopilotMutation, CreateCopilotMutation,
EditCopilotMutation, EditCopilotMutation,
EditUserMutation, EditUserMutation,
CreateEventMutation,
EditEventMutation,
CreateOrganizationMutation,
EditOrganizationMutation,
OrganizationUserMutation,
)) ))
+4
View File
@@ -2,7 +2,9 @@ from strawberry.tools import merge_types
from .entities.aircraft import AircraftQueries from .entities.aircraft import AircraftQueries
from .entities.airport import AirportQueries from .entities.airport import AirportQueries
from .entities.copilot import CopilotQueries from .entities.copilot import CopilotQueries
from .entities.event import EventQueries
from .entities.flight import FlightQueries from .entities.flight import FlightQueries
from .entities.organization import OrganizationQueries
from .entities.photo import PhotoQueries from .entities.photo import PhotoQueries
from .entities.poi import PointOfInterestQueries from .entities.poi import PointOfInterestQueries
from .entities.poi_type import PointOfInterestTypeQueries from .entities.poi_type import PointOfInterestTypeQueries
@@ -20,4 +22,6 @@ Query = merge_types('Query', (
PhotoQueries, PhotoQueries,
PointOfInterestQueries, PointOfInterestQueries,
PointOfInterestTypeQueries, PointOfInterestTypeQueries,
EventQueries,
OrganizationQueries,
)) ))
+3
View File
@@ -1,4 +1,6 @@
import dataclasses import dataclasses
from typing import Set
import strawberry import strawberry
from fastapi_jwt import JwtAuthorizationCredentials from fastapi_jwt import JwtAuthorizationCredentials
from fastapi_jwt.jwt import JwtAccessBearerCookie from fastapi_jwt.jwt import JwtAccessBearerCookie
@@ -30,6 +32,7 @@ class LoggingExtension(SchemaExtension):
@dataclasses.dataclass @dataclasses.dataclass
class GraphQLContext(BaseContext): class GraphQLContext(BaseContext):
user_id: int user_id: int
organization_ids: Set[int]
jwt_auth_credentials: JwtAuthorizationCredentials jwt_auth_credentials: JwtAuthorizationCredentials
jwt: JwtAccessBearerCookie jwt: JwtAccessBearerCookie
background_tasks: BackgroundTasks background_tasks: BackgroundTasks
+18 -2
View File
@@ -1,12 +1,15 @@
from datetime import timedelta from datetime import timedelta
from fastapi import FastAPI, APIRouter, Depends, Security from fastapi import FastAPI, APIRouter, Depends, Security
from fastapi_jwt import JwtAuthorizationCredentials, JwtAccessBearerCookie, JwtRefreshBearerCookie from fastapi_jwt import JwtAuthorizationCredentials, JwtAccessBearerCookie, JwtRefreshBearerCookie
from sqlalchemy import select
from starlette.background import BackgroundTasks from starlette.background import BackgroundTasks
from starlette.middleware.cors import CORSMiddleware from starlette.middleware.cors import CORSMiddleware
from starlette.responses import RedirectResponse, Response from starlette.responses import RedirectResponse, Response
from starlette.staticfiles import StaticFiles from starlette.staticfiles import StaticFiles
from strawberry.fastapi import GraphQLRouter from strawberry.fastapi import GraphQLRouter
from config import APP_SECRET_KEY, GRAPHIQL, APP_DEBUG, ALLOW_CORS_ORIGINS 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.login import LoginEndpoint, LoginInput, RefreshEndpoint, LogoutEndpoint
from endpoints.registration import RegistrationInput, RegistrationEndpoint from endpoints.registration import RegistrationInput, RegistrationEndpoint
from graphql_schema.schema import schema, GraphQLContext from graphql_schema.schema import schema, GraphQLContext
@@ -55,9 +58,22 @@ class App:
app.mount("/static", StaticFiles(directory="/app/static"), name="static") app.mount("/static", StaticFiles(directory="/app/static"), name="static")
def setup_graphql_endpoint(self, app: FastAPI): 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( return GraphQLContext(
user_id=credentials['id'] if credentials else None, user_id=user_id,
organization_ids=organization_ids,
jwt_auth_credentials=credentials, jwt_auth_credentials=credentials,
jwt=self.access_security, jwt=self.access_security,
background_tasks=Depends(BackgroundTasks) background_tasks=Depends(BackgroundTasks)