Refaktoring

This commit is contained in:
Michal Kváček
2023-09-26 06:49:15 +02:00
parent dc49cd347c
commit 3f03021f5c
29 changed files with 294 additions and 513 deletions
@@ -5,7 +5,7 @@ from external.elevation import elevation_api
from external.gpx_parser import GPXParser
async def add_terrain_elevation(flight: dict, gpx_filename: str):
async def add_terrain_elevation_to_flight(flight: dict, gpx_filename: str):
path = "/app/uploads/tracks" # TODO vytahnout do configu
gpx_parser = GPXParser(f"{path}/{gpx_filename}")
@@ -24,3 +24,19 @@ async def add_terrain_elevation(flight: dict, gpx_filename: str):
except ClientResponseError as e:
print(e)
async def add_terrain_elevation_to_photo(photo):
async with get_session() as db:
try:
elevation = await elevation_api.get_elevation_for_points([
{"lat": photo.gps_latitude, "lng": photo.gps_longitude}
])
if not elevation:
print("Cannot get elevation")
return
terrain_elevation = elevation[0]['elevation']
await models.Photo.update(db_session=db, obj=photo, data={"terrain_elevation": terrain_elevation})
except Exception as e:
print(f"Cannot get elevation: {e}")
-19
View File
@@ -1,19 +0,0 @@
from database import models
from dependencies.db import get_session
from external.elevation import elevation_api
async def add_terrain_elevation(photo):
async with get_session() as db:
try:
elevation = await elevation_api.get_elevation_for_points([
{"lat": photo.gps_latitude, "lng": photo.gps_longitude}
])
if not elevation:
print("Cannot get elevation")
return
terrain_elevation = elevation[0]['elevation']
await models.Photo.update(db_session=db, obj=photo, data={"terrain_elevation": terrain_elevation})
except Exception as e:
print(f"Cannot get elevation: {e}")
@@ -1,50 +0,0 @@
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, Organization
async def load(ids: List[int]):
async with async_session() as session:
models = (await session.scalars(select(Aircraft).filter(Aircraft.id.in_(ids)))).all()
models_by_id = {model.id: model for model in models}
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)
-16
View File
@@ -1,16 +0,0 @@
from typing import List
from sqlalchemy import select
from strawberry.dataloader import DataLoader
from database import async_session
from database.models import Airport
async def load(ids: List[int]):
async with async_session() as session:
models = (await session.scalars(select(Airport).filter(Airport.id.in_(ids)))).all()
models_by_id = {model.id: model for model in models}
return [models_by_id[id_] for id_ in ids]
airport_dataloader = DataLoader(load_fn=load, cache=False)
+75
View File
@@ -0,0 +1,75 @@
from collections import defaultdict
from sqlalchemy import select
from typing import Type, List, Optional
from database import models, async_session
class SingleModelByIdDataloader:
def __init__(
self, model: Type[models.BaseModel], relationship_column=None, filters: Optional[list] = None
) -> None:
super().__init__()
self.model = model
self.relationship_column = relationship_column if relationship_column else model.id
if filters is None:
filters = []
self.filters = filters
async def load(self, ids: List[int]):
async with async_session() as session:
query = select(self.model, self.relationship_column).filter(self.relationship_column.in_(ids))
for filter_ in self.filters:
query = query.filter(filter_)
items = (await session.execute(query)).all()
items_by_id = {rel_id: item for item, rel_id in items}
return [items_by_id.get(id_) for id_ in ids]
class MultiModelsDataloader:
def __init__(
self,
model: Type[models.BaseModel],
relationship_column,
extra_join: Optional[list] = None,
filters: Optional[list] = None,
order_by: Optional[list] = None,
):
if extra_join is None:
extra_join = []
if filters is None:
filters = []
if order_by is None:
order_by = [model.id.desc()] # defaultne radit od nejnovejsich zaznamu
self.model = model
self.relationship_column = relationship_column
self.extra_join = extra_join
self.order_by = order_by
self.filters = filters
async def load(self, ids: List[int]):
async with async_session() as session:
rel_column = self.relationship_column
query = (
select(self.model, rel_column)
.filter(rel_column.in_(ids))
.order_by(*self.order_by)
)
for table in self.extra_join:
query = query.join(table)
for filter_ in self.filters:
query = query.filter(filter_)
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]
@@ -1,24 +0,0 @@
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 Copilot, Flight
async def load(ids: List[int]):
async with async_session() as session:
models = (await session.execute(
select(Copilot, Flight.id)
.join(Copilot.flights)
.filter(Flight.id.in_(ids))
)).all()
copilots_by_flight_id = defaultdict(list)
for copilot, flight_id in models:
copilots_by_flight_id[flight_id].append(copilot)
return [copilots_by_flight_id[id_] for id_ in ids]
flight_copilots_dataloader = DataLoader(load_fn=load, cache=False)
-23
View File
@@ -1,23 +0,0 @@
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)
-51
View File
@@ -1,51 +0,0 @@
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 Flight, Copilot, PointOfInterest, Event
class FlightsLoader:
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(Flight, rel_column)
.filter(rel_column.in_(ids))
.order_by(Flight.takeoff_datetime.desc())
)
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]
flights_by_copilot_dataloader = DataLoader(
load_fn=FlightsLoader(Copilot.id, extra_join=[Flight.copilots]).load,
cache=False
)
flights_by_aircraft_dataloader = DataLoader(load_fn=FlightsLoader(Flight.aircraft_id).load, cache=False)
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,89 @@
from strawberry.dataloader import DataLoader
from database import models
from graphql_schema.dataloaders.base import MultiModelsDataloader
# TODO: doresit razeni modelu!
aircrafts_from_organization_dataloader = DataLoader(
load_fn=MultiModelsDataloader(
models.Aircraft,
relationship_column=models.Organization.id,
extra_join=[models.Aircraft.organization]
).load,
cache=False
)
flight_copilots_dataloader = DataLoader(
load_fn=MultiModelsDataloader(
models.Copilot,
relationship_column=models.Flight.id,
extra_join=[models.Copilot.flights]).load,
cache=False)
flights_by_copilot_dataloader = DataLoader(
load_fn=MultiModelsDataloader(
models.Flight,
relationship_column=models.Copilot.id,
extra_join=[models.Flight.copilots]).load,
cache=False
)
flights_by_aircraft_dataloader = DataLoader(
load_fn=MultiModelsDataloader(
models.Flight,
relationship_column=models.Flight.aircraft_id
).load, cache=False)
flight_by_poi_dataloader = DataLoader(
load_fn=MultiModelsDataloader(
models.Flight,
relationship_column=models.PointOfInterest.id,
extra_join=[models.Flight.track, models.PointOfInterest]).load,
cache=False
)
flights_by_event_dataloader = DataLoader(
load_fn=MultiModelsDataloader(
models.Flight,
relationship_column=models.Event.id,
extra_join=[models.Flight.event]).load,
cache=False
)
user_organizations_dataloader = DataLoader(load_fn=MultiModelsDataloader(
models.Organization,
relationship_column=models.user_is_in_organization.c.user_id,
extra_join=[models.user_is_in_organization]
).load, cache=False)
users_in_organization_dataloader = DataLoader(
load_fn=MultiModelsDataloader(
models.User,
relationship_column=models.Organization.id,
extra_join=[models.Organization.users]
).load,
cache=False
)
photos_dataloader = DataLoader(
load_fn=MultiModelsDataloader(
models.Photo,
relationship_column=models.Photo.flight_id,
order_by=[models.Photo.exposed_at]
).load,
cache=False
)
poi_photos_dataloader = DataLoader(
load_fn=MultiModelsDataloader(
models.Photo,
relationship_column=models.Photo.point_of_interest_id,
order_by=[models.Photo.exposed_at]
).load,
cache=False
)
flight_track_dataloader = DataLoader(
load_fn=MultiModelsDataloader(
models.FlightTrack,
relationship_column=models.FlightTrack.flight_id,
order_by=[models.FlightTrack.order]
).load,
cache=False
)
@@ -1,39 +0,0 @@
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)
-43
View File
@@ -1,43 +0,0 @@
from collections import defaultdict
from typing import List, Literal
from sqlalchemy import select
from strawberry.dataloader import DataLoader
from database import async_session
from database.models import Photo
class PhotoDataloader:
def __init__(self, relationship_column: Literal['flight_id', 'point_of_interest_id']) -> None:
super().__init__()
self.relationship_column = relationship_column
async def load_collection(self, ids: List[int]):
async with async_session() as session:
models = (await session.scalars(
select(Photo)
.filter(getattr(Photo, self.relationship_column).in_(ids))
.order_by(Photo.exposed_at)
)).all()
photos_by_relationship_id = defaultdict(list)
for photo in models:
photos_by_relationship_id[getattr(photo, self.relationship_column)].append(photo)
return [photos_by_relationship_id[id_] for id_ in ids]
async def flight_cover_photo_load(ids: List[int]):
async with async_session() as session:
models = (
await session.scalars(
select(Photo)
.filter(Photo.is_flight_cover.is_(True))
.filter(Photo.flight_id.in_(ids)))
).all()
photos = {p.flight_id: p for p in models}
return [photos.get(id_) for id_ in ids]
cover_photo_loader = DataLoader(load_fn=flight_cover_photo_load, cache=False)
photos_dataloader = DataLoader(load_fn=PhotoDataloader("flight_id").load_collection, cache=False)
poi_photos_dataloader = DataLoader(load_fn=PhotoDataloader("point_of_interest_id").load_collection, cache=False)
-45
View File
@@ -1,45 +0,0 @@
from collections import defaultdict
from typing import List
from sqlalchemy import select
from strawberry.dataloader import DataLoader
from database import async_session
from database.models import PointOfInterest, FlightTrack, PointOfInterestType
async def load_flight_track(flight_ids: List[int]):
async with async_session() as session:
query = (
select(FlightTrack)
.filter(FlightTrack.flight_id.in_(flight_ids))
.order_by(FlightTrack.order)
)
data = (await session.scalars(query)).all()
pois_by_flight_id = defaultdict(list)
for poi in data:
pois_by_flight_id[poi.flight_id].append(poi)
return [pois_by_flight_id[id_] for id_ in flight_ids]
flight_track_dataloader = DataLoader(load_fn=load_flight_track, cache=False)
async def load_poi(ids: List[int]):
async with async_session() as session:
models = (await session.scalars(select(PointOfInterest).filter(PointOfInterest.id.in_(ids)))).all()
models_by_id = {model.id: model for model in models}
return [models_by_id.get(id_) for id_ in ids]
async def load_poi_type(ids: List[int]):
async with async_session() as session:
models = (await session.scalars(select(PointOfInterestType).filter(PointOfInterestType.id.in_(ids)))).all()
models_by_id = {model.id: model for model in models}
return [models_by_id.get(id_) for id_ in ids]
poi_dataloader = DataLoader(load_fn=load_poi, cache=False)
poi_type_dataloader = DataLoader(load_fn=load_poi_type, cache=False)
@@ -0,0 +1,24 @@
from typing import Optional, Type
from strawberry.dataloader import DataLoader
from database import models
from graphql_schema.dataloaders.base import SingleModelByIdDataloader
def create_dataloader(model: Type[models.BaseModel], relationship_column=None, filters: Optional[list] = None):
loader = SingleModelByIdDataloader(model, relationship_column, filters).load
return DataLoader(load_fn=loader, cache=False)
airport_dataloader = create_dataloader(models.Airport)
aircraft_dataloader = create_dataloader(models.Aircraft)
event_dataloader = create_dataloader(models.Event)
organizations_dataloader = create_dataloader(models.Organization)
airport_weather_info_loader = create_dataloader(models.WeatherInfo)
poi_dataloader = create_dataloader(models.PointOfInterest)
poi_type_dataloader = create_dataloader(models.PointOfInterestType)
cover_photo_loader = create_dataloader(
models.Photo,
relationship_column=models.Photo.flight_id,
filters=[models.Photo.is_flight_cover.is_(True)]
)
-40
View File
@@ -1,40 +0,0 @@
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
)
-16
View File
@@ -1,16 +0,0 @@
from typing import List
from sqlalchemy import select
from strawberry.dataloader import DataLoader
from database import async_session
from database.models import WeatherInfo
async def load(ids: List[int]):
async with async_session() as session:
models = (await session.scalars(select(WeatherInfo).filter(WeatherInfo.id.in_(ids)))).all()
models_by_id = {model.id: model for model in models}
return [models_by_id.get(id_) for id_ in ids]
airport_weather_info_loader = DataLoader(load_fn=load, cache=False)
+7 -11
View File
@@ -1,4 +1,4 @@
from typing import List, Optional, Annotated, TYPE_CHECKING, Set
from typing import List, Optional, Annotated, TYPE_CHECKING
import strawberry
from strawberry.file_uploads import Upload
from database import models
@@ -9,8 +9,8 @@ from upload_utils import handle_file_upload, delete_file, get_public_url
from graphql_schema.entities.helpers.combobox import handle_combobox_save
from .resolvers.aircraft import get_aircraft_resolver
from .resolvers.base import get_list, get_one
from ..dataloaders.flight import flights_by_aircraft_dataloader
from ..dataloaders.organizations import organizations_dataloader
from ..dataloaders.multi_models import flights_by_aircraft_dataloader
from ..dataloaders.single_model import organizations_dataloader
from ..types import ComboboxInput
if TYPE_CHECKING:
@@ -22,18 +22,14 @@ AIRCRAFT_UPLOAD_DEST_PATH = "/app/uploads/aircrafts/"
@strawberry_sqlalchemy_type(models.Aircraft)
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)
flights: List[Annotated["Flight", strawberry.lazy('.flight')]] = strawberry.field(
resolver=lambda root: flights_by_aircraft_dataloader.load(root.id)
)
organization: Optional[Annotated["Organization", strawberry.lazy(".organization")]] = strawberry.field(
resolver=load_organization
resolver=lambda root: organizations_dataloader.load(root.organization_id)
)
+4 -6
View File
@@ -1,12 +1,11 @@
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_copilot_dataloader
from graphql_schema.sqlalchemy_to_strawberry_type import strawberry_sqlalchemy_type, strawberry_sqlalchemy_input
from .resolvers.base import get_base_resolver, get_list, get_one
from ..dataloaders.multi_models import flights_by_copilot_dataloader
if TYPE_CHECKING:
from .flight import Flight
@@ -14,10 +13,9 @@ if TYPE_CHECKING:
@strawberry_sqlalchemy_type(models.Copilot)
class Copilot:
async def load_flights(root):
return await flights_by_copilot_dataloader.load(root.id)
flights: List[Annotated["Flight", strawberry.lazy('.flight')]] = strawberry.field(resolver=load_flights)
flights: List[Annotated["Flight", strawberry.lazy('.flight')]] = strawberry.field(
resolver=lambda root: flights_by_copilot_dataloader.load(root.id)
)
@strawberry.type
+4 -6
View File
@@ -1,12 +1,11 @@
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
from .resolvers.base import get_base_resolver, get_list, get_one
from ..dataloaders.multi_models import flights_by_event_dataloader
if TYPE_CHECKING:
from .flight import Flight
@@ -14,10 +13,9 @@ if TYPE_CHECKING:
@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)
flights: List[Annotated["Flight", strawberry.lazy('.flight')]] = strawberry.field(
resolver=lambda root: flights_by_event_dataloader.load(root.id)
)
@strawberry.type
+23 -45
View File
@@ -6,18 +6,13 @@ from fastapi import HTTPException
from sqlalchemy import select, insert, delete
from starlette.status import HTTP_401_UNAUTHORIZED
from strawberry.file_uploads import Upload
from background_jobs.elevation import add_terrain_elevation_to_flight
from database import models
from database.models import flight_has_copilot
from decorators.endpoints import authenticated_user_only
from decorators.error_logging import error_logging
from dependencies.db import get_session
from external.gpx_parser import GPXParser
from graphql_schema.dataloaders.aircraft import aircraft_dataloader
from graphql_schema.dataloaders.airport import airport_dataloader
from graphql_schema.dataloaders.copilots import flight_copilots_dataloader
from graphql_schema.dataloaders.photos import photos_dataloader, cover_photo_loader
from graphql_schema.dataloaders.poi import flight_track_dataloader, poi_dataloader
from graphql_schema.dataloaders.weather import airport_weather_info_loader
from graphql_schema.entities.aircraft import Aircraft
from graphql_schema.entities.airport import Airport
from graphql_schema.entities.photo import Photo
@@ -26,9 +21,13 @@ 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_combobox_save
handle_upload_gpx, handle_airport_changed, handle_combobox_save
)
from ..dataloaders.multi_models import flight_copilots_dataloader, flight_track_dataloader, photos_dataloader
from ..dataloaders.single_model import (
poi_dataloader, event_dataloader, aircraft_dataloader, airport_dataloader, cover_photo_loader,
airport_weather_info_loader
)
from ..dataloaders.event import event_dataloader
from ..types import ComboboxInput
if TYPE_CHECKING:
@@ -38,10 +37,9 @@ if TYPE_CHECKING:
@strawberry_sqlalchemy_type(models.FlightTrack)
class FlightTrack:
async def load_poi(root):
return await poi_dataloader.load(root.point_of_interest_id)
point_of_interest: PointOfInterest = strawberry.field(resolver=load_poi)
point_of_interest: PointOfInterest = strawberry.field(
resolver=lambda root: poi_dataloader.load(root.point_of_interest_id)
)
@strawberry_sqlalchemy_type(models.WeatherInfo)
@@ -71,27 +69,6 @@ class GPXTrack:
@strawberry_sqlalchemy_type(models.Flight)
class Flight:
async def load_takeoff_airport(root):
return await airport_dataloader.load(root.takeoff_airport_id)
async def load_landing_airport(root):
return await airport_dataloader.load(root.landing_airport_id)
async def load_aircraft(root):
return await aircraft_dataloader.load(root.aircraft_id)
async def load_photos(root):
return await photos_dataloader.load(root.id)
async def load_cover_photo(root):
return await cover_photo_loader.load(root.id)
async def load_takeoff_weather_info(root):
return await airport_weather_info_loader.load(root.takeoff_weather_info_id)
async def load_landing_weather_info(root):
return await airport_weather_info_loader.load(root.landing_weather_info_id)
def duration_min_calculated(root):
if root.duration_total:
return root.duration_total
@@ -102,9 +79,6 @@ class Flight:
return 0
async def load_track(root):
return await flight_track_dataloader.load(root.id)
async def load_gpx_track(root):
if not root.gpx_track_filename:
return None
@@ -144,14 +118,18 @@ class Flight:
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)
cover_photo: Optional[Photo] = strawberry.field(resolver=load_cover_photo)
track: List[FlightTrack] = strawberry.field(resolver=load_track)
takeoff_weather_info: Optional[WeatherInfo] = strawberry.field(resolver=load_takeoff_weather_info)
landing_weather_info: Optional[WeatherInfo] = strawberry.field(resolver=load_landing_weather_info)
photos: List[Photo] = strawberry.field(resolver=load_photos)
aircraft: Aircraft = strawberry.field(resolver=lambda root: aircraft_dataloader.load(root.aircraft_id))
takeoff_airport: Airport = strawberry.field(resolver=lambda root: airport_dataloader.load(root.takeoff_airport_id))
landing_airport: Airport = strawberry.field(resolver=lambda root: airport_dataloader.load(root.landing_airport_id))
cover_photo: Optional[Photo] = strawberry.field(resolver=lambda root: cover_photo_loader.load(root.id))
track: List[FlightTrack] = strawberry.field(resolver=lambda root: flight_track_dataloader.load(root.id))
takeoff_weather_info: Optional[WeatherInfo] = strawberry.field(
resolver=lambda root: airport_weather_info_loader.load(root.takeoff_weather_info_id)
)
landing_weather_info: Optional[WeatherInfo] = strawberry.field(
resolver=lambda root: airport_weather_info_loader.load(root.landing_weather_info_id)
)
photos: List[Photo] = strawberry.field(resolver=lambda root: photos_dataloader.load(root.id))
gpx_track_url: Optional[str] = strawberry.field(resolver=load_gpx_track_url) # TODO: odstranit
gpx_track: Optional[GPXTrack] = strawberry.field(resolver=load_gpx_track)
@@ -284,7 +262,7 @@ class EditFlightMutation:
if input.gpx_track is not None:
data['gpx_track_filename'] = await handle_upload_gpx(flight, input.gpx_track)
info.context.background_tasks.add_task(
add_terrain_elevation, flight=flight.as_dict(), gpx_filename=data['gpx_track_filename']
add_terrain_elevation_to_flight, flight=flight.as_dict(), gpx_filename=data['gpx_track_filename']
)
if input.takeoff_airport and input.landing_airport:
@@ -1,6 +1,6 @@
import asyncio
from datetime import datetime
from typing import List, Type, Literal, Optional, Tuple
from typing import List, Literal, Optional, Tuple
from sqlalchemy import select, delete
from sqlalchemy.ext.asyncio import AsyncSession
from strawberry.file_uploads import Upload
@@ -13,7 +13,9 @@ from upload_utils import delete_file, handle_file_upload
weather_api = Weather()
async def handle_weather_info(db: AsyncSession, date_time: datetime, airport: models.Airport) -> Optional[models.WeatherInfo]:
async def handle_weather_info(
db: AsyncSession, date_time: datetime, airport: models.Airport
) -> Optional[models.WeatherInfo]:
if not airport.gps_latitude or not airport.gps_longitude:
return None
@@ -89,7 +91,9 @@ async def handle_aircraft_save(db: AsyncSession, user_id: int, aircraft: Combobo
})
async def get_airports(db, takeoff_airport: ComboboxInput, landing_airport: ComboboxInput, user_id: int) -> Tuple[models.Airport, models.Airport]:
async def get_airports(
db, takeoff_airport: ComboboxInput, landing_airport: ComboboxInput, user_id: int
) -> Tuple[models.Airport, models.Airport]:
takeoff_airport_id = await handle_combobox_save(
db, models.Airport, takeoff_airport, user_id,
name_column="icao_code", extra_data={"name": takeoff_airport.name}
@@ -147,8 +151,6 @@ async def handle_upload_gpx(flight: models.Flight, gpx_track: Upload):
return await handle_file_upload(gpx_track, path)
async def handle_copilots_edit(db: AsyncSession, copilots: List[ComboboxInput], user_id: int) -> List[int]:
async def handle_copilots_edit(db: AsyncSession, copilots: List[ComboboxInput], user_id: int) -> tuple:
cors = [handle_combobox_save(db, models.Copilot, copilot, user_id) for copilot in copilots]
return await asyncio.gather(*cors)
+8 -12
View File
@@ -1,6 +1,6 @@
from typing import List, Annotated, TYPE_CHECKING
import strawberry
from sqlalchemy import select, delete
from sqlalchemy import delete
from sqlalchemy.dialects.mysql import insert
from sqlalchemy.exc import IntegrityError
from database import models
@@ -8,8 +8,7 @@ from decorators.endpoints import authenticated_user_only
from dependencies.db import get_session
from graphql_schema.sqlalchemy_to_strawberry_type import strawberry_sqlalchemy_type, strawberry_sqlalchemy_input
from .resolvers.base import get_base_resolver, get_list, get_one
from ..dataloaders.aircraft import aircrafts_from_organization_dataloader
from ..dataloaders.users import users_in_organization_dataloader
from ..dataloaders.multi_models import users_in_organization_dataloader, aircrafts_from_organization_dataloader
if TYPE_CHECKING:
from .user import User
@@ -18,15 +17,12 @@ if TYPE_CHECKING:
@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)
users: List[Annotated["User", strawberry.lazy(".user")]] = strawberry.field(
resolver=lambda root: users_in_organization_dataloader.load(root.id)
)
aircrafts: List[Annotated["Aircraft", strawberry.lazy(".aircraft")]] = strawberry.field(
resolver=lambda root: aircrafts_from_organization_dataloader.load(root.id)
)
@strawberry.type
+10 -18
View File
@@ -3,18 +3,19 @@ from typing import List, Optional, Annotated, TYPE_CHECKING
import strawberry
from sqlalchemy import select, update
from strawberry.file_uploads import Upload
from background_jobs.photo import add_terrain_elevation
from background_jobs.elevation import add_terrain_elevation_to_photo
from database import models
from decorators.endpoints import authenticated_user_only
from dependencies.db import get_session
from graphql_schema.dataloaders.poi import poi_dataloader
from graphql_schema.sqlalchemy_to_strawberry_type import strawberry_sqlalchemy_type
from graphql_schema.types import ComboboxInput
from upload_utils import (
get_public_url, handle_file_upload, delete_file, parse_exif_info, generate_thumbnail, file_exists, resize_image, rotate_image
get_public_url, handle_file_upload, delete_file, parse_exif_info, generate_thumbnail, file_exists, resize_image,
rotate_image
)
from graphql_schema.entities.helpers.combobox import handle_combobox_save
from .resolvers.base import get_base_resolver, get_list
from ..dataloaders.single_model import poi_dataloader
if TYPE_CHECKING:
from .poi import PointOfInterest
@@ -22,9 +23,6 @@ if TYPE_CHECKING:
@strawberry_sqlalchemy_type(models.Photo)
class Photo:
def resolve_url(root):
return get_public_url(f"photos/{root.flight_id}/{root.filename}")
def resolve_thumb_url(root):
thumbnail = get_photo_basepath(root.flight_id) + "/thumbs/" + root.filename
if not file_exists(thumbnail):
@@ -32,15 +30,11 @@ class Photo:
return get_public_url(f"photos/{root.flight_id}/thumbs/{root.filename}")
async def load_poi(root):
if not root.point_of_interest_id:
return None
return await poi_dataloader.load(root.point_of_interest_id)
url: str = strawberry.field(resolver=resolve_url)
url: str = strawberry.field(resolver=lambda root: get_public_url(f"photos/{root.flight_id}/{root.filename}"))
thumbnail_url: str = strawberry.field(resolver=resolve_thumb_url)
point_of_interest: Optional[Annotated["PointOfInterest", strawberry.lazy('.poi')]] = strawberry.field(resolver=load_poi) # noqa
point_of_interest: Optional[Annotated["PointOfInterest", strawberry.lazy('.poi')]] = strawberry.field(
resolver=lambda root: poi_dataloader.load(root.point_of_interest_id)
)
def get_base_query(user_id: int):
@@ -100,7 +94,7 @@ class UploadPhotoMutation:
info.context.background_tasks.add_task(generate_thumbnail, path=path, filename=filename)
if exif_info.get("gps_latitude") and exif_info.get("gps_longitude"):
info.context.background_tasks.add_task(add_terrain_elevation, photo=photo)
info.context.background_tasks.add_task(add_terrain_elevation_to_photo, photo=photo)
return photo
@@ -165,7 +159,7 @@ class EditPhotoMutation:
angle=angle,
),
rotate_image(
path=get_photo_basepath(photo.flight_id)+"/thumbs",
path=get_photo_basepath(photo.flight_id) + "/thumbs",
filename=photo.filename,
angle=angle,
),
@@ -174,8 +168,6 @@ class EditPhotoMutation:
return Photo(**photo.as_dict())
@strawberry.type
class DeletePhotoMutation:
@strawberry.mutation()
+10 -30
View File
@@ -4,14 +4,13 @@ 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.dataloaders.flight import flight_by_poi_dataloader
from graphql_schema.dataloaders.photos import poi_photos_dataloader
from graphql_schema.dataloaders.poi import poi_type_dataloader
from graphql_schema.entities.helpers.combobox import handle_combobox_save
from graphql_schema.entities.poi_type import PointOfInterestType
from graphql_schema.sqlalchemy_to_strawberry_type import strawberry_sqlalchemy_type, strawberry_sqlalchemy_input
from graphql_schema.types import ComboboxInput
from .resolvers.base import get_base_resolver, get_list, get_one
from ..dataloaders.multi_models import flight_by_poi_dataloader, poi_photos_dataloader
from ..dataloaders.single_model import poi_type_dataloader
if TYPE_CHECKING:
from .flight import Flight
@@ -20,34 +19,15 @@ if TYPE_CHECKING:
@strawberry_sqlalchemy_type(models.PointOfInterest)
class PointOfInterest:
async def load_photos(root):
return await poi_photos_dataloader.load(root.id)
async def load_type(root):
return await poi_type_dataloader.load(root.type_id)
async def load_flights(root):
return await flight_by_poi_dataloader.load(root.id)
type: Optional[PointOfInterestType] = strawberry.field(resolver=load_type)
photos: List[Annotated["Photo", strawberry.lazy('.photo')]] = strawberry.field(resolver=load_photos)
flights: List[Annotated["Flight", strawberry.lazy('.flight')]] = strawberry.field(resolver=load_flights)
def get_base_query(user_id: int, only_my: bool = False):
query = (
select(models.PointOfInterest)
.filter(models.PointOfInterest.deleted.is_(False))
type: Optional[PointOfInterestType] = strawberry.field(
resolver=lambda root: poi_type_dataloader.load(root.type_id)
)
photos: List[Annotated["Photo", strawberry.lazy('.photo')]] = strawberry.field(
resolver=lambda root: poi_photos_dataloader.load(root.id)
)
flights: List[Annotated["Flight", strawberry.lazy('.flight')]] = strawberry.field(
resolver=lambda root: flight_by_poi_dataloader.load(root.id)
)
if only_my:
query = query.filter(models.PointOfInterest.created_by_id == user_id)
else:
query = query.filter(or_(
models.PointOfInterest.created_by_id == user_id,
models.PointOfInterest.is_public.is_(True)
))
return query
@strawberry.type
-1
View File
@@ -1,6 +1,5 @@
from typing import List
import strawberry
from sqlalchemy import select, or_
from database import models
from decorators.endpoints import authenticated_user_only
from graphql_schema.entities.resolvers.base import get_base_resolver, get_list, get_one
@@ -21,4 +21,4 @@ def get_aircraft_resolver(user_id: int, organization_ids: Optional[Set[int]] = N
)
)
return query
return query
+10 -6
View File
@@ -9,13 +9,14 @@ 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
from ..dataloaders.multi_models import user_organizations_dataloader
if TYPE_CHECKING:
from .organization import Organization
@strawberry_sqlalchemy_type(models.User, exclude_fields=['password_hashed'])
class User:
async def load_avatar_image_url(root):
@@ -30,12 +31,11 @@ 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)
organizations: List[Annotated['Organization', strawberry.lazy(".organization")]] = strawberry.field(
resolver=lambda root: user_organizations_dataloader.load(root.id)
)
@strawberry.type
@@ -85,7 +85,11 @@ class EditUserMutation:
)).one()
user_image_path = f"/app/uploads/profile/{user.id}"
data = {key: getattr(input, key) for key in ("name", "description", "public_username") if getattr(input, key) is not None}
data = {
key: getattr(input, key)
for key in ("name", "description", "public_username")
if getattr(input, key) is not None
}
if input.avatar_image:
if user.avatar_image_filename:
delete_file(f"{user_image_path}/{user.avatar_image_filename}", silent=True)
+3 -1
View File
@@ -3,7 +3,9 @@ from graphql_schema.entities.aircraft import CreateAircraftMutation, EditAircraf
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.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
@@ -41,6 +41,7 @@ def strawberry_sqlalchemy_type(model, exclude_fields: Optional[typing.Union[List
return wrapper
def strawberry_sqlalchemy_input(
model,
exclude_fields: Optional[typing.Union[List, typing.Tuple]] = None,
-3
View File
@@ -9,7 +9,6 @@ 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
@@ -69,8 +68,6 @@ class App:
.filter(models.user_is_in_organization.c.user_id == user_id)
)).all())
print("ORGANIZATION IDS", organization_ids)
return GraphQLContext(
user_id=user_id,
organization_ids=organization_ids,