Uprava queries a dotazovani se pres base query builder, fix nahledu fotky pro editor, cache na fotky

This commit is contained in:
Michal Kváček
2024-02-13 22:59:51 +01:00
parent 51316c492b
commit e28a6fce57
14 changed files with 209 additions and 90 deletions
+31 -8
View File
@@ -1,5 +1,5 @@
from typing import Optional, Type from typing import Optional, Type
from sqlalchemy import select, or_ from sqlalchemy import select, or_, and_
from database import models from database import models
@@ -13,6 +13,8 @@ class QueryBuilder:
created_by_id: Optional[int] = None, created_by_id: Optional[int] = None,
order_by: Optional[list] = None, order_by: Optional[list] = None,
only_public: Optional[bool] = False, only_public: Optional[bool] = False,
only_my: Optional[bool] = False,
# include_others_public: Optional[bool] = False,
url_slug: Optional[str] = None, url_slug: Optional[str] = None,
include_deleted: bool = False include_deleted: bool = False
): ):
@@ -24,21 +26,42 @@ class QueryBuilder:
if not include_deleted and hasattr(self.model, "deleted"): if not include_deleted and hasattr(self.model, "deleted"):
query = query.filter(self.model.deleted.is_(False)) query = query.filter(self.model.deleted.is_(False))
ownership_filters = []
my_filters = []
if only_public and hasattr(self.model, "is_public"): 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"): if hasattr(self.model, "url_slug"):
query = query.filter(self.model.url_slug != '') my_filters.append(self.model.url_slug != '')
elif hasattr(self.model, "created_by_id") and created_by_id:
query = query.filter(or_( if only_my:
self.model.created_by_id.is_(None), if created_by_id and hasattr(self.model, "created_by_id"):
self.model.created_by_id == 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'): if url_slug is not None and hasattr(self.model, 'url_slug'):
query = query.filter(self.model.url_slug == url_slug) query = query.filter(self.model.url_slug == url_slug)
if order_by: if order_by:
query = query.order_by(*order_by) query = query.order_by(*order_by)
elif hasattr(self.model, "name"): elif hasattr(self.model, "name"):
query = query.order_by(self.model.name) query = query.order_by(self.model.name)
+15 -11
View File
@@ -10,7 +10,21 @@ from database import async_session, models
from graphql_schema.schema import GraphQLContext, schema 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)): async def setup_graphql_context(credentials: JwtAuthorizationCredentials = Security(access_security)):
user_id = credentials['id'] if credentials else None user_id = credentials['id'] if credentials else None
organization_ids = set() 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) .filter(models.user_is_in_organization.c.user_id == user_id)
)).all()) )).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( return GraphQLContext(
user_id=user_id, user_id=user_id,
organization_ids=organization_ids, organization_ids=organization_ids,
+3 -2
View File
@@ -9,7 +9,8 @@ from utils.image import PhotoEditor
class PhotoEditorEndpoint(AuthEndpoint): 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: async with get_session() as db:
photo = (await db.scalars( photo = (await db.scalars(
select(models.Photo) select(models.Photo)
@@ -17,7 +18,7 @@ class PhotoEditorEndpoint(AuthEndpoint):
)).one() )).one()
basepath = get_photo_basepath(photo.flight_id) basepath = get_photo_basepath(photo.flight_id)
filename = photo.filename filename = f"{photo.filename}.{photo.filename_extension}"
original_filename = '_original_' + filename original_filename = '_original_' + filename
if os.path.exists(f"{basepath}/{original_filename}"): if os.path.exists(f"{basepath}/{original_filename}"):
+1 -2
View File
@@ -36,7 +36,6 @@ class AircraftQueries:
public: Optional[bool] = False public: Optional[bool] = False
) -> Aircraft: ) -> Aircraft:
filter_params = {} filter_params = {}
if id: if id:
filter_params['object_id'] = id filter_params['object_id'] = id
@@ -45,7 +44,7 @@ class AircraftQueries:
return await AircraftQueryResolver().get_one( return await AircraftQueryResolver().get_one(
user_id=info.context.user_id, 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, only_public=public,
**filter_params **filter_params
) )
+6 -11
View File
@@ -1,6 +1,6 @@
from typing import List from typing import List
import strawberry import strawberry
from sqlalchemy import delete from sqlalchemy import delete, select
from sqlalchemy.dialects.mysql import insert from sqlalchemy.dialects.mysql import insert
from sqlalchemy.exc import IntegrityError from sqlalchemy.exc import IntegrityError
from database import models from database import models
@@ -8,6 +8,7 @@ from decorators.endpoints import authenticated_user_only
from database.transaction import get_session from database.transaction import get_session
from decorators.error_logging import error_logging from decorators.error_logging import error_logging
from graphql_schema.entities.resolvers.base import BaseQueryResolver, BaseMutationResolver 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.mutation_input import CreateOrganizationInput, EditOrganizationInput
from graphql_schema.entities.types.types import Organization from graphql_schema.entities.types.types import Organization
@@ -18,19 +19,13 @@ class OrganizationQueries:
@error_logging @error_logging
@authenticated_user_only() @authenticated_user_only()
async def organizations(root, info) -> List[Organization]: async def organizations(root, info) -> List[Organization]:
return await BaseQueryResolver(Organization, models.Organization).get_list( return await OrganizationQueryResolver().get_list()
info.context.user_id,
order_by=[models.Organization.name]
)
@strawberry.field() @strawberry.field()
@error_logging @error_logging
@authenticated_user_only() @authenticated_user_only()
async def organization(root, info, id: int) -> Organization: async def organization(root, info, id: int) -> Organization:
return await BaseQueryResolver(Organization, models.Organization).get_one( return await OrganizationQueryResolver().get_one(object_id=id)
object_id=id,
user_id=info.context.user_id
)
@strawberry.type @strawberry.type
@@ -62,7 +57,7 @@ class OrganizationUserMutation:
async def add_to_organization(root, info, organization_id: int) -> Organization: async def add_to_organization(root, info, organization_id: int) -> Organization:
async with get_session() as db: async with get_session() as db:
organization = (await db.scalars( organization = (await db.scalars(
BaseQueryResolver(Organization, models.Organization).get_query(object_id=organization_id) OrganizationQueryResolver().get_query(object_id=organization_id)
)).one() )).one()
try: try:
@@ -83,7 +78,7 @@ class OrganizationUserMutation:
async def remove_from_organization(root, info, organization_id: int) -> Organization: async def remove_from_organization(root, info, organization_id: int) -> Organization:
async with get_session() as db: async with get_session() as db:
organization = (await db.scalars( organization = (await db.scalars(
BaseQueryResolver(Organization, models.Organization).get_query(object_id=organization_id) OrganizationQueryResolver().get_query(object_id=organization_id)
)).one() )).one()
await db.execute( await db.execute(
@@ -1,5 +1,8 @@
from operator import or_ from operator import or_
from typing import Set, Optional from typing import Set, Optional
from sqlalchemy import and_
from database import models from database import models
from database.transaction import get_session from database.transaction import get_session
from graphql_schema.entities.helpers.combobox import handle_combobox_save from graphql_schema.entities.helpers.combobox import handle_combobox_save
@@ -21,16 +24,20 @@ class AircraftQueryResolver(BaseQueryResolver):
*args, *args,
**kwargs, **kwargs,
): ):
call_sign = kwargs.get("call_sign")
filters = {} filters = {}
if object_id: if object_id:
filters['object_id'] = object_id filters['object_id'] = object_id
if call_sign: if not organization_ids:
filters['call_sign'] = call_sign 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( query = super().get_query(
order_by=[models.Aircraft.id.desc()], only_my=False,
user_id=user_id if not organization_ids else None, only_public=kwargs.get("only_public", False),
order_by=order_by,
**filters **filters
) )
if organization_ids: if organization_ids:
@@ -38,7 +45,10 @@ class AircraftQueryResolver(BaseQueryResolver):
query.filter( query.filter(
or_( or_(
models.Aircraft.created_by_id == user_id, 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)
)
) )
) )
) )
+27 -14
View File
@@ -35,38 +35,51 @@ class BaseQueryResolver(BaseResolver):
object_id: Optional[int] = None, object_id: Optional[int] = None,
order_by: Optional[list] = None, order_by: Optional[list] = None,
only_public: Optional[bool] = False, only_public: Optional[bool] = False,
only_my: Optional[bool] = False,
url_slug: Optional[str] = None,
**kwargs, **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 object_id:
if not hasattr(self.model, "id"): if not hasattr(self.model, "id"):
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) query = query.filter(self.model.id == object_id)
search = kwargs.pop("search", None) query = self._handle_search(query, 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))
if kwargs: if kwargs:
query = query.filter_by(**kwargs) query = query.filter_by(**kwargs)
return query 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]: 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) return await self._get_list(query)
async def get_one(self, user_id: Optional[int] = None, **kwargs) -> GQL_TYPE: 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) return await self._get_one(query)
@@ -14,10 +14,16 @@ class CopilotQueryResolver(BaseQueryResolver):
object_id: Optional[int] = None, object_id: Optional[int] = None,
order_by: Optional[list] = None, order_by: Optional[list] = None,
only_public: Optional[bool] = False, only_public: Optional[bool] = False,
*args, **kwargs **kwargs
): ):
pilot_username = kwargs.pop("pilot_username", None) 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: if pilot_username:
query = ( query = (
@@ -21,7 +21,8 @@ class EventQueryResolver(BaseQueryResolver):
user_id, object_id, user_id, object_id,
order_by=[models.Event.date_from.desc()], order_by=[models.Event.date_from.desc()],
url_slug=kwargs.get("url_slug"), url_slug=kwargs.get("url_slug"),
only_public=not bool(user_id) only_public=only_public,
only_my=True
) )
if kwargs.get('username'): if kwargs.get('username'):
@@ -41,7 +41,8 @@ class FlightQueryResolver(BaseQueryResolver):
user_id, user_id,
**filters, **filters,
order_by=[models.Flight.takeoff_datetime.desc()], 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"): if kwargs.get("event_id"):
@@ -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
+1 -9
View File
@@ -22,6 +22,7 @@ from graphql_schema.dataloaders.single_model import (
airport_weather_info_loader, organizations_dataloader, flight_dataloader, photo_adjustment_dataloader, airport_weather_info_loader, organizations_dataloader, flight_dataloader, photo_adjustment_dataloader,
photo_dataloader, user_dataloader photo_dataloader, user_dataloader
) )
from graphql_schema.permissions import IsAuthenticated
from graphql_schema.sqlalchemy_to_strawberry_type import strawberry_sqlalchemy_type 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 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) @strawberry_sqlalchemy_type(models.Flight)
class Flight: class Flight:
def __init__(self, **kwargs): def __init__(self, **kwargs):
+13
View File
@@ -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)
+52 -23
View File
@@ -18,6 +18,17 @@ from endpoints.photo_editor_preview import PhotoEditorEndpoint
from endpoints.registration import RegistrationInput, RegistrationEndpoint 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: class App:
api_router = APIRouter(dependencies=[]) api_router = APIRouter(dependencies=[])
access_security = JwtAccessBearerCookie( access_security = JwtAccessBearerCookie(
@@ -38,7 +49,21 @@ class App:
ignore_errors=[GraphQLError, HTTPException] 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_exception_handlers(app)
self.setup_middleware(app) self.setup_middleware(app)
@@ -65,22 +90,14 @@ class App:
@staticmethod @staticmethod
def setup_static_paths(app: FastAPI): def setup_static_paths(app: FastAPI):
app.mount("/uploads", StaticFiles(directory="/app/uploads"), name="uploads") app.mount("/uploads", app=StaticFilesCache(directory="/app/uploads"), name="uploads")
app.mount("/static", StaticFiles(directory="/app/static"), name="static") app.mount(
"/static",
app=StaticFilesCache(directory="/app/static", cachecontrol=f"public, max-age={24 * 3600}"),
name="static"
)
def setup_routes(self, app: FastAPI): 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) @self.api_router.post("/registration", status_code=201)
async def registration(user: RegistrationInput, background_tasks: BackgroundTasks): async def registration(user: RegistrationInput, background_tasks: BackgroundTasks):
return await RegistrationEndpoint().on_post(user, background_tasks) return await RegistrationEndpoint().on_post(user, background_tasks)
@@ -92,6 +109,23 @@ class App:
refresh_token=self.refresh_security refresh_token=self.refresh_security
).on_post(user, resp) ).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}") @self.api_router.get("/forgotten-password/token/{token}")
async def token_info(token: str): async def token_info(token: str):
return await ForgottenPasswordEndpoint().token_info(token) return await ForgottenPasswordEndpoint().token_info(token)
@@ -104,17 +138,10 @@ class App:
async def password_reset(input: ChangeForgottenPassword): async def password_reset(input: ChangeForgottenPassword):
return await ForgottenPasswordEndpoint().change_password(input) 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): async def contact(input: ContactInput, background_tasks: BackgroundTasks):
return await ContactEndpoint().on_post(input, background_tasks) 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}") @self.api_router.get("/photo/editor-preview/{photo_id}")
async def photo_editor_preview( async def photo_editor_preview(
photo_id: int, photo_id: int,
@@ -145,5 +172,7 @@ class App:
rotate=rotate, rotate=rotate,
) )
setup_graphql_endpoint(app, self.access_security)
# musi byt na konci # musi byt na konci
app.include_router(self.api_router) app.include_router(self.api_router)