From 10bc0847b61b9e77c791982b498aa3bd84b5a5dd Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Michal=20Kv=C3=A1=C4=8Dek?= Date: Thu, 25 Jul 2024 19:50:09 +0200 Subject: [PATCH] Trasu z GPX nahravat do DB --- ...levation_altitude_to_track_03de3f7f8fdd.py | 46 +++++++++++ ...cated_column_with_altitude_e3620deb41b6.py | 33 ++++++++ src/database/models.py | 40 ++++----- src/external/gpx_parser.py | 11 ++- .../dataloaders/flight_duration.py | 4 +- .../dataloaders/multi_models.py | 19 +++-- .../dataloaders/single_model.py | 1 + .../entities/resolvers/flight.py | 8 +- src/graphql_schema/entities/types/types.py | 68 +++++---------- src/scripts/migrate_gpx_to_db.py | 82 +++++++++++++++++++ 10 files changed, 231 insertions(+), 81 deletions(-) create mode 100644 alembic/versions/20240711-074559_add_speed_elevation_altitude_to_track_03de3f7f8fdd.py create mode 100644 alembic/versions/20240715-063513_delete_duplicated_column_with_altitude_e3620deb41b6.py create mode 100644 src/scripts/migrate_gpx_to_db.py diff --git a/alembic/versions/20240711-074559_add_speed_elevation_altitude_to_track_03de3f7f8fdd.py b/alembic/versions/20240711-074559_add_speed_elevation_altitude_to_track_03de3f7f8fdd.py new file mode 100644 index 0000000..1730851 --- /dev/null +++ b/alembic/versions/20240711-074559_add_speed_elevation_altitude_to_track_03de3f7f8fdd.py @@ -0,0 +1,46 @@ +"""add speed, elevation, altitude to track + +Revision ID: 03de3f7f8fdd +Revises: 8cc01c03e980 +Create Date: 2024-07-11 07:45:59.376941 + +""" +from alembic import op +import sqlalchemy as sa + + +# revision identifiers, used by Alembic. +revision = '03de3f7f8fdd' +down_revision = '8cc01c03e980' +branch_labels = None +depends_on = None + + +def upgrade() -> None: + # ### commands auto generated by Alembic - please adjust! ### + op.add_column('track', sa.Column('min_speed', sa.Float(), nullable=True)) + op.add_column('track', sa.Column('max_speed', sa.Float(), nullable=True)) + op.add_column('track', sa.Column('avg_speed', sa.Float(), nullable=True)) + op.add_column('track', sa.Column('max_altitude', sa.Float(), nullable=True)) + op.add_column('track', sa.Column('avg_altitude', sa.Float(), nullable=True)) + op.add_column('track', sa.Column('total_duration', sa.Integer(), nullable=True, comment='Total duration in seconds')) + op.add_column('track_point', sa.Column('terrain_elevation', sa.Float(), nullable=True)) + op.add_column('track_point', sa.Column('speed', sa.Float(), nullable=True)) + op.add_column('track_point', sa.Column('altitude', sa.Float(), nullable=True)) + op.add_column('track_point', sa.Column('magnetic_variation', sa.Float(), nullable=True)) + # ### end Alembic commands ### + + +def downgrade() -> None: + # ### commands auto generated by Alembic - please adjust! ### + op.drop_column('track_point', 'magnetic_variation') + op.drop_column('track_point', 'altitude') + op.drop_column('track_point', 'speed') + op.drop_column('track_point', 'terrain_elevation') + op.drop_column('track', 'total_duration') + op.drop_column('track', 'avg_altitude') + op.drop_column('track', 'max_altitude') + op.drop_column('track', 'avg_speed') + op.drop_column('track', 'max_speed') + op.drop_column('track', 'min_speed') + # ### end Alembic commands ### diff --git a/alembic/versions/20240715-063513_delete_duplicated_column_with_altitude_e3620deb41b6.py b/alembic/versions/20240715-063513_delete_duplicated_column_with_altitude_e3620deb41b6.py new file mode 100644 index 0000000..6e68b9b --- /dev/null +++ b/alembic/versions/20240715-063513_delete_duplicated_column_with_altitude_e3620deb41b6.py @@ -0,0 +1,33 @@ +"""delete duplicated column with altitude + +Revision ID: e3620deb41b6 +Revises: 03de3f7f8fdd +Create Date: 2024-07-15 06:35:13.536320 + +""" +from alembic import op +import sqlalchemy as sa +from sqlalchemy.dialects import mysql + +# revision identifiers, used by Alembic. +revision = 'e3620deb41b6' +down_revision = '03de3f7f8fdd' +branch_labels = None +depends_on = None + + +def upgrade() -> None: + # ### commands auto generated by Alembic - please adjust! ### + + op.rename_table("flight_track", "flight_turn_point") + op.create_index(op.f('ix_track_point_timestamp'), 'track_point', ['timestamp'], unique=False) + op.drop_column('track_point', 'elevation') + # ### end Alembic commands ### + + +def downgrade() -> None: + # ### commands auto generated by Alembic - please adjust! ### + op.add_column('track_point', sa.Column('elevation', mysql.FLOAT(), nullable=True)) + op.drop_index(op.f('ix_track_point_timestamp'), table_name='track_point') + op.rename_table("flight_turn_point", "flight_track") + # ### end Alembic commands ### diff --git a/src/database/models.py b/src/database/models.py index a90ab8b..298b9a9 100644 --- a/src/database/models.py +++ b/src/database/models.py @@ -124,20 +124,30 @@ class Track(BaseModel): id: Mapped[int] = mapped_column(primary_key=True) bounds: Mapped[list[tuple[float, float]]] = mapped_column(JSON()) + min_speed: Mapped[float] = mapped_column(Float, nullable=True) + max_speed: Mapped[float] = mapped_column(Float, nullable=True) + avg_speed: Mapped[float] = mapped_column(Float, nullable=True) + max_altitude: Mapped[float] = mapped_column(Float, nullable=True) + avg_altitude: Mapped[float] = mapped_column(Float, nullable=True) + total_duration: Mapped[int] = mapped_column(Integer, nullable=True, comment="Total duration in seconds") created_by_id: Mapped[int] = mapped_column(Integer, ForeignKey('user.id')) created_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now()) - created_by: Mapped['User'] = relationship() + + # created_by: Mapped['User'] = relationship() class TrackPoint(BaseModel): __tablename__ = "track_point" id: Mapped[int] = mapped_column(primary_key=True) - timestamp: Mapped[datetime] = mapped_column(DateTime, nullable=False) + timestamp: Mapped[datetime] = mapped_column(DateTime, nullable=False, index=True) track_id: Mapped[id] = mapped_column(Integer, ForeignKey('track.id')) gps_latitude: Mapped[float] = mapped_column(Float, nullable=False) gps_longitude: Mapped[float] = mapped_column(Float, nullable=False) - elevation: Mapped[float] = mapped_column(Float, nullable=True) + terrain_elevation: Mapped[float] = mapped_column(Float, nullable=True) + speed: Mapped[float] = mapped_column(Float, nullable=True) + altitude: Mapped[float] = mapped_column(Float, nullable=True) + magnetic_variation: Mapped[float] = mapped_column(Float, nullable=True) class FlightPlanMarker(BaseModel): @@ -313,22 +323,6 @@ class Aircraft(BaseModel): organization: Mapped['Organization'] = relationship() flights: Mapped[Set['Flight']] = relationship() created_by: Mapped['User'] = relationship() - # notes: Mapped['AircraftNotes'] = relationship() - - -# class AircraftNotes(BaseModel): -# __tablename__ = "aircraft_notes" -# -# id: Mapped[int] = mapped_column(primary_key=True) -# aircraft_id: Mapped[int] = mapped_column(Integer, ForeignKey("aircraft.id"), nullable=False) -# name: Mapped[str] = mapped_column(String(128), nullable=False) -# description: Mapped[str] = mapped_column(Text, nullable=False) -# is_public: Mapped[bool] = mapped_column(Boolean, server_default='0') -# created_by_id: Mapped[int] = mapped_column(Integer, ForeignKey('user.id')) -# created_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now()) -# -# created_by: Mapped['User'] = relationship() -# aircraft: Mapped['Aircraft'] = relationship() class Organization(BaseModel): @@ -344,15 +338,15 @@ class Organization(BaseModel): created_by: Mapped['User'] = relationship() -class FlightTrack(BaseModel): - __tablename__ = "flight_track" +class FlightTurnPoint(BaseModel): + __tablename__ = "flight_turn_point" id: Mapped[int] = mapped_column(primary_key=True) flight_id: Mapped[int] = mapped_column(Integer, ForeignKey("flight.id"), nullable=False) point_of_interest_id: Mapped[int] = mapped_column(Integer, ForeignKey("point_of_interest.id"), nullable=True) airport_id: Mapped[int] = mapped_column(Integer, ForeignKey("airport.id"), nullable=True) landing_duration: Mapped[int] = mapped_column(Integer, nullable=True) - order: Mapped[int] = mapped_column(Integer) + order: Mapped[int] = mapped_column(Integer, index=True) flight: Mapped['Flight'] = relationship() point_of_interest: Mapped['PointOfInterest'] = relationship() @@ -426,7 +420,7 @@ class Flight(BaseModel): landing_airport: Mapped['Airport'] = relationship(foreign_keys=[landing_airport_id]) weather_info_landing: Mapped[WeatherInfo] = relationship(foreign_keys=[landing_weather_info_id]) weather_info_takeoff: Mapped[WeatherInfo] = relationship(foreign_keys=[takeoff_weather_info_id]) - track: Mapped['FlightTrack'] = relationship() + turn_points: Mapped[list['FlightTurnPoint']] = relationship() event: Mapped['Event'] = relationship() copilots: Mapped[List['Copilot']] = relationship(secondary=flight_has_copilot) aircraft: Mapped['Aircraft'] = relationship() diff --git a/src/external/gpx_parser.py b/src/external/gpx_parser.py index 9c344b6..c025bf5 100644 --- a/src/external/gpx_parser.py +++ b/src/external/gpx_parser.py @@ -1,5 +1,5 @@ from collections import defaultdict -from datetime import datetime +from datetime import datetime, timedelta from typing import List, Dict, Any from aiocache import cached from lxml import etree @@ -78,6 +78,11 @@ class GPXParser: 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()) @@ -111,6 +116,10 @@ class GPXParser: 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() diff --git a/src/graphql_schema/dataloaders/flight_duration.py b/src/graphql_schema/dataloaders/flight_duration.py index 96fc02b..2d5f51e 100644 --- a/src/graphql_schema/dataloaders/flight_duration.py +++ b/src/graphql_schema/dataloaders/flight_duration.py @@ -10,8 +10,8 @@ async def load_flight_durations(ids: List[int]): select( models.Flight.id, func.timediff(models.Flight.landing_datetime, models.Flight.takeoff_datetime).label("diff"), - func.coalesce(func.sum(models.FlightTrack.landing_duration), 0).label("landing_duration") - ).join(models.Flight.track, isouter=True) + func.coalesce(func.sum(models.FlightTurnPoint.landing_duration), 0).label("landing_duration") + ).join(models.Flight.turn_points, isouter=True) .group_by(models.Flight.id) .filter(models.Flight.id.in_(ids)) diff --git a/src/graphql_schema/dataloaders/multi_models.py b/src/graphql_schema/dataloaders/multi_models.py index b5cde55..1564bba 100644 --- a/src/graphql_schema/dataloaders/multi_models.py +++ b/src/graphql_schema/dataloaders/multi_models.py @@ -98,7 +98,7 @@ flight_by_poi_dataloader = DataLoader( models.Flight, relationship_column=models.PointOfInterest.id, order_by=[models.Flight.takeoff_datetime.desc()], - extra_join=[models.Flight.track, models.PointOfInterest] + extra_join=[models.Flight.turn_points, models.PointOfInterest] ).load, cache=False ) @@ -155,11 +155,20 @@ poi_photos_dataloader = DataLoader( ).load, cache=False ) -flight_track_dataloader = DataLoader( +flight_turn_points_dataloader = DataLoader( load_fn=MultiModelsDataloader( - models.FlightTrack, - relationship_column=models.FlightTrack.flight_id, - order_by=[models.FlightTrack.order] + models.FlightTurnPoint, + relationship_column=models.FlightTurnPoint.flight_id, + order_by=[models.FlightTurnPoint.order] + ).load, + cache=False +) + +track_points_dataloder = DataLoader( + load_fn=MultiModelsDataloader( + models.TrackPoint, + relationship_column=models.TrackPoint.track_id, + order_by=[models.TrackPoint.timestamp, models.TrackPoint.id] ).load, cache=False ) diff --git a/src/graphql_schema/dataloaders/single_model.py b/src/graphql_schema/dataloaders/single_model.py index 3dd669b..68dba5f 100644 --- a/src/graphql_schema/dataloaders/single_model.py +++ b/src/graphql_schema/dataloaders/single_model.py @@ -15,6 +15,7 @@ aircraft_dataloader = create_dataloader(models.Aircraft) event_dataloader = create_dataloader(models.Event) organizations_dataloader = create_dataloader(models.Organization) airport_weather_info_loader = create_dataloader(models.WeatherInfo) +track_dataloader = create_dataloader(models.Track) poi_dataloader = create_dataloader(models.PointOfInterest) poi_type_dataloader = create_dataloader(models.PointOfInterestType) flight_dataloader = create_dataloader(models.Flight) diff --git a/src/graphql_schema/entities/resolvers/flight.py b/src/graphql_schema/entities/resolvers/flight.py index e40fb9d..91a644b 100644 --- a/src/graphql_schema/entities/resolvers/flight.py +++ b/src/graphql_schema/entities/resolvers/flight.py @@ -58,8 +58,8 @@ class FlightQueryResolver(BaseQueryResolver): if kwargs.get("point_of_interest_id"): query = ( - query.join(models.Flight.track) - .filter(models.FlightTrack.point_of_interest_id == kwargs["point_of_interest_id"]) + query.join(models.Flight.turn_points) + .filter(models.FlightTurnPoint.point_of_interest_id == kwargs["point_of_interest_id"]) ) if kwargs.get('username'): @@ -252,7 +252,7 @@ async def handle_upload_gpx(gpx_track: Upload, context, original_gpx_filename: O async def handle_track_edit(db: AsyncSession, flight_id: int, track: List[TrackItemInput], user_id: int): - await db.execute(delete(models.FlightTrack).filter(models.FlightTrack.flight_id == flight_id)) + await db.execute(delete(models.FlightTurnPoint).filter(models.FlightTurnPoint.flight_id == flight_id)) order = 0 for item in track: @@ -280,7 +280,7 @@ async def handle_track_edit(db: AsyncSession, flight_id: int, track: List[TrackI } ) - await models.FlightTrack.create( + await models.FlightTurnPoint.create( db, data={ "flight_id": flight_id, diff --git a/src/graphql_schema/entities/types/types.py b/src/graphql_schema/entities/types/types.py index 08dbeb8..30473fe 100644 --- a/src/graphql_schema/entities/types/types.py +++ b/src/graphql_schema/entities/types/types.py @@ -1,7 +1,5 @@ from __future__ import annotations - import math -from datetime import datetime from typing import Optional, List import strawberry from database import models @@ -10,17 +8,17 @@ from utils.gps import get_bearing, get_distance from external.gpx_parser import GPXParser from graphql_schema.dataloaders.flight_duration import flight_duration_dataloader from graphql_schema.dataloaders.multi_models import ( - poi_photos_dataloader, flight_by_poi_dataloader, flight_copilots_dataloader, flight_track_dataloader, + poi_photos_dataloader, flight_by_poi_dataloader, flight_copilots_dataloader, flight_turn_points_dataloader, photos_dataloader, flights_by_aircraft_dataloader, users_in_organization_dataloader, aircrafts_from_organization_dataloader, user_organizations_dataloader, flights_by_event_dataloader, flights_by_copilot_dataloader, public_flights_by_event_dataloader, public_flights_by_copilot_dataloader, photo_copilots_dataloader, photos_aircraft_dataloader, copilots_in_photo_dataloader, flight_plan_markers_dataloader, - reporting_points_dataloader, flight_plan_copilots_dataloader, runways_dataloader, frequencies_dataloader + reporting_points_dataloader, flight_plan_copilots_dataloader, runways_dataloader, frequencies_dataloader, track_points_dataloder ) from graphql_schema.dataloaders.single_model import ( poi_dataloader, poi_type_dataloader, event_dataloader, aircraft_dataloader, airport_dataloader, airport_weather_info_loader, organizations_dataloader, flight_dataloader, photo_adjustment_dataloader, - photo_dataloader, user_dataloader + photo_dataloader, user_dataloader, track_dataloader ) from graphql_schema.permissions import IsAuthenticated from graphql_schema.sqlalchemy_to_strawberry_type import strawberry_sqlalchemy_type @@ -33,18 +31,6 @@ class Point: lng: float -@strawberry.type -class GPXTrack: - coordinates: List[Point] - speed: List[float] - altitude: List[float] - magnetic_variation: List[float] - terrain_elevation: List[float] - time: List[datetime] - max_speed: float - avg_speed: float - max_altitude: float - avg_altitude: float @strawberry_sqlalchemy_type(models.ReportingPoint) @@ -88,8 +74,8 @@ class Airspace: ) -@strawberry_sqlalchemy_type(models.FlightTrack) -class FlightTrack: +@strawberry_sqlalchemy_type(models.FlightTurnPoint) +class FlightTurnPoint: point_of_interest: Optional[PointOfInterest] = strawberry.field( resolver=lambda root: poi_dataloader.load(root.point_of_interest_id) ) @@ -133,6 +119,20 @@ class Photo: ) +@strawberry_sqlalchemy_type(models.TrackPoint) +class TrackPoint: + coordinates: Point = strawberry.field(resolver=lambda root: Point(lat=root.gps_latitude, lng=root.gps_longitude)) + + +@strawberry_sqlalchemy_type(models.Track, exclude_fields=['bounds']) +class Track: + bounds: list[tuple[float, float]] + map_bounds: list[Point] = strawberry.field( + resolver=lambda root: [Point(lat=point[1], lng=point[0]) for point in root.bounds] + ) + points: list[TrackPoint] = strawberry.field(resolver=lambda root: track_points_dataloder.load(root.id)) + + @strawberry_sqlalchemy_type(models.Flight) class Flight: def __init__(self, **kwargs): @@ -141,28 +141,6 @@ class Flight: for key, value in kwargs.items(): setattr(self, key, value) - async def load_gpx_track(root): - if not root.gpx_track_filename: - return None - - try: - gpx_parser = GPXParser(f"{FLIGHT_GPX_TRACK_PATH}/{root.gpx_track_filename}") - except OSError: - return None - - return GPXTrack( - coordinates=[Point(**point) for point in await gpx_parser.get_coordinates()], - speed=await gpx_parser.get_speed(), - altitude=await gpx_parser.get_altitude(), - terrain_elevation=await gpx_parser.get_terrain_elevation(), - time=await gpx_parser.get_times(), - max_speed=await gpx_parser.get_max_speed(), - avg_speed=await gpx_parser.get_avg_speed(), - max_altitude=await gpx_parser.get_max_altitude(), - avg_altitude=await gpx_parser.get_avg_altitude(), - magnetic_variation=await gpx_parser.get_magnetic_variation(), - ) - @authenticated_user_only(raise_when_unauthorized=False, return_value_unauthorized=[]) async def load_copilots(root): return await flight_copilots_dataloader.load(root.id) @@ -182,7 +160,7 @@ class Flight: resolver=lambda root: airport_dataloader.load(root.landing_airport_id) ) title_photo: Optional[Photo] = strawberry.field(resolver=lambda root: photo_dataloader.load(root.title_photo_id)) - track: List[FlightTrack] = strawberry.field(resolver=lambda root: flight_track_dataloader.load(root.id)) + turn_points: List[FlightTurnPoint] = strawberry.field(resolver=lambda root: flight_turn_points_dataloader.load(root.id)) takeoff_weather_info: Optional[WeatherInfo] = strawberry.field( resolver=lambda root: airport_weather_info_loader.load(root.takeoff_weather_info_id) ) @@ -190,10 +168,8 @@ class Flight: resolver=lambda root: airport_weather_info_loader.load(root.landing_weather_info_id) ) photos: List[Photo] = strawberry.field(resolver=lambda root: photos_dataloader.load(root.id)) - gpx_track: Optional[GPXTrack] = strawberry.field(resolver=load_gpx_track) - duration_min_calculated: int = strawberry.field( - resolver=lambda root: flight_duration_dataloader.load(root.id) - ) + track: Optional[Track] = strawberry.field(resolver=lambda root: track_dataloader.load(root.track_id)) + duration_min_calculated: int = strawberry.field(resolver=lambda root: flight_duration_dataloader.load(root.id)) social_image_url: Optional[str] = strawberry.field( resolver=lambda root: get_public_url(f'photos/{root.id}/title_photo.jpg') ) diff --git a/src/scripts/migrate_gpx_to_db.py b/src/scripts/migrate_gpx_to_db.py new file mode 100644 index 0000000..2797692 --- /dev/null +++ b/src/scripts/migrate_gpx_to_db.py @@ -0,0 +1,82 @@ +import asyncio +import sys +from sqlalchemy import select, delete + +sys.path.insert(0, "/app/src") +from database import async_session, models # noqa +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() + + for flight in flights: + await migrate_flight(db, flight) + + +if __name__ == "__main__": + loop = asyncio.get_event_loop() + loop.run_until_complete(migrate_gpx())