Refaktoring resolveru

This commit is contained in:
Michal Kváček
2023-09-25 10:04:03 +02:00
parent 22a28ae70e
commit d820b92964
16 changed files with 253 additions and 252 deletions
+26
View File
@@ -0,0 +1,26 @@
from aiohttp import ClientResponseError
from database import models
from dependencies.db import get_session
from external.elevation import elevation_api
from external.gpx_parser import GPXParser
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()
try:
elevation = await elevation_api.get_elevation_for_points(coordinates)
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}")
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)
+9 -39
View File
@@ -1,13 +1,14 @@
from typing import List, Optional, Annotated, TYPE_CHECKING, Set
import strawberry
from strawberry.file_uploads import Upload
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, strawberry_sqlalchemy_input
from upload_utils import handle_file_upload, delete_file, get_public_url
from .helpers.flight import handle_combobox_save
from graphql_schema.entities.helpers.combobox import handle_combobox_save
from .resolvers.aircraft import get_aircraft_resolver
from .resolvers.base import get_list, get_one
from ..dataloaders.flight import flights_by_aircraft_dataloader
from ..dataloaders.organizations import organizations_dataloader
from ..types import ComboboxInput
@@ -36,43 +37,19 @@ class Aircraft:
)
def get_base_query(user_id: int, organization_ids: Set[int]):
return (
select(models.Aircraft)
.filter(
or_(
models.Aircraft.created_by_id == user_id,
models.Aircraft.organization_id.in_(organization_ids)
)
)
.filter(models.Aircraft.deleted.is_(False))
.order_by(models.Aircraft.id.desc())
)
@strawberry.type
class AircraftQueries:
@strawberry.field()
@authenticated_user_only()
async def aircrafts(root, info) -> List[Aircraft]:
async with get_session() as db:
aircrafts = (await db.scalars(
get_base_query(info.context.user_id, info.context.organization_ids)
)).all()
return [Aircraft(**a.as_dict()) for a in aircrafts]
query = get_aircraft_resolver(info.context.user_id, info.context.organization_ids)
return await get_list(models.Aircraft, query)
@strawberry.field()
@authenticated_user_only()
async def aircraft(root, info, id: int) -> Aircraft:
query = (
get_base_query(info.context.user_id, info.context.organization_ids)
.filter(models.Aircraft.id == id)
)
async with get_session() as db:
aircraft = (await db.scalars(query)).one()
return Aircraft(**aircraft.as_dict())
query = get_aircraft_resolver(info.context.user_id, info.context.organization_ids, id)
return await get_one(models.Aircraft, query)
@strawberry.type
@@ -92,7 +69,6 @@ class CreateAircraftMutation:
input_data['photo_filename'] = await handle_file_upload(input.photo, AIRCRAFT_UPLOAD_DEST_PATH)
async with get_session() as db:
if input.organization:
input_data['organization_id'] = await handle_combobox_save(
db,
@@ -134,10 +110,7 @@ class EditAircraftMutation:
user_id=info.context.user_id,
)
aircraft = (await db.scalars(
get_base_query(info.context.user_id, set()) # TODO: bude fungovat prazdny set?
.filter(models.Aircraft.id == id)
)).one()
aircraft = (await db.scalars(get_aircraft_resolver(user_id=info.context.user_id, aircraft_id=id))).one()
if input.photo:
if aircraft.photo_filename:
@@ -156,10 +129,7 @@ class DeleteAircraftMutation:
@authenticated_user_only()
async def delete_aircraft(self, info, id: int) -> Aircraft:
async with get_session() as db:
aircraft = (await db.scalars(
get_base_query(info.context.user_id)
.filter(models.Aircraft.id == id)
)).one()
aircraft = (await db.scalars(get_aircraft_resolver(info.context.user_id, aircraft_id=id))).one()
aircraft = await models.Aircraft.update(db, obj=aircraft, data=dict(deleted=True))
return Aircraft(**aircraft.as_dict())
+6 -27
View File
@@ -1,9 +1,9 @@
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.entities.resolvers.airport import get_airport_resolver
from graphql_schema.entities.resolvers.base import get_list, get_one
from graphql_schema.sqlalchemy_to_strawberry_type import strawberry_sqlalchemy_type
@@ -12,37 +12,16 @@ class Airport:
pass
def get_base_query(user_id: int):
return (
select(models.Airport)
.filter(models.Airport.deleted.is_(False))
.filter(or_(
models.Airport.created_by_id == user_id,
models.Airport.created_by_id.is_(None),
))
.order_by(models.Airport.icao_code)
)
@strawberry.type
class AirportQueries:
@strawberry.field()
@authenticated_user_only()
async def airports(root, info) -> List[Airport]:
query = get_base_query(info.context.user_id)
async with get_session() as db:
airports = (await db.scalars(query)).all()
return [Airport(**a.as_dict()) for a in airports]
query = get_airport_resolver(info.context.user_id)
return await get_list(models.Airport, query)
@strawberry.field()
@authenticated_user_only()
async def airport(root, info, id: int) -> Airport:
query = (
get_base_query(info.context.user_id)
.filter(models.Airport.id == id)
)
async with get_session() as db:
airport = (await db.scalars(query)).one()
return Airport(**airport.as_dict())
query = get_airport_resolver(info.context.user_id, id)
return await get_one(models.Airport, query)
+6 -22
View File
@@ -6,6 +6,7 @@ from decorators.endpoints import authenticated_user_only
from dependencies.db import get_session
from graphql_schema.dataloaders.flight import flights_by_copilot_dataloader
from graphql_schema.sqlalchemy_to_strawberry_type import strawberry_sqlalchemy_type, strawberry_sqlalchemy_input
from .resolvers.base import get_base_resolver, get_list, get_one
if TYPE_CHECKING:
from .flight import Flight
@@ -19,36 +20,19 @@ class Copilot:
flights: List[Annotated["Flight", strawberry.lazy('.flight')]] = strawberry.field(resolver=load_flights)
def get_base_query(user_id: int):
return (
select(models.Copilot)
.filter(models.Copilot.created_by_id == user_id)
.filter(models.Copilot.deleted.is_(False))
.order_by(models.Copilot.name)
)
@strawberry.type
class CopilotQueries:
@strawberry.field()
@authenticated_user_only()
async def copilots(root, info) -> List[Copilot]:
async with get_session() as db:
copilots = (await db.scalars(
get_base_query(info.context.user_id)
)).all()
return [Copilot(**c.as_dict()) for c in copilots]
query = get_base_resolver(models.Copilot, user_id=info.context.user_id, order_by=[models.Copilot.name])
return await get_list(models.Copilot, query)
@strawberry.field()
@authenticated_user_only()
async def copilot(root, info, id: int) -> Copilot:
async with get_session() as db:
copilot = (await db.scalars(
get_base_query(info.context.user_id)
.filter(models.Copilot.id == id)
)).one()
return Copilot(**copilot.as_dict())
query = get_base_resolver(models.Copilot, object_id=id, user_id=info.context.user_id)
return await get_one(models.Copilot, query)
@strawberry.type
@@ -84,7 +68,7 @@ class EditCopilotMutation:
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)
query=get_base_resolver(models.Copilot, object_id=id, user_id=info.context.user_id)
)).one()
updated_copilot = await models.Copilot.update(db, obj=copilot, data=input.to_dict())
+9 -22
View File
@@ -6,6 +6,7 @@ from decorators.endpoints import authenticated_user_only
from dependencies.db import get_session
from graphql_schema.dataloaders.flight import flights_by_event_dataloader
from graphql_schema.sqlalchemy_to_strawberry_type import strawberry_sqlalchemy_type, strawberry_sqlalchemy_input
from .resolvers.base import get_base_resolver, get_list, get_one
if TYPE_CHECKING:
from .flight import Flight
@@ -19,36 +20,22 @@ class Event:
flights: List[Annotated["Flight", strawberry.lazy('.flight')]] = strawberry.field(resolver=load_flights)
def get_base_query(user_id: int):
return (
select(models.Event)
.filter(models.Event.created_by_id == user_id)
.filter(models.Event.deleted.is_(False))
.order_by(models.Event.date_from.desc(), models.Event.id.desc())
)
@strawberry.type
class EventQueries:
@strawberry.field()
@authenticated_user_only()
async def events(root, info) -> List[Event]:
async with get_session() as db:
events = (await db.scalars(
get_base_query(info.context.user_id)
)).all()
return [Event(**c.as_dict()) for c in events]
query = get_base_resolver(
models.Event, user_id=info.context.user_id,
order_by=[models.Event.date_from.desc(), models.Event.id.desc()]
)
return await get_list(models.Event, query)
@strawberry.field()
@authenticated_user_only()
async def event(root, info, id: int) -> Event:
async with get_session() as db:
event = (await db.scalars(
get_base_query(info.context.user_id)
.filter(models.Event.id == id)
)).one()
return Event(**event.as_dict())
query = get_base_resolver(models.Event, user_id=info.context.user_id, object_id=id)
return await get_one(models.Event, query)
@strawberry.type
@@ -84,7 +71,7 @@ class EditEventMutation:
async def edit_event(root, info, id: int, input: EditEventInput) -> Event:
async with get_session() as db:
event = (await db.scalars(
get_base_query(info.context.user_id).filter(models.Event.id == id)
get_base_resolver(models.Event, user_id=info.context.user_id, object_id=id)
)).one()
updated_event = await models.Event.update(db, obj=event, data=input.to_dict())
@@ -0,0 +1,28 @@
from typing import Type, Optional
from sqlalchemy.ext.asyncio import AsyncSession
from database import models
from graphql_schema.types import ComboboxInput
async def handle_combobox_save(
db: AsyncSession,
model: Type[models.BaseModel],
input: ComboboxInput,
user_id: int,
name_column: str = "name",
extra_data: Optional[dict] = None
) -> int:
if input.id:
return input.id
else:
if not extra_data:
extra_data = {}
data = {name_column: input.name, **extra_data}
if hasattr(model, "created_by_id"):
data["created_by_id"] = user_id
obj = await model.create(db, data)
await db.flush()
return obj.id
+1 -47
View File
@@ -1,15 +1,12 @@
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 strawberry.file_uploads import Upload
from database import models
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.entities.helpers.combobox import handle_combobox_save
from graphql_schema.types import ComboboxInput
from upload_utils import delete_file, handle_file_upload
@@ -141,27 +138,6 @@ async def handle_airport_changed(
setattr(flight, f"{type_}_datetime", input_datetime)
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()
try:
elevation = await elevation_api.get_elevation_for_points(coordinates)
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}")
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):
path = "/app/uploads/tracks"
@@ -176,25 +152,3 @@ async def handle_copilots_edit(db: AsyncSession, copilots: List[ComboboxInput],
return await asyncio.gather(*cors)
async def handle_combobox_save(
db: AsyncSession,
model: Type[models.BaseModel],
input: ComboboxInput,
user_id: int,
name_column: str = "name",
extra_data: Optional[dict] = None
) -> int:
if input.id:
return input.id
else:
if not extra_data:
extra_data = {}
data = {name_column: input.name, **extra_data}
if hasattr(model, "created_by_id"):
data["created_by_id"] = user_id
obj = await model.create(db, data)
await db.flush()
return obj.id
+9 -26
View File
@@ -1,13 +1,13 @@
from typing import List, Annotated, TYPE_CHECKING
import strawberry
from sqlalchemy import select, or_, delete
from sqlalchemy import select, delete
from sqlalchemy.dialects.mysql import insert
from sqlalchemy.exc import IntegrityError
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, strawberry_sqlalchemy_input
from .resolvers.base import get_base_resolver, get_list, get_one
from ..dataloaders.aircraft import aircrafts_from_organization_dataloader
from ..dataloaders.users import users_in_organization_dataloader
@@ -29,33 +29,19 @@ class Organization:
aircrafts: List[Annotated["Aircraft", strawberry.lazy(".aircraft")]] = strawberry.field(resolver=load_aircrafts)
def get_base_query():
return (
select(models.Organization)
.filter(models.Organization.deleted.is_(False))
.order_by(models.Organization.name)
)
@strawberry.type
class OrganizationQueries:
@strawberry.field()
@authenticated_user_only()
async def organizations(root, info) -> List[Organization]:
async with get_session() as db:
organizations = (await db.scalars(get_base_query())).all()
return [Organization(**c.as_dict()) for c in organizations]
query = get_base_resolver(models.Organization, order_by=[models.Organization.name])
return await get_list(models.Organization, query)
@strawberry.field()
@authenticated_user_only()
async def organization(root, info, id: int) -> Organization:
async with get_session() as db:
organization = (await db.scalars(
get_base_query()
.filter(models.Organization.id == id)
)).one()
return Organization(**organization.as_dict())
query = get_base_resolver(models.Organization, object_id=id)
return await get_one(models.Organization, query)
@strawberry.type
@@ -87,7 +73,7 @@ class OrganizationUserMutation:
@authenticated_user_only()
async def add_to_organization(root, info, organization_id: int) -> Organization:
async with get_session() as db:
organization = (await db.scalars(get_base_query().filter(models.Organization.id == organization_id))).one()
organization = (await db.scalars(get_base_resolver(models.Organization, object_id=organization_id))).one()
try:
await db.execute(
@@ -106,7 +92,7 @@ class OrganizationUserMutation:
@authenticated_user_only()
async def remove_from_organization(root, info, organization_id: int) -> Organization:
async with get_session() as db:
organization = (await db.scalars(get_base_query().filter(models.Organization.id == organization_id))).one()
organization = (await db.scalars(get_base_resolver(models.Organization, object_id=organization_id))).one()
await db.execute(
delete(models.user_is_in_organization).filter_by(
@@ -118,7 +104,6 @@ class OrganizationUserMutation:
return Organization(**organization.as_dict())
@strawberry.type
class EditOrganizationMutation:
@strawberry_sqlalchemy_input(model=models.Organization, exclude_fields=["id"])
@@ -130,9 +115,7 @@ class EditOrganizationMutation:
async def edit_organization(root, info, id: int, input: EditOrganizationInput) -> Organization:
async with get_session() as db:
organization = (await db.scalars(
get_base_query()
.filter(models.Organization.created_by_id == info.context.user_id)
.filter(models.Organization.id == id)
get_base_resolver(models.Organization, user_id=info.context.user_id, object_id=id)
)).one()
updated_organization = await models.Organization.update(db, obj=organization, data=input.to_dict())
+10 -10
View File
@@ -13,7 +13,8 @@ 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, rotate_image
)
from .helpers.flight import handle_combobox_save
from graphql_schema.entities.helpers.combobox import handle_combobox_save
from .resolvers.base import get_base_resolver, get_list
if TYPE_CHECKING:
from .poi import PointOfInterest
@@ -58,11 +59,8 @@ def get_photo_basepath(flight_id: int) -> str:
class PhotoQueries:
@strawberry.field()
async def photos(root, info) -> List[Photo]:
query = get_base_query(info.context.user_id)
async with get_session() as db:
photos = (await db.scalars(query)).all()
return [Photo(**photo.as_dict()) for photo in photos]
query = get_base_resolver(models.Photo, user_id=info.context.user_id)
return await get_list(models.Photo, query)
@strawberry.type
@@ -97,6 +95,7 @@ class UploadPhotoMutation:
}, db_session=db)
photo = Photo(**photo_model.as_dict())
# TODO: udelat primo konkretni bg joby na resize a thumbnaily
info.context.background_tasks.add_task(resize_image, path=path, filename=filename, new_width=2500, quality=85)
info.context.background_tasks.add_task(generate_thumbnail, path=path, filename=filename)
@@ -118,15 +117,16 @@ class EditPhotoMutation:
@strawberry.mutation()
@authenticated_user_only()
async def edit_photo(self, info, id: int, input: EditPhotoInput) -> Photo:
query = get_base_query(info.context.user_id)
# TODO: base trida pro inputy s definici to_dict/as_dict?
data = {
key: getattr(input, key) for key in ('name', 'description', 'is_flight_cover')
if getattr(input, key) is not None
}
query = get_base_resolver(models.Photo, user_id=info.context.user_id, object_id=id)
async with get_session() as db:
photo = (await db.scalars(query.filter(models.Photo.id == id))).one()
photo = (await db.scalars(query)).one()
if input.point_of_interest:
data['point_of_interest_id'] = await handle_combobox_save(
@@ -155,8 +155,8 @@ class EditPhotoMutation:
async def rotate_photo(self, info, id: int, angle: int) -> Photo:
async with get_session() as db:
query = get_base_query(info.context.user_id)
photo = (await db.scalars(query.filter(models.Photo.id == id))).one()
query = get_base_resolver(models.Photo, user_id=info.context.user_id, object_id=id)
photo = (await db.scalars(query)).one()
await asyncio.gather(
rotate_image(
+16 -26
View File
@@ -7,10 +7,11 @@ from dependencies.db import get_session
from graphql_schema.dataloaders.flight import flight_by_poi_dataloader
from graphql_schema.dataloaders.photos import poi_photos_dataloader
from graphql_schema.dataloaders.poi import poi_type_dataloader
from graphql_schema.entities.helpers.flight import handle_combobox_save
from graphql_schema.entities.helpers.combobox import handle_combobox_save
from graphql_schema.entities.poi_type import PointOfInterestType
from graphql_schema.sqlalchemy_to_strawberry_type import strawberry_sqlalchemy_type, strawberry_sqlalchemy_input
from graphql_schema.types import ComboboxInput
from .resolvers.base import get_base_resolver, get_list, get_one
if TYPE_CHECKING:
from .flight import Flight
@@ -51,28 +52,17 @@ def get_base_query(user_id: int, only_my: bool = False):
@strawberry.type
class PointOfInterestQueries:
@strawberry.field()
@authenticated_user_only()
async def points_of_interest(root, info) -> List[PointOfInterest]:
query = (
get_base_query(info.context.user_id)
.order_by(models.PointOfInterest.id.desc())
)
async with get_session() as db:
pois = (await db.scalars(query)).all()
return [PointOfInterest(**poi.as_dict()) for poi in pois]
query = get_base_resolver(models.PointOfInterest, user_id=info.context.user_id)
return await get_list(models.PointOfInterest, query)
@strawberry.field()
@authenticated_user_only()
async def point_of_interest(root, info, id: int) -> PointOfInterest:
query = (
get_base_query(info.context.user_id)
.filter(models.PointOfInterest.id == id)
)
async with get_session() as db:
poi = (await db.scalars(query)).one()
return PointOfInterest(**poi.as_dict())
query = get_base_resolver(models.PointOfInterest, user_id=info.context.user_id, object_id=id)
return await get_one(models.PointOfInterest, query)
@strawberry.type
@@ -105,20 +95,19 @@ class EditPointOfInterestMutation:
@strawberry.mutation
@authenticated_user_only()
async def edit_point_of_interest(root, info, id: int, input: EditPointOfInterestInput) -> PointOfInterest:
# TODO: kontrola organizace
input_data = input.to_dict()
query = get_base_resolver(
models.PointOfInterest, user_id=info.context.user_id, object_id=id, include_public=False
)
async with get_session() as db:
if input.type is not None:
input_data['type_id'] = await handle_combobox_save(
db, models.PointOfInterestType, input.type, info.context.user_id
)
poi = (
await db.scalars(
get_base_query(info.context.user_id, only_my=True)
.filter(models.PointOfInterest.id == id)
)).one()
poi = (await db.scalars(query)).one()
updated_poi = await models.PointOfInterest.update(db, obj=poi, data=input_data)
return PointOfInterest(**updated_poi.as_dict())
@@ -129,11 +118,12 @@ class DeletePointOfInterestMutation:
@strawberry.mutation
@authenticated_user_only()
async def delete_point_of_interest(self, info, id: int) -> PointOfInterest:
query = get_base_resolver(
models.PointOfInterest, user_id=info.context.user_id, object_id=id, include_public=False
)
async with get_session() as db:
poi = (await db.scalars(
get_base_query(info.context.user_id, only_my=True)
.filter(models.PointOfInterest.id == id)
)).one()
poi = (await db.scalars(query)).one()
updated_poi = await models.PointOfInterest.update(db, obj=poi, data=dict(deleted=True))
return PointOfInterest(**updated_poi.as_dict())
+5 -33
View File
@@ -3,7 +3,7 @@ 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.entities.resolvers.base import get_base_resolver, get_list, get_one
from graphql_schema.sqlalchemy_to_strawberry_type import strawberry_sqlalchemy_type
@@ -12,48 +12,20 @@ class PointOfInterestType:
pass
def get_base_query(user_id: int, only_my: bool = False):
query = (
select(models.PointOfInterestType)
.filter(models.PointOfInterestType.deleted.is_(False))
)
if only_my:
query = query.filter(models.PointOfInterestType.created_by_id == user_id)
else:
query = query.filter(or_(
models.PointOfInterestType.created_by_id == user_id,
models.PointOfInterestType.is_public.is_(True)
))
return query
@strawberry.type
class PointOfInterestTypeQueries:
@strawberry.field()
@authenticated_user_only()
async def point_of_interest_types(root, info) -> List[PointOfInterestType]:
query = (
get_base_query(info.context.user_id, only_my=False)
.order_by(models.PointOfInterestType.id.desc())
)
async with get_session() as db:
poi_types = (await db.scalars(query)).all()
return [PointOfInterestType(**poi_type.as_dict()) for poi_type in poi_types]
query = get_base_resolver(models.PointOfInterestType, user_id=info.context.user_id)
return await get_list(models.PointOfInterestType, query)
@strawberry.field()
@authenticated_user_only()
async def point_of_interest_type(root, info, id: int) -> PointOfInterestType:
query = (
get_base_query(info.context.user_id)
.filter(models.PointOfInterestType.id == id)
)
async with get_session() as db:
poi_type = (await db.scalars(query)).one()
return PointOfInterestType(**poi_type.as_dict())
query = get_base_resolver(models.PointOfInterestType, user_id=info.context.user_id, object_id=id)
return await get_one(models.PointOfInterestType, query)
#
# @strawberry.type
@@ -0,0 +1,24 @@
from operator import or_
from typing import Set, Optional
from database import models
from graphql_schema.entities.resolvers.base import get_base_resolver
def get_aircraft_resolver(user_id: int, organization_ids: Optional[Set[int]] = None, aircraft_id: Optional[int] = None):
query = get_base_resolver(
model=models.Aircraft,
order_by=[models.Aircraft.id.desc()],
object_id=aircraft_id,
user_id=user_id if not organization_ids else None
)
if organization_ids:
query = (
query.filter(
or_(
models.Aircraft.created_by_id == user_id,
models.Aircraft.organization_id.in_(organization_ids)
)
)
)
return query
@@ -0,0 +1,17 @@
from typing import Optional
from sqlalchemy import or_
from database import models
from graphql_schema.entities.resolvers.base import get_base_resolver
def get_airport_resolver(user_id: int, airport_id: Optional[int] = None):
if airport_id:
return get_base_resolver(model=models.Airport, object_id=airport_id, user_id=user_id)
return (
get_base_resolver(model=models.Airport)
.filter(or_(
models.Airport.created_by_id == user_id,
models.Airport.created_by_id.is_(None),
))
)
@@ -0,0 +1,53 @@
from typing import Optional, Type
from sqlalchemy import select, or_
from database import models
from dependencies.db import get_session
def get_base_resolver(
model: Type[models.BaseModel],
object_id: Optional[int] = None,
user_id: Optional[int] = None,
order_by: Optional[list] = None,
include_public: Optional[bool] = True,
):
query = select(model)
if object_id:
if hasattr(model, "id"):
query = query.filter(model.id == object_id)
else:
raise AssertionError(f"Model {model} has no ID column! Cannot query by ID!")
if hasattr(model, "deleted"):
query = query.filter(model.deleted.is_(False))
ownership_clause = []
if hasattr(model, "is_public") and include_public:
ownership_clause.append(model.is_public.is_(True))
if hasattr(model, "created_by_id") and user_id:
ownership_clause.append(model.created_by_id == user_id)
if len(ownership_clause) > 1:
query = query.filter(or_(*ownership_clause))
elif len(ownership_clause) == 1:
query = query.filter(*ownership_clause)
if order_by:
query = query.order_by(*order_by)
return query
async def get_list(model, query):
async with get_session() as db:
items = (await db.scalars(query)).all()
return [model(**m.as_dict()) for m in items]
async def get_one(model, query):
async with get_session() as db:
data = (await db.scalars(query)).one()
return model(**data.as_dict())