diff --git a/src/background_jobs/elevation.py b/src/background_jobs/elevation.py index cb146dd..463170b 100644 --- a/src/background_jobs/elevation.py +++ b/src/background_jobs/elevation.py @@ -1,4 +1,3 @@ -from aiohttp import ClientResponseError from sqlalchemy import select from database import models from database.transaction import get_session @@ -16,18 +15,15 @@ async def add_terrain_elevation_to_flight(flight_id: int): .join(models.Track.flight) .filter(models.Flight.id == flight_id)) ).all() - try: - await update_track_points_elevation(db, track_points) - except ClientResponseError as e: - print(e) + await update_track_points_elevation(db, track_points) @retryable async def add_terrain_elevation_to_photo(photo): 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 diff --git a/src/external/gpx_parser.py b/src/external/gpx_parser.py index 6cb89fd..216b967 100644 --- a/src/external/gpx_parser.py +++ b/src/external/gpx_parser.py @@ -1,7 +1,6 @@ from collections import defaultdict from datetime import datetime, timedelta from typing import List, Dict, Any -from aiocache import cached from lxml import etree @@ -67,59 +66,47 @@ class GPXParser: def run_xpath(self, path: str): return self.gpx.xpath(path, namespaces=self.namespace) - @cached() async def get_times_all(self): nodes = self.run_xpath("//gpx:trkpt/gpx:time") return [datetime.fromisoformat(node.text).astimezone() for node in nodes] - @cached() async def get_times(self): times = await self.get_times_all() return await self.sample_times(times) - @cached() async def get_total_duration(self) -> timedelta: times = await self.get_times_all() return times[-1] - times[0] - @cached() async def get_coordinates(self) -> List[Dict[str, float]]: return await self.average_coordinates(await self.get_coordinates_all()) - @cached() async def get_coordinates_all(self) -> List[Dict[str, float]]: nodes = self.run_xpath("//gpx:trkpt") return [{"lat": float(node.attrib["lat"]), "lng": float(node.attrib['lon'])} for node in nodes] - @cached() async def get_speed(self) -> List[float]: nodes = self.run_xpath("//gpx:speed") return await self.average_sample_numbers([float(node.text) for node in nodes]) - @cached() async def get_magnetic_variation(self) -> List[float]: nodes = self.run_xpath("//gpx:magvar") return await self.average_sample_numbers([int(node.text) for node in nodes]) - @cached() async def get_altitude(self) -> List[float]: nodes = self.run_xpath("//gpx:ele") return await self.average_sample_numbers([float(node.text) for node in nodes]) - @cached() async def get_terrain_elevation(self) -> List[float]: nodes = self.run_xpath("//gpx:terrain_elevation") return await self.average_sample_numbers([float(node.text) for node in nodes]) - @cached() async def get_max_speed(self): return max(await self.get_speed()) or 0 - @cached() async def get_min_speed(self): return min(await self.get_speed()) or 0 - @cached() async def get_avg_speed(self): speeds = await self.get_speed() if not speeds: @@ -127,11 +114,9 @@ class GPXParser: return round(sum(speeds) / len(speeds), 2) - @cached() async def get_max_altitude(self): return max(await self.get_altitude()) or 0 - @cached() async def get_avg_altitude(self): altitudes = await self.get_altitude() return round(sum(altitudes) / len(altitudes), 2) diff --git a/src/graphql_schema/entities/resolvers/flight.py b/src/graphql_schema/entities/resolvers/flight.py index 91a644b..cd18fbd 100644 --- a/src/graphql_schema/entities/resolvers/flight.py +++ b/src/graphql_schema/entities/resolvers/flight.py @@ -1,16 +1,14 @@ import asyncio import random -import uuid -from typing import List, Optional -from sqlalchemy import delete, insert, select, func, text +from typing import Optional +from sqlalchemy import delete, insert from sqlalchemy.ext.asyncio import AsyncSession -from strawberry.file_uploads import Upload from background_jobs.elevation import add_terrain_elevation_to_flight +from utils.flight_track_helpers import handle_upload_gpx, save_track_from_gpx_to_db, extract_basic_flight_info_from_gpx from background_jobs.weather import download_weather_for_flight from database import models from database.models import flight_has_copilot from database.transaction import get_session -from external.gpx_parser import GPXParser from graphql_schema.entities.helpers.combobox import handle_combobox_save from graphql_schema.entities.resolvers.base import BaseMutationResolver, BaseQueryResolver from graphql_schema.entities.types.mutation_input import ( @@ -18,9 +16,8 @@ from graphql_schema.entities.types.mutation_input import ( ) from graphql_schema.entities.types.types import Flight from paths import FLIGHT_GPX_TRACK_PATH -from utils.file import delete_file -from utils.str_utils import random_str from utils.file import handle_file_upload +from utils.str_utils import random_str class FlightQueryResolver(BaseQueryResolver): @@ -75,55 +72,13 @@ class FlightMutationResolver(BaseMutationResolver): def __init__(self): super().__init__(Flight, models.Flight) - async def get_airport_id_by_gps(self, gps_lat: float, gps_lng: float) -> Optional[int]: - async with (get_session() as db): - query = ( - select(models.Airport, func.coalesce(6371 * func.acos( - func.cos(func.radians(gps_lat)) * - func.cos(func.radians(models.Airport.gps_latitude)) * - func.cos(func.radians(models.Airport.gps_longitude) - func.radians(gps_lng)) + - func.sin(func.radians(gps_lat)) * - func.sin(func.radians(models.Airport.gps_latitude)) - ), 9999).label("distance")) - .filter(models.Airport.use_in_gpx_guess.is_(True)) - .order_by("distance") - .having(text("distance < 1")) - .limit(1) - ) - - data = (await db.execute(query)).one_or_none() - if data: - airport, distance = data - return airport.id - - return None - - async def extract_data_from_gpx(self, gpx_filename: str) -> dict: - data = GPXParser(f"{FLIGHT_GPX_TRACK_PATH}/{gpx_filename}") - - times, coordinates = await asyncio.gather( - data.get_times(), - data.get_coordinates() - ) - takeoff_airport_id, landing_airport_id = await asyncio.gather( - self.get_airport_id_by_gps(coordinates[0]['lat'], coordinates[0]['lng']), - self.get_airport_id_by_gps(coordinates[-1]['lat'], coordinates[-1]['lng']), - ) - - return { - "takeoff_airport_id": takeoff_airport_id, - "landing_airport_id": landing_airport_id, - "takeoff_datetime": times[0], - "landing_datetime": times[-1], - } - async def create(self, context, input: CreateFlightInput) -> Flight: data = input.to_dict() user_id = context.user_id if input.gpx_track_file: - data['gpx_track_filename'] = await handle_upload_gpx(gpx_track=input.gpx_track_file, context=context) - data_from_gpx = await self.extract_data_from_gpx(data['gpx_track_filename']) + data['gpx_track_filename'] = await handle_file_upload(input.gpx_track_file, FLIGHT_GPX_TRACK_PATH) + data_from_gpx = await extract_basic_flight_info_from_gpx(data['gpx_track_filename']) data.update(data_from_gpx) else: async with get_session() as db: @@ -153,16 +108,10 @@ class FlightMutationResolver(BaseMutationResolver): if input.track is not None: await handle_track_edit(db=db, flight_id=flight.id, track=input.track, user_id=user_id) - context.background_tasks.add_task( - download_weather_for_flight, - flight_id=flight.id, airport_id=flight.takeoff_airport_id, date_time=flight.takeoff_datetime, - type_="takeoff" - ) - context.background_tasks.add_task( - download_weather_for_flight, - flight_id=flight.id, airport_id=flight.landing_airport_id, date_time=flight.landing_datetime, - type_="landing" - ) + if data['gpx_track_filename']: + await save_track_from_gpx_to_db(gpx_filename=data['gpx_track_filename'], flight_id=flight.id) + + schedule_background_tasks(flight.id, data, context) return flight @@ -176,40 +125,23 @@ class FlightMutationResolver(BaseMutationResolver): data = input.to_dict() if input.gpx_track_file is not None: - data['gpx_track_filename'] = await handle_upload_gpx( - gpx_track=input.gpx_track_file, - context=context, - original_gpx_filename=flight_data['gpx_track_filename'] - ) + data['gpx_track_filename'] = await handle_upload_gpx(gpx_track=input.gpx_track_file, flight_id=flight_id) async with get_session() as db: if input.takeoff_airport: - takeoff_airport_id = await handle_combobox_save( + data['takeoff_airport_id'] = await handle_combobox_save( db, models.Airport, input.takeoff_airport, user_id, name_column="icao_code", extra_data={"name": input.takeoff_airport.name} ) - data['takeoff_airport_id'] = takeoff_airport_id data['takeoff_datetime'] = input.takeoff_datetime or flight_data['takeoff_datetime'] - context.background_tasks.add_task( - download_weather_for_flight, flight_id=id, airport_id=takeoff_airport_id, date_time=data['takeoff_datetime'], - type_="takeoff" - ) - if input.landing_airport: - landing_airport_id = await handle_combobox_save( + data['landing_airport_id'] = await handle_combobox_save( db, models.Airport, input.landing_airport, user_id, name_column="icao_code", extra_data={"name": input.landing_airport.name} ) - - data['landing_airport_id'] = landing_airport_id data['landing_datetime'] = input.landing_datetime or flight_data['landing_datetime'] - context.background_tasks.add_task( - download_weather_for_flight, flight_id=id, airport_id=landing_airport_id, date_time=data['landing_datetime'], - type_="landing" - ) - if input.aircraft is not None: data['aircraft_id'] = await handle_aircraft_save(db, user_id, input.aircraft) @@ -238,20 +170,26 @@ class FlightMutationResolver(BaseMutationResolver): token=random_str(64) )) - return await self._do_update(db, flight_data, data) + flight_model = await self._do_update(db, flight_data, data) + + schedule_background_tasks(flight_id, data, context) + + return flight_model -async def handle_upload_gpx(gpx_track: Upload, context, original_gpx_filename: Optional[str] = None): - if original_gpx_filename: - delete_file(FLIGHT_GPX_TRACK_PATH + "/" + original_gpx_filename, silent=True) - - filename = await handle_file_upload(gpx_track, FLIGHT_GPX_TRACK_PATH) - context.background_tasks.add_task(add_terrain_elevation_to_flight, flight_id=id, gpx_filename=filename) - - return filename +def schedule_background_tasks(flight_id: int, flight_data: dict, context) -> None: + context.background_tasks.add_task(add_terrain_elevation_to_flight, flight_id=flight_id) + context.background_tasks.add_task( + download_weather_for_flight, flight_id=id, airport_id=flight_data['takeoff_airport_id'], + date_time=flight_data['takeoff_datetime'], type_="takeoff" + ) + context.background_tasks.add_task( + download_weather_for_flight, flight_id=id, airport_id=flight_data['landing_airport_id'], + date_time=flight_data['landing_datetime'], type_="landing" + ) -async def handle_track_edit(db: AsyncSession, flight_id: int, track: List[TrackItemInput], user_id: int): +async def handle_track_edit(db: AsyncSession, flight_id: int, track: list[TrackItemInput], user_id: int): await db.execute(delete(models.FlightTurnPoint).filter(models.FlightTurnPoint.flight_id == flight_id)) order = 0 diff --git a/src/scripts/elevation.py b/src/scripts/elevation.py index 8009366..ac2fa2f 100644 --- a/src/scripts/elevation.py +++ b/src/scripts/elevation.py @@ -11,31 +11,27 @@ from database.transaction import get_session async def add_elevation_to_photos(): - async with async_session() as session: - photos = (await session.scalars( + async with get_session() as db: + photos = (await db.scalars( select(models.Photo) .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 + {"lat": p.gps_latitude, "lng": p.gps_longitude, "id": p.id} for p in photos if p.gps_latitude or p.gps_longitude ] - photos_by_coordinates = {(p.gps_latitude, p.gps_longitude): p for p in photos} if not coordinates: print("all done") return points = await elevation_api.get_elevation_for_points(coordinates) for point in points: - photo = photos_by_coordinates[point.lat, point.lng] - await models.Photo.update(db_session=session, obj=photo, data={"terrain_elevation": point.elevation}) - await session.flush() - await session.commit() + await models.Photo.update(db_session=db, id=point.id, data={"terrain_elevation": point.elevation}) async def add_elevation_to_tracks(): async with get_session() as db: - track_points = (await db.session.scalars( + track_points = (await db.scalars( select(models.TrackPoint) .filter(models.TrackPoint.terrain_elevation.is_(None)) )).all() diff --git a/src/scripts/migrate_gpx_to_db.py b/src/scripts/migrate_gpx_to_db.py index 2797692..6301cd5 100644 --- a/src/scripts/migrate_gpx_to_db.py +++ b/src/scripts/migrate_gpx_to_db.py @@ -1,80 +1,24 @@ import asyncio import sys -from sqlalchemy import select, delete +from sqlalchemy import select sys.path.insert(0, "/app/src") -from database import async_session, models # noqa +from utils.flight_track_helpers import save_track_from_gpx_to_db +from database import models from database.transaction import get_session -from external.gpx_parser import GPXParser -from paths import FLIGHT_GPX_TRACK_PATH - - -def get_bounds(coordinates) -> list[tuple[float, float]]: - latitudes = [c['lat'] for c in coordinates] - longitudes = [c['lng'] for c in coordinates] - - return [ - (min(latitudes), min(longitudes)), - (min(latitudes), max(longitudes)), - (max(latitudes), min(longitudes)), - (max(latitudes), max(latitudes)) - ] - - -async def migrate_flight(db, flight: models.Flight): - try: - gpx_parser = GPXParser(file=f"{FLIGHT_GPX_TRACK_PATH}/{flight.gpx_track_filename}") - except OSError as e: - print(f"Cannot process {flight.id=}: {flight.gpx_track_filename}: {e}") - return - - altitudes = await gpx_parser.get_altitude() - terrain_elevations = await gpx_parser.get_terrain_elevation() - coordinates = await gpx_parser.get_coordinates() - speeds = await gpx_parser.get_speed() - magnetic_variations = await gpx_parser.get_magnetic_variation() - times = await gpx_parser.get_times() - - track_data = { - "bounds": get_bounds(coordinates), - "min_speed": await gpx_parser.get_min_speed(), - "avg_speed": await gpx_parser.get_avg_speed(), - "max_speed": await gpx_parser.get_max_speed(), - "total_duration": (await gpx_parser.get_total_duration()).seconds, - "max_altitude": await gpx_parser.get_max_altitude(), - "avg_altitude": await gpx_parser.get_avg_altitude(), - } - - if flight.track_id is None: - track = await models.Track.create(db, { - **track_data, - "created_by_id": flight.created_by_id - }) - flight.track_id = track.id - else: - await db.execute(delete(models.TrackPoint).filter(models.TrackPoint.track_id == flight.track_id)) - track = await models.Track.get_one(db, id=flight.track_id) - await models.Track.update(db, data=track_data, obj=track) - - for i in range(len(coordinates)): - await models.TrackPoint.create(db, { - "track_id": track.id, - "altitude": altitudes[i] if i < len(altitudes) else None, - "magnetic_variation": magnetic_variations[i] if i < len(magnetic_variations) else None, - "terrain_elevation": terrain_elevations[i] if i < len(terrain_elevations) else None, - "speed": speeds[i] if i < len(speeds) else None, - "gps_latitude": coordinates[i]['lat'], - "gps_longitude": coordinates[i]['lng'], - "timestamp": times[i] - }) async def migrate_gpx(): async with get_session() as db: - flights = (await db.scalars(select(models.Flight).filter(models.Flight.gpx_track_filename.is_not(None)))).all() + flights = ( + await db.execute( + select(models.Flight.id, models.Flight.gpx_track_filename) + .filter(models.Flight.gpx_track_filename.is_not(None))) + ).all() + flight_tracks = {f.id: f.gpx_track_filename for f in flights} - for flight in flights: - await migrate_flight(db, flight) + for id, gpx_filename in flight_tracks.items(): + await save_track_from_gpx_to_db(flight_id=id, gpx_filename=gpx_filename) if __name__ == "__main__": diff --git a/src/utils/flight_track_helpers.py b/src/utils/flight_track_helpers.py new file mode 100644 index 0000000..c951fc6 --- /dev/null +++ b/src/utils/flight_track_helpers.py @@ -0,0 +1,125 @@ +import asyncio +from typing import Optional + +from sqlalchemy import delete, func, select, text +from sqlalchemy import delete +from strawberry.file_uploads import Upload + +from database import models +from database.transaction import get_session +from external.gpx_parser import GPXParser +from paths import FLIGHT_GPX_TRACK_PATH +from utils.file import handle_file_upload, delete_file + + +def get_bounds(coordinates: list[dict[str, float]]) -> list[tuple[float, float]]: + latitudes = [c['lat'] for c in coordinates] + longitudes = [c['lng'] for c in coordinates] + + return [ + (min(latitudes), min(longitudes)), + (min(latitudes), max(longitudes)), + (max(latitudes), min(longitudes)), + (max(latitudes), max(latitudes)) + ] + + +async def get_airport_id_by_gps(gps_lat: float, gps_lng: float) -> Optional[int]: + async with get_session() as db: + query = ( + select(models.Airport, func.coalesce(6371 * func.acos( + func.cos(func.radians(gps_lat)) * + func.cos(func.radians(models.Airport.gps_latitude)) * + func.cos(func.radians(models.Airport.gps_longitude) - func.radians(gps_lng)) + + func.sin(func.radians(gps_lat)) * + func.sin(func.radians(models.Airport.gps_latitude)) + ), 9999).label("distance")) + .filter(models.Airport.use_in_gpx_guess.is_(True)) + .order_by("distance") + .having(text("distance < 1")) + .limit(1) + ) + + data = (await db.execute(query)).one_or_none() + if data: + airport, distance = data + return airport.id + + return None + + +async def extract_basic_flight_info_from_gpx(gpx_filename: str) -> dict: + data = GPXParser(f"{FLIGHT_GPX_TRACK_PATH}/{gpx_filename}") + + times, coordinates = await asyncio.gather( + data.get_times(), + data.get_coordinates() + ) + takeoff_airport_id, landing_airport_id = await asyncio.gather( + get_airport_id_by_gps(coordinates[0]['lat'], coordinates[0]['lng']), + get_airport_id_by_gps(coordinates[-1]['lat'], coordinates[-1]['lng']), + ) + + return { + "takeoff_airport_id": takeoff_airport_id, + "landing_airport_id": landing_airport_id, + "takeoff_datetime": times[0], + "landing_datetime": times[-1], + } + + +async def handle_upload_gpx(gpx_track: Upload, flight_id: int): + filename = await handle_file_upload(gpx_track, FLIGHT_GPX_TRACK_PATH) + await save_track_from_gpx_to_db(gpx_filename=filename, flight_id=flight_id) + return filename + + +async def save_track_from_gpx_to_db(gpx_filename: str, flight_id: int | None = None): + try: + gpx_parser = GPXParser(file=f"{FLIGHT_GPX_TRACK_PATH}/{gpx_filename}") + except OSError as e: + print(f"Cannot process {gpx_filename}: {e}") + return + + altitudes = await gpx_parser.get_altitude() + terrain_elevations = await gpx_parser.get_terrain_elevation() + coordinates = await gpx_parser.get_coordinates() + speeds = await gpx_parser.get_speed() + magnetic_variations = await gpx_parser.get_magnetic_variation() + times = await gpx_parser.get_times() + + track_data = { + "bounds": get_bounds(coordinates), + "min_speed": await gpx_parser.get_min_speed(), + "avg_speed": await gpx_parser.get_avg_speed(), + "max_speed": await gpx_parser.get_max_speed(), + "total_duration": (await gpx_parser.get_total_duration()).seconds, + "max_altitude": await gpx_parser.get_max_altitude(), + "avg_altitude": await gpx_parser.get_avg_altitude(), + } + + async with get_session() as db: + flight = await models.Flight.get_one(db, id=flight_id) + + if flight.track_id is None: + track = await models.Track.create(db, { + **track_data, + "created_by_id": flight.created_by_id + }) + flight.track_id = track.id + else: + await db.execute(delete(models.TrackPoint).filter(models.TrackPoint.track_id == flight.track_id)) + track = await models.Track.get_one(db, id=flight.track_id) + await models.Track.update(db, data=track_data, obj=track) + + for i in range(len(coordinates)): + await models.TrackPoint.create(db, { + "track_id": track.id, + "altitude": altitudes[i] if i < len(altitudes) else None, + "magnetic_variation": magnetic_variations[i] if i < len(magnetic_variations) else None, + "terrain_elevation": terrain_elevations[i] if i < len(terrain_elevations) else None, + "speed": speeds[i] if i < len(speeds) else None, + "gps_latitude": coordinates[i]['lat'], + "gps_longitude": coordinates[i]['lng'], + "timestamp": times[i] + })