From e28a6fce577ec0019f37bb557f20232cc88428ac Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Michal=20Kv=C3=A1=C4=8Dek?= Date: Tue, 13 Feb 2024 22:59:51 +0100 Subject: [PATCH] Uprava queries a dotazovani se pres base query builder, fix nahledu fotky pro editor, cache na fotky --- src/database/query_builder.py | 39 ++++++++-- src/endpoints/graphql.py | 26 ++++--- src/endpoints/photo_editor_preview.py | 5 +- src/graphql_schema/entities/aircraft.py | 3 +- src/graphql_schema/entities/organization.py | 17 ++--- .../entities/resolvers/aircraft.py | 22 ++++-- src/graphql_schema/entities/resolvers/base.py | 41 ++++++---- .../entities/resolvers/copilot.py | 10 ++- .../entities/resolvers/event.py | 3 +- .../entities/resolvers/flight.py | 3 +- .../entities/resolvers/organization.py | 32 ++++++++ src/graphql_schema/entities/types/types.py | 10 +-- src/graphql_schema/permissions.py | 13 ++++ src/main.py | 75 +++++++++++++------ 14 files changed, 209 insertions(+), 90 deletions(-) create mode 100644 src/graphql_schema/entities/resolvers/organization.py create mode 100644 src/graphql_schema/permissions.py diff --git a/src/database/query_builder.py b/src/database/query_builder.py index b653c89..a408e8b 100644 --- a/src/database/query_builder.py +++ b/src/database/query_builder.py @@ -1,5 +1,5 @@ from typing import Optional, Type -from sqlalchemy import select, or_ +from sqlalchemy import select, or_, and_ from database import models @@ -13,6 +13,8 @@ class QueryBuilder: created_by_id: Optional[int] = None, order_by: Optional[list] = None, only_public: Optional[bool] = False, + only_my: Optional[bool] = False, + # include_others_public: Optional[bool] = False, url_slug: Optional[str] = None, include_deleted: bool = False ): @@ -24,21 +26,42 @@ class QueryBuilder: if not include_deleted and hasattr(self.model, "deleted"): query = query.filter(self.model.deleted.is_(False)) + ownership_filters = [] + my_filters = [] if only_public and hasattr(self.model, "is_public"): - query = query.filter(self.model.is_public.is_(True)) + my_filters.append(self.model.is_public.is_(True)) if hasattr(self.model, "url_slug"): - query = query.filter(self.model.url_slug != '') - elif hasattr(self.model, "created_by_id") and created_by_id: - query = query.filter(or_( - self.model.created_by_id.is_(None), - self.model.created_by_id == created_by_id - )) + my_filters.append(self.model.url_slug != '') + + if only_my: + if created_by_id and hasattr(self.model, "created_by_id"): + my_filters.append(self.model.created_by_id == created_by_id) + + # others_filters = [] + # if include_others_public: + # # TODO: Toto jeste neni implementovane nikde v resolverech! + # if hasattr(self.model, "is_public"): + # others_filters.append(self.model.is_public.is_(True)) + # + # if hasattr(self.model, "created_by_id"): + # others_filters.append(or_( + # self.model.created_by_id != created_by_id, + # self.model.created_by_id.is_(None) + # )) + + if my_filters: + ownership_filters.append(and_(*my_filters)) + # if others_filters: + # ownership_filters.append(and_(*others_filters)) + + query = query.filter(or_(*ownership_filters)) if url_slug is not None and hasattr(self.model, 'url_slug'): query = query.filter(self.model.url_slug == url_slug) if order_by: query = query.order_by(*order_by) + elif hasattr(self.model, "name"): query = query.order_by(self.model.name) diff --git a/src/endpoints/graphql.py b/src/endpoints/graphql.py index de04a24..6922050 100644 --- a/src/endpoints/graphql.py +++ b/src/endpoints/graphql.py @@ -10,7 +10,21 @@ from database import async_session, models from graphql_schema.schema import GraphQLContext, schema -def setup_graphql_endpoint(app: FastAPI, access_security: JwtAccessBearerCookie, api_router: APIRouter): +def setup_graphql_endpoint(app: FastAPI, access_security: JwtAccessBearerCookie): + if not APP_DEBUG: + return + debug_router = APIRouter() + + @debug_router.get("/graphql/autologin") + async def autologin(): + access_token = access_security.create_access_token(subject={"id": 7, "name": "Franta Vomacka"}) + response = RedirectResponse(url="/graphql") + access_security.set_access_cookie(response, access_token, expires_delta=timedelta(days=14)) + + return response + + app.include_router(debug_router) + async def setup_graphql_context(credentials: JwtAuthorizationCredentials = Security(access_security)): user_id = credentials['id'] if credentials else None organization_ids = set() @@ -22,16 +36,6 @@ def setup_graphql_endpoint(app: FastAPI, access_security: JwtAccessBearerCookie, .filter(models.user_is_in_organization.c.user_id == user_id) )).all()) - if APP_DEBUG: - @api_router.get("/graphql/autologin") - async def autologin(): - access_token = access_security.create_access_token(subject={"id": 6, "name": "Franta Vomacka"}) - - response = RedirectResponse(url="/graphql") - access_security.set_access_cookie(response, access_token, expires_delta=timedelta(days=14)) - - return response - return GraphQLContext( user_id=user_id, organization_ids=organization_ids, diff --git a/src/endpoints/photo_editor_preview.py b/src/endpoints/photo_editor_preview.py index f4a7438..cacb3b6 100644 --- a/src/endpoints/photo_editor_preview.py +++ b/src/endpoints/photo_editor_preview.py @@ -9,7 +9,8 @@ from utils.image import PhotoEditor class PhotoEditorEndpoint(AuthEndpoint): - async def show_preview(self, photo_id: int, **kwargs): + @staticmethod + async def show_preview(photo_id: int, **kwargs): async with get_session() as db: photo = (await db.scalars( select(models.Photo) @@ -17,7 +18,7 @@ class PhotoEditorEndpoint(AuthEndpoint): )).one() basepath = get_photo_basepath(photo.flight_id) - filename = photo.filename + filename = f"{photo.filename}.{photo.filename_extension}" original_filename = '_original_' + filename if os.path.exists(f"{basepath}/{original_filename}"): diff --git a/src/graphql_schema/entities/aircraft.py b/src/graphql_schema/entities/aircraft.py index 1534bc6..0d0fd7d 100644 --- a/src/graphql_schema/entities/aircraft.py +++ b/src/graphql_schema/entities/aircraft.py @@ -36,7 +36,6 @@ class AircraftQueries: public: Optional[bool] = False ) -> Aircraft: filter_params = {} - if id: filter_params['object_id'] = id @@ -45,7 +44,7 @@ class AircraftQueries: return await AircraftQueryResolver().get_one( user_id=info.context.user_id, - organization_ids=info.context.organization_ids, + organization_ids=info.context.organization_ids if not public else None, only_public=public, **filter_params ) diff --git a/src/graphql_schema/entities/organization.py b/src/graphql_schema/entities/organization.py index bcc0f96..f50a291 100644 --- a/src/graphql_schema/entities/organization.py +++ b/src/graphql_schema/entities/organization.py @@ -1,6 +1,6 @@ from typing import List import strawberry -from sqlalchemy import delete +from sqlalchemy import delete, select from sqlalchemy.dialects.mysql import insert from sqlalchemy.exc import IntegrityError from database import models @@ -8,6 +8,7 @@ from decorators.endpoints import authenticated_user_only from database.transaction import get_session from decorators.error_logging import error_logging from graphql_schema.entities.resolvers.base import BaseQueryResolver, BaseMutationResolver +from graphql_schema.entities.resolvers.organization import OrganizationQueryResolver from graphql_schema.entities.types.mutation_input import CreateOrganizationInput, EditOrganizationInput from graphql_schema.entities.types.types import Organization @@ -18,19 +19,13 @@ class OrganizationQueries: @error_logging @authenticated_user_only() async def organizations(root, info) -> List[Organization]: - return await BaseQueryResolver(Organization, models.Organization).get_list( - info.context.user_id, - order_by=[models.Organization.name] - ) + return await OrganizationQueryResolver().get_list() @strawberry.field() @error_logging @authenticated_user_only() async def organization(root, info, id: int) -> Organization: - return await BaseQueryResolver(Organization, models.Organization).get_one( - object_id=id, - user_id=info.context.user_id - ) + return await OrganizationQueryResolver().get_one(object_id=id) @strawberry.type @@ -62,7 +57,7 @@ class OrganizationUserMutation: async def add_to_organization(root, info, organization_id: int) -> Organization: async with get_session() as db: organization = (await db.scalars( - BaseQueryResolver(Organization, models.Organization).get_query(object_id=organization_id) + OrganizationQueryResolver().get_query(object_id=organization_id) )).one() try: @@ -83,7 +78,7 @@ class OrganizationUserMutation: async def remove_from_organization(root, info, organization_id: int) -> Organization: async with get_session() as db: organization = (await db.scalars( - BaseQueryResolver(Organization, models.Organization).get_query(object_id=organization_id) + OrganizationQueryResolver().get_query(object_id=organization_id) )).one() await db.execute( diff --git a/src/graphql_schema/entities/resolvers/aircraft.py b/src/graphql_schema/entities/resolvers/aircraft.py index fb69672..e27df91 100644 --- a/src/graphql_schema/entities/resolvers/aircraft.py +++ b/src/graphql_schema/entities/resolvers/aircraft.py @@ -1,5 +1,8 @@ from operator import or_ from typing import Set, Optional + +from sqlalchemy import and_ + from database import models from database.transaction import get_session from graphql_schema.entities.helpers.combobox import handle_combobox_save @@ -21,16 +24,20 @@ class AircraftQueryResolver(BaseQueryResolver): *args, **kwargs, ): - call_sign = kwargs.get("call_sign") filters = {} if object_id: filters['object_id'] = object_id - if call_sign: - filters['call_sign'] = call_sign + if not organization_ids: + filters['user_id'] = user_id + if kwargs.get("call_sign"): + filters['call_sign'] = kwargs["call_sign"] + if kwargs.get("search"): + filters['search'] = kwargs.pop("search", None) query = super().get_query( - order_by=[models.Aircraft.id.desc()], - user_id=user_id if not organization_ids else None, + only_my=False, + only_public=kwargs.get("only_public", False), + order_by=order_by, **filters ) if organization_ids: @@ -38,7 +45,10 @@ class AircraftQueryResolver(BaseQueryResolver): query.filter( or_( models.Aircraft.created_by_id == user_id, - models.Aircraft.organization_id.in_(organization_ids) + and_( + models.Aircraft.organization_id.in_(organization_ids), + models.Aircraft.is_public.is_(True) + ) ) ) ) diff --git a/src/graphql_schema/entities/resolvers/base.py b/src/graphql_schema/entities/resolvers/base.py index 7b5155a..61713ad 100644 --- a/src/graphql_schema/entities/resolvers/base.py +++ b/src/graphql_schema/entities/resolvers/base.py @@ -35,38 +35,51 @@ class BaseQueryResolver(BaseResolver): object_id: Optional[int] = None, order_by: Optional[list] = None, only_public: Optional[bool] = False, + only_my: Optional[bool] = False, + url_slug: Optional[str] = None, **kwargs, ): - query = self.query_builder.get_simple_query(created_by_id=user_id, order_by=order_by, only_public=only_public) + query = self.query_builder.get_simple_query( + created_by_id=user_id, + order_by=order_by, + only_public=only_public, + only_my=only_my, + url_slug=url_slug + ) if object_id: 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) - search = kwargs.pop("search", None) - if search: - search_clauses = [] - if hasattr(self.model, "name"): - search_clauses.append(self.model.name.contains(search)) - - if hasattr(self.model, "description"): - search_clauses.append(self.model.description.contains(search)) - - if search_clauses: - query = query.filter(or_(*search_clauses)) + query = self._handle_search(query, kwargs.pop("search", None)) if kwargs: query = query.filter_by(**kwargs) return query + def _handle_search(self, query, search): + if not search: + return query + + search_clauses = [] + if hasattr(self.model, "name"): + search_clauses.append(self.model.name.contains(search)) + + if hasattr(self.model, "description"): + search_clauses.append(self.model.description.contains(search)) + + if search_clauses: + query = query.filter(or_(*search_clauses)) + return query + async def get_list(self, user_id: Optional[int] = None, **kwargs) -> List[GQL_TYPE]: - query = self.get_query(user_id, **kwargs) + query = self.get_query(user_id=user_id, **kwargs) return await self._get_list(query) async def get_one(self, user_id: Optional[int] = None, **kwargs) -> GQL_TYPE: - query = self.get_query(user_id, **kwargs) + query = self.get_query(user_id=user_id, **kwargs) return await self._get_one(query) diff --git a/src/graphql_schema/entities/resolvers/copilot.py b/src/graphql_schema/entities/resolvers/copilot.py index 1e0918c..a736a4f 100644 --- a/src/graphql_schema/entities/resolvers/copilot.py +++ b/src/graphql_schema/entities/resolvers/copilot.py @@ -14,10 +14,16 @@ class CopilotQueryResolver(BaseQueryResolver): object_id: Optional[int] = None, order_by: Optional[list] = None, only_public: Optional[bool] = False, - *args, **kwargs + **kwargs ): pilot_username = kwargs.pop("pilot_username", None) - query = super().get_query(user_id, object_id, order_by, only_public, *args, **kwargs) + + query = super().get_query( + user_id=user_id, object_id=object_id, order_by=order_by, + only_public=only_public, + only_my=bool(user_id) and not only_public, + **kwargs + ) if pilot_username: query = ( diff --git a/src/graphql_schema/entities/resolvers/event.py b/src/graphql_schema/entities/resolvers/event.py index 5c7f0e7..9289711 100644 --- a/src/graphql_schema/entities/resolvers/event.py +++ b/src/graphql_schema/entities/resolvers/event.py @@ -21,7 +21,8 @@ class EventQueryResolver(BaseQueryResolver): user_id, object_id, order_by=[models.Event.date_from.desc()], url_slug=kwargs.get("url_slug"), - only_public=not bool(user_id) + only_public=only_public, + only_my=True ) if kwargs.get('username'): diff --git a/src/graphql_schema/entities/resolvers/flight.py b/src/graphql_schema/entities/resolvers/flight.py index 648e500..2b21975 100644 --- a/src/graphql_schema/entities/resolvers/flight.py +++ b/src/graphql_schema/entities/resolvers/flight.py @@ -41,7 +41,8 @@ class FlightQueryResolver(BaseQueryResolver): user_id, **filters, order_by=[models.Flight.takeoff_datetime.desc()], - only_public=only_public + only_public=only_public, + only_my=not only_public ) if kwargs.get("event_id"): diff --git a/src/graphql_schema/entities/resolvers/organization.py b/src/graphql_schema/entities/resolvers/organization.py new file mode 100644 index 0000000..c606066 --- /dev/null +++ b/src/graphql_schema/entities/resolvers/organization.py @@ -0,0 +1,32 @@ +from typing import Optional +from sqlalchemy import select +from database import models +from graphql_schema.entities.resolvers.base import BaseQueryResolver +from graphql_schema.entities.types.types import Organization + + +class OrganizationQueryResolver(BaseQueryResolver): + + def __init__(self): + super().__init__(Organization, models.Organization) + + def get_query( + self, + object_id: Optional[int] = None, + order_by: Optional[list] = None, + **kwargs): + + query = ( + select(models.Organization) + .filter(models.Organization.deleted.is_(False)) + ) + + if not order_by: + order_by = [models.Organization.name] + + query = query.order_by(*order_by) + + if object_id: + query = query.where(models.Organization.id == object_id) + + return query diff --git a/src/graphql_schema/entities/types/types.py b/src/graphql_schema/entities/types/types.py index c4f9beb..743ec63 100644 --- a/src/graphql_schema/entities/types/types.py +++ b/src/graphql_schema/entities/types/types.py @@ -22,6 +22,7 @@ from graphql_schema.dataloaders.single_model import ( airport_weather_info_loader, organizations_dataloader, flight_dataloader, photo_adjustment_dataloader, photo_dataloader, user_dataloader ) +from graphql_schema.permissions import IsAuthenticated from graphql_schema.sqlalchemy_to_strawberry_type import strawberry_sqlalchemy_type from paths import get_avatar_url, get_title_image_url, get_photo_thumbnail_url, get_photo_url, FLIGHT_GPX_TRACK_PATH @@ -96,15 +97,6 @@ class Photo: ) -class IsAuthenticated(BasePermission): - message = "User is not authenticated" - error_class = GraphQLError - error_extensions = {"code": "UNAUTHORIZED"} - - def has_permission(self, source: Any, info: Info, **kwargs) -> bool: - return bool(info.context.user_id) - - @strawberry_sqlalchemy_type(models.Flight) class Flight: def __init__(self, **kwargs): diff --git a/src/graphql_schema/permissions.py b/src/graphql_schema/permissions.py new file mode 100644 index 0000000..1ee4ebb --- /dev/null +++ b/src/graphql_schema/permissions.py @@ -0,0 +1,13 @@ +from typing import Any +from graphql import GraphQLError +from strawberry import BasePermission +from strawberry.types import Info + + +class IsAuthenticated(BasePermission): + message = "User is not authenticated" + error_class = GraphQLError + error_extensions = {"code": "UNAUTHORIZED"} + + def has_permission(self, source: Any, info: Info, **kwargs) -> bool: + return bool(info.context.user_id) diff --git a/src/main.py b/src/main.py index 564ac49..9db531a 100644 --- a/src/main.py +++ b/src/main.py @@ -18,6 +18,17 @@ from endpoints.photo_editor_preview import PhotoEditorEndpoint from endpoints.registration import RegistrationInput, RegistrationEndpoint +class StaticFilesCache(StaticFiles): + def __init__(self, *args, cachecontrol="public, max-age=31536000, s-maxage=31536000, immutable", **kwargs): + self.cachecontrol = cachecontrol + super().__init__(*args, **kwargs) + + def file_response(self, *args, **kwargs) -> Response: + resp: Response = super().file_response(*args, **kwargs) + resp.headers.setdefault("Cache-Control", self.cachecontrol) + return resp + + class App: api_router = APIRouter(dependencies=[]) access_security = JwtAccessBearerCookie( @@ -38,7 +49,21 @@ class App: ignore_errors=[GraphQLError, HTTPException] ) - app = FastAPI() + app = FastAPI( + title="Polétání.cz (API)", + version="1.0.0", + summary="Nejenom GraphQL API pro aplikaci Polétání.cz", + terms_of_service="https://poletani.cz/podminky", + contact={ + "name": "Michal Kváček", + "url": "https://poletani.cz/michal/", + "email": "michal@kvacek.cz", + }, + license_info={ + "name": "Apache 2.0", + "url": "https://www.apache.org/licenses/LICENSE-2.0.html", + }, + ) self.setup_exception_handlers(app) self.setup_middleware(app) @@ -65,22 +90,14 @@ class App: @staticmethod def setup_static_paths(app: FastAPI): - app.mount("/uploads", StaticFiles(directory="/app/uploads"), name="uploads") - app.mount("/static", StaticFiles(directory="/app/static"), name="static") + app.mount("/uploads", app=StaticFilesCache(directory="/app/uploads"), name="uploads") + app.mount( + "/static", + app=StaticFilesCache(directory="/app/static", cachecontrol=f"public, max-age={24 * 3600}"), + name="static" + ) def setup_routes(self, app: FastAPI): - setup_graphql_endpoint(app, self.access_security, self.api_router) - - @self.api_router.post("/refresh") - async def refresh( - resp: Response, - credentials: JwtAuthorizationCredentials = Security(self.refresh_security) - ): - return await RefreshEndpoint( - access_token=self.access_security, - refresh_token=self.refresh_security - ).on_post(resp, credentials) - @self.api_router.post("/registration", status_code=201) async def registration(user: RegistrationInput, background_tasks: BackgroundTasks): return await RegistrationEndpoint().on_post(user, background_tasks) @@ -92,6 +109,23 @@ class App: refresh_token=self.refresh_security ).on_post(user, resp) + @self.api_router.post("/refresh", summary="Refresh access token") + async def refresh( + resp: Response, + credentials: JwtAuthorizationCredentials = Security(self.refresh_security) + ): + return await RefreshEndpoint( + access_token=self.access_security, + refresh_token=self.refresh_security + ).on_post(resp, credentials) + + @self.api_router.post("/logout") + async def logout(resp: Response): + return await LogoutEndpoint( + access_token=self.access_security, + refresh_token=self.refresh_security + ).on_post(resp) + @self.api_router.get("/forgotten-password/token/{token}") async def token_info(token: str): return await ForgottenPasswordEndpoint().token_info(token) @@ -104,17 +138,10 @@ class App: async def password_reset(input: ChangeForgottenPassword): return await ForgottenPasswordEndpoint().change_password(input) - @self.api_router.post("/contact") + @self.api_router.post("/contact", summary="Send email from contact form") async def contact(input: ContactInput, background_tasks: BackgroundTasks): return await ContactEndpoint().on_post(input, background_tasks) - @self.api_router.post("/logout") - async def logout(resp: Response): - return await LogoutEndpoint( - access_token=self.access_security, - refresh_token=self.refresh_security - ).on_post(resp) - @self.api_router.get("/photo/editor-preview/{photo_id}") async def photo_editor_preview( photo_id: int, @@ -145,5 +172,7 @@ class App: rotate=rotate, ) + setup_graphql_endpoint(app, self.access_security) + # musi byt na konci app.include_router(self.api_router)