From f4d709d621253375f6170c2b3429e308e1fb8c8c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Michal=20Kv=C3=A1=C4=8Dek?= Date: Tue, 12 Dec 2023 23:41:44 +0100 Subject: [PATCH] Strankovani --- src/config.py | 1 - src/decorators/endpoints.py | 14 +++++- src/endpoints/login.py | 3 +- .../dataloaders/single_model.py | 2 +- src/graphql_schema/entities/aircraft.py | 28 +++++++----- src/graphql_schema/entities/copilot.py | 8 +--- src/graphql_schema/entities/event.py | 37 ++++++++++------ src/graphql_schema/entities/flight.py | 38 +++++++++------- .../entities/helpers/pagination.py | 43 +++++++++++++++++++ src/graphql_schema/entities/photo.py | 8 ++-- src/graphql_schema/entities/poi.py | 30 +++++++------ .../entities/resolvers/aircraft.py | 3 -- .../entities/resolvers/photo.py | 7 ++- src/graphql_schema/entities/types/types.py | 7 ++- src/main.py | 10 ++--- 15 files changed, 154 insertions(+), 85 deletions(-) create mode 100644 src/graphql_schema/entities/helpers/pagination.py diff --git a/src/config.py b/src/config.py index f0db6b9..5d3e42f 100644 --- a/src/config.py +++ b/src/config.py @@ -3,7 +3,6 @@ import os APP_DEBUG = True GRAPHIQL = True -ACCESS_TOKEN_VALIDITY_MINUTES = 20 REFRESH_TOKEN_VALIDITY_DAYS = 30 API_URL = os.environ.get("API_URL") or "http://localhost:8000" diff --git a/src/decorators/endpoints.py b/src/decorators/endpoints.py index ace7e88..71a6bed 100644 --- a/src/decorators/endpoints.py +++ b/src/decorators/endpoints.py @@ -3,8 +3,18 @@ from fastapi import HTTPException from starlette.status import HTTP_401_UNAUTHORIZED -def public_endpoint(func): - pass +def allow_public(func): + @wraps(func) + async def decorator(*args, **kwargs): + if 'info' in kwargs: + user_id = kwargs['info'].context.user_id + public = kwargs.get('public') + if not user_id and not public: + raise HTTPException(HTTP_401_UNAUTHORIZED) + + return await func(*args, **kwargs) + + return decorator def authenticated_user_only(raise_when_unauthorized: bool = True, return_value_unauthorized=None): diff --git a/src/endpoints/login.py b/src/endpoints/login.py index ad68379..216c7d3 100644 --- a/src/endpoints/login.py +++ b/src/endpoints/login.py @@ -41,8 +41,7 @@ class LoginEndpoint(AuthEndpoint): return { "user": user, - "access_token": access_token, - "access_token_validity": self.access_security.access_expires_delta.total_seconds(), + "access_token": access_token } diff --git a/src/graphql_schema/dataloaders/single_model.py b/src/graphql_schema/dataloaders/single_model.py index b14a42e..3dd669b 100644 --- a/src/graphql_schema/dataloaders/single_model.py +++ b/src/graphql_schema/dataloaders/single_model.py @@ -21,4 +21,4 @@ flight_dataloader = create_dataloader(models.Flight) photo_adjustment_dataloader = create_dataloader( models.PhotoAdjustment, relationship_column=models.PhotoAdjustment.photo_id ) -photo_dataloader = create_dataloader(models.Photo) \ No newline at end of file +photo_dataloader = create_dataloader(models.Photo) diff --git a/src/graphql_schema/entities/aircraft.py b/src/graphql_schema/entities/aircraft.py index 20d5e9a..9f1ab62 100644 --- a/src/graphql_schema/entities/aircraft.py +++ b/src/graphql_schema/entities/aircraft.py @@ -1,10 +1,8 @@ -from typing import List, Optional +from typing import Optional import strawberry -from fastapi import HTTPException -from starlette.status import HTTP_401_UNAUTHORIZED - -from decorators.endpoints import authenticated_user_only +from decorators.endpoints import authenticated_user_only, allow_public from decorators.error_logging import error_logging +from .helpers.pagination import get_pagination_window, PaginationWindow from .resolvers.aircraft import AircraftMutationResolver, AircraftQueryResolver from graphql_schema.entities.types.mutation_input import CreateAircraftInput, EditAircraftInput from graphql_schema.entities.types.types import Aircraft @@ -15,20 +13,28 @@ class AircraftQueries: @strawberry.field() @error_logging @authenticated_user_only() - async def aircrafts(root, info) -> List[Aircraft]: - return await AircraftQueryResolver().get_list( + async def aircrafts(root, info, limit: int, offset: int = 0) -> PaginationWindow[Aircraft]: + query = AircraftQueryResolver().get_query( info.context.user_id, organization_ids=info.context.organization_ids ) + return await get_pagination_window( + query=query, + item_type=Aircraft, + limit=limit, + offset=offset + ) + @strawberry.field() @error_logging + @allow_public async def aircraft(root, info, id: int, public: Optional[bool] = False) -> Aircraft: - if not info.context.user_id and not public: - raise HTTPException(HTTP_401_UNAUTHORIZED) - return await AircraftQueryResolver().get_one( - id, user_id=info.context.user_id, organization_ids=info.context.organization_ids, public=public + id, + user_id=info.context.user_id, + organization_ids=info.context.organization_ids, + public=public ) diff --git a/src/graphql_schema/entities/copilot.py b/src/graphql_schema/entities/copilot.py index cf3b995..e9d585d 100644 --- a/src/graphql_schema/entities/copilot.py +++ b/src/graphql_schema/entities/copilot.py @@ -1,11 +1,9 @@ from typing import List, Optional import strawberry -from fastapi import HTTPException -from starlette.status import HTTP_401_UNAUTHORIZED from strawberry.types import Info from database import models from decorators.error_logging import error_logging -from decorators.endpoints import authenticated_user_only +from decorators.endpoints import authenticated_user_only, allow_public from graphql_schema.entities.resolvers.base import BaseMutationResolver from graphql_schema.entities.resolvers.copilot import CopilotQueryResolver from graphql_schema.entities.types.mutation_input import CreateCopilotInput, EditCopilotInput @@ -22,10 +20,8 @@ class CopilotQueries: @strawberry.field() @error_logging + @allow_public async def copilot(root, info: Info, id: int, pilot_username: Optional[str] = None) -> Copilot: - if not info.context.user_id and not pilot_username: - raise HTTPException(HTTP_401_UNAUTHORIZED) - return await CopilotQueryResolver().get_one( id, user_id=info.context.user_id, diff --git a/src/graphql_schema/entities/event.py b/src/graphql_schema/entities/event.py index 2341127..a744108 100644 --- a/src/graphql_schema/entities/event.py +++ b/src/graphql_schema/entities/event.py @@ -1,10 +1,9 @@ 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 decorators.endpoints import authenticated_user_only, allow_public from decorators.error_logging import error_logging +from graphql_schema.entities.helpers.pagination import PaginationWindow, get_pagination_window 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 @@ -15,25 +14,37 @@ from graphql_schema.entities.types.types import Event class EventQueries: @strawberry.field() @error_logging - 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( + @allow_public + async def events( + root, + info, + limit: int, + offset: int = 0, + username: Optional[str] = None, + public: Optional[bool] = False, + ) -> PaginationWindow[Event]: + query = EventQueryResolver().get_query( 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()], + public=public + ) + + return await get_pagination_window( + query=query, + item_type=Event, + limit=limit, + offset=offset ) @strawberry.field() @error_logging - 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) - + @allow_public + async def event(root, info, id: int, username: Optional[str] = None, public: Optional[bool] = False) -> Event: return await EventQueryResolver().get_one( id, username=username, + public=public, user_id=info.context.user_id ) diff --git a/src/graphql_schema/entities/flight.py b/src/graphql_schema/entities/flight.py index 218af40..af494fa 100644 --- a/src/graphql_schema/entities/flight.py +++ b/src/graphql_schema/entities/flight.py @@ -1,17 +1,16 @@ import asyncio -from typing import List, Optional +from typing import Optional import strawberry -from fastapi import HTTPException -from starlette.status import HTTP_401_UNAUTHORIZED from background_jobs.weather import download_weather from database import models -from decorators.endpoints import authenticated_user_only -from decorators.error_logging import error_logging from database.transaction import get_session +from decorators.endpoints import authenticated_user_only, allow_public +from decorators.error_logging import error_logging from graphql_schema.entities.resolvers.flight import handle_aircraft_save, FlightMutationResolver, FlightQueryResolver from graphql_schema.entities.types.mutation_input import EditFlightInput, CreateFlightInput -from .helpers.combobox import handle_combobox_save from graphql_schema.entities.types.types import Flight +from .helpers.combobox import handle_combobox_save +from .helpers.pagination import PaginationWindow, get_pagination_window @strawberry.type @@ -19,34 +18,43 @@ class FlightQueries: @strawberry.field() @error_logging + @allow_public async def flights( root, info, + limit: int, + offset: int = 0, username: Optional[str] = None, public: Optional[bool] = False, copilot_id: Optional[int] = None, point_of_interest_id: Optional[int] = None, aircraft_id: Optional[int] = None, - ) -> List[Flight]: - if not info.context.user_id and not public: - raise HTTPException(HTTP_401_UNAUTHORIZED) - - return await FlightQueryResolver().get_list( + ) -> PaginationWindow[Flight]: + query = FlightQueryResolver().get_query( user_id=info.context.user_id, username=username, only_public=public, copilot_id=copilot_id, aircraft_id=aircraft_id, point_of_interest_id=point_of_interest_id + ) + return await get_pagination_window( + query=query, + item_type=Flight, + limit=limit, + offset=offset, ) @strawberry.field() @error_logging + @allow_public async def flight(root, info, id: int, username: Optional[str] = None, public: Optional[bool] = False) -> Flight: - if not info.context.user_id and not public: - raise HTTPException(HTTP_401_UNAUTHORIZED) - - return await FlightQueryResolver().get_one(id, user_id=info.context.user_id, username=username, public=public) + return await FlightQueryResolver().get_one( + id, + user_id=info.context.user_id, + username=username, + public=public + ) @strawberry.type diff --git a/src/graphql_schema/entities/helpers/pagination.py b/src/graphql_schema/entities/helpers/pagination.py new file mode 100644 index 0000000..06e9868 --- /dev/null +++ b/src/graphql_schema/entities/helpers/pagination.py @@ -0,0 +1,43 @@ +from typing import TypeVar, Generic, List +import strawberry +from sqlalchemy import func, Select +from database.transaction import get_session + +Item = TypeVar("Item") + + +@strawberry.type +class PaginationWindow(Generic[Item]): + items: List[Item] = strawberry.field( + description="The list of items in this pagination window." + ) + + total_items_count: int = strawberry.field( + description="Total number of items in the filtered dataset." + ) + + +async def get_pagination_window( + query: Select, + item_type: type, + limit: int, + offset: int = 0, +) -> PaginationWindow: + if limit <= 0: + raise Exception(f"limit ({limit}) must be > 0") + + async with get_session() as db: + cnt_query = query.with_only_columns(func.count()) + total_items_count = (await db.scalars(cnt_query)).one() + + # if offset != 0 and not 0 <= offset < total_items_count: + # raise Exception(f"offset ({offset}) is out of range " f"(0-{total_items_count - 1})") + + async with get_session() as db: + data = (await db.scalars(query.limit(limit).offset(offset))).all() + dataset = [item_type(**i.as_dict()) for i in data] + + return PaginationWindow( + items=dataset, + total_items_count=total_items_count + ) diff --git a/src/graphql_schema/entities/photo.py b/src/graphql_schema/entities/photo.py index 9880623..28687c4 100644 --- a/src/graphql_schema/entities/photo.py +++ b/src/graphql_schema/entities/photo.py @@ -1,7 +1,7 @@ from typing import List, Optional import strawberry from database import models -from decorators.endpoints import authenticated_user_only +from decorators.endpoints import authenticated_user_only, allow_public from decorators.error_logging import error_logging from graphql_schema.entities.resolvers.base import BaseQueryResolver from graphql_schema.entities.resolvers.photo import PhotoMutationResolver, PhotoQueryResolver @@ -13,6 +13,7 @@ from graphql_schema.entities.types.mutation_input import EditPhotoInput, UploadP class PhotoQueries: @strawberry.field() @error_logging + @allow_public async def photos( root, info, flight_id: Optional[int] = None, @@ -32,8 +33,9 @@ class PhotoQueries: @strawberry.field() @error_logging - async def photo(root, info, id: int) -> Photo: - return await BaseQueryResolver(Photo, models.Photo).get_one(id, user_id=info.context.user_id) + @allow_public + async def photo(root, info, id: int, public: Optional[bool] = False,) -> Photo: + return await BaseQueryResolver(Photo, models.Photo).get_one(id, user_id=info.context.user_id, public=public) @strawberry.type diff --git a/src/graphql_schema/entities/poi.py b/src/graphql_schema/entities/poi.py index 4b35454..7a2062a 100644 --- a/src/graphql_schema/entities/poi.py +++ b/src/graphql_schema/entities/poi.py @@ -1,12 +1,11 @@ from typing import List 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 decorators.endpoints import authenticated_user_only, allow_public from database.transaction import get_session from decorators.error_logging import error_logging from graphql_schema.entities.helpers.combobox import handle_combobox_save +from graphql_schema.entities.helpers.pagination import get_pagination_window, PaginationWindow from graphql_schema.entities.resolvers.base import BaseQueryResolver, BaseMutationResolver from graphql_schema.entities.types.types import PointOfInterest from graphql_schema.entities.types.mutation_input import CreatePointOfInterestInput, EditPointOfInterestInput @@ -16,20 +15,25 @@ from graphql_schema.entities.types.mutation_input import CreatePointOfInterestIn class PointOfInterestQueries: @strawberry.field() @error_logging - async def points_of_interest(root, info, public: bool = False) -> List[PointOfInterest]: - if not info.context.user_id and not public: - raise HTTPException(HTTP_401_UNAUTHORIZED) - - return await BaseQueryResolver(PointOfInterest, models.PointOfInterest).get_list( - info.context.user_id, - only_public=public + @allow_public + async def points_of_interest( + root, info, + limit: int, offset: int = 0, + public: bool = False + ) -> PaginationWindow[PointOfInterest]: + query = BaseQueryResolver(PointOfInterest, models.PointOfInterest).get_query( + info.context.user_id, only_public=public + ) + return await get_pagination_window( + query=query, + item_type=PointOfInterest, + limit=limit, + offset=offset ) @strawberry.field() + @allow_public async def point_of_interest(root, info, id: int, public: bool = False) -> PointOfInterest: - if not info.context.user_id and not public: - raise HTTPException(HTTP_401_UNAUTHORIZED) - return await BaseQueryResolver(PointOfInterest, models.PointOfInterest).get_one( id, info.context.user_id, only_public=public diff --git a/src/graphql_schema/entities/resolvers/aircraft.py b/src/graphql_schema/entities/resolvers/aircraft.py index 8e5bedf..a7fafb5 100644 --- a/src/graphql_schema/entities/resolvers/aircraft.py +++ b/src/graphql_schema/entities/resolvers/aircraft.py @@ -6,9 +6,6 @@ from graphql_schema.entities.helpers.combobox import handle_combobox_save from graphql_schema.entities.resolvers.base import BaseMutationResolver, BaseQueryResolver from graphql_schema.entities.types.mutation_input import EditAircraftInput, CreateAircraftInput from graphql_schema.entities.types.types import Aircraft -from paths import AIRCRAFT_UPLOAD_DEST_PATH -from utils.file import delete_file -from utils.upload import handle_file_upload class AircraftQueryResolver(BaseQueryResolver): diff --git a/src/graphql_schema/entities/resolvers/photo.py b/src/graphql_schema/entities/resolvers/photo.py index b3db684..8627244 100644 --- a/src/graphql_schema/entities/resolvers/photo.py +++ b/src/graphql_schema/entities/resolvers/photo.py @@ -1,16 +1,15 @@ import os import shutil -from typing import Type, Optional, List - +from typing import Optional from PIL import Image from pydantic import BaseModel -from sqlalchemy import select, delete, insert +from sqlalchemy import delete, insert from background_jobs.elevation import add_terrain_elevation_to_photo from background_jobs.photo import generate_thumbnail, resize_photo 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, BaseQueryResolver, GQL_TYPE +from graphql_schema.entities.resolvers.base import BaseMutationResolver, BaseQueryResolver from graphql_schema.entities.types.mutation_input import EditPhotoInput, UploadPhotoInput, AdjustmentInput from graphql_schema.entities.types.types import Photo from paths import get_photo_basepath diff --git a/src/graphql_schema/entities/types/types.py b/src/graphql_schema/entities/types/types.py index 6a4efd5..a3fff48 100644 --- a/src/graphql_schema/entities/types/types.py +++ b/src/graphql_schema/entities/types/types.py @@ -10,7 +10,8 @@ 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, public_flights_by_event_dataloader, public_flights_by_copilot_dataloader, photo_copilots_dataloader, photos_aircraft_dataloader, copilots_in_photo_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 ) from graphql_schema.dataloaders.single_model import ( poi_dataloader, poi_type_dataloader, event_dataloader, aircraft_dataloader, airport_dataloader, @@ -18,9 +19,7 @@ from graphql_schema.dataloaders.single_model import ( photo_dataloader, user_dataloader ) from graphql_schema.sqlalchemy_to_strawberry_type import strawberry_sqlalchemy_type -from paths import ( - get_public_url, get_avatar_url, get_title_image_url, get_photo_thumbnail_url, get_photo_url, FLIGHT_GPX_TRACK_PATH -) +from paths import get_avatar_url, get_title_image_url, get_photo_thumbnail_url, get_photo_url, FLIGHT_GPX_TRACK_PATH @strawberry.type diff --git a/src/main.py b/src/main.py index 240c60d..6ef67ed 100644 --- a/src/main.py +++ b/src/main.py @@ -1,7 +1,7 @@ import sentry_sdk from datetime import timedelta from typing import Optional -from fastapi import FastAPI, APIRouter, Depends, Security +from fastapi import FastAPI, APIRouter, Depends, Security, HTTPException from fastapi_jwt import JwtAuthorizationCredentials, JwtAccessBearerCookie, JwtRefreshBearerCookie from graphql import GraphQLError from sqlalchemy import select @@ -10,10 +10,7 @@ from starlette.middleware.cors import CORSMiddleware from starlette.responses import RedirectResponse, Response from starlette.staticfiles import StaticFiles from strawberry.fastapi import GraphQLRouter -from config import ( - APP_SECRET_KEY, GRAPHIQL, APP_DEBUG, ALLOW_CORS_ORIGINS, SENTRY_DSN, ACCESS_TOKEN_VALIDITY_MINUTES, - REFRESH_TOKEN_VALIDITY_DAYS -) +from config import APP_SECRET_KEY, GRAPHIQL, APP_DEBUG, ALLOW_CORS_ORIGINS, SENTRY_DSN, REFRESH_TOKEN_VALIDITY_DAYS from database import models, async_session from endpoints.login import LoginEndpoint, LoginInput, RefreshEndpoint, LogoutEndpoint from endpoints.photo_editor_preview import PhotoEditorEndpoint @@ -26,7 +23,6 @@ class App: access_security = JwtAccessBearerCookie( secret_key=APP_SECRET_KEY, auto_error=False, - access_expires_delta=timedelta(minutes=ACCESS_TOKEN_VALIDITY_MINUTES), ) refresh_security = JwtRefreshBearerCookie( secret_key=APP_SECRET_KEY, @@ -39,7 +35,7 @@ class App: sentry_sdk.init( dsn=SENTRY_DSN, enable_tracing=True, - ignore_errors = [GraphQLError] + ignore_errors=[GraphQLError, HTTPException] ) app = FastAPI()