Uprava prace s DB, refaktoring

This commit is contained in:
Michal Kváček
2023-09-13 10:04:06 +02:00
parent 64cb544e2d
commit 1760dba99d
27 changed files with 195 additions and 182 deletions
+8 -3
View File
@@ -34,7 +34,12 @@ class FlightsLoader:
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_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)
flight_by_poi_dataloader = DataLoader(
load_fn=FlightsLoader(PointOfInterest.id, extra_join=[Flight.track, PointOfInterest]).load,
cache=False
)
+9 -6
View File
@@ -77,13 +77,16 @@ class CreateAircraftMutation:
if input.photo:
input_data['photo_filename'] = await handle_file_upload(input.photo, AIRCRAFT_UPLOAD_DEST_PATH)
return await models.Aircraft.create(
db,
data=dict(
**input_data,
created_by_id=info.context.user_id,
async with get_session() as db:
aircraft = await models.Aircraft.create(
db,
data=dict(
**input_data,
created_by_id=info.context.user_id,
)
)
)
return Aircraft(**aircraft.as_dict())
@strawberry.type
+1 -5
View File
@@ -53,7 +53,6 @@ class CopilotQueries:
@strawberry.type
class CreateCopilotMutation:
@strawberry_sqlalchemy_input(model=models.Copilot, exclude_fields=["id"])
class CreateCopilotInput:
pass
@@ -76,7 +75,6 @@ class CreateCopilotMutation:
@strawberry.type
class EditCopilotMutation:
@strawberry_sqlalchemy_input(model=models.Copilot, exclude_fields=["id"])
class EditCopilotInput:
pass
@@ -84,12 +82,10 @@ class EditCopilotMutation:
@strawberry.mutation
@authenticated_user_only()
async def edit_copilot(root, info, id: int, input: EditCopilotInput) -> Copilot:
async with get_session() as db:
copilot = (await db.scalars(
get_base_query(info.context.user_id).filter(models.Copilot.id == id)
)).one()
updated_copilot = await models.Copilot.update(db, obj=copilot, data=input.to_dict())
return Copilot(**updated_copilot.as_dict())
return Copilot(**updated_copilot.as_dict())
+13 -11
View File
@@ -1,9 +1,8 @@
import asyncio
from datetime import timedelta, datetime
from typing import List, Optional, Annotated, TYPE_CHECKING, Tuple
from typing import List, Optional, Annotated, TYPE_CHECKING
import strawberry
from fastapi import HTTPException
from lxml import etree
from sqlalchemy import select, insert, delete
from starlette.status import HTTP_401_UNAUTHORIZED
from strawberry.file_uploads import Upload
@@ -25,7 +24,10 @@ from graphql_schema.entities.photo import Photo
from graphql_schema.entities.poi import PointOfInterest
from graphql_schema.sqlalchemy_to_strawberry_type import strawberry_sqlalchemy_type, strawberry_sqlalchemy_input
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
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
)
from ..types import ComboboxInput
if TYPE_CHECKING:
@@ -134,7 +136,7 @@ class Flight:
return await flight_copilots_dataloader.load(root.id)
duration_min_calculated: int = strawberry.field(resolver=duration_min_calculated)
copilots: Optional[List[Annotated["Copilot", strawberry.lazy(".copilot")]]] = strawberry.field(resolver=load_copilots)
copilots: Optional[List[Annotated["Copilot", strawberry.lazy(".copilot")]]] = strawberry.field(resolver=load_copilots) # noqa
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)
@@ -220,7 +222,9 @@ class CreateFlightMutation:
data = input.to_dict()
async with get_session() as db:
takeoff_airport, landing_airport = await get_airports(db, input.takeoff_airport.id, input.landing_airport.id)
takeoff_airport, landing_airport = await get_airports(
db, input.takeoff_airport, input.landing_airport, info.context.user_id,
)
aircraft_id = await handle_aircraft_save(db, info.context.user_id, input.aircraft)
weather_takeoff, weather_landing = await asyncio.gather(
@@ -231,8 +235,8 @@ class CreateFlightMutation:
flight = await models.Flight.create(db, data={
**data,
"takeoff_weather_info_id": weather_takeoff.id,
"landing_weather_info_id": weather_landing.id,
"takeoff_weather_info_id": weather_takeoff.id if weather_takeoff else None,
"landing_weather_info_id": weather_landing.id if weather_landing else None,
"takeoff_airport_id": takeoff_airport.id,
"landing_airport_id": landing_airport.id,
"has_terrain_elevation": False,
@@ -268,9 +272,7 @@ class EditFlightMutation:
)).one()
takeoff_airport, landing_airport = await get_airports(
db,
takeoff_airport_id=input.takeoff_airport.id if input.takeoff_airport else flight.takeoff_airport_id,
landing_airport_id=input.landing_airport.id if input.landing_airport else flight.landing_airport_id,
db, input.takeoff_airport, input.landing_airport, info.context.user_id,
)
data = input.to_dict()
@@ -278,7 +280,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, gpx_filename=data['gpx_track_filename'], db=db
add_terrain_elevation, flight=flight.as_dict(), gpx_filename=data['gpx_track_filename']
)
if (
+44 -20
View File
@@ -1,26 +1,33 @@
import asyncio
from datetime import datetime
from typing import List, Type, Literal, Optional, Tuple
from aiohttp import ClientResponseError
from sqlalchemy import select, delete
from sqlalchemy.ext.asyncio import AsyncSession
from starlette.background import BackgroundTasks
from strawberry.file_uploads import Upload
from database import models
from external.elevation import ElevationAPI, elevation_api
from dependencies.db import get_session
from external.elevation import elevation_api
from external.gpx_parser import GPXParser
from external.weather import Weather
from graphql_schema.types import ComboboxInput
from upload_utils import delete_file, file_exists, handle_file_upload
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) -> models.WeatherInfo:
weather = await weather_api.get_weather_for_hour(date_time.astimezone(), (airport.gps_latitude, airport.gps_longitude))
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
try:
weather = await weather_api.get_weather_for_hour(
date_time.astimezone(),
gps=(airport.gps_latitude, airport.gps_longitude)
)
except Exception as e:
print(e)
return None
model = models.WeatherInfo(**{
"datetime": weather['datetime'],
"qnh": weather['pressure_msl'],
@@ -56,7 +63,10 @@ async def handle_track_edit(db: AsyncSession, flight: models.Flight, track: List
poi_object = poi_map.get(item.id)
if not poi_object:
poi_object = await models.PointOfInterest.create(db, data=dict(created_by_id=user_id, name=item.name, description=""))
poi_object = await models.PointOfInterest.create(
db,
data=dict(created_by_id=user_id, name=item.name, description="")
)
await db.flush()
await models.FlightTrack.create(
@@ -82,7 +92,21 @@ async def handle_aircraft_save(db: AsyncSession, user_id: int, aircraft: Combobo
})
async def get_airports(db, takeoff_airport_id: int, landing_airport_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}
)
if landing_airport.id != takeoff_airport_id or landing_airport.name != takeoff_airport.name:
landing_airport_id = await handle_combobox_save(
db, models.Airport, landing_airport, user_id,
name_column="icao_code",
extra_data={"name": landing_airport.name}
)
else:
landing_airport_id = takeoff_airport_id
takeoff_airport = (await db.scalars(
select(models.Airport).filter(models.Airport.id == takeoff_airport_id)
)).one()
@@ -103,39 +127,39 @@ async def handle_airport_changed(
):
flight_datetime = getattr(flight, f"{type_}_datetime")
if input_datetime and input_datetime != flight_datetime:
weather = await handle_weather_info(db, input_datetime, airport)
await db.flush()
existing_weather_id = getattr(flight, f"{type_}_weather_info_id")
if existing_weather_id:
# db.delete(delete())
pass
setattr(flight, f"{type_}_weather_info_id", weather.id)
weather = await handle_weather_info(db, input_datetime, airport)
await db.flush()
if weather:
setattr(flight, f"{type_}_weather_info_id", weather.id)
setattr(flight, f"{type_}_airport_id", airport.id)
setattr(flight, f"{type_}_datetime", input_datetime)
async def add_terrain_elevation(db: AsyncSession, flight: models.Flight, gpx_filename: str):
async def add_terrain_elevation(flight: dict, gpx_filename: str):
path = "/app/uploads/tracks" # TODO vytahnout do configu
gpx_parser = GPXParser(f"{path}/{gpx_filename}")
coordinates = await gpx_parser.get_coordinates()
print("AAAAAAAAAAAAAAAAAAAAAAAAA", coordinates)
try:
elevation = await elevation_api.get_elevation_for_points(coordinates)
print("ELEVATION", elevation)
tree_with_elevation = gpx_parser.add_terrain_elevation(elevation)
output_name = f"terrain_{gpx_filename}"
gpx_parser.write(tree_with_elevation, f"{path}/{output_name}")
await models.Flight.update(db, {"gpx_track_filename": output_name, "has_terrain_elevation": True}, obj=flight)
except ClientResponseError:
print("NEumim elevation!")
async with get_session() as db:
await models.Flight.update(
db, {"gpx_track_filename": output_name, "has_terrain_elevation": True},
id=flight['id'])
except ClientResponseError as e:
print(e)
async def handle_upload_gpx(flight: models.Flight, gpx_track: Upload):
+6 -8
View File
@@ -7,10 +7,11 @@ 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.entities.poi import PointOfInterest
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
from upload_utils import (
get_public_url, handle_file_upload, delete_file, parse_exif_info, generate_thumbnail, file_exists, resize_image
)
from .helpers.flight import handle_combobox_save
if TYPE_CHECKING:
@@ -37,7 +38,7 @@ class Photo:
url: str = strawberry.field(resolver=resolve_url)
thumbnail_url: str = strawberry.field(resolver=resolve_thumb_url)
point_of_interest: Optional[Annotated["PointOfInterest", strawberry.lazy('.poi')]] = strawberry.field(resolver=load_poi)
point_of_interest: Optional[Annotated["PointOfInterest", strawberry.lazy('.poi')]] = strawberry.field(resolver=load_poi) # noqa
def get_base_query(user_id: int):
@@ -161,10 +162,7 @@ class DeletePhotoMutation:
photo = Photo(**photo_model.as_dict())
base_path = get_photo_basepath(photo.flight_id)
try:
delete_file(f"{base_path}/{photo.filename}")
delete_file(f"{base_path}/thumbs/{photo.filename}")
except Exception as e:
print(e)
delete_file(f"{base_path}/{photo.filename}", silent=True)
delete_file(f"{base_path}/thumbs/{photo.filename}", silent=True)
return photo_model
+1 -2
View File
@@ -1,11 +1,10 @@
from typing import List, Optional
from typing import List
import strawberry
from sqlalchemy import select, or_
from database import models
from decorators.endpoints import authenticated_user_only
from dependencies.db import get_session
from graphql_schema.sqlalchemy_to_strawberry_type import strawberry_sqlalchemy_type
from graphql_schema.types import ComboboxInput
@strawberry_sqlalchemy_type(models.PointOfInterestType)
+1 -2
View File
@@ -1,4 +1,3 @@
from functools import wraps
from typing import Optional
import strawberry
from graphql import GraphQLError
@@ -79,7 +78,7 @@ class EditUserMutation:
)).one()
user_image_path = f"/app/uploads/profile/{user.id}"
data = input.to_dict()
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)
-1
View File
@@ -2,7 +2,6 @@ import dataclasses
import strawberry
from fastapi_jwt import JwtAuthorizationCredentials
from fastapi_jwt.jwt import JwtAccessBearerCookie
from sqlalchemy.ext.asyncio import AsyncSession
from starlette.background import BackgroundTasks
from strawberry.extensions import SchemaExtension
from strawberry.fastapi import BaseContext
@@ -41,7 +41,6 @@ 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,