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)
frac = val % 1
# https://www.pgc.umn.edu/apps/convert/ N47 17.6 E012 47.3
print(frac, frac * 60, frac * 3600)
+17 -13
View File
@@ -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()
+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 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)
+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 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
)
@@ -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
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
+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 .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)
+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
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
+7
View File
@@ -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,
))
+4
View File
@@ -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,
))
+3
View File
@@ -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
+18 -2
View File
@@ -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)