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 sqlalchemy import select
from typing import Type, List, Optional
from database import models, async_session
from database.query_builder import QueryBuilder
class SingleModelByIdDataloader:
class BaseDataloader:
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:
super().__init__()
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:
filters = []
self.filters = filters
class SingleModelByIdDataloader(BaseDataloader):
async def load(self, ids: List[int]):
async with async_session() as session:
query = select(self.model, self.relationship_column).filter(self.relationship_column.in_(ids))
for filter_ in self.filters:
query = query.filter(filter_)
query = (
self.query_builder.get_simple_query(extra_select=[self.relationship_column], include_deleted=True)
.filter(self.relationship_column.in_(ids))
.filter(*self.filters)
)
items = (await session.execute(query)).all()
items_by_id = {rel_id: item for item, rel_id in items}
return [items_by_id.get(id_) for id_ in ids]
class MultiModelsDataloader:
class MultiModelsDataloader(BaseDataloader):
def __init__(
self,
model: Type[models.BaseModel],
relationship_column,
extra_join: Optional[list] = None,
relationship_column=None,
filters: Optional[list] = None,
extra_join: Optional[list] = None,
order_by: Optional[list] = None,
):
super().__init__(model, relationship_column, filters)
if extra_join is None:
extra_join = []
if filters is None:
filters = []
self.extra_join = extra_join
if order_by is None:
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.filters = filters
async def load(self, ids: List[int]):
async with async_session() as session:
rel_column = self.relationship_column
async with async_session() as db:
query = (
select(self.model, rel_column)
.filter(rel_column.in_(ids))
.order_by(*self.order_by)
self.query_builder.get_simple_query(
extra_select=[self.relationship_column],
order_by=self.order_by
)
.filter(self.relationship_column.in_(ids))
.filter(*self.filters)
)
for table in self.extra_join:
query = query.join(table)
for joined_table in self.extra_join:
query = query.join(joined_table)
for filter_ in self.filters:
query = query.filter(filter_)
if self.filters:
query = query.filter(*self.filters)
data = (await session.execute(query)).all()
data = (await db.execute(query)).all()
result_data = defaultdict(list)
for item, rel_id in data:
@@ -48,6 +48,15 @@ flights_by_event_dataloader = DataLoader(
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(
models.Organization,
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
from fastapi import HTTPException
from starlette.status import HTTP_401_UNAUTHORIZED
from database import models
from decorators.endpoints import authenticated_user_only
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.types import Event
@@ -11,17 +15,26 @@ from graphql_schema.entities.types.types import Event
@strawberry.type
class EventQueries:
@strawberry.field()
@authenticated_user_only()
async def events(root, info) -> List[Event]:
return await BaseQueryResolver(Event, models.Event).get_list(
async def events(root, info, username: Optional[str] = None) -> List[Event]:
if not info.context.user_id and not username:
raise HTTPException(HTTP_401_UNAUTHORIZED)
return await EventQueryResolver().get_list(
info.context.user_id,
username=username,
order_by=[models.Event.date_from.desc(), models.Event.id.desc()]
)
@strawberry.field()
@authenticated_user_only()
async def event(root, info, id: int) -> Event:
return await BaseQueryResolver(Event, models.Event).get_one(id, user_id=info.context.user_id)
async def event(root, info, id: int, username: Optional[str] = None) -> Event:
if not info.context.user_id and not username:
raise HTTPException(HTTP_401_UNAUTHORIZED)
return await EventQueryResolver().get_one(
id,
username=username,
user_id=info.context.user_id
)
@strawberry.type
@@ -36,7 +49,7 @@ class EventMutation:
async def edit_event(root, info, id: int, input: EditEventInput) -> Event:
async with get_session() as db:
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()
updated_event = await models.Event.update(db, obj=event, data=input.to_dict())
+1 -1
View File
@@ -63,7 +63,7 @@ class FlightMutation:
"name": "",
"description": ""
})
flight = FlightMutationResolver().create(data, info.context.user_id)
flight = await FlightMutationResolver().create(data, info.context.user_id)
info.context.background_tasks.add_task(
download_weather,
+12 -67
View File
@@ -1,21 +1,19 @@
import asyncio
from typing import List, Optional
from typing import List
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.photo import generate_thumbnail, resize_photo
from database import models
from decorators.endpoints import authenticated_user_only
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 paths import get_photo_basepath
from graphql_schema.entities.helpers.combobox import handle_combobox_save
from utils.file import delete_file
from utils.image import parse_exif_info, rotate_image
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
@@ -26,15 +24,7 @@ class PhotoQueries:
@strawberry.type
class UploadPhotoMutation:
@strawberry.input
class UploadPhotoInput:
photo: Upload
flight_id: int
name: Optional[str] = None
description: Optional[str] = None
point_of_interest: Optional[ComboboxInput] = None
class PhotoMutation:
@strawberry.mutation
@authenticated_user_only()
async def upload_photo(self, info, input: UploadPhotoInput) -> Photo:
@@ -42,8 +32,9 @@ class UploadPhotoMutation:
filename = await handle_file_upload(input.photo, path)
exif_info = await parse_exif_info(path, filename)
async with get_session() as db:
photo_model = await models.Photo.create(data={
photo = await PhotoMutationResolver().create(
user_id=info.context.user_id,
data={
"flight_id": input.flight_id,
"name": input.name,
"filename": filename,
@@ -53,9 +44,8 @@ class UploadPhotoMutation:
"gps_longitude": exif_info.get("gps_longitude"),
"gps_altitude": exif_info.get("gps_altitude"),
"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(generate_thumbnail, path=path, filename=filename)
@@ -65,52 +55,10 @@ class UploadPhotoMutation:
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()
@authenticated_user_only()
async def edit_photo(self, info, id: int, input: EditPhotoInput) -> Photo:
data = input.to_dict()
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())
return await PhotoMutationResolver().update(id, input, info.context.user_id)
@strawberry.mutation()
@authenticated_user_only()
@@ -137,13 +85,10 @@ class EditPhotoMutation:
return Photo(**photo_as_dict)
@strawberry.type
class DeletePhotoMutation:
@strawberry.mutation()
@authenticated_user_only()
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)
delete_file(f"{base_path}/{photo.filename}", silent=True)
+1 -1
View File
@@ -45,7 +45,7 @@ class PointOfInterestMutation:
input_data = input.to_dict()
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:
+52 -65
View File
@@ -1,103 +1,72 @@
from typing import Optional, Type, T
from sqlalchemy import select, or_
from typing import Optional, Type, TypeVar, Generic, List
from sqlalchemy.ext.asyncio import AsyncSession
from database import models
from database.query_builder import QueryBuilder
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.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(
self,
user_id: Optional[int] = None,
object_id: Optional[int] = None,
order_by: Optional[list] = None,
include_public: Optional[bool] = True,
only_public: Optional[bool] = False,
*args,
**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 hasattr(self.model, "id"):
query = query.filter(self.model.id == object_id)
else:
if not hasattr(self.model, "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"):
query = query.filter(self.model.deleted.is_(False))
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)
if only_public and hasattr(self.model, "is_public"):
query = query.filter(self.model.is_public.is_(True))
return query
async def _get_list(self, query):
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:
async def get_list(self, user_id: Optional[int] = None, **kwargs) -> List[GQL_TYPE]:
query = self.get_query(user_id, **kwargs)
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)
return await self._get_one(query)
class BaseMutationResolver:
model: Type[models.BaseModel]
graphql_type: Type[T] = None
class BaseMutationResolver(BaseResolver):
async def _get_one(self, db: AsyncSession, id: int, created_by_id: int) -> models.BaseModel:
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]):
self.graphql_type = graphql_type
self.model = model
async def delete(self, user_id: int, id: int) -> T:
async def _do_create(self, data) -> GQL_TYPE:
async with get_session() as db:
query = BaseQueryResolver(self.graphql_type, self.model).get_query(user_id, object_id=id)
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)
model = await self.model.create(db, data=data)
return self.graphql_type(**model.as_dict())
async def create(self, data: dict, user_id: Optional[int] = None) -> T:
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:
async def _do_update(self, db: AsyncSession, obj: models.BaseModel | dict, data: dict) -> GQL_TYPE:
update_where = {}
if isinstance(obj, models.BaseModel):
update_where['obj'] = obj
@@ -106,3 +75,21 @@ class BaseMutationResolver:
model = await self.model.update(db, data=data, **update_where)
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(
user_id, object_id,
order_by=[models.Flight.takeoff_datetime.desc()],
include_public=bool(user_id)
only_public=not bool(user_id)
)
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
@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=[
"id", "aircraft_id", "deleted", "landing_airport_id", "takeoff_airport_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 datetime import datetime
from typing import Optional, Annotated, List
from typing import Optional, List
import strawberry
from database import models
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,
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
flights_by_copilot_dataloader, public_flights_by_event_dataloader
)
from graphql_schema.dataloaders.single_model import (
poi_dataloader, poi_type_dataloader, event_dataloader, aircraft_dataloader, airport_dataloader, cover_photo_loader,
@@ -161,13 +161,16 @@ class Organization:
class User:
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))
organizations: List[Annotated['Organization', strawberry.lazy(".organization")]] = strawberry.field(
organizations: List[Organization] = strawberry.field(
resolver=lambda root: user_organizations_dataloader.load(root.id)
)
@strawberry_sqlalchemy_type(models.Event)
class Event:
flights: List[Annotated["Flight", strawberry.lazy('.flight')]] = strawberry.field(
resolver=lambda root: flights_by_event_dataloader.load(root.id)
)
async def load_flights(root, info):
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.flight import FlightMutation
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.user import EditUserMutation
Mutation = merge_types("Mutation", (
AircraftMutation,
FlightMutation,
UploadPhotoMutation,
EditPhotoMutation,
DeletePhotoMutation,
PhotoMutation,
PointOfInterestMutation,
CreateCopilotMutation,
EditCopilotMutation,
@@ -3,6 +3,7 @@ from typing import List, Optional
import strawberry
import sqlalchemy
from sqlalchemy import Column
from logger import log
from database.models import BaseModel
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
annotations_[name] = type_
except NotImplementedError as e:
print(f"Neimplementovano: {e}, {name=}")
log.warning(f"Cannot annotate {name} in {model} for GQL type. Exception: {e}")
return annotations_
+1 -1
View File
@@ -90,7 +90,7 @@ class App:
if APP_DEBUG:
@self.api_router.get("/graphql/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")
self.access_security.set_access_cookie(response, access_token, expires_delta=timedelta(days=14))