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
@@ -0,0 +1,38 @@
"""make gps in airport nullable
Revision ID: a09857ec9721
Revises: 6e5cc5123a2b
Create Date: 2023-09-13 09:32:06.444179
"""
from alembic import op
import sqlalchemy as sa
from sqlalchemy.dialects import mysql
# revision identifiers, used by Alembic.
revision = 'a09857ec9721'
down_revision = '6e5cc5123a2b'
branch_labels = None
depends_on = None
def upgrade() -> None:
# ### commands auto generated by Alembic - please adjust! ###
op.alter_column('airport', 'gps_latitude',
existing_type=mysql.FLOAT(),
nullable=True)
op.alter_column('airport', 'gps_longitude',
existing_type=mysql.FLOAT(),
nullable=True)
# ### end Alembic commands ###
def downgrade() -> None:
# ### commands auto generated by Alembic - please adjust! ###
op.alter_column('airport', 'gps_longitude',
existing_type=mysql.FLOAT(),
nullable=False)
op.alter_column('airport', 'gps_latitude',
existing_type=mysql.FLOAT(),
nullable=False)
# ### end Alembic commands ###
+17 -4
View File
@@ -3,7 +3,16 @@ import csv
def deg_to_dec(val: str): def deg_to_dec(val: str):
val = val[:-1].lstrip("0") val = float(val[:-1]) / 100
deg = int(val)
frac = val % 1
print(frac, frac * 60, frac * 3600)
return deg + (frac * 60)
# h_m, s = val.split(".") # h_m, s = val.split(".")
# h_m = h_m.strip("0") # h_m = h_m.strip("0")
# h = int(h_m[0:2]) # h = int(h_m[0:2])
@@ -16,15 +25,19 @@ def deg_to_dec(val: str):
# print("--konec-----------------------------------") # print("--konec-----------------------------------")
# return h + (m / 60.0) + (int(s) / 3600.0) # return h + (m / 60.0) + (int(s) / 3600.0)
return len(val or "") # return val
# return len(val or "")
# ZDAKOV: @49.504378,14.1808905 # ZDAKOV: @49.504378,14.1808905
# 4930.250N / 100 -> cele cislo stupne, desetinne prevest do sedesatkove soustavy
with open("./poi.csv") as f: with open("./poi.csv") as f:
reader = csv.DictReader(f) reader = csv.DictReader(f)
for row in reader: for row in reader:
# if row['name'] != 'ZDAKOV': if row['name'] != 'ZDAKOV':
# continue continue
print(row['name'], row['lat'], row['lon'], deg_to_dec(row['lat']), deg_to_dec(row['lon'])) print(row['name'], row['lat'], row['lon'], deg_to_dec(row['lat']), deg_to_dec(row['lon']))
+3 -3
View File
@@ -6,7 +6,9 @@ from external.elevation import elevation_api
async def add_terrain_elevation(photo): async def add_terrain_elevation(photo):
async with get_session() as db: async with get_session() as db:
try: try:
elevation = await elevation_api.get_elevation_for_points([{"lat": photo.gps_latitude, "lng": photo.gps_longitude}]) elevation = await elevation_api.get_elevation_for_points([
{"lat": photo.gps_latitude, "lng": photo.gps_longitude}
])
if not elevation: if not elevation:
print("Cannot get elevation") print("Cannot get elevation")
return return
@@ -15,5 +17,3 @@ async def add_terrain_elevation(photo):
await models.Photo.update(db_session=db, obj=photo, data={"terrain_elevation": terrain_elevation}) await models.Photo.update(db_session=db, obj=photo, data={"terrain_elevation": terrain_elevation})
except Exception as e: except Exception as e:
print(f"Cannot get elevation: {e}") print(f"Cannot get elevation: {e}")
+3 -3
View File
@@ -64,12 +64,12 @@ class Airport(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)
icao_code: Mapped[str] = mapped_column(String(8), nullable=False) icao_code: Mapped[str] = mapped_column(String(8), nullable=False)
gps_latitude: Mapped[float] = mapped_column(Float, nullable=False) gps_latitude: Mapped[float] = mapped_column(Float, nullable=True)
gps_longitude: Mapped[float] = mapped_column(Float, nullable=False) gps_longitude: Mapped[float] = mapped_column(Float, nullable=True)
elevation: Mapped[int] = mapped_column(Integer, nullable=True) elevation: Mapped[int] = mapped_column(Integer, nullable=True)
is_public: Mapped[bool] = mapped_column(Boolean, server_default='0') is_public: Mapped[bool] = mapped_column(Boolean, server_default='0')
created_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now()) created_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now())
created_by_id: Mapped[int] = mapped_column(Integer, ForeignKey('user.id'), nullable=True) # automaticky import nebude mit ID created_by_id: Mapped[int] = mapped_column(Integer, ForeignKey('user.id'), nullable=True) # automaticky import nebude mit ID # noqa
deleted: Mapped[bool] = mapped_column(Boolean, nullable=False, server_default='0') deleted: Mapped[bool] = mapped_column(Boolean, nullable=False, server_default='0')
metars: Mapped['Metar'] = relationship(back_populates="airport") metars: Mapped['Metar'] = relationship(back_populates="airport")
-35
View File
@@ -1,35 +0,0 @@
import asyncio
import functools
from typing import Callable
from database import async_session
def transactional(func: Callable) -> Callable:
@functools.wraps(func)
async def _wrapper(*args, **kwargs):
# db_session = db_session_context.get()
# if db_session:
# return func(*args, **kwargs)
db_session = async_session()
print("STARTUJI TRANSAKCI")
db_session.begin()
# db_session_context.set(db_session)
try:
kwargs['db'] = db_session
result = await func(*args, **kwargs)
await db_session.commit()
print("KONEC TRANSAKCE V DEKORATORU")
except Exception as e:
await db_session.rollback()
raise
finally:
await db_session.close()
# db_session_context.set(None)
return result
return _wrapper
+1 -1
View File
@@ -9,6 +9,6 @@ def error_logging(func):
try: try:
return await func(*args, **kwargs) return await func(*args, **kwargs)
except NoResultFound as e: except NoResultFound as e:
raise GraphQLError(f"Not found", original_error=e) raise GraphQLError("Not found", original_error=e)
return decorator return decorator
+2 -14
View File
@@ -1,8 +1,4 @@
from contextlib import asynccontextmanager from contextlib import asynccontextmanager
from typing import AsyncGenerator
from sqlalchemy.ext.asyncio import AsyncSession
from database import async_session from database import async_session
@@ -13,18 +9,10 @@ async def get_session():
try: try:
yield session yield session
await session.commit() await session.commit()
except: except Exception as e:
await session.rollback() await session.rollback()
print(f"ERROR: {e}")
raise raise
finally: finally:
session.expunge_all() session.expunge_all()
await session.close() await session.close()
async def db_session():
async with async_session() as session:
async with session.begin():
yield session
await session.flush()
await session.commit()
print("CCCCCCCCCCCCCOOOOOOOOOOOOOOMMMMMMMMMMIIIIIIIIIIIITTTTTTTTTTTT")
-6
View File
@@ -1,12 +1,6 @@
from fastapi_jwt.jwt import JwtAccess, JwtRefresh from fastapi_jwt.jwt import JwtAccess, JwtRefresh
class BaseEndpoint(object):
def __init__(self, db):
self.db = db
super().__init__()
class AuthEndpoint: class AuthEndpoint:
def __init__(self, *args, **kwargs): def __init__(self, *args, **kwargs):
self.access_security: JwtAccess = kwargs.pop("access_token") self.access_security: JwtAccess = kwargs.pop("access_token")
+12 -8
View File
@@ -4,7 +4,8 @@ from passlib.hash import bcrypt
from sqlalchemy import select from sqlalchemy import select
from starlette.responses import Response from starlette.responses import Response
from database.models import User from database.models import User
from endpoints.base import BaseEndpoint, AuthEndpoint from dependencies.db import get_session
from endpoints.base import AuthEndpoint
from pydantic import BaseModel from pydantic import BaseModel
@@ -13,18 +14,21 @@ class LoginInput(BaseModel):
password: str password: str
class LoginEndpoint(AuthEndpoint, BaseEndpoint): class LoginEndpoint(AuthEndpoint):
async def on_post(self, user_data: LoginInput, resp: Response) -> dict: async def on_post(self, user_data: LoginInput, resp: Response) -> dict:
query = select(User).filter_by(email=user_data.email) query = select(User).filter_by(email=user_data.email)
logged_user = (await self.db.scalars(query)).first()
if not logged_user: async with get_session() as db:
raise HTTPException(status_code=401, detail="Invalid user") logged_user = (await db.scalars(query)).first()
if not logged_user:
raise HTTPException(status_code=401, detail="Invalid user")
if not bcrypt.verify(user_data.password, logged_user.password_hashed): user = logged_user.as_dict()
if not bcrypt.verify(user_data.password, user['password_hashed']):
raise HTTPException(status_code=401, detail="Bad username or password") raise HTTPException(status_code=401, detail="Bad username or password")
subject = {"id": logged_user.id, "email": logged_user.email} subject = {"id": user['id'], "email": user['email']}
access_token = self.access_security.create_access_token(subject=subject) access_token = self.access_security.create_access_token(subject=subject)
refresh_token = self.refresh_security.create_refresh_token(subject=subject) refresh_token = self.refresh_security.create_refresh_token(subject=subject)
@@ -35,7 +39,7 @@ class LoginEndpoint(AuthEndpoint, BaseEndpoint):
) )
return { return {
"user": logged_user.as_dict(), "user": user,
"access_token": access_token, "access_token": access_token,
"access_token_validity": self.access_security.access_expires_delta.total_seconds(), "access_token_validity": self.access_security.access_expires_delta.total_seconds(),
} }
+13 -16
View File
@@ -4,7 +4,7 @@ from sqlalchemy import select
from typing import Optional from typing import Optional
from pydantic import BaseModel, root_validator, Field from pydantic import BaseModel, root_validator, Field
from database.models import User from database.models import User
from endpoints.base import BaseEndpoint from dependencies.db import get_session
from passlib.hash import bcrypt from passlib.hash import bcrypt
@@ -23,22 +23,19 @@ class RegistrationInput(BaseModel):
return values return values
class RegistrationEndpoint(BaseEndpoint): class RegistrationEndpoint:
async def on_post(self, user_data: RegistrationInput) -> User: async def on_post(self, user_data: RegistrationInput) -> User:
query = select(User).filter_by(email=user_data.email) query = select(User).filter_by(email=user_data.email)
existing_user = (await self.db.scalars(query)).first() async with get_session() as db:
existing_user = (await db.scalars(query)).first()
if existing_user: if existing_user:
raise HTTPException(status_code=422, detail="User already exists") raise HTTPException(status_code=422, detail="User already exists")
model = await User.create(self.db, { model = await User.create(db, {
"name": user_data.name, "name": user_data.name,
"email": user_data.email, "email": user_data.email,
"password_hashed": bcrypt.hash(user_data.password), "password_hashed": bcrypt.hash(user_data.password),
"description": "" "description": ""
}) })
return model.as_dict()
await self.db.commit()
return model.as_dict()
+1 -2
View File
@@ -1,5 +1,4 @@
from typing import List, Tuple, Dict from typing import List, Dict
import aiohttp import aiohttp
+3 -2
View File
@@ -37,7 +37,7 @@ class Weather:
query_string = urllib.parse.urlencode(params) query_string = urllib.parse.urlencode(params)
return f"{url}{query_string}" return f"{url}{query_string}"
@cached(ttl=6*3600) @cached(ttl=6 * 3600)
async def download_weather_for_day(self, date: datetime.date, gps: Tuple[float, float]): async def download_weather_for_day(self, date: datetime.date, gps: Tuple[float, float]):
url = self.get_weather_info_url(start_date=date, end_date=date, gps=gps) url = self.get_weather_info_url(start_date=date, end_date=date, gps=gps)
@@ -46,7 +46,8 @@ class Weather:
resp.raise_for_status() resp.raise_for_status()
return await resp.json() return await resp.json()
async def get_weather_for_hour(self, date_time: datetime.datetime, gps: Tuple[float, float]) -> Dict[str, float|str]: async def get_weather_for_hour(
self, date_time: datetime.datetime, gps: Tuple[float, float]) -> Dict[str, float | str]:
data = await self.download_weather_for_day(date_time.date(), gps) data = await self.download_weather_for_day(date_time.date(), gps)
# TODO: kontrola timezone! # TODO: kontrola timezone!
-13
View File
@@ -1,13 +0,0 @@
import asyncio
from external.elevation import ElevationAPI
from external.gpx_parser import GPXParser
async def test():
elevation_api = ElevationAPI()
parser = GPXParser("./uploads/tracks/37b979cb-71c0-4d09-a2c1-cfdad5f7a0cf-OK-AUR_28_AUR_Bristell_NG5_Zapisnik_letu_2023-08-04-00 00_2023-08-04-12 00(1).gpx")
points = await parser.get_coordinates()
elevation = await elevation_api.get_elevation_for_points(points)
parser.add_terrain_elevation(elevation)
asyncio.run(test())
+8 -3
View File
@@ -34,7 +34,12 @@ class FlightsLoader:
return [result_data[id_] for id_ in ids] 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) flights_by_aircraft_dataloader = DataLoader(load_fn=FlightsLoader(Flight.aircraft_id).load, cache=False)
flight_by_poi_dataloader = DataLoader(
flight_by_poi_dataloader = DataLoader(load_fn=FlightsLoader(PointOfInterest.id, extra_join=[Flight.track, PointOfInterest]).load, cache=False) 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: if input.photo:
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)
return await models.Aircraft.create( async with get_session() as db:
db, aircraft = await models.Aircraft.create(
data=dict( db,
**input_data, data=dict(
created_by_id=info.context.user_id, **input_data,
created_by_id=info.context.user_id,
)
) )
)
return Aircraft(**aircraft.as_dict())
@strawberry.type @strawberry.type
-4
View File
@@ -53,7 +53,6 @@ class CopilotQueries:
@strawberry.type @strawberry.type
class CreateCopilotMutation: class CreateCopilotMutation:
@strawberry_sqlalchemy_input(model=models.Copilot, exclude_fields=["id"]) @strawberry_sqlalchemy_input(model=models.Copilot, exclude_fields=["id"])
class CreateCopilotInput: class CreateCopilotInput:
pass pass
@@ -76,7 +75,6 @@ class CreateCopilotMutation:
@strawberry.type @strawberry.type
class EditCopilotMutation: class EditCopilotMutation:
@strawberry_sqlalchemy_input(model=models.Copilot, exclude_fields=["id"]) @strawberry_sqlalchemy_input(model=models.Copilot, exclude_fields=["id"])
class EditCopilotInput: class EditCopilotInput:
pass pass
@@ -84,12 +82,10 @@ class EditCopilotMutation:
@strawberry.mutation @strawberry.mutation
@authenticated_user_only() @authenticated_user_only()
async def edit_copilot(root, info, id: int, input: EditCopilotInput) -> Copilot: async def edit_copilot(root, info, id: int, input: EditCopilotInput) -> Copilot:
async with get_session() as db: async with get_session() as db:
copilot = (await db.scalars( copilot = (await db.scalars(
get_base_query(info.context.user_id).filter(models.Copilot.id == id) get_base_query(info.context.user_id).filter(models.Copilot.id == id)
)).one() )).one()
updated_copilot = await models.Copilot.update(db, obj=copilot, data=input.to_dict()) 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 import asyncio
from datetime import timedelta, datetime from datetime import timedelta, datetime
from typing import List, Optional, Annotated, TYPE_CHECKING, Tuple from typing import List, Optional, Annotated, TYPE_CHECKING
import strawberry import strawberry
from fastapi import HTTPException from fastapi import HTTPException
from lxml import etree
from sqlalchemy import select, insert, delete from sqlalchemy import select, insert, delete
from starlette.status import HTTP_401_UNAUTHORIZED from starlette.status import HTTP_401_UNAUTHORIZED
from strawberry.file_uploads import Upload 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.entities.poi import PointOfInterest
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 get_public_url 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 from ..types import ComboboxInput
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -134,7 +136,7 @@ class Flight:
return await flight_copilots_dataloader.load(root.id) return await flight_copilots_dataloader.load(root.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) copilots: Optional[List[Annotated["Copilot", strawberry.lazy(".copilot")]]] = strawberry.field(resolver=load_copilots) # noqa
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)
@@ -220,7 +222,9 @@ class CreateFlightMutation:
data = input.to_dict() data = input.to_dict()
async with get_session() as db: 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) aircraft_id = await handle_aircraft_save(db, info.context.user_id, input.aircraft)
weather_takeoff, weather_landing = await asyncio.gather( weather_takeoff, weather_landing = await asyncio.gather(
@@ -231,8 +235,8 @@ class CreateFlightMutation:
flight = await models.Flight.create(db, data={ flight = await models.Flight.create(db, data={
**data, **data,
"takeoff_weather_info_id": weather_takeoff.id, "takeoff_weather_info_id": weather_takeoff.id if weather_takeoff else None,
"landing_weather_info_id": weather_landing.id, "landing_weather_info_id": weather_landing.id if weather_landing else None,
"takeoff_airport_id": takeoff_airport.id, "takeoff_airport_id": takeoff_airport.id,
"landing_airport_id": landing_airport.id, "landing_airport_id": landing_airport.id,
"has_terrain_elevation": False, "has_terrain_elevation": False,
@@ -268,9 +272,7 @@ class EditFlightMutation:
)).one() )).one()
takeoff_airport, landing_airport = await get_airports( takeoff_airport, landing_airport = await get_airports(
db, db, input.takeoff_airport, input.landing_airport, info.context.user_id,
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,
) )
data = input.to_dict() data = input.to_dict()
@@ -278,7 +280,7 @@ class EditFlightMutation:
if input.gpx_track is not None: if input.gpx_track is not None:
data['gpx_track_filename'] = await handle_upload_gpx(flight, input.gpx_track) data['gpx_track_filename'] = await handle_upload_gpx(flight, input.gpx_track)
info.context.background_tasks.add_task( 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 ( if (
+44 -20
View File
@@ -1,26 +1,33 @@
import asyncio import asyncio
from datetime import datetime from datetime import datetime
from typing import List, Type, Literal, Optional, Tuple from typing import List, Type, Literal, Optional, Tuple
from aiohttp import ClientResponseError from aiohttp import ClientResponseError
from sqlalchemy import select, delete from sqlalchemy import select, delete
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from starlette.background import BackgroundTasks
from strawberry.file_uploads import Upload from strawberry.file_uploads import Upload
from database import models 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.gpx_parser import GPXParser
from external.weather import Weather from external.weather import Weather
from graphql_schema.types import ComboboxInput 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() weather_api = Weather()
async def handle_weather_info(db: AsyncSession, date_time: datetime, airport: models.Airport) -> models.WeatherInfo: async def handle_weather_info(db: AsyncSession, date_time: datetime, airport: models.Airport) -> Optional[models.WeatherInfo]:
weather = await weather_api.get_weather_for_hour(date_time.astimezone(), (airport.gps_latitude, airport.gps_longitude)) 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(**{ model = models.WeatherInfo(**{
"datetime": weather['datetime'], "datetime": weather['datetime'],
"qnh": weather['pressure_msl'], "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) poi_object = poi_map.get(item.id)
if not poi_object: 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 db.flush()
await models.FlightTrack.create( 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( takeoff_airport = (await db.scalars(
select(models.Airport).filter(models.Airport.id == takeoff_airport_id) select(models.Airport).filter(models.Airport.id == takeoff_airport_id)
)).one() )).one()
@@ -103,39 +127,39 @@ async def handle_airport_changed(
): ):
flight_datetime = getattr(flight, f"{type_}_datetime") flight_datetime = getattr(flight, f"{type_}_datetime")
if input_datetime and input_datetime != flight_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") existing_weather_id = getattr(flight, f"{type_}_weather_info_id")
if existing_weather_id: if existing_weather_id:
# db.delete(delete()) # db.delete(delete())
pass 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_}_airport_id", airport.id)
setattr(flight, f"{type_}_datetime", input_datetime) 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 path = "/app/uploads/tracks" # TODO vytahnout do configu
gpx_parser = GPXParser(f"{path}/{gpx_filename}") gpx_parser = GPXParser(f"{path}/{gpx_filename}")
coordinates = await gpx_parser.get_coordinates() coordinates = await gpx_parser.get_coordinates()
print("AAAAAAAAAAAAAAAAAAAAAAAAA", coordinates)
try: try:
elevation = await elevation_api.get_elevation_for_points(coordinates) elevation = await elevation_api.get_elevation_for_points(coordinates)
print("ELEVATION", elevation)
tree_with_elevation = gpx_parser.add_terrain_elevation(elevation) tree_with_elevation = gpx_parser.add_terrain_elevation(elevation)
output_name = f"terrain_{gpx_filename}" output_name = f"terrain_{gpx_filename}"
gpx_parser.write(tree_with_elevation, f"{path}/{output_name}") 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: async with get_session() as db:
print("NEumim elevation!") 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): 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 decorators.endpoints import authenticated_user_only
from dependencies.db import get_session from dependencies.db import get_session
from graphql_schema.dataloaders.poi import poi_dataloader 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.sqlalchemy_to_strawberry_type import strawberry_sqlalchemy_type
from graphql_schema.types import ComboboxInput 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 from .helpers.flight import handle_combobox_save
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -37,7 +38,7 @@ class Photo:
url: str = strawberry.field(resolver=resolve_url) url: str = strawberry.field(resolver=resolve_url)
thumbnail_url: str = strawberry.field(resolver=resolve_thumb_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): def get_base_query(user_id: int):
@@ -161,10 +162,7 @@ class DeletePhotoMutation:
photo = Photo(**photo_model.as_dict()) photo = Photo(**photo_model.as_dict())
base_path = get_photo_basepath(photo.flight_id) base_path = get_photo_basepath(photo.flight_id)
try: delete_file(f"{base_path}/{photo.filename}", silent=True)
delete_file(f"{base_path}/{photo.filename}") delete_file(f"{base_path}/thumbs/{photo.filename}", silent=True)
delete_file(f"{base_path}/thumbs/{photo.filename}")
except Exception as e:
print(e)
return photo_model return photo_model
+1 -2
View File
@@ -1,11 +1,10 @@
from typing import List, Optional from typing import List
import strawberry import strawberry
from sqlalchemy import select, or_ 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 from graphql_schema.sqlalchemy_to_strawberry_type import strawberry_sqlalchemy_type
from graphql_schema.types import ComboboxInput
@strawberry_sqlalchemy_type(models.PointOfInterestType) @strawberry_sqlalchemy_type(models.PointOfInterestType)
+1 -2
View File
@@ -1,4 +1,3 @@
from functools import wraps
from typing import Optional from typing import Optional
import strawberry import strawberry
from graphql import GraphQLError from graphql import GraphQLError
@@ -79,7 +78,7 @@ class EditUserMutation:
)).one() )).one()
user_image_path = f"/app/uploads/profile/{user.id}" 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 input.avatar_image:
if user.avatar_image_filename: if user.avatar_image_filename:
delete_file(f"{user_image_path}/{user.avatar_image_filename}", silent=True) delete_file(f"{user_image_path}/{user.avatar_image_filename}", silent=True)
-1
View File
@@ -2,7 +2,6 @@ import dataclasses
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
from sqlalchemy.ext.asyncio import AsyncSession
from starlette.background import BackgroundTasks from starlette.background import BackgroundTasks
from strawberry.extensions import SchemaExtension from strawberry.extensions import SchemaExtension
from strawberry.fastapi import BaseContext from strawberry.fastapi import BaseContext
@@ -41,7 +41,6 @@ def strawberry_sqlalchemy_type(model, exclude_fields: Optional[typing.Union[List
return wrapper return wrapper
def strawberry_sqlalchemy_input( def strawberry_sqlalchemy_input(
model, model,
exclude_fields: Optional[typing.Union[List, typing.Tuple]] = None, exclude_fields: Optional[typing.Union[List, typing.Tuple]] = None,
+4 -7
View File
@@ -1,14 +1,12 @@
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.ext.asyncio import AsyncSession
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 dependencies.db import db_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
@@ -97,23 +95,22 @@ class App:
).on_post(resp, credentials) ).on_post(resp, credentials)
@self.api_router.post("/login") @self.api_router.post("/login")
async def login(resp: Response, user: LoginInput, db: AsyncSession = Depends(db_session)): async def login(resp: Response, user: LoginInput):
return await LoginEndpoint( return await LoginEndpoint(
db=db,
access_token=self.access_security, access_token=self.access_security,
refresh_token=self.refresh_security refresh_token=self.refresh_security
).on_post(user, resp) ).on_post(user, resp)
@self.api_router.post("/logout") @self.api_router.post("/logout")
async def login(resp: Response): async def logout(resp: Response):
return await LogoutEndpoint( return await LogoutEndpoint(
access_token=self.access_security, access_token=self.access_security,
refresh_token=self.refresh_security refresh_token=self.refresh_security
).on_post(resp) ).on_post(resp)
@self.api_router.post("/registration", status_code=201) @self.api_router.post("/registration", status_code=201)
async def registration(user: RegistrationInput, db: AsyncSession = Depends(db_session)): async def registration(user: RegistrationInput):
return await RegistrationEndpoint(db=db).on_post(user) return await RegistrationEndpoint().on_post(user)
# musi byt na konci # musi byt na konci
app.include_router(self.api_router) app.include_router(self.api_router)
+3
View File
@@ -0,0 +1,3 @@
PHOTO_BASE_PATH = ""
AIRCRAFT_BASE_PATH = ""
FLIGHT_BASE_PATH = ""
+7 -6
View File
@@ -1,13 +1,12 @@
import asyncio import asyncio
import sys import sys
from sqlalchemy import select from sqlalchemy import select
sys.path.insert(0, "/app/src") sys.path.insert(0, "/app/src")
from database import async_session, models from database import async_session, models # noqa
from external.elevation import elevation_api from external.elevation import elevation_api # noqa
from external.gpx_parser import GPXParser from external.gpx_parser import GPXParser # noqa
async def add_elevation_to_photos(): async def add_elevation_to_photos():
@@ -17,7 +16,9 @@ async def add_elevation_to_photos():
.filter(models.Photo.terrain_elevation.is_(None)) .filter(models.Photo.terrain_elevation.is_(None))
)).all() )).all()
coordinates = [{"lat": p.gps_latitude, "lng": p.gps_longitude} for p in photos if p.gps_latitude or p.gps_longitude] coordinates = [
{"lat": p.gps_latitude, "lng": p.gps_longitude} for p in photos if p.gps_latitude or p.gps_longitude
]
photos_by_corrdinates = {(p.gps_latitude, p.gps_longitude): p for p in photos} photos_by_corrdinates = {(p.gps_latitude, p.gps_longitude): p for p in photos}
if not coordinates: if not coordinates:
print("all done") print("all done")
@@ -35,7 +36,7 @@ async def add_elevation_to_tracks():
async with async_session() as session: async with async_session() as session:
flights = (await session.scalars( flights = (await session.scalars(
select(models.Flight) select(models.Flight)
.filter(models.Flight.has_terrain_elevation == False) .filter(models.Flight.has_terrain_elevation.is_(False))
.filter(models.Flight.gpx_track_filename.isnot(None)) .filter(models.Flight.gpx_track_filename.isnot(None))
)).all() )).all()
+4 -2
View File
@@ -62,7 +62,9 @@ async def parse_exif_info(path: str, filename: str) -> dict:
return exif_info return exif_info
async def resize_image(path: str, filename: str, new_width: int, quality: int = 90, dest_path: str = None, dest_filename: str = None): async def resize_image(
path: str, filename: str, new_width: int, quality: int = 90, dest_path: str = None, dest_filename: str = None
):
if not dest_path: if not dest_path:
dest_path = path dest_path = path
@@ -78,7 +80,7 @@ async def resize_image(path: str, filename: str, new_width: int, quality: int =
image = image.resize((new_width, new_height), Image.LANCZOS) image = image.resize((new_width, new_height), Image.LANCZOS)
check_directories(dest_path) check_directories(dest_path)
image.save(f"{dest_path}/{dest_filename}", 'JPEG', quality=quality) image.save(f"{dest_path}/{dest_filename}", 'JPEG', quality=quality)
except UnidentifiedImageError as e: except UnidentifiedImageError:
pass pass