From 1760dba99deb68a9b8eac7956b118086a622a761 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Michal=20Kv=C3=A1=C4=8Dek?= Date: Wed, 13 Sep 2023 10:04:06 +0200 Subject: [PATCH] Uprava prace s DB, refaktoring --- ...ke_gps_in_airport_nullable_a09857ec9721.py | 38 +++++++++++ db/csvToDb.py | 21 ++++-- src/background_jobs/photo.py | 6 +- src/database/models.py | 6 +- src/decorators/db.py | 35 ---------- src/decorators/error_logging.py | 2 +- src/dependencies/db.py | 16 +---- src/endpoints/base.py | 6 -- src/endpoints/login.py | 20 +++--- src/endpoints/registration.py | 29 ++++----- src/external/elevation.py | 3 +- src/external/weather.py | 7 +- src/gpx_test.py | 13 ---- src/graphql_schema/dataloaders/flight.py | 11 +++- src/graphql_schema/entities/aircraft.py | 15 +++-- src/graphql_schema/entities/copilot.py | 6 +- src/graphql_schema/entities/flight.py | 24 +++---- src/graphql_schema/entities/helpers/flight.py | 64 +++++++++++++------ src/graphql_schema/entities/photo.py | 14 ++-- src/graphql_schema/entities/poi_type.py | 3 +- src/graphql_schema/entities/user.py | 3 +- src/graphql_schema/schema.py | 1 - .../sqlalchemy_to_strawberry_type.py | 1 - src/main.py | 11 ++-- src/paths.py | 3 + src/scripts/elevation.py | 13 ++-- src/upload_utils.py | 6 +- 27 files changed, 195 insertions(+), 182 deletions(-) create mode 100644 alembic/versions/20230913-093206_make_gps_in_airport_nullable_a09857ec9721.py delete mode 100644 src/decorators/db.py delete mode 100644 src/gpx_test.py diff --git a/alembic/versions/20230913-093206_make_gps_in_airport_nullable_a09857ec9721.py b/alembic/versions/20230913-093206_make_gps_in_airport_nullable_a09857ec9721.py new file mode 100644 index 0000000..d021bc9 --- /dev/null +++ b/alembic/versions/20230913-093206_make_gps_in_airport_nullable_a09857ec9721.py @@ -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 ### diff --git a/db/csvToDb.py b/db/csvToDb.py index 9b95670..bf8ffe5 100644 --- a/db/csvToDb.py +++ b/db/csvToDb.py @@ -3,7 +3,16 @@ import csv 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 = h_m.strip("0") # h = int(h_m[0:2]) @@ -16,15 +25,19 @@ def deg_to_dec(val: str): # print("--konec-----------------------------------") # return h + (m / 60.0) + (int(s) / 3600.0) - return len(val or "") + # return val + + # return len(val or "") # ZDAKOV: @49.504378,14.1808905 +# 4930.250N / 100 -> cele cislo stupne, desetinne prevest do sedesatkove soustavy + with open("./poi.csv") as f: reader = csv.DictReader(f) for row in reader: - # if row['name'] != 'ZDAKOV': - # continue + if row['name'] != 'ZDAKOV': + continue print(row['name'], row['lat'], row['lon'], deg_to_dec(row['lat']), deg_to_dec(row['lon'])) \ No newline at end of file diff --git a/src/background_jobs/photo.py b/src/background_jobs/photo.py index 5fd1f36..6ab9593 100644 --- a/src/background_jobs/photo.py +++ b/src/background_jobs/photo.py @@ -6,7 +6,9 @@ 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}]) + elevation = await elevation_api.get_elevation_for_points([ + {"lat": photo.gps_latitude, "lng": photo.gps_longitude} + ]) if not elevation: print("Cannot get elevation") 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}) except Exception as e: print(f"Cannot get elevation: {e}") - - diff --git a/src/database/models.py b/src/database/models.py index 27de7bd..7e2827d 100644 --- a/src/database/models.py +++ b/src/database/models.py @@ -64,12 +64,12 @@ class Airport(BaseModel): id: Mapped[int] = mapped_column(primary_key=True) name: Mapped[str] = mapped_column(String(128), nullable=False) icao_code: Mapped[str] = mapped_column(String(8), nullable=False) - gps_latitude: Mapped[float] = mapped_column(Float, nullable=False) - gps_longitude: Mapped[float] = mapped_column(Float, nullable=False) + gps_latitude: Mapped[float] = mapped_column(Float, nullable=True) + gps_longitude: Mapped[float] = mapped_column(Float, nullable=True) elevation: Mapped[int] = mapped_column(Integer, nullable=True) is_public: Mapped[bool] = mapped_column(Boolean, server_default='0') 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') metars: Mapped['Metar'] = relationship(back_populates="airport") diff --git a/src/decorators/db.py b/src/decorators/db.py deleted file mode 100644 index e16a966..0000000 --- a/src/decorators/db.py +++ /dev/null @@ -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 \ No newline at end of file diff --git a/src/decorators/error_logging.py b/src/decorators/error_logging.py index 137225e..fc08b38 100644 --- a/src/decorators/error_logging.py +++ b/src/decorators/error_logging.py @@ -9,6 +9,6 @@ def error_logging(func): try: return await func(*args, **kwargs) except NoResultFound as e: - raise GraphQLError(f"Not found", original_error=e) + raise GraphQLError("Not found", original_error=e) return decorator diff --git a/src/dependencies/db.py b/src/dependencies/db.py index 2f7e79c..bcb9db7 100644 --- a/src/dependencies/db.py +++ b/src/dependencies/db.py @@ -1,8 +1,4 @@ from contextlib import asynccontextmanager -from typing import AsyncGenerator - -from sqlalchemy.ext.asyncio import AsyncSession - from database import async_session @@ -13,18 +9,10 @@ async def get_session(): try: yield session await session.commit() - except: + except Exception as e: await session.rollback() + print(f"ERROR: {e}") raise finally: session.expunge_all() 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") diff --git a/src/endpoints/base.py b/src/endpoints/base.py index 034652e..ab95653 100644 --- a/src/endpoints/base.py +++ b/src/endpoints/base.py @@ -1,12 +1,6 @@ from fastapi_jwt.jwt import JwtAccess, JwtRefresh -class BaseEndpoint(object): - def __init__(self, db): - self.db = db - super().__init__() - - class AuthEndpoint: def __init__(self, *args, **kwargs): self.access_security: JwtAccess = kwargs.pop("access_token") diff --git a/src/endpoints/login.py b/src/endpoints/login.py index fd2ea5d..99c0c61 100644 --- a/src/endpoints/login.py +++ b/src/endpoints/login.py @@ -4,7 +4,8 @@ from passlib.hash import bcrypt from sqlalchemy import select from starlette.responses import Response 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 @@ -13,18 +14,21 @@ class LoginInput(BaseModel): password: str -class LoginEndpoint(AuthEndpoint, BaseEndpoint): +class LoginEndpoint(AuthEndpoint): async def on_post(self, user_data: LoginInput, resp: Response) -> dict: query = select(User).filter_by(email=user_data.email) - logged_user = (await self.db.scalars(query)).first() - if not logged_user: - raise HTTPException(status_code=401, detail="Invalid user") + async with get_session() as db: + 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") - 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) refresh_token = self.refresh_security.create_refresh_token(subject=subject) @@ -35,7 +39,7 @@ class LoginEndpoint(AuthEndpoint, BaseEndpoint): ) return { - "user": logged_user.as_dict(), + "user": user, "access_token": access_token, "access_token_validity": self.access_security.access_expires_delta.total_seconds(), } diff --git a/src/endpoints/registration.py b/src/endpoints/registration.py index 71f0d61..53fda73 100644 --- a/src/endpoints/registration.py +++ b/src/endpoints/registration.py @@ -4,7 +4,7 @@ from sqlalchemy import select from typing import Optional from pydantic import BaseModel, root_validator, Field from database.models import User -from endpoints.base import BaseEndpoint +from dependencies.db import get_session from passlib.hash import bcrypt @@ -23,22 +23,19 @@ class RegistrationInput(BaseModel): return values -class RegistrationEndpoint(BaseEndpoint): - +class RegistrationEndpoint: async def on_post(self, user_data: RegistrationInput) -> User: 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: - raise HTTPException(status_code=422, detail="User already exists") + if existing_user: + raise HTTPException(status_code=422, detail="User already exists") - model = await User.create(self.db, { - "name": user_data.name, - "email": user_data.email, - "password_hashed": bcrypt.hash(user_data.password), - "description": "" - }) - - await self.db.commit() - - return model.as_dict() + model = await User.create(db, { + "name": user_data.name, + "email": user_data.email, + "password_hashed": bcrypt.hash(user_data.password), + "description": "" + }) + return model.as_dict() diff --git a/src/external/elevation.py b/src/external/elevation.py index 5b4003c..c7bc2ba 100644 --- a/src/external/elevation.py +++ b/src/external/elevation.py @@ -1,5 +1,4 @@ -from typing import List, Tuple, Dict - +from typing import List, Dict import aiohttp diff --git a/src/external/weather.py b/src/external/weather.py index c4163af..33f7a55 100644 --- a/src/external/weather.py +++ b/src/external/weather.py @@ -37,7 +37,7 @@ class Weather: query_string = urllib.parse.urlencode(params) 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]): url = self.get_weather_info_url(start_date=date, end_date=date, gps=gps) @@ -46,7 +46,8 @@ class Weather: resp.raise_for_status() 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) # TODO: kontrola timezone! @@ -55,4 +56,4 @@ class Weather: result_data = {metric: data['hourly'][metric][idx] for metric in self.METRICS} result_data['datetime'] = datetime.datetime.strptime(data['hourly']['time'][idx], "%Y-%m-%dT%H:%M") - return result_data \ No newline at end of file + return result_data diff --git a/src/gpx_test.py b/src/gpx_test.py deleted file mode 100644 index 3c6917d..0000000 --- a/src/gpx_test.py +++ /dev/null @@ -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()) \ No newline at end of file diff --git a/src/graphql_schema/dataloaders/flight.py b/src/graphql_schema/dataloaders/flight.py index c7560dc..be15066 100644 --- a/src/graphql_schema/dataloaders/flight.py +++ b/src/graphql_schema/dataloaders/flight.py @@ -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) \ No newline at end of file +flight_by_poi_dataloader = DataLoader( + load_fn=FlightsLoader(PointOfInterest.id, extra_join=[Flight.track, PointOfInterest]).load, + cache=False +) diff --git a/src/graphql_schema/entities/aircraft.py b/src/graphql_schema/entities/aircraft.py index 6bb0c5b..a6e4183 100644 --- a/src/graphql_schema/entities/aircraft.py +++ b/src/graphql_schema/entities/aircraft.py @@ -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 diff --git a/src/graphql_schema/entities/copilot.py b/src/graphql_schema/entities/copilot.py index f8cd8e6..4915465 100644 --- a/src/graphql_schema/entities/copilot.py +++ b/src/graphql_schema/entities/copilot.py @@ -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()) \ No newline at end of file + return Copilot(**updated_copilot.as_dict()) diff --git a/src/graphql_schema/entities/flight.py b/src/graphql_schema/entities/flight.py index b538d61..3ffdc95 100644 --- a/src/graphql_schema/entities/flight.py +++ b/src/graphql_schema/entities/flight.py @@ -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 ( diff --git a/src/graphql_schema/entities/helpers/flight.py b/src/graphql_schema/entities/helpers/flight.py index 3d32c52..d38591f 100644 --- a/src/graphql_schema/entities/helpers/flight.py +++ b/src/graphql_schema/entities/helpers/flight.py @@ -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): diff --git a/src/graphql_schema/entities/photo.py b/src/graphql_schema/entities/photo.py index 795b77d..8808553 100644 --- a/src/graphql_schema/entities/photo.py +++ b/src/graphql_schema/entities/photo.py @@ -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 diff --git a/src/graphql_schema/entities/poi_type.py b/src/graphql_schema/entities/poi_type.py index 52913ac..94394cd 100644 --- a/src/graphql_schema/entities/poi_type.py +++ b/src/graphql_schema/entities/poi_type.py @@ -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) diff --git a/src/graphql_schema/entities/user.py b/src/graphql_schema/entities/user.py index 84bd05a..8858226 100644 --- a/src/graphql_schema/entities/user.py +++ b/src/graphql_schema/entities/user.py @@ -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) diff --git a/src/graphql_schema/schema.py b/src/graphql_schema/schema.py index 6554007..bfaa002 100644 --- a/src/graphql_schema/schema.py +++ b/src/graphql_schema/schema.py @@ -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 diff --git a/src/graphql_schema/sqlalchemy_to_strawberry_type.py b/src/graphql_schema/sqlalchemy_to_strawberry_type.py index d1c6b89..7764885 100644 --- a/src/graphql_schema/sqlalchemy_to_strawberry_type.py +++ b/src/graphql_schema/sqlalchemy_to_strawberry_type.py @@ -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, diff --git a/src/main.py b/src/main.py index 91cc43e..de9d614 100644 --- a/src/main.py +++ b/src/main.py @@ -1,14 +1,12 @@ from datetime import timedelta from fastapi import FastAPI, APIRouter, Depends, Security from fastapi_jwt import JwtAuthorizationCredentials, JwtAccessBearerCookie, JwtRefreshBearerCookie -from sqlalchemy.ext.asyncio import AsyncSession 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 dependencies.db import db_session from endpoints.login import LoginEndpoint, LoginInput, RefreshEndpoint, LogoutEndpoint from endpoints.registration import RegistrationInput, RegistrationEndpoint from graphql_schema.schema import schema, GraphQLContext @@ -97,23 +95,22 @@ class App: ).on_post(resp, credentials) @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( - db=db, access_token=self.access_security, refresh_token=self.refresh_security ).on_post(user, resp) @self.api_router.post("/logout") - async def login(resp: Response): + async def logout(resp: Response): return await LogoutEndpoint( access_token=self.access_security, refresh_token=self.refresh_security ).on_post(resp) @self.api_router.post("/registration", status_code=201) - async def registration(user: RegistrationInput, db: AsyncSession = Depends(db_session)): - return await RegistrationEndpoint(db=db).on_post(user) + async def registration(user: RegistrationInput): + return await RegistrationEndpoint().on_post(user) # musi byt na konci app.include_router(self.api_router) diff --git a/src/paths.py b/src/paths.py index e69de29..9a29fee 100644 --- a/src/paths.py +++ b/src/paths.py @@ -0,0 +1,3 @@ +PHOTO_BASE_PATH = "" +AIRCRAFT_BASE_PATH = "" +FLIGHT_BASE_PATH = "" diff --git a/src/scripts/elevation.py b/src/scripts/elevation.py index fe45ffc..46f316f 100644 --- a/src/scripts/elevation.py +++ b/src/scripts/elevation.py @@ -1,13 +1,12 @@ import asyncio import sys - from sqlalchemy import select sys.path.insert(0, "/app/src") -from database import async_session, models -from external.elevation import elevation_api -from external.gpx_parser import GPXParser +from database import async_session, models # noqa +from external.elevation import elevation_api # noqa +from external.gpx_parser import GPXParser # noqa async def add_elevation_to_photos(): @@ -17,7 +16,9 @@ async def add_elevation_to_photos(): .filter(models.Photo.terrain_elevation.is_(None)) )).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} if not coordinates: print("all done") @@ -35,7 +36,7 @@ async def add_elevation_to_tracks(): async with async_session() as session: flights = (await session.scalars( 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)) )).all() diff --git a/src/upload_utils.py b/src/upload_utils.py index 7e5d74f..91f9698 100644 --- a/src/upload_utils.py +++ b/src/upload_utils.py @@ -62,7 +62,9 @@ async def parse_exif_info(path: str, filename: str) -> dict: 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: 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) check_directories(dest_path) image.save(f"{dest_path}/{dest_filename}", 'JPEG', quality=quality) - except UnidentifiedImageError as e: + except UnidentifiedImageError: pass