import asyncio from typing import Optional from sqlalchemy import delete, select from sqlalchemy.dialects.mysql import insert from database import models from database.models import flight_plan_has_copilot from database.transaction import get_session from graphql_schema.entities.helpers.combobox import handle_combobox_save from graphql_schema.entities.resolvers.base import BaseMutationResolver, BaseQueryResolver from graphql_schema.entities.resolvers.flight import handle_aircraft_save from graphql_schema.entities.types.mutation_input import CreateFlightPlanInput, EditFlightPlanInput from graphql_schema.entities.types.types import FlightPlan from utils.str_utils import random_str class FlightPlanQueryResolver(BaseQueryResolver): def __init__(self): super().__init__(graphql_type=FlightPlan, model=models.FlightPlan) def get_query( self, user_id: Optional[int] = None, only_public: Optional[bool] = False, object_id: Optional[int] = None, *args, **kwargs ): filters = {"object_id": object_id} if object_id else {} query = super().get_query( user_id, **filters, order_by=[models.FlightPlan.planned_takeoff_datetime.desc(), models.FlightPlan.id.desc()], only_public=only_public, only_my=not only_public ) if kwargs.get('username'): query = ( query.join(models.FlightPlan.created_by) .filter(models.User.public_username == kwargs['username']) ) return query class FlightPlanMutationResolver(BaseMutationResolver): def __init__(self): super().__init__(graphql_type=FlightPlan, model=models.FlightPlan) @staticmethod async def save_markers(db, flight_plan: models.FlightPlan, markers: list): position = 0 for marker in markers: if marker.type == 'poi': assert bool(marker.point_of_interest_id) if marker.type == 'airport': assert bool(marker.airport_id) await models.FlightPlanMarker.create(db, data={ "position": position, "flight_plan_id": flight_plan.id, "airport_id": marker.airport_id, "point_of_interest_id": marker.point_of_interest_id, "type": marker.type, "name": marker.name, "gps_latitude": marker.gps_latitude, "gps_longitude": marker.gps_longitude }) position += 1 @staticmethod async def reset_plan_markers(db, flight_plan: models.FlightPlan): await db.execute( delete(models.FlightPlanMarker) .filter(models.FlightPlanMarker.flight_plan_id == flight_plan.id) ) async def create(self, context, data: CreateFlightPlanInput) -> FlightPlan: input_data = data.to_dict() input_data['created_by_id'] = context.user_id async with get_session() as db: flight_plan = await self._do_create(db, data=input_data) await self.save_markers(db, flight_plan, data.markers) return flight_plan async def save_copilots(self, db, flight_plan_id: int, copilots: list, user_id: int): await db.execute(delete(flight_plan_has_copilot).filter_by(flight_plan_id=flight_plan_id)) copilots = await asyncio.gather(*[ handle_combobox_save(db, models.Copilot, copilot, user_id) for copilot in copilots ]) for copilot_id in copilots: await db.execute(insert(flight_plan_has_copilot).values( flight_plan_id=flight_plan_id, copilot_id=copilot_id, token=random_str(64) )) async def update(self, id: int, data: EditFlightPlanInput, user_id: int) -> FlightPlan: input_data = data.to_dict() async with get_session() as db: flight_plan_model = await self._get_one(db, id=id, created_by_id=user_id) if data.aircraft is not None: input_data['aircraft_id'] = await handle_aircraft_save(db, user_id, data.aircraft) if data.markers is not None: await self.reset_plan_markers(db, flight_plan_model) await self.save_markers(db, flight_plan_model, data.markers) if flight_plan_model.is_default_name: if data.markers: markers = data.markers else: markers = (await db.scalars(select(models.FlightPlanMarker).filter(models.FlightPlanMarker.flight_plan_id == id))).all() used_markers = evenly_spaced_elements(markers, 5) input_data['name'] = " - ".join(m.name for m in used_markers) if data.copilots is not None: await self.save_copilots( db, flight_plan_id=id, copilots=data.copilots, user_id=user_id ) flight_plan = await self._do_update(db, obj=flight_plan_model, data=input_data) return flight_plan def evenly_spaced_elements(lst: list, count: int) -> list: if count > len(lst): return lst interval = (len(lst) - 1) / (count - 1) if count > 1 else 0 return [lst[int(round(i * interval))] for i in range(count)]