Refaktoring

This commit is contained in:
Michal Kváček
2023-10-14 22:42:51 +02:00
parent 7f7f79dad6
commit bc258efa3a
16 changed files with 274 additions and 184 deletions
+31
View File
@@ -0,0 +1,31 @@
from typing import Optional, Type
from sqlalchemy import select
from database import models
class QueryBuilder:
def __init__(self, model: Type[models.BaseModel]):
self.model = model
def get_simple_query(
self,
extra_select: Optional[list] = None,
created_by_id: Optional[int] = None,
order_by: Optional[list] = None,
include_deleted: bool = False
):
if not extra_select:
extra_select = []
query = select(self.model, *extra_select)
if not include_deleted and hasattr(self.model, "deleted"):
query = query.filter(self.model.deleted.is_(False))
if hasattr(self.model, "created_by_id") and created_by_id:
query = query.filter(self.model.created_by_id == created_by_id)
if order_by:
query = query.order_by(*order_by)
return query
+35 -28
View File
@@ -1,72 +1,79 @@
from collections import defaultdict from collections import defaultdict
from sqlalchemy import select
from typing import Type, List, Optional from typing import Type, List, Optional
from database import models, async_session from database import models, async_session
from database.query_builder import QueryBuilder
class SingleModelByIdDataloader: class BaseDataloader:
def __init__( def __init__(
self, model: Type[models.BaseModel], relationship_column=None, filters: Optional[list] = None self,
model: Type[models.BaseModel],
relationship_column, filters: Optional[list] = None
) -> None: ) -> None:
super().__init__() super().__init__()
self.model = model self.model = model
self.relationship_column = relationship_column if relationship_column else model.id self.query_builder = QueryBuilder(self.model)
if relationship_column is None:
relationship_column = model.id
self.relationship_column = relationship_column
if filters is None: if filters is None:
filters = [] filters = []
self.filters = filters self.filters = filters
class SingleModelByIdDataloader(BaseDataloader):
async def load(self, ids: List[int]): async def load(self, ids: List[int]):
async with async_session() as session: async with async_session() as session:
query = select(self.model, self.relationship_column).filter(self.relationship_column.in_(ids)) query = (
for filter_ in self.filters: self.query_builder.get_simple_query(extra_select=[self.relationship_column], include_deleted=True)
query = query.filter(filter_) .filter(self.relationship_column.in_(ids))
.filter(*self.filters)
)
items = (await session.execute(query)).all() items = (await session.execute(query)).all()
items_by_id = {rel_id: item for item, rel_id in items} items_by_id = {rel_id: item for item, rel_id in items}
return [items_by_id.get(id_) for id_ in ids] return [items_by_id.get(id_) for id_ in ids]
class MultiModelsDataloader: class MultiModelsDataloader(BaseDataloader):
def __init__( def __init__(
self, self,
model: Type[models.BaseModel], model: Type[models.BaseModel],
relationship_column, relationship_column=None,
extra_join: Optional[list] = None,
filters: Optional[list] = None, filters: Optional[list] = None,
extra_join: Optional[list] = None,
order_by: Optional[list] = None, order_by: Optional[list] = None,
): ):
super().__init__(model, relationship_column, filters)
if extra_join is None: if extra_join is None:
extra_join = [] extra_join = []
self.extra_join = extra_join
if filters is None:
filters = []
if order_by is None: if order_by is None:
order_by = [model.id.desc()] # defaultne radit od nejnovejsich zaznamu order_by = [model.id.desc()] # defaultne radit od nejnovejsich zaznamu
self.model = model
self.relationship_column = relationship_column
self.extra_join = extra_join
self.order_by = order_by self.order_by = order_by
self.filters = filters
async def load(self, ids: List[int]): async def load(self, ids: List[int]):
async with async_session() as session: async with async_session() as db:
rel_column = self.relationship_column
query = ( query = (
select(self.model, rel_column) self.query_builder.get_simple_query(
.filter(rel_column.in_(ids)) extra_select=[self.relationship_column],
.order_by(*self.order_by) order_by=self.order_by
)
.filter(self.relationship_column.in_(ids))
.filter(*self.filters)
) )
for table in self.extra_join: for joined_table in self.extra_join:
query = query.join(table) query = query.join(joined_table)
for filter_ in self.filters: if self.filters:
query = query.filter(filter_) query = query.filter(*self.filters)
data = (await session.execute(query)).all() data = (await db.execute(query)).all()
result_data = defaultdict(list) result_data = defaultdict(list)
for item, rel_id in data: for item, rel_id in data:
@@ -48,6 +48,15 @@ flights_by_event_dataloader = DataLoader(
cache=False cache=False
) )
public_flights_by_event_dataloader = DataLoader(
load_fn=MultiModelsDataloader(
models.Flight,
relationship_column=models.Event.id,
filters=[models.Flight.is_public.is_(True)],
extra_join=[models.Flight.event]).load,
cache=False
)
user_organizations_dataloader = DataLoader(load_fn=MultiModelsDataloader( user_organizations_dataloader = DataLoader(load_fn=MultiModelsDataloader(
models.Organization, models.Organization,
relationship_column=models.user_is_in_organization.c.user_id, relationship_column=models.user_is_in_organization.c.user_id,
+22 -9
View File
@@ -1,9 +1,13 @@
from typing import List from typing import List, Optional
import strawberry import strawberry
from fastapi import HTTPException
from starlette.status import HTTP_401_UNAUTHORIZED
from database import models from database import models
from decorators.endpoints import authenticated_user_only from decorators.endpoints import authenticated_user_only
from database.transaction import get_session from database.transaction import get_session
from graphql_schema.entities.resolvers.base import BaseQueryResolver, BaseMutationResolver from graphql_schema.entities.resolvers.base import BaseMutationResolver
from graphql_schema.entities.resolvers.event import EventQueryResolver
from graphql_schema.entities.types.mutation_input import CreateEventInput, EditEventInput from graphql_schema.entities.types.mutation_input import CreateEventInput, EditEventInput
from graphql_schema.entities.types.types import Event from graphql_schema.entities.types.types import Event
@@ -11,17 +15,26 @@ from graphql_schema.entities.types.types import Event
@strawberry.type @strawberry.type
class EventQueries: class EventQueries:
@strawberry.field() @strawberry.field()
@authenticated_user_only() async def events(root, info, username: Optional[str] = None) -> List[Event]:
async def events(root, info) -> List[Event]: if not info.context.user_id and not username:
return await BaseQueryResolver(Event, models.Event).get_list( raise HTTPException(HTTP_401_UNAUTHORIZED)
return await EventQueryResolver().get_list(
info.context.user_id, info.context.user_id,
username=username,
order_by=[models.Event.date_from.desc(), models.Event.id.desc()] order_by=[models.Event.date_from.desc(), models.Event.id.desc()]
) )
@strawberry.field() @strawberry.field()
@authenticated_user_only() async def event(root, info, id: int, username: Optional[str] = None) -> Event:
async def event(root, info, id: int) -> Event: if not info.context.user_id and not username:
return await BaseQueryResolver(Event, models.Event).get_one(id, user_id=info.context.user_id) raise HTTPException(HTTP_401_UNAUTHORIZED)
return await EventQueryResolver().get_one(
id,
username=username,
user_id=info.context.user_id
)
@strawberry.type @strawberry.type
@@ -36,7 +49,7 @@ class EventMutation:
async def edit_event(root, info, id: int, input: EditEventInput) -> Event: async def edit_event(root, info, id: int, input: EditEventInput) -> Event:
async with get_session() as db: async with get_session() as db:
event = (await db.scalars( event = (await db.scalars(
BaseQueryResolver(Event, models.Event).get_query(user_id=info.context.user_id, object_id=id) EventQueryResolver().get_query(user_id=info.context.user_id, object_id=id)
)).one() )).one()
updated_event = await models.Event.update(db, obj=event, data=input.to_dict()) updated_event = await models.Event.update(db, obj=event, data=input.to_dict())
+1 -1
View File
@@ -63,7 +63,7 @@ class FlightMutation:
"name": "", "name": "",
"description": "" "description": ""
}) })
flight = FlightMutationResolver().create(data, info.context.user_id) flight = await FlightMutationResolver().create(data, info.context.user_id)
info.context.background_tasks.add_task( info.context.background_tasks.add_task(
download_weather, download_weather,
+12 -67
View File
@@ -1,21 +1,19 @@
import asyncio import asyncio
from typing import List, Optional from typing import List
import strawberry import strawberry
from sqlalchemy import update
from strawberry.file_uploads import Upload
from background_jobs.elevation import add_terrain_elevation_to_photo from background_jobs.elevation import add_terrain_elevation_to_photo
from background_jobs.photo import generate_thumbnail, resize_photo from background_jobs.photo import generate_thumbnail, resize_photo
from database import models from database import models
from decorators.endpoints import authenticated_user_only from decorators.endpoints import authenticated_user_only
from database.transaction import get_session from database.transaction import get_session
from graphql_schema.entities.resolvers.base import BaseQueryResolver, BaseMutationResolver from graphql_schema.entities.resolvers.base import BaseQueryResolver
from graphql_schema.entities.resolvers.photo import PhotoMutationResolver
from graphql_schema.entities.types.types import Photo from graphql_schema.entities.types.types import Photo
from paths import get_photo_basepath from paths import get_photo_basepath
from graphql_schema.entities.helpers.combobox import handle_combobox_save
from utils.file import delete_file from utils.file import delete_file
from utils.image import parse_exif_info, rotate_image from utils.image import parse_exif_info, rotate_image
from utils.upload import handle_file_upload from utils.upload import handle_file_upload
from graphql_schema.entities.types.mutation_input import ComboboxInput from graphql_schema.entities.types.mutation_input import EditPhotoInput, UploadPhotoInput
@strawberry.type @strawberry.type
@@ -26,15 +24,7 @@ class PhotoQueries:
@strawberry.type @strawberry.type
class UploadPhotoMutation: class PhotoMutation:
@strawberry.input
class UploadPhotoInput:
photo: Upload
flight_id: int
name: Optional[str] = None
description: Optional[str] = None
point_of_interest: Optional[ComboboxInput] = None
@strawberry.mutation @strawberry.mutation
@authenticated_user_only() @authenticated_user_only()
async def upload_photo(self, info, input: UploadPhotoInput) -> Photo: async def upload_photo(self, info, input: UploadPhotoInput) -> Photo:
@@ -42,8 +32,9 @@ class UploadPhotoMutation:
filename = await handle_file_upload(input.photo, path) filename = await handle_file_upload(input.photo, path)
exif_info = await parse_exif_info(path, filename) exif_info = await parse_exif_info(path, filename)
async with get_session() as db: photo = await PhotoMutationResolver().create(
photo_model = await models.Photo.create(data={ user_id=info.context.user_id,
data={
"flight_id": input.flight_id, "flight_id": input.flight_id,
"name": input.name, "name": input.name,
"filename": filename, "filename": filename,
@@ -53,9 +44,8 @@ class UploadPhotoMutation:
"gps_longitude": exif_info.get("gps_longitude"), "gps_longitude": exif_info.get("gps_longitude"),
"gps_altitude": exif_info.get("gps_altitude"), "gps_altitude": exif_info.get("gps_altitude"),
"is_flight_cover": False, "is_flight_cover": False,
"created_by_id": info.context.user_id, },
}, db_session=db) )
photo = Photo(**photo_model.as_dict())
info.context.background_tasks.add_task(resize_photo, path=path, filename=filename) info.context.background_tasks.add_task(resize_photo, path=path, filename=filename)
info.context.background_tasks.add_task(generate_thumbnail, path=path, filename=filename) info.context.background_tasks.add_task(generate_thumbnail, path=path, filename=filename)
@@ -65,52 +55,10 @@ class UploadPhotoMutation:
return photo return photo
@strawberry.type
class EditPhotoMutation:
@strawberry.input
class EditPhotoInput:
name: Optional[str] = None
description: Optional[str] = None
point_of_interest: Optional[ComboboxInput] = None
is_flight_cover: Optional[bool] = None
def to_dict(self):
return {
key: getattr(self, key) for key in ('name', 'description', 'is_flight_cover')
if getattr(self, key) is not None
}
@strawberry.mutation() @strawberry.mutation()
@authenticated_user_only() @authenticated_user_only()
async def edit_photo(self, info, id: int, input: EditPhotoInput) -> Photo: async def edit_photo(self, info, id: int, input: EditPhotoInput) -> Photo:
data = input.to_dict() return await PhotoMutationResolver().update(id, input, info.context.user_id)
async with get_session() as db:
photo = (await db.scalars(
BaseQueryResolver(Photo, models.Photo).get_query(user_id=info.context.user_id, object_id=id)
)).one()
if input.point_of_interest:
data['point_of_interest_id'] = await handle_combobox_save(
db,
models.PointOfInterest,
input.point_of_interest,
info.context.user_id,
extra_data={
"description": ""
}
)
if input.is_flight_cover:
# reset other covers
(await db.execute(
update(models.Photo)
.filter(models.Photo.flight_id == photo.flight_id)
.filter(models.Photo.id != id).values(is_flight_cover=False))
)
updated_model = await models.Photo.update(db, obj=photo, data=data)
return Photo(**updated_model.as_dict())
@strawberry.mutation() @strawberry.mutation()
@authenticated_user_only() @authenticated_user_only()
@@ -137,13 +85,10 @@ class EditPhotoMutation:
return Photo(**photo_as_dict) return Photo(**photo_as_dict)
@strawberry.type
class DeletePhotoMutation:
@strawberry.mutation() @strawberry.mutation()
@authenticated_user_only() @authenticated_user_only()
async def delete_photo(self, info, id: int) -> Photo: async def delete_photo(self, info, id: int) -> Photo:
photo = await BaseMutationResolver(Photo, models.Photo).delete(user_id=info.context.user_id, id=id) photo = await PhotoMutationResolver().delete(user_id=info.context.user_id, id=id)
base_path = get_photo_basepath(photo.flight_id) base_path = get_photo_basepath(photo.flight_id)
delete_file(f"{base_path}/{photo.filename}", silent=True) delete_file(f"{base_path}/{photo.filename}", silent=True)
+1 -1
View File
@@ -45,7 +45,7 @@ class PointOfInterestMutation:
input_data = input.to_dict() input_data = input.to_dict()
query = BaseQueryResolver(PointOfInterest, models.PointOfInterest).get_query( query = BaseQueryResolver(PointOfInterest, models.PointOfInterest).get_query(
user_id=info.context.user_id, object_id=id, include_public=False user_id=info.context.user_id, object_id=id, only_public=False
) )
async with get_session() as db: async with get_session() as db:
+52 -65
View File
@@ -1,103 +1,72 @@
from typing import Optional, Type, T from typing import Optional, Type, TypeVar, Generic, List
from sqlalchemy import select, or_
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from database import models from database import models
from database.query_builder import QueryBuilder
from database.transaction import get_session from database.transaction import get_session
GQL_TYPE = TypeVar('GQL_TYPE')
class BaseQueryResolver:
def __init__(self, graphql_type, model): class BaseResolver(Generic[GQL_TYPE]):
def __init__(self, graphql_type: GQL_TYPE, model: Type[models.BaseModel]):
self.graphql_type = graphql_type self.graphql_type = graphql_type
self.model = model self.model = model
self.query_builder = QueryBuilder(self.model)
class BaseQueryResolver(BaseResolver):
async def _get_list(self, query) -> List[GQL_TYPE]:
async with get_session() as db:
items = (await db.scalars(query)).all()
return [self.graphql_type(**m.as_dict()) for m in items]
async def _get_one(self, query) -> GQL_TYPE:
async with get_session() as db:
data = (await db.scalars(query)).one()
return self.graphql_type(**data.as_dict())
def get_query( def get_query(
self, self,
user_id: Optional[int] = None, user_id: Optional[int] = None,
object_id: Optional[int] = None, object_id: Optional[int] = None,
order_by: Optional[list] = None, order_by: Optional[list] = None,
include_public: Optional[bool] = True, only_public: Optional[bool] = False,
*args, *args,
**kwargs, **kwargs,
): ):
query = select(self.model) query = self.query_builder.get_simple_query(created_by_id=user_id, order_by=order_by)
if object_id: if object_id:
if hasattr(self.model, "id"): if not hasattr(self.model, "id"):
query = query.filter(self.model.id == object_id)
else:
raise AssertionError(f"Model {self.model} has no ID column! Cannot query by ID!") raise AssertionError(f"Model {self.model} has no ID column! Cannot query by ID!")
query = query.filter(self.model.id == object_id)
if hasattr(self.model, "deleted"): if only_public and hasattr(self.model, "is_public"):
query = query.filter(self.model.deleted.is_(False)) query = query.filter(self.model.is_public.is_(True))
ownership_clause = []
if hasattr(self.model, "is_public") and include_public:
ownership_clause.append(self.model.is_public.is_(True))
if hasattr(self.model, "created_by_id") and user_id:
ownership_clause.append(self.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 return query
async def _get_list(self, query): async def get_list(self, user_id: Optional[int] = None, **kwargs) -> List[GQL_TYPE]:
async with get_session() as db:
items = (await db.scalars(query)).all()
return [self.model(**m.as_dict()) for m in items]
async def _get_one(self, query):
async with get_session() as db:
data = (await db.scalars(query)).one()
return self.model(**data.as_dict())
async def get_list(self, user_id: Optional[int] = None, **kwargs) -> list:
query = self.get_query(user_id, **kwargs) query = self.get_query(user_id, **kwargs)
return await self._get_list(query) return await self._get_list(query)
async def get_one(self, id: int, user_id: Optional[int] = None, **kwargs): async def get_one(self, id: int, user_id: Optional[int] = None, **kwargs) -> GQL_TYPE:
query = self.get_query(user_id, object_id=id, **kwargs) query = self.get_query(user_id, object_id=id, **kwargs)
return await self._get_one(query) return await self._get_one(query)
class BaseMutationResolver: class BaseMutationResolver(BaseResolver):
model: Type[models.BaseModel] async def _get_one(self, db: AsyncSession, id: int, created_by_id: int) -> models.BaseModel:
graphql_type: Type[T] = None query = self.query_builder.get_simple_query(created_by_id=created_by_id).filter(self.model.id == id)
return (await db.scalars(query)).one()
def __init__(self, graphql_type: Type[T], model: Type[models.BaseModel]): async def _do_create(self, data) -> GQL_TYPE:
self.graphql_type = graphql_type
self.model = model
async def delete(self, user_id: int, id: int) -> T:
async with get_session() as db: async with get_session() as db:
query = BaseQueryResolver(self.graphql_type, self.model).get_query(user_id, object_id=id) model = await self.model.create(db, data=data)
model = (await db.scalars(query)).one()
if hasattr(self.model, "deleted"):
model = await self.model.update(db, obj=model, data=dict(deleted=True))
else:
await db.delete(model)
return self.graphql_type(**model.as_dict()) return self.graphql_type(**model.as_dict())
async def create(self, data: dict, user_id: Optional[int] = None) -> T: async def _do_update(self, db: AsyncSession, obj: models.BaseModel | dict, data: dict) -> GQL_TYPE:
input_data = {**data}
if hasattr(self.model, "created_by_id"):
input_data['created_by_id'] = user_id
async with get_session() as db:
model = await self.model.create(db, data=input_data)
return self.graphql_type(**model.as_dict())
async def _do_update(self, db: AsyncSession, obj: models.BaseModel | dict, data: dict) -> T:
update_where = {} update_where = {}
if isinstance(obj, models.BaseModel): if isinstance(obj, models.BaseModel):
update_where['obj'] = obj update_where['obj'] = obj
@@ -106,3 +75,21 @@ class BaseMutationResolver:
model = await self.model.update(db, data=data, **update_where) model = await self.model.update(db, data=data, **update_where)
return self.graphql_type(**model.as_dict()) return self.graphql_type(**model.as_dict())
async def delete(self, user_id: int, id: int) -> GQL_TYPE:
async with get_session() as db:
model = await self._get_one(db, id, user_id)
if hasattr(self.model, "deleted"):
model = await self.model.update(db, obj=model, data=dict(deleted=True))
else:
await db.delete(model)
return self.graphql_type(**model.as_dict())
async def create(self, data: dict, user_id: Optional[int] = None) -> GQL_TYPE:
input_data = {**data}
if hasattr(self.model, "created_by_id"):
input_data['created_by_id'] = user_id
return await self._do_create(input_data)
@@ -0,0 +1,33 @@
from typing import Optional
from database import models
from graphql_schema.entities.resolvers.base import BaseQueryResolver
from graphql_schema.entities.types.types import Event
class EventQueryResolver(BaseQueryResolver):
def __init__(self):
super().__init__(graphql_type=Event, model=models.Event)
def get_query(
self,
user_id: Optional[int] = None,
object_id: Optional[int] = None,
order_by: Optional[list] = None,
only_public: Optional[bool] = True,
*args,
**kwargs,
):
query = super().get_query(
user_id, object_id,
order_by=[models.Event.date_from.desc()],
only_public=not bool(user_id)
)
if kwargs.get('username'):
query = (
query
.join(models.Event.created_by)
.filter(models.User.public_username == kwargs['username'])
)
return query
@@ -25,7 +25,7 @@ class FlightQueryResolver(BaseQueryResolver):
query = super().get_query( query = super().get_query(
user_id, object_id, user_id, object_id,
order_by=[models.Flight.takeoff_datetime.desc()], order_by=[models.Flight.takeoff_datetime.desc()],
include_public=bool(user_id) only_public=not bool(user_id)
) )
if kwargs.get('username'): if kwargs.get('username'):
@@ -0,0 +1,40 @@
from sqlalchemy import update
from sqlalchemy.ext.asyncio import AsyncSession
from database import models
from database.transaction import get_session
from graphql_schema.entities.helpers.combobox import handle_combobox_save
from graphql_schema.entities.resolvers.base import BaseMutationResolver
from graphql_schema.entities.types.mutation_input import EditPhotoInput
from graphql_schema.entities.types.types import Photo
class PhotoMutationResolver(BaseMutationResolver):
def __init__(self):
super().__init__(Photo, models.Photo)
async def reset_flight_cover(self, db: AsyncSession, flight_id: int):
(await db.execute(
update(models.Photo)
.filter(models.Photo.flight_id == flight_id)
.filter(models.Photo.id != id).values(is_flight_cover=False))
)
async def update(self, id: int, input: EditPhotoInput, user_id: int) -> Photo:
data = input.to_dict()
async with get_session() as db:
photo = await self._get_one(db, id, created_by_id=user_id)
if input.point_of_interest:
data['point_of_interest_id'] = await handle_combobox_save(
db,
models.PointOfInterest,
input.point_of_interest,
user_id,
extra_data={"description": ""}
)
if input.is_flight_cover:
# reset other covers
await self.reset_flight_cover(db, photo.flight_id)
return await self._do_update(db, obj=photo, data=data)
@@ -54,6 +54,29 @@ class EditEventInput(BaseGraphqlInputType):
pass pass
@strawberry.input
class UploadPhotoInput:
photo: Upload
flight_id: int
name: Optional[str] = None
description: Optional[str] = None
point_of_interest: Optional[ComboboxInput] = None
@strawberry.input
class EditPhotoInput:
name: Optional[str] = None
description: Optional[str] = None
point_of_interest: Optional[ComboboxInput] = None
is_flight_cover: Optional[bool] = None
def to_dict(self):
return {
key: getattr(self, key) for key in ('name', 'description', 'is_flight_cover')
if getattr(self, key) is not None
}
@strawberry_sqlalchemy_input(models.Flight, exclude_fields=[ @strawberry_sqlalchemy_input(models.Flight, exclude_fields=[
"id", "aircraft_id", "deleted", "landing_airport_id", "takeoff_airport_id", "id", "aircraft_id", "deleted", "landing_airport_id", "takeoff_airport_id",
"takeoff_weather_info_id", "landing_weather_info_id", "gpx_track_filename", "event_id" "takeoff_weather_info_id", "landing_weather_info_id", "gpx_track_filename", "event_id"
+9 -6
View File
@@ -1,6 +1,6 @@
from __future__ import annotations from __future__ import annotations
from datetime import datetime from datetime import datetime
from typing import Optional, Annotated, List from typing import Optional, List
import strawberry import strawberry
from database import models from database import models
from decorators.endpoints import authenticated_user_only from decorators.endpoints import authenticated_user_only
@@ -10,7 +10,7 @@ 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_track_dataloader,
photos_dataloader, flights_by_aircraft_dataloader, users_in_organization_dataloader, photos_dataloader, flights_by_aircraft_dataloader, users_in_organization_dataloader,
aircrafts_from_organization_dataloader, user_organizations_dataloader, flights_by_event_dataloader, aircrafts_from_organization_dataloader, user_organizations_dataloader, flights_by_event_dataloader,
flights_by_copilot_dataloader flights_by_copilot_dataloader, public_flights_by_event_dataloader
) )
from graphql_schema.dataloaders.single_model import ( from graphql_schema.dataloaders.single_model import (
poi_dataloader, poi_type_dataloader, event_dataloader, aircraft_dataloader, airport_dataloader, cover_photo_loader, poi_dataloader, poi_type_dataloader, event_dataloader, aircraft_dataloader, airport_dataloader, cover_photo_loader,
@@ -161,13 +161,16 @@ class Organization:
class User: class User:
avatar_image_url: Optional[str] = strawberry.field(resolver=lambda root: get_avatar_url(root)) avatar_image_url: Optional[str] = strawberry.field(resolver=lambda root: get_avatar_url(root))
title_image_url: str = strawberry.field(resolver=lambda root: get_title_image_url(root)) title_image_url: str = strawberry.field(resolver=lambda root: get_title_image_url(root))
organizations: List[Annotated['Organization', strawberry.lazy(".organization")]] = strawberry.field( organizations: List[Organization] = strawberry.field(
resolver=lambda root: user_organizations_dataloader.load(root.id) resolver=lambda root: user_organizations_dataloader.load(root.id)
) )
@strawberry_sqlalchemy_type(models.Event) @strawberry_sqlalchemy_type(models.Event)
class Event: class Event:
flights: List[Annotated["Flight", strawberry.lazy('.flight')]] = strawberry.field( async def load_flights(root, info):
resolver=lambda root: flights_by_event_dataloader.load(root.id) is_user_logged_in = bool(info.context.user_id)
) dataloader = flights_by_event_dataloader if is_user_logged_in else public_flights_by_event_dataloader
return await dataloader.load(root.id)
flights: List[Flight] = strawberry.field(resolver=load_flights)
+2 -4
View File
@@ -4,16 +4,14 @@ from graphql_schema.entities.copilot import CreateCopilotMutation, EditCopilotMu
from graphql_schema.entities.event import EventMutation from graphql_schema.entities.event import EventMutation
from graphql_schema.entities.flight import FlightMutation from graphql_schema.entities.flight import FlightMutation
from graphql_schema.entities.organization import OrganizationUserMutation, OrganizationMutation from graphql_schema.entities.organization import OrganizationUserMutation, OrganizationMutation
from graphql_schema.entities.photo import UploadPhotoMutation, DeletePhotoMutation, EditPhotoMutation from graphql_schema.entities.photo import PhotoMutation
from graphql_schema.entities.poi import PointOfInterestMutation from graphql_schema.entities.poi import PointOfInterestMutation
from graphql_schema.entities.user import EditUserMutation from graphql_schema.entities.user import EditUserMutation
Mutation = merge_types("Mutation", ( Mutation = merge_types("Mutation", (
AircraftMutation, AircraftMutation,
FlightMutation, FlightMutation,
UploadPhotoMutation, PhotoMutation,
EditPhotoMutation,
DeletePhotoMutation,
PointOfInterestMutation, PointOfInterestMutation,
CreateCopilotMutation, CreateCopilotMutation,
EditCopilotMutation, EditCopilotMutation,
@@ -3,6 +3,7 @@ from typing import List, Optional
import strawberry import strawberry
import sqlalchemy import sqlalchemy
from sqlalchemy import Column from sqlalchemy import Column
from logger import log
from database.models import BaseModel from database.models import BaseModel
from graphql_schema.entities.types.base import BaseGraphqlInputType from graphql_schema.entities.types.base import BaseGraphqlInputType
@@ -22,7 +23,7 @@ def get_annotations_for_scalars(model: BaseModel, exclude_fields=None, force_opt
type_ = typing.Optional[column.type.python_type] if is_optional else column.type.python_type type_ = typing.Optional[column.type.python_type] if is_optional else column.type.python_type
annotations_[name] = type_ annotations_[name] = type_
except NotImplementedError as e: except NotImplementedError as e:
print(f"Neimplementovano: {e}, {name=}") log.warning(f"Cannot annotate {name} in {model} for GQL type. Exception: {e}")
return annotations_ return annotations_
+1 -1
View File
@@ -90,7 +90,7 @@ class App:
if APP_DEBUG: if APP_DEBUG:
@self.api_router.get("/graphql/autologin") @self.api_router.get("/graphql/autologin")
async def autologin(): async def autologin():
access_token = self.access_security.create_access_token(subject={"id": 1, "name": "Franta Vomacka"}) access_token = self.access_security.create_access_token(subject={"id": 6, "name": "Franta Vomacka"})
response = RedirectResponse(url="/graphql") response = RedirectResponse(url="/graphql")
self.access_security.set_access_cookie(response, access_token, expires_delta=timedelta(days=14)) self.access_security.set_access_cookie(response, access_token, expires_delta=timedelta(days=14))