diff --git a/src/database/models.py b/src/database/models.py index c6a65f9..475b663 100644 --- a/src/database/models.py +++ b/src/database/models.py @@ -3,12 +3,11 @@ import datetime from typing import Set, List from sqlalchemy import String, DateTime, ForeignKey, Text, Integer, func, Table, Column, Boolean, select, Float, Enum from sqlalchemy.dialects.mysql import JSON -from sqlalchemy.orm import Mapped, relationship, as_declarative, mapped_column +from sqlalchemy.orm import Mapped, relationship, mapped_column, DeclarativeBase from sqlalchemy.ext.asyncio import AsyncSession -@as_declarative() -class BaseModel: +class BaseModel(DeclarativeBase): excluded_columns_in_dict = ("deleted",) @classmethod @@ -119,6 +118,7 @@ class FlightPlan(BaseModel): created_by: Mapped['User'] = relationship() + class Track(BaseModel): __tablename__ = "track" @@ -134,7 +134,7 @@ class Track(BaseModel): created_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now()) flight: Mapped['Flight'] = relationship() - # created_by: Mapped['User'] = relationship() + track_points: Mapped[list['TrackPoint']] = relationship() class TrackPoint(BaseModel): @@ -142,7 +142,7 @@ class TrackPoint(BaseModel): id: Mapped[int] = mapped_column(primary_key=True) timestamp: Mapped[datetime] = mapped_column(DateTime, nullable=False, index=True) - track_id: Mapped[id] = mapped_column(Integer, ForeignKey('track.id')) + track_id: Mapped[int] = mapped_column(Integer, ForeignKey('track.id')) gps_latitude: Mapped[float] = mapped_column(Float, nullable=False) gps_longitude: Mapped[float] = mapped_column(Float, nullable=False) terrain_elevation: Mapped[float] = mapped_column(Float, nullable=True) diff --git a/src/database/query_builder.py b/src/database/query_builder.py index 6b06a04..9a909ee 100644 --- a/src/database/query_builder.py +++ b/src/database/query_builder.py @@ -1,10 +1,10 @@ -from typing import Optional, Type +from typing import Type from sqlalchemy import select, or_, and_ from database import models -class QueryBuilder: - def __init__(self, model: Type[models.BaseModel]): +class QueryBuilder[ModelType: models.BaseModel]: + def __init__(self, model: Type[ModelType]): self.model = model def get_simple_query( diff --git a/src/decorators/endpoints.py b/src/decorators/endpoints.py deleted file mode 100644 index 8c5b0b9..0000000 --- a/src/decorators/endpoints.py +++ /dev/null @@ -1,38 +0,0 @@ -from functools import wraps -from fastapi import HTTPException -from starlette.status import HTTP_401_UNAUTHORIZED - - -def raise_unauthorized(): - raise HTTPException(HTTP_401_UNAUTHORIZED, "Not authorized") - - -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_unauthorized() - - return await func(*args, **kwargs) - - return decorator - - -def authenticated_user_only(raise_when_unauthorized: bool = True, return_value_unauthorized=None): - def wrapper(func): - @wraps(func) - async def decorator(*args, **kwargs): - if 'info' in kwargs: - if not kwargs['info'].context.user_id: - if raise_when_unauthorized: - raise_unauthorized() - else: - return return_value_unauthorized - return await func(*args, **kwargs) - - return decorator - - return wrapper diff --git a/src/decorators/error_logging.py b/src/decorators/error_logging.py deleted file mode 100644 index 31bafa5..0000000 --- a/src/decorators/error_logging.py +++ /dev/null @@ -1,18 +0,0 @@ -from functools import wraps -from fastapi import HTTPException -from graphql import GraphQLError -from sqlalchemy.exc import NoResultFound - - -def error_logging(func): - @wraps(func) - async def decorator(*args, **kwargs): - try: - return await func(*args, **kwargs) - except NoResultFound as e: - raise GraphQLError("Not found", original_error=e) - except HTTPException as e: - if e.status_code == 401: - raise GraphQLError("Not authorized", original_error=e) - - return decorator diff --git a/src/endpoints/graphql.py b/src/endpoints/graphql.py deleted file mode 100644 index ea50baf..0000000 --- a/src/endpoints/graphql.py +++ /dev/null @@ -1,52 +0,0 @@ -from datetime import timedelta -from fastapi import FastAPI, Security, Depends, BackgroundTasks, APIRouter -from fastapi_jwt import JwtAuthorizationCredentials -from fastapi_jwt.jwt import JwtAccessBearerCookie -from sqlalchemy import select -from starlette.responses import RedirectResponse -from strawberry.fastapi import GraphQLRouter -from config import GRAPHIQL, APP_DEBUG -from database import async_session, models -from graphql_schema.schema import GraphQLContext, schema - - -def setup_graphql_endpoint(app: FastAPI, access_security: JwtAccessBearerCookie): - if APP_DEBUG: - debug_router = APIRouter() - - @debug_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 - - 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() - - if user_id: - async with async_session() as db: - organization_ids = set((await db.scalars( - select(models.user_is_in_organization.c.organization_id) - .filter(models.user_is_in_organization.c.user_id == user_id) - )).all()) - - return GraphQLContext( - user_id=user_id, - organization_ids=organization_ids, - jwt_auth_credentials=credentials, - jwt=access_security, - background_tasks=Depends(BackgroundTasks) - ) - - graphql_app = GraphQLRouter( - schema, - graphiql=GRAPHIQL, - debug=APP_DEBUG, - context_getter=setup_graphql_context - ) - app.include_router(graphql_app, prefix="/graphql") diff --git a/src/external/openair_parser.py b/src/external/openair_parser.py index 0b5bc38..a820b68 100644 --- a/src/external/openair_parser.py +++ b/src/external/openair_parser.py @@ -53,7 +53,7 @@ class Airspace: upper_limit: str = None lower_limit: str = None center: Optional[Coordinates] = None - radius_nm: Optional[float] = None + radius_nm: float | None = None bounds: list[Coordinates] = dataclasses.field(default_factory=lambda: []) diff --git a/src/graphql_schema/context.py b/src/graphql_schema/context.py new file mode 100644 index 0000000..746abb4 --- /dev/null +++ b/src/graphql_schema/context.py @@ -0,0 +1,33 @@ +import dataclasses +from fastapi import BackgroundTasks, Depends, Security +from fastapi_jwt import JwtAuthorizationCredentials +from sqlalchemy import select +from strawberry.fastapi import BaseContext +from database import async_session, models +from jwt import access_security + + +@dataclasses.dataclass +class GraphQLContext(BaseContext): + user_id: int + organization_ids: set[int] + jwt_auth_credentials: JwtAuthorizationCredentials + background_tasks: BackgroundTasks + +async def setup_graphql_context(credentials: JwtAuthorizationCredentials = Security(access_security)): + user_id = credentials['id'] if credentials else None + organization_ids = set() + + if user_id: + async with async_session() as db: + organization_ids = set((await db.scalars( + select(models.user_is_in_organization.c.organization_id) + .filter(models.user_is_in_organization.c.user_id == user_id) + )).all()) + + return GraphQLContext( + user_id=user_id, + organization_ids=organization_ids, + jwt_auth_credentials=credentials, + background_tasks=Depends(BackgroundTasks) + ) diff --git a/src/graphql_schema/dataloaders/base.py b/src/graphql_schema/dataloaders/base.py index 84999a0..301dcc2 100644 --- a/src/graphql_schema/dataloaders/base.py +++ b/src/graphql_schema/dataloaders/base.py @@ -1,5 +1,5 @@ from collections import defaultdict -from typing import Type, List, Optional +from typing import Type from logger import log from database import models, async_session from database.query_builder import QueryBuilder @@ -25,7 +25,7 @@ class BaseDataloader: class SingleModelByIdDataloader(BaseDataloader): - async def load(self, ids: List[int]): + async def load(self, ids: list[int]): async with async_session() as session: query = ( self.query_builder.get_simple_query(extra_select=[self.relationship_column], include_deleted=True) @@ -80,7 +80,7 @@ class MultiModelsDataloader(BaseDataloader): return query - async def load(self, ids: List[int]): + async def load(self, ids: list[int]): query = self.get_query(ids) async with async_session() as db: diff --git a/src/graphql_schema/entities/aircraft.py b/src/graphql_schema/entities/aircraft.py index db15ed6..84eb954 100644 --- a/src/graphql_schema/entities/aircraft.py +++ b/src/graphql_schema/entities/aircraft.py @@ -1,70 +1,47 @@ -from typing import Optional import strawberry -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 .helpers.filters import get_filters +from .helpers.pagination import 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 +from ..extensions.field.auth import AllowPublicAccess, AuthenticatedOnly +from ..extensions.field.pagination import OffsetPagination @strawberry.type class AircraftQueries: - @strawberry.field() - @error_logging - @authenticated_user_only() - async def aircrafts(root, info, limit: int, offset: int = 0) -> PaginationWindow[Aircraft]: - query = AircraftQueryResolver().get_query( + @strawberry.field(extensions=[OffsetPagination(item_type=Aircraft), AuthenticatedOnly()]) + async def aircrafts(root, info) -> PaginationWindow[Aircraft]: + return 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 + @strawberry.field(extensions=[AllowPublicAccess()]) async def aircraft( root, info, id: int | None = None, call_sign: str | None = None, public: bool | None = False ) -> Aircraft: - filter_params = {} - if id: - filter_params['object_id'] = id - - if call_sign: - filter_params['call_sign'] = call_sign - return await AircraftQueryResolver().get_one( user_id=info.context.user_id, organization_ids=info.context.organization_ids if not public else None, only_public=public, - **filter_params + **get_filters(object_id=id, call_sign=call_sign) ) @strawberry.type class AircraftMutation: - @strawberry.mutation - @error_logging - @authenticated_user_only() + @strawberry.mutation(extensions=[AuthenticatedOnly()]) async def create_aircraft(root, info, input: CreateAircraftInput) -> Aircraft: return await AircraftMutationResolver().create(info.context, input) - @strawberry.mutation - @error_logging - @authenticated_user_only() + @strawberry.mutation(extensions=[AuthenticatedOnly()]) async def edit_aircraft(root, info, id: int, input: EditAircraftInput) -> Aircraft: return await AircraftMutationResolver().update(id, info.context, data=input) - @strawberry.mutation - @authenticated_user_only() + @strawberry.mutation(extensions=[AuthenticatedOnly()]) async def delete_aircraft(self, info, id: int) -> Aircraft: return await AircraftMutationResolver().delete(info.context, id) diff --git a/src/graphql_schema/entities/airport.py b/src/graphql_schema/entities/airport.py index 1637aa3..9df880f 100644 --- a/src/graphql_schema/entities/airport.py +++ b/src/graphql_schema/entities/airport.py @@ -1,22 +1,17 @@ -from typing import List import strawberry from database import models -from decorators.error_logging import error_logging -from decorators.endpoints import authenticated_user_only from graphql_schema.entities.resolvers.base import BaseQueryResolver from graphql_schema.entities.types.types import Airport +from graphql_schema.extensions.field.auth import AuthenticatedOnly @strawberry.type class AirportQueries: @strawberry.field() - @error_logging - async def airports(root, info) -> List[Airport]: + async def airports(root, info) -> list[Airport]: return await BaseQueryResolver(Airport, models.Airport).get_list(info.context.user_id) - @strawberry.field() - @error_logging - @authenticated_user_only() + @strawberry.field(extensions=[AuthenticatedOnly()]) async def airport(root, info, id: int) -> Airport: return await BaseQueryResolver(Airport, models.Airport).get_one( object_id=id, diff --git a/src/graphql_schema/entities/airspace.py b/src/graphql_schema/entities/airspace.py index f5f074e..174b1c8 100644 --- a/src/graphql_schema/entities/airspace.py +++ b/src/graphql_schema/entities/airspace.py @@ -1,7 +1,5 @@ -from typing import List, Optional import strawberry from database import models -from decorators.error_logging import error_logging from graphql_schema.entities.resolvers.base import BaseQueryResolver from graphql_schema.entities.types.types import Airspace @@ -9,10 +7,9 @@ from graphql_schema.entities.types.types import Airspace @strawberry.type class AirspaceQueries: @strawberry.field() - @error_logging async def airspaces( - root, info, country: str | None = None, types: Optional[list[str]] = None - ) -> List[Airspace]: + root, info, country: str | None = None, types: list[str] | None = None + ) -> list[Airspace]: filters = [] if country: diff --git a/src/graphql_schema/entities/copilot.py b/src/graphql_schema/entities/copilot.py index 0a12f8f..e8b85d3 100644 --- a/src/graphql_schema/entities/copilot.py +++ b/src/graphql_schema/entities/copilot.py @@ -1,28 +1,22 @@ -from typing import List, Optional import strawberry from graphql import GraphQLError from strawberry.types import Info from database import models -from decorators.error_logging import error_logging -from decorators.endpoints import authenticated_user_only, allow_public -from graphql_schema.entities.helpers.detail import get_detail_filters +from graphql_schema.entities.helpers.filters import get_filters 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 from graphql_schema.entities.types.types import Copilot +from graphql_schema.extensions.field.auth import AllowPublicAccess, AuthenticatedOnly @strawberry.type class CopilotQueries: - @strawberry.field() - @error_logging - @authenticated_user_only() - async def copilots(root, info: Info) -> List[Copilot]: + @strawberry.field(extensions=[AuthenticatedOnly()]) + async def copilots(root, info: Info) -> list[Copilot]: return await CopilotQueryResolver().get_list(info.context.user_id) - @strawberry.field() - @error_logging - @allow_public + @strawberry.field(extensions=[AllowPublicAccess()]) async def copilot( root, info: Info, id: int | None = None, @@ -32,17 +26,10 @@ class CopilotQueries: upload_flight_slug: str | None = None, public: bool | None = False ) -> Copilot: - filter_params = {} - if id: - filter_params['object_id'] = id - if url_slug is not None: - filter_params['url_slug'] = url_slug - if upload_token and upload_flight_slug: - filter_params['upload_token'] = upload_token - filter_params['upload_flight_slug'] = upload_flight_slug - if pilot_username: - filter_params['pilot_username'] = pilot_username - + filter_params = get_filters( + object_id=id, url_slug=url_slug, pilot_username=pilot_username, upload_token=upload_token, + upload_flight_slug=upload_flight_slug, + ) if not filter_params: raise GraphQLError(f"Invalid identification supplied: {filter_params}") @@ -55,14 +42,10 @@ class CopilotQueries: @strawberry.type class CopilotMutation: - @strawberry.mutation - @error_logging - @authenticated_user_only() - async def create_copilot(root, info, input: CreateCopilotInput) -> Copilot: + @strawberry.mutation(extensions=[AuthenticatedOnly()]) + async def create_copilot(root, info: Info, input: CreateCopilotInput) -> Copilot: return await BaseMutationResolver(Copilot, models.Copilot).create(info.context, data=input) - @strawberry.mutation - @error_logging - @authenticated_user_only() - async def edit_copilot(root, info, id: int, input: EditCopilotInput) -> Copilot: - return await BaseMutationResolver(Copilot, models.Copilot).update(id, input, info.context.user_id) + @strawberry.mutation(extensions=[AuthenticatedOnly()]) + async def edit_copilot(root, info: Info, id: int, input: EditCopilotInput) -> Copilot: + return await BaseMutationResolver(Copilot, models.Copilot).update(info.context, id, input, info.context.user_id) diff --git a/src/graphql_schema/entities/event.py b/src/graphql_schema/entities/event.py index 949a01e..ddc76ce 100644 --- a/src/graphql_schema/entities/event.py +++ b/src/graphql_schema/entities/event.py @@ -1,46 +1,33 @@ -from typing import Optional import strawberry from database import models -from decorators.endpoints import authenticated_user_only, allow_public -from decorators.error_logging import error_logging -from graphql_schema.entities.helpers.detail import get_detail_filters -from graphql_schema.entities.helpers.pagination import PaginationWindow, get_pagination_window +from graphql_schema.entities.helpers.filters import get_detail_filters +from graphql_schema.entities.helpers.pagination import PaginationWindow 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 +from graphql_schema.extensions.field.auth import AllowPublicAccess, AuthenticatedOnly +from graphql_schema.extensions.field.pagination import OffsetPagination @strawberry.type class EventQueries: - @strawberry.field() - @error_logging - @allow_public + @strawberry.field(extensions=[OffsetPagination(item_type=Event), AllowPublicAccess()]) async def events( root, info, - limit: int, - offset: int = 0, username: str | None = None, public: bool | None = False, ) -> PaginationWindow[Event]: - query = EventQueryResolver().get_query( + return EventQueryResolver().get_query( user_id=info.context.user_id, username=username, order_by=[models.Event.date_from.desc(), models.Event.name.desc()], only_public=public, ) - return await get_pagination_window( - query=query, - item_type=Event, - limit=limit, - offset=offset - ) - @strawberry.field() - @error_logging - @allow_public + @strawberry.field(extensions=[AllowPublicAccess()]) async def event( root, info, id: int | None = None, @@ -61,14 +48,10 @@ class EventQueries: @strawberry.type class EventMutation: - @strawberry.mutation - @error_logging - @authenticated_user_only() + @strawberry.mutation(extensions=[AuthenticatedOnly()]) async def create_event(root, info, input: CreateEventInput) -> Event: return await BaseMutationResolver(Event, models.Event).create(info.context, input) - @strawberry.mutation - @error_logging - @authenticated_user_only() + @strawberry.mutation(extensions=[AuthenticatedOnly()]) async def edit_event(root, info, id: int, input: EditEventInput) -> Event: - return await BaseMutationResolver(Event, models.Event).update(id, input, info.context.user_id) + return await BaseMutationResolver(Event, models.Event).update(info.context, id, input, info.context.user_id) diff --git a/src/graphql_schema/entities/flight.py b/src/graphql_schema/entities/flight.py index 620f94a..ac08e75 100644 --- a/src/graphql_schema/entities/flight.py +++ b/src/graphql_schema/entities/flight.py @@ -1,24 +1,19 @@ -from typing import Optional import strawberry -from decorators.endpoints import authenticated_user_only, allow_public -from decorators.error_logging import error_logging from graphql_schema.entities.resolvers.flight import FlightMutationResolver, FlightQueryResolver from graphql_schema.entities.types.mutation_input import EditFlightInput, CreateFlightInput from graphql_schema.entities.types.types import Flight -from .helpers.detail import get_detail_filters -from .helpers.pagination import PaginationWindow, get_pagination_window +from .helpers.filters import get_detail_filters +from .helpers.pagination import PaginationWindow +from ..extensions.field.auth import AllowPublicAccess, AuthenticatedOnly +from ..extensions.field.pagination import OffsetPagination @strawberry.type class FlightQueries: - @strawberry.field() - @error_logging - @allow_public + @strawberry.field(extensions=[OffsetPagination(item_type=Flight), AllowPublicAccess()]) async def flights( root, info, - limit: int, - offset: int = 0, username: str | None = None, event_id: int | None = None, public: bool | None = False, @@ -26,7 +21,7 @@ class FlightQueries: point_of_interest_id: int | None = None, aircraft_id: int | None = None, ) -> PaginationWindow[Flight]: - query = FlightQueryResolver().get_query( + return FlightQueryResolver().get_query( user_id=info.context.user_id, username=username, event_id=event_id, @@ -36,16 +31,7 @@ class FlightQueries: 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 + @strawberry.field(extensions=[AllowPublicAccess()]) async def flight( root, info, id: int | None = None, @@ -66,20 +52,14 @@ class FlightQueries: @strawberry.type class FlightMutation: - @strawberry.mutation - @error_logging - @authenticated_user_only() + @strawberry.mutation(extensions=[AuthenticatedOnly()]) async def create_flight(self, info, input: CreateFlightInput) -> Flight: return await FlightMutationResolver().create(info.context, input) - @strawberry.mutation - @error_logging - @authenticated_user_only() + @strawberry.mutation(extensions=[AuthenticatedOnly()]) async def edit_flight(self, info, id: int, input: EditFlightInput) -> Flight: return await FlightMutationResolver().update(info.context, id, input) - @strawberry.mutation - @error_logging - @authenticated_user_only() + @strawberry.mutation(extensions=[AuthenticatedOnly()]) async def delete_flight(self, info, id: int) -> Flight: return await FlightMutationResolver().delete(info.context.user_id, id) diff --git a/src/graphql_schema/entities/flight_plan.py b/src/graphql_schema/entities/flight_plan.py index 83ffd8b..851bbf4 100644 --- a/src/graphql_schema/entities/flight_plan.py +++ b/src/graphql_schema/entities/flight_plan.py @@ -1,25 +1,19 @@ -from typing import List, Optional import strawberry from strawberry.types import Info -from decorators.endpoints import authenticated_user_only, allow_public -from decorators.error_logging import error_logging -from graphql_schema.entities.helpers.detail import get_detail_filters +from graphql_schema.entities.helpers.filters import get_detail_filters from graphql_schema.entities.resolvers.flight_plan import FlightPlanMutationResolver, FlightPlanQueryResolver from graphql_schema.entities.types.mutation_input import CreateFlightPlanInput, EditFlightPlanInput from graphql_schema.entities.types.types import FlightPlan +from graphql_schema.extensions.field.auth import AuthenticatedOnly, AllowPublicAccess @strawberry.type class FlightPlanQueries: - @strawberry.field() - @error_logging - @authenticated_user_only() - async def flight_plans(root, info: Info) -> List[FlightPlan]: + @strawberry.field(extensions=[AuthenticatedOnly()]) + async def flight_plans(root, info: Info) -> list[FlightPlan]: return await FlightPlanQueryResolver().get_list(info.context.user_id) - @strawberry.field() - @error_logging - @allow_public + @strawberry.field(extensions=[AllowPublicAccess()]) async def flight_plan( root, info: Info, @@ -40,14 +34,10 @@ class FlightPlanQueries: @strawberry.type class FlightPlanMutation: - @strawberry.mutation - @error_logging - @authenticated_user_only() - async def create_flight_plan(root, info, input: CreateFlightPlanInput) -> FlightPlan: + @strawberry.mutation(extensions=[AuthenticatedOnly()]) + async def create_flight_plan(root, info: Info, input: CreateFlightPlanInput) -> FlightPlan: return await FlightPlanMutationResolver().create(info.context, data=input) - @strawberry.mutation - @error_logging - @authenticated_user_only() - async def edit_flight_plan(root, info, id: int, input: EditFlightPlanInput) -> FlightPlan: + @strawberry.mutation(extensions=[AuthenticatedOnly()]) + async def edit_flight_plan(root, info: Info, id: int, input: EditFlightPlanInput) -> FlightPlan: return await FlightPlanMutationResolver().update(info.context, id, input) diff --git a/src/graphql_schema/entities/helpers/detail.py b/src/graphql_schema/entities/helpers/filters.py similarity index 55% rename from src/graphql_schema/entities/helpers/detail.py rename to src/graphql_schema/entities/helpers/filters.py index 2acab0d..040c988 100644 --- a/src/graphql_schema/entities/helpers/detail.py +++ b/src/graphql_schema/entities/helpers/filters.py @@ -1,13 +1,14 @@ -from typing import Optional +from typing import Any + from graphql import GraphQLError +def get_filters(**kwargs) -> dict[str, Any]: + return {k: v for k, v in kwargs.items() if v is not None} + + def get_detail_filters(id: int | None = None, url_slug: str | None = None) -> dict: - filter_params = {} - if id: - filter_params['object_id'] = id - if url_slug is not None: - filter_params['url_slug'] = url_slug + filter_params = get_filters(object_id=id, url_slug=url_slug) if not filter_params: raise GraphQLError("You must specifiy either urlSlug or id!") diff --git a/src/graphql_schema/entities/helpers/pagination.py b/src/graphql_schema/entities/helpers/pagination.py index a6014e4..20dd23c 100644 --- a/src/graphql_schema/entities/helpers/pagination.py +++ b/src/graphql_schema/entities/helpers/pagination.py @@ -30,10 +30,6 @@ async def get_pagination_window( 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] diff --git a/src/graphql_schema/entities/organization.py b/src/graphql_schema/entities/organization.py index f0c505c..33ca238 100644 --- a/src/graphql_schema/entities/organization.py +++ b/src/graphql_schema/entities/organization.py @@ -3,27 +3,23 @@ import strawberry from sqlalchemy import delete from sqlalchemy.dialects.mysql import insert from sqlalchemy.exc import IntegrityError +from strawberry import Info from database import models -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 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 +from graphql_schema.extensions.field.auth import AuthenticatedOnly @strawberry.type class OrganizationQueries: - @strawberry.field() - @error_logging - @authenticated_user_only() + @strawberry.field(extensions=[AuthenticatedOnly()]) async def organizations(root, info) -> List[Organization]: return await OrganizationQueryResolver().get_list() - @strawberry.field() - @error_logging - @authenticated_user_only() + @strawberry.field(extensions=[AuthenticatedOnly()]) async def organization(root, info, id: int) -> Organization: return await OrganizationQueryResolver().get_one(object_id=id) @@ -31,17 +27,14 @@ class OrganizationQueries: @strawberry.type class OrganizationMutation: - @strawberry.mutation - @error_logging - @authenticated_user_only() - async def create_organization(root, info, input: CreateOrganizationInput) -> Organization: + @strawberry.mutation(extensions=[AuthenticatedOnly()]) + async def create_organization(root, info: Info, input: CreateOrganizationInput) -> Organization: return await BaseMutationResolver(Organization, models.Organization).create(info.context, data=input) - @strawberry.mutation - @error_logging - @authenticated_user_only() - async def edit_organization(root, info, id: int, input: EditOrganizationInput) -> Organization: + @strawberry.mutation(extensions=[AuthenticatedOnly()]) + async def edit_organization(root, info: Info, id: int, input: EditOrganizationInput) -> Organization: return await BaseMutationResolver(Organization, models.Organization).update( + info.context, id, data=input, user_id=info.context.user_id @@ -51,9 +44,7 @@ class OrganizationMutation: @strawberry.type class OrganizationUserMutation: - @strawberry.mutation - @error_logging - @authenticated_user_only() + @strawberry.mutation(extensions=[AuthenticatedOnly()]) async def add_to_organization(root, info, organization_id: int) -> Organization: async with get_session() as db: organization = (await db.scalars( @@ -72,9 +63,7 @@ class OrganizationUserMutation: return Organization(**organization.as_dict()) - @strawberry.mutation - @error_logging - @authenticated_user_only() + @strawberry.mutation(extensions=[AuthenticatedOnly()]) async def remove_from_organization(root, info, organization_id: int) -> Organization: async with get_session() as db: organization = (await db.scalars( diff --git a/src/graphql_schema/entities/photo.py b/src/graphql_schema/entities/photo.py index 7dcfc4e..5e0e624 100644 --- a/src/graphql_schema/entities/photo.py +++ b/src/graphql_schema/entities/photo.py @@ -1,28 +1,31 @@ -from typing import List, Optional import strawberry +from fastapi import HTTPException +from starlette.status import HTTP_401_UNAUTHORIZED +from strawberry import Info from database import models -from decorators.endpoints import authenticated_user_only, allow_public, raise_unauthorized -from decorators.error_logging import error_logging from graphql_schema.entities.resolvers.base import BaseQueryResolver from graphql_schema.entities.resolvers.photo import PhotoMutationResolver, PhotoQueryResolver from graphql_schema.entities.types.types import Photo from graphql_schema.entities.types.mutation_input import EditPhotoInput, UploadPhotoInput, AdjustmentInput +from graphql_schema.extensions.field.auth import AuthenticatedOnly, AllowPublicAccess + + +def raise_unauthorized(): + raise HTTPException(HTTP_401_UNAUTHORIZED, "Not authorized") @strawberry.type class PhotoQueries: - @strawberry.field() - @error_logging - @allow_public + @strawberry.field(extensions=[AllowPublicAccess()]) async def photos( - root, info, + root, info: Info, flight_id: int | None = None, copilot_id: int | None = None, uploaded_by_copilot_id: int | None = None, point_of_interest_id: int | None = None, aircraft_id: int | None = None, public: bool | None = False, - ) -> List[Photo]: + ) -> list[Photo]: return await PhotoQueryResolver().get_list( public=public, flight_id=flight_id, @@ -34,10 +37,8 @@ class PhotoQueries: order_by=[models.Photo.exposed_at] ) - @strawberry.field() - @error_logging - @allow_public - async def photo(root, info, id: int, public: bool | None = False, ) -> Photo: + @strawberry.field(extensions=[AllowPublicAccess()]) + async def photo(root, info: Info, id: int, public: bool | None = False) -> Photo: return await BaseQueryResolver(Photo, models.Photo).get_one( object_id=id, user_id=info.context.user_id, @@ -47,22 +48,18 @@ class PhotoQueries: @strawberry.type class PhotoMutation: - @strawberry.mutation - @error_logging + @strawberry.mutation() async def upload_photo(self, info, input: UploadPhotoInput) -> Photo: if info.context.user_id is None and not input.copilot_upload_token: raise_unauthorized() return await PhotoMutationResolver().upload(info, input) - @strawberry.mutation() - @error_logging - @authenticated_user_only() + @strawberry.mutation(extensions=[AuthenticatedOnly()]) async def edit_photo(self, info, id: int, input: EditPhotoInput) -> Photo: return await PhotoMutationResolver().update(info.context, id, input, info.context.user_id) @strawberry.mutation() - @error_logging async def change_orientation(self, info, id: int, direction: str, copilot_upload_token: str | None = None) -> Photo: if info.context.user_id is None and not copilot_upload_token: raise_unauthorized() @@ -75,14 +72,11 @@ class PhotoMutation: info=info ) - @strawberry.mutation() - @error_logging - @authenticated_user_only() + @strawberry.mutation(extensions=[AuthenticatedOnly()]) async def adjust_photo(self, info, id: int, adjustment: AdjustmentInput) -> Photo: return await PhotoMutationResolver().adjust(id, info=info, user_id=info.context.user_id, adjustment=adjustment) @strawberry.mutation() - @error_logging async def delete_photo(self, info, id: int, copilot_upload_token: str | None = None) -> Photo: if info.context.user_id is None and not copilot_upload_token: raise_unauthorized() diff --git a/src/graphql_schema/entities/poi.py b/src/graphql_schema/entities/poi.py index b12a56b..a07b0de 100644 --- a/src/graphql_schema/entities/poi.py +++ b/src/graphql_schema/entities/poi.py @@ -1,44 +1,33 @@ -from typing import Optional import strawberry from database import models -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.detail import get_detail_filters -from graphql_schema.entities.helpers.pagination import get_pagination_window, PaginationWindow +from graphql_schema.entities.helpers.filters import get_detail_filters +from graphql_schema.entities.helpers.pagination import 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 +from graphql_schema.extensions.field.auth import AuthenticatedOnly, AllowPublicAccess +from graphql_schema.extensions.field.pagination import OffsetPagination @strawberry.type class PointOfInterestQueries: - @strawberry.field() - @error_logging - @allow_public + @strawberry.field(extensions=[OffsetPagination(item_type=PointOfInterest), AllowPublicAccess()]) async def points_of_interest( root, info, - limit: int, offset: int = 0, search: str | None = None, - public: bool = False + public: bool = False, ) -> PaginationWindow[PointOfInterest]: - query = BaseQueryResolver(PointOfInterest, models.PointOfInterest).get_query( + return BaseQueryResolver(PointOfInterest, models.PointOfInterest).get_query( info.context.user_id, only_my=bool(info.context.user_id), include_others_public=True, only_public=public, search=search, ) - return await get_pagination_window( - query=query, - item_type=PointOfInterest, - limit=limit, - offset=offset - ) - @strawberry.field() - @allow_public + @strawberry.field(extensions=[AllowPublicAccess()]) async def point_of_interest( root, info, url_slug: str | None = None, @@ -56,9 +45,7 @@ class PointOfInterestQueries: @strawberry.type class PointOfInterestMutation: - @strawberry.mutation - @error_logging - @authenticated_user_only() + @strawberry.mutation(extensions=[AuthenticatedOnly()]) async def create_point_of_interest(root, info, input: CreatePointOfInterestInput) -> PointOfInterest: input_data = input.to_dict() @@ -73,9 +60,7 @@ class PointOfInterestMutation: db, input_data ) - @strawberry.mutation - @error_logging - @authenticated_user_only() + @strawberry.mutation(extensions=[AuthenticatedOnly()]) async def edit_point_of_interest(root, info, id: int, input: EditPointOfInterestInput) -> PointOfInterest: input_data = input.to_dict() @@ -93,8 +78,6 @@ class PointOfInterestMutation: updated_poi = await models.PointOfInterest.update(db, obj=poi, data=input_data) return PointOfInterest(**updated_poi.as_dict()) - @strawberry.mutation - @error_logging - @authenticated_user_only() + @strawberry.mutation(extensions=[AuthenticatedOnly()]) async def delete_point_of_interest(self, info, id: int) -> PointOfInterest: return await BaseMutationResolver(PointOfInterest, models.PointOfInterest).delete(info.context.user_id, id=id) diff --git a/src/graphql_schema/entities/poi_type.py b/src/graphql_schema/entities/poi_type.py index 2cb9076..8261d89 100644 --- a/src/graphql_schema/entities/poi_type.py +++ b/src/graphql_schema/entities/poi_type.py @@ -1,92 +1,20 @@ -from typing import List import strawberry from database import models -from decorators.endpoints import authenticated_user_only -from decorators.error_logging import error_logging from graphql_schema.entities.resolvers.base import BaseQueryResolver from graphql_schema.entities.types.types import PointOfInterestType +from graphql_schema.extensions.field.auth import AuthenticatedOnly @strawberry.type class PointOfInterestTypeQueries: - @strawberry.field() - @error_logging - @authenticated_user_only() - async def point_of_interest_types(root, info) -> List[PointOfInterestType]: + @strawberry.field(extensions=[AuthenticatedOnly()]) + async def point_of_interest_types(root, info) -> list[PointOfInterestType]: return await BaseQueryResolver(PointOfInterestType, models.PointOfInterestType).get_list(info.context.user_id) - @strawberry.field() - @error_logging - @authenticated_user_only() + @strawberry.field(extensions=[AuthenticatedOnly()]) async def point_of_interest_type(root, info, id: int) -> PointOfInterestType: return await BaseQueryResolver(PointOfInterestType, models.PointOfInterestType).get_one( object_id=id, user_id=info.context.user_id ) - -# -# @strawberry.type -# class CreatePointOfInterestMutation: -# @strawberry_sqlalchemy_input(models.PointOfInterest, exclude_fields=['id', 'type_id']) -# class CreatePointOfInterestInput: -# type: # Optional[ComboboxInput] = None -# -# @strawberry.mutation -# @authenticated_user_only() -# async def create_point_of_interest(root, info, input: CreatePointOfInterestInput) -> PointOfInterest: -# input_data = input.to_dict() -# -# input_data['type_id'] = await handle_combobox_save( -# info.context.db, -# models.PointOfInterestType, -# input.type, -# info.context.user_id -# ) -# -# return await models.PointOfInterest.create( -# info.context.db, -# data=dict( -# **input_data, -# created_by_id=info.context.user_id, -# ) -# ) -# -# -# @strawberry.type -# class EditPointOfInterestMutation: -# @strawberry_sqlalchemy_input(models.PointOfInterest, exclude_fields=['id', 'type_id']) -# class EditPointOfInterestInput: -# type: Optional[ComboboxInput] = None -# -# @strawberry.mutation -# @authenticated_user_only() -# async def edit_point_of_interest(root, info, id: int, input: EditPointOfInterestInput) -> PointOfInterest: -# # TODO: kontrola organizace -# input_data = input.to_dict() -# -# if 'type' in input: -# input_data['type_id'] = await handle_combobox_save( -# info.context.db, -# models.PointOfInterestType, -# input.type, -# info.context.user_id -# ) -# -# poi = ( -# await info.context.db.scalars( -# get_base_query(info.context.user_id, only_my=True) -# .filter(models.PointOfInterest.id == id)) -# ).one() -# return await models.PointOfInterest.update(info.context.db, obj=poi, data=input_data) -# -# -# @strawberry.type -# class DeletePointOfInterestMutation: -# -# @strawberry.mutation -# @authenticated_user_only() -# async def delete_point_of_interest(self, info, id: int) -> PointOfInterest: -# poi = get_base_query(info.context.user_id, only_my=True).filter(models.PointOfInterest.id == id).one() -# -# return await models.PointOfInterest.update(info.context.db, obj=poi, data=dict(deleted=True)) diff --git a/src/graphql_schema/entities/resolvers/aircraft.py b/src/graphql_schema/entities/resolvers/aircraft.py index 271e37e..a7c6db3 100644 --- a/src/graphql_schema/entities/resolvers/aircraft.py +++ b/src/graphql_schema/entities/resolvers/aircraft.py @@ -3,6 +3,7 @@ from typing import Set, Optional from sqlalchemy import and_ from database import models from database.transaction import get_session +from graphql_schema.context import GraphQLContext 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 @@ -60,7 +61,7 @@ class AircraftMutationResolver(BaseMutationResolver): def __init__(self): super().__init__(graphql_type=Aircraft, model=models.Aircraft) - async def create(self, context, data: CreateAircraftInput) -> Aircraft: + async def create(self, context: GraphQLContext, data: CreateAircraftInput) -> Aircraft: input_data = data.to_dict() async with get_session() as db: @@ -75,7 +76,7 @@ class AircraftMutationResolver(BaseMutationResolver): return await self._do_create(db, data=input_data) - async def update(self, context, id: int, data: EditAircraftInput) -> Aircraft: + async def update(self, context: GraphQLContext, id: int, data: EditAircraftInput) -> Aircraft: update_data = data.to_dict() async with get_session() as db: if data.organization: diff --git a/src/graphql_schema/entities/resolvers/base.py b/src/graphql_schema/entities/resolvers/base.py index b03932c..706066b 100644 --- a/src/graphql_schema/entities/resolvers/base.py +++ b/src/graphql_schema/entities/resolvers/base.py @@ -1,12 +1,13 @@ -from typing import Optional, Type, TypeVar, Generic, List - +from typing import Type, TypeVar, Generic from sqlalchemy import or_ from sqlalchemy.ext.asyncio import AsyncSession from database import models from database.query_builder import QueryBuilder from database.transaction import get_session +from graphql_schema.context import GraphQLContext from graphql_schema.entities.types.base import BaseGraphqlInputType + GQL_TYPE = TypeVar('GQL_TYPE') @@ -18,7 +19,7 @@ class BaseResolver(Generic[GQL_TYPE]): class BaseQueryResolver(BaseResolver): - async def _get_list(self, query) -> List[GQL_TYPE]: + async def _get_list(self, query) -> list[GQL_TYPE]: async with get_session() as db: items = (await db.scalars(query)).all() @@ -77,7 +78,7 @@ class BaseQueryResolver(BaseResolver): query = query.filter(or_(*search_clauses)) return query - async def get_list(self, user_id: int | None = None, **kwargs) -> List[GQL_TYPE]: + async def get_list(self, user_id: int | None = None, **kwargs) -> list[GQL_TYPE]: query = self.get_query(user_id=user_id, **kwargs) return await self._get_list(query) @@ -107,7 +108,7 @@ class BaseMutationResolver(BaseResolver): model = await self.model.update(db, data=data, **update_where) return self.graphql_type(**model.as_dict()) - async def create(self, context, data: BaseGraphqlInputType) -> GQL_TYPE: + async def create(self, context: GraphQLContext, data: BaseGraphqlInputType) -> GQL_TYPE: input_data = data.to_dict() if hasattr(self.model, "created_by_id"): @@ -116,12 +117,12 @@ class BaseMutationResolver(BaseResolver): async with get_session() as db: return await self._do_create(db, input_data) - async def update(self, context, id: int, data: BaseGraphqlInputType, user_id: int) -> GQL_TYPE: + async def update(self, context: GraphQLContext, id: int, data: BaseGraphqlInputType) -> GQL_TYPE: async with get_session() as db: - item = await self._get_one(db, id, user_id) + item = await self._get_one(db, id, context.user_id) return await self._do_update(db, item, data.to_dict()) - async def delete(self, context, id: int, **kwargs) -> GQL_TYPE: + async def delete(self, context: GraphQLContext, id: int, **kwargs) -> GQL_TYPE: async with get_session() as db: model = await self._get_one(db, id, context.user_id) diff --git a/src/graphql_schema/entities/resolvers/copilot.py b/src/graphql_schema/entities/resolvers/copilot.py index 67e06d4..7825aa0 100644 --- a/src/graphql_schema/entities/resolvers/copilot.py +++ b/src/graphql_schema/entities/resolvers/copilot.py @@ -1,7 +1,4 @@ -from typing import Optional - from sqlalchemy import and_ - from database import models from graphql_schema.entities.resolvers.base import BaseQueryResolver from graphql_schema.entities.types.types import Copilot diff --git a/src/graphql_schema/entities/resolvers/flight.py b/src/graphql_schema/entities/resolvers/flight.py index 69be5f1..35921af 100644 --- a/src/graphql_schema/entities/resolvers/flight.py +++ b/src/graphql_schema/entities/resolvers/flight.py @@ -5,6 +5,7 @@ from sqlalchemy import delete, insert, or_, select from sqlalchemy.ext.asyncio import AsyncSession from background_jobs.elevation import add_terrain_elevation_to_flight from background_jobs.flight_title_photo import add_circular_avatar, generate_flight_title_photo +from graphql_schema.context import GraphQLContext from utils.flight_track_helpers import handle_upload_gpx, save_track_from_gpx_to_db, extract_basic_flight_info_from_gpx from background_jobs.weather import download_weather_for_flight from database import models @@ -82,7 +83,7 @@ class FlightMutationResolver(BaseMutationResolver): def __init__(self): super().__init__(Flight, models.Flight) - async def create(self, context, input: CreateFlightInput) -> Flight: + async def create(self, context: GraphQLContext, input: CreateFlightInput) -> Flight: data = input.to_dict() user_id = context.user_id @@ -125,7 +126,7 @@ class FlightMutationResolver(BaseMutationResolver): return flight - async def update(self, context, id: int, input: EditFlightInput) -> Flight: + async def update(self, context: GraphQLContext, id: int, input: EditFlightInput) -> Flight: user_id = context.user_id async with get_session() as db: flight = await self._get_one(db, id, user_id) @@ -197,13 +198,13 @@ def schedule_background_tasks(flight_id: int, flight_data: dict, context) -> Non if flight_data.get("takeoff_airport_id"): context.background_tasks.add_task( - download_weather_for_flight, flight_id=id, airport_id=flight_data['takeoff_airport_id'], + download_weather_for_flight, flight_id=flight_id, airport_id=flight_data['takeoff_airport_id'], date_time=flight_data['takeoff_datetime'], type_="takeoff" ) if flight_data.get("landing_airport_id"): context.background_tasks.add_task( - download_weather_for_flight, flight_id=id, airport_id=flight_data['landing_airport_id'], + download_weather_for_flight, flight_id=flight_id, airport_id=flight_data['landing_airport_id'], date_time=flight_data['landing_datetime'], type_="landing" ) diff --git a/src/graphql_schema/entities/resolvers/flight_plan.py b/src/graphql_schema/entities/resolvers/flight_plan.py index 9848465..38c34af 100644 --- a/src/graphql_schema/entities/resolvers/flight_plan.py +++ b/src/graphql_schema/entities/resolvers/flight_plan.py @@ -1,10 +1,10 @@ import asyncio -from typing import Optional from sqlalchemy import delete, select from sqlalchemy.dialects.mysql import insert from database import models from database.models import flight_plan_has_copilot from database.transaction import get_session +from graphql_schema.context import GraphQLContext from graphql_schema.entities.helpers.combobox import handle_combobox_save from graphql_schema.entities.resolvers.base import BaseMutationResolver, BaseQueryResolver from graphql_schema.entities.resolvers.flight import handle_aircraft_save @@ -49,7 +49,7 @@ class FlightPlanMutationResolver(BaseMutationResolver): def __init__(self): super().__init__(graphql_type=FlightPlan, model=models.FlightPlan) - async def create(self, context, data: CreateFlightPlanInput) -> FlightPlan: + async def create(self, context: GraphQLContext, data: CreateFlightPlanInput) -> FlightPlan: input_data = data.to_dict() input_data['created_by_id'] = context.user_id @@ -64,7 +64,7 @@ class FlightPlanMutationResolver(BaseMutationResolver): ) return flight_plan - async def update(self, context, id: int, data: EditFlightPlanInput) -> FlightPlan: + async def update(self, context: GraphQLContext, id: int, data: EditFlightPlanInput) -> FlightPlan: input_data = data.to_dict() user_id = context.user_id @@ -87,10 +87,10 @@ class FlightPlanMutationResolver(BaseMutationResolver): if flight_plan_model.is_default_name: if not markers: - markers = await db.scalars( + markers = (await db.scalars( select(models.FlightPlanMarker) .filter(models.FlightPlanMarker.flight_plan_id == id) - ).all() + )).all() used_markers = evenly_spaced_elements(markers, 5) input_data['name'] = " - ".join(m.name for m in used_markers) diff --git a/src/graphql_schema/entities/resolvers/photo.py b/src/graphql_schema/entities/resolvers/photo.py index 60b8abb..a0f9072 100644 --- a/src/graphql_schema/entities/resolvers/photo.py +++ b/src/graphql_schema/entities/resolvers/photo.py @@ -11,6 +11,7 @@ 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.context import GraphQLContext 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 EditPhotoInput, UploadPhotoInput, AdjustmentInput @@ -172,17 +173,17 @@ class PhotoMutationResolver(BaseMutationResolver): return photo - async def update(self, context, id: int, input: EditPhotoInput, user_id: int) -> Photo: + async def update(self, context: GraphQLContext, id: int, input: EditPhotoInput) -> Photo: data = input.to_dict() async with get_session() as db: - photo = await self._get_one(db, id, created_by_id=user_id) + photo = await self._get_one(db, id, created_by_id=context.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, + context.user_id, extra_data={"description": ""} ) @@ -268,7 +269,7 @@ class PhotoMutationResolver(BaseMutationResolver): "cache_key": int(time()) }) - async def delete(self, context, id: int, **kwargs) -> Photo: + async def delete(self, context: GraphQLContext, id: int, **kwargs) -> Photo: copilot_upload_token = kwargs.get("copilot_upload_token") await self._get_photo_details(id, context.user_id, copilot_upload_token, copy_original=False) diff --git a/src/graphql_schema/entities/types/mutation_input.py b/src/graphql_schema/entities/types/mutation_input.py index b28efec..394d7dd 100644 --- a/src/graphql_schema/entities/types/mutation_input.py +++ b/src/graphql_schema/entities/types/mutation_input.py @@ -112,12 +112,12 @@ class CropInput(BaseGraphqlInputType): @strawberry.input class AdjustmentInput: - rotate: Optional[float] = 0 + rotate: float | None = 0 crop_after_rotate: bool | None = True, - brightness: Optional[float] = 1 - contrast: Optional[float] = 1 - saturation: Optional[float] = 1 - sharpness: Optional[float] = 1 + brightness: float | None = 1 + contrast: float | None = 1 + saturation: float | None = 1 + sharpness: float | None = 1 crop: Optional[CropInput] = None @@ -153,8 +153,8 @@ class TrackItemInput: point_of_interest: Optional[ComboboxInput] = None airport: Optional[ComboboxInput] = None landing_duration: int | None = None - gps_latitude: Optional[float] = None - gps_longitude: Optional[float] = None + gps_latitude: float | None = None + gps_longitude: float | None = None @strawberry_sqlalchemy_input(models.Aircraft, exclude_fields=['id', 'photo_filename']) diff --git a/src/graphql_schema/entities/types/types.py b/src/graphql_schema/entities/types/types.py index 222bcb6..cdfd4be 100644 --- a/src/graphql_schema/entities/types/types.py +++ b/src/graphql_schema/entities/types/types.py @@ -3,7 +3,7 @@ import math from typing import Optional, List import strawberry from database import models -from decorators.endpoints import authenticated_user_only +# from decorators.endpoints import authenticated_user_only from utils.gps import get_bearing, get_distance from external.gpx_parser import GPXParser from graphql_schema.dataloaders.flight_duration import flight_duration_dataloader @@ -141,11 +141,11 @@ class Flight: for key, value in kwargs.items(): setattr(self, key, value) - @authenticated_user_only(raise_when_unauthorized=False, return_value_unauthorized=[]) + # @authenticated_user_only(raise_when_unauthorized=False, return_value_unauthorized=[]) async def load_copilots(root): return await flight_copilots_dataloader.load(root.id) - @authenticated_user_only(raise_when_unauthorized=False, return_value_unauthorized=[]) + # @authenticated_user_only(raise_when_unauthorized=False, return_value_unauthorized=[]) async def load_event(root): return await event_dataloader.load(root.event_id) @@ -189,7 +189,7 @@ class FlightPlanMarker: @strawberry.type class FlightPlanTrack: bearing: int | None - distance: Optional[float] + distance: float | None from_: FlightPlanMarker = strawberry.field(name="from") to: Optional[FlightPlanMarker] @@ -218,7 +218,7 @@ class FlightPlan: ) return navigation - @authenticated_user_only(raise_when_unauthorized=False, return_value_unauthorized=[]) + # @authenticated_user_only(raise_when_unauthorized=False, return_value_unauthorized=[]) async def load_copilots(root): return await flight_plan_copilots_dataloader.load(root.id) diff --git a/src/graphql_schema/entities/user.py b/src/graphql_schema/entities/user.py index c1ab297..a5fad0a 100644 --- a/src/graphql_schema/entities/user.py +++ b/src/graphql_schema/entities/user.py @@ -3,13 +3,13 @@ import strawberry from graphql import GraphQLError from passlib.hash import bcrypt from sqlalchemy import select +from strawberry import Info from strawberry.file_uploads import Upload from background_jobs.photo import resize_photo 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 graphql_schema.entities.types.types import User +from graphql_schema.extensions.field.auth import AuthenticatedOnly from utils.file import delete_file from utils.file import handle_file_upload @@ -17,8 +17,7 @@ from utils.file import handle_file_upload @strawberry.type class UserQueries: @strawberry.field() - @error_logging - async def user(root, info, username: str) -> User: + async def user(root, info: Info, username: str) -> User: if len(username) == 0: raise GraphQLError("Username not set!") @@ -28,10 +27,8 @@ class UserQueries: return user - @strawberry.field() - @authenticated_user_only() - @error_logging - async def logged_user(root, info) -> User: + @strawberry.field(extensions=[AuthenticatedOnly()]) + async def logged_user(root, info: Info) -> User: async with get_session() as db: user_model = (await db.scalars( select(models.User).filter_by(id=info.context.user_id) @@ -52,9 +49,8 @@ class EditUserMutation: avatar_image: Optional[Upload] = None title_image: Optional[Upload] = None - @strawberry.mutation - @authenticated_user_only() - async def edit_logged_user(root, info, input: EditUserInput) -> User: + @strawberry.mutation(extensions=[AuthenticatedOnly()]) + async def edit_logged_user(root, info: Info, input: EditUserInput) -> User: async with get_session() as db: user = (await db.scalars( select(models.User).filter_by(id=info.context.user_id) diff --git a/src/graphql_schema/extensions/__init__.py b/src/graphql_schema/extensions/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/graphql_schema/extensions/field/__init__.py b/src/graphql_schema/extensions/field/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/graphql_schema/extensions/field/auth.py b/src/graphql_schema/extensions/field/auth.py new file mode 100644 index 0000000..bc2693c --- /dev/null +++ b/src/graphql_schema/extensions/field/auth.py @@ -0,0 +1,25 @@ +from typing import Any +from fastapi import HTTPException +from starlette.status import HTTP_401_UNAUTHORIZED +from strawberry import Info +from strawberry.extensions import FieldExtension +from strawberry.extensions.field_extension import AsyncExtensionResolver + + +class AuthenticatedOnly(FieldExtension): + async def resolve_async(self, next_: AsyncExtensionResolver, source: Any, info: Info, **kwargs: Any) -> Any: + if not info.context.user_id: + raise HTTPException(HTTP_401_UNAUTHORIZED, "Not authorized") + + return await next_(source, info, **kwargs) + + +class AllowPublicAccess(FieldExtension): + async def resolve_async(self, next_: AsyncExtensionResolver, source: Any, info: Info, **kwargs: Any) -> Any: + user_id = info.context.user_id + public = kwargs.get('public') + + if not user_id and not public: + raise HTTPException(HTTP_401_UNAUTHORIZED, "Not authorized") + + return await next_(source, info, **kwargs) diff --git a/src/graphql_schema/extensions/field/pagination.py b/src/graphql_schema/extensions/field/pagination.py new file mode 100644 index 0000000..b3c84fa --- /dev/null +++ b/src/graphql_schema/extensions/field/pagination.py @@ -0,0 +1,48 @@ +from typing import Callable, Any, Type +import strawberry +from strawberry.annotation import StrawberryAnnotation +from strawberry.extensions import FieldExtension +from strawberry.types.arguments import StrawberryArgument +from strawberry.types.field import StrawberryField +from graphql_schema.entities.helpers.pagination import get_pagination_window, PaginationWindow + +class OffsetPagination[Item](FieldExtension): + + def __init__(self, item_type: Type[Item]): + super().__init__() + + self.item_type = item_type + + def apply(self, field: StrawberryField) -> StrawberryField: + offset_arg = StrawberryArgument( + python_name="offset", + graphql_name="offset", + type_annotation=StrawberryAnnotation(annotation=int | None), + default=0, + ) + + limit_arg = StrawberryArgument( + python_name="limit", + graphql_name="limit", + type_annotation=StrawberryAnnotation(annotation=int), + default=10, + ) + + field.arguments.append(offset_arg) + field.arguments.append(limit_arg) + + return field + + async def resolve_async( + self, next_: Callable[..., Any], source: Any, info: strawberry.Info, + limit: int, offset: int = 0, + **kwargs + ) -> PaginationWindow[Item]: + query = await next_(source, info, **kwargs) + + return await get_pagination_window( + query=query, + item_type=self.item_type, + limit=limit, + offset=offset, + ) diff --git a/src/graphql_schema/extensions/schema/__init__.py b/src/graphql_schema/extensions/schema/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/graphql_schema/extensions/schema/error_logging.py b/src/graphql_schema/extensions/schema/error_logging.py new file mode 100644 index 0000000..d8e9d8b --- /dev/null +++ b/src/graphql_schema/extensions/schema/error_logging.py @@ -0,0 +1,19 @@ +from typing import Callable, Any +from fastapi import HTTPException +from graphql import GraphQLResolveInfo, GraphQLError +from sqlalchemy.exc import NoResultFound +from strawberry.extensions import SchemaExtension +from strawberry.utils.await_maybe import AwaitableOrValue + + +class ErrorLogging(SchemaExtension): + async def resolve_async(self, _next: Callable, root: Any, info: GraphQLResolveInfo, *args: str, **kwargs: Any) -> AwaitableOrValue[object]: + try: + return await _next(root, info, *args, **kwargs) + except NoResultFound as e: + raise GraphQLError("Not found", original_error=e) + except HTTPException as e: + if e.status_code == 401: + raise GraphQLError("Not authorized", original_error=e) + except Exception as e: + raise GraphQLError(f"Unknown error: {e}", original_error=e) diff --git a/src/graphql_schema/schema.py b/src/graphql_schema/schema.py index cd59084..912abec 100644 --- a/src/graphql_schema/schema.py +++ b/src/graphql_schema/schema.py @@ -1,11 +1,7 @@ -import dataclasses -from typing import Set import strawberry -from fastapi_jwt import JwtAuthorizationCredentials -from fastapi_jwt.jwt import JwtAccessBearerCookie -from starlette.background import BackgroundTasks -from strawberry.extensions import SchemaExtension -from strawberry.fastapi import BaseContext +from strawberry.extensions import SchemaExtension, ValidationCache + +from graphql_schema.extensions.schema.error_logging import ErrorLogging from .mutation import Mutation from .query import Query @@ -28,17 +24,14 @@ class LoggingExtension(SchemaExtension): print("request end") -@dataclasses.dataclass -class GraphQLContext(BaseContext): - user_id: int - organization_ids: Set[int] - jwt_auth_credentials: JwtAuthorizationCredentials - jwt: JwtAccessBearerCookie - background_tasks: BackgroundTasks schema = strawberry.Schema( query=Query, mutation=Mutation, - extensions=[LoggingExtension] + extensions=[ + ErrorLogging(), + LoggingExtension(), + ValidationCache() + ], ) diff --git a/src/graphql_schema/sqlalchemy_to_strawberry_type.py b/src/graphql_schema/sqlalchemy_to_strawberry_type.py index e62e6c2..de72cee 100644 --- a/src/graphql_schema/sqlalchemy_to_strawberry_type.py +++ b/src/graphql_schema/sqlalchemy_to_strawberry_type.py @@ -28,7 +28,7 @@ def get_annotations_for_scalars(model: BaseModel, exclude_fields=None, force_opt return annotations_ -def strawberry_sqlalchemy_type(model, exclude_fields: Optional[typing.Union[List, typing.Tuple]] = None): +def strawberry_sqlalchemy_type(model: BaseModel, exclude_fields: list | tuple | None = None): if exclude_fields is None: exclude_fields = [] @@ -47,7 +47,7 @@ def strawberry_sqlalchemy_input( model, exclude_fields: Optional[typing.Union[List, typing.Tuple]] = None, all_optional: bool = False -) -> typing.Callable[[...], strawberry.object_type]: +) -> typing.Callable[[...], strawberry.type]: if exclude_fields is None: exclude_fields = [] diff --git a/src/jwt.py b/src/jwt.py new file mode 100644 index 0000000..17199df --- /dev/null +++ b/src/jwt.py @@ -0,0 +1,14 @@ +from datetime import timedelta +from fastapi_jwt import JwtAccessBearerCookie, JwtRefreshBearerCookie +from .config import APP_SECRET_KEY, APP_DEBUG, REFRESH_TOKEN_VALIDITY_DAYS + +access_security = JwtAccessBearerCookie( + secret_key=APP_SECRET_KEY, + auto_error=False, + access_expires_delta=timedelta(days=1) if APP_DEBUG else timedelta(minutes=20) +) +refresh_security = JwtRefreshBearerCookie( + secret_key=APP_SECRET_KEY, + auto_error=True, + refresh_expires_delta=timedelta(days=REFRESH_TOKEN_VALIDITY_DAYS), +) diff --git a/src/main.py b/src/main.py index 0b040d6..1badabd 100644 --- a/src/main.py +++ b/src/main.py @@ -1,22 +1,15 @@ import sentry_sdk -from datetime import timedelta -from typing import Optional -from fastapi import FastAPI, APIRouter, Security, HTTPException -from fastapi_jwt import JwtAuthorizationCredentials, JwtAccessBearerCookie, JwtRefreshBearerCookie +from fastapi import FastAPI, HTTPException, APIRouter from graphql import GraphQLError from sqlalchemy.exc import NoResultFound from starlette.background import BackgroundTasks from starlette.middleware.cors import CORSMiddleware from starlette.responses import Response, JSONResponse from starlette.staticfiles import StaticFiles -from config import APP_SECRET_KEY, ALLOW_CORS_ORIGINS, SENTRY_DSN, REFRESH_TOKEN_VALIDITY_DAYS, APP_DEBUG +from config import ALLOW_CORS_ORIGINS, SENTRY_DSN from endpoints.contact import ContactEndpoint, ContactInput -from endpoints.forgotten_password import ForgottenPasswordRequest, ForgottenPasswordEndpoint, ChangeForgottenPassword -from endpoints.graphql import setup_graphql_endpoint -from endpoints.login import LoginEndpoint, LoginInput, RefreshEndpoint, LogoutEndpoint -from endpoints.photo_editor_preview import PhotoEditorEndpoint -from endpoints.registration import RegistrationInput, RegistrationEndpoint from endpoints.sitemap import SitemapEndpoint +from routers import forgotten_password, auth, photo_preview, graphql class StaticFilesCache(StaticFiles): @@ -31,17 +24,6 @@ class StaticFilesCache(StaticFiles): class App: - api_router = APIRouter(dependencies=[]) - access_security = JwtAccessBearerCookie( - secret_key=APP_SECRET_KEY, - auto_error=False, - access_expires_delta=timedelta(days=1) if APP_DEBUG else timedelta(minutes=20) - ) - refresh_security = JwtRefreshBearerCookie( - secret_key=APP_SECRET_KEY, - auto_error=True, - refresh_expires_delta=timedelta(days=REFRESH_TOKEN_VALIDITY_DAYS), - ) def create_app(self): if SENTRY_DSN: @@ -100,94 +82,20 @@ class App: ) def setup_routes(self, app: FastAPI): - @self.api_router.post("/registration", status_code=201) - async def registration(user: RegistrationInput, background_tasks: BackgroundTasks): - return await RegistrationEndpoint().on_post(user, background_tasks) + api_router = APIRouter() - @self.api_router.post("/login") - async def login(resp: Response, user: LoginInput): - return await LoginEndpoint( - access_token=self.access_security, - 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}", - summary="Info about token used for resetting password" - ) - async def token_info(token: str): - return await ForgottenPasswordEndpoint().token_info(token) - - @self.api_router.post( - "/forgotten-password/request", - summary="Request password change, e-mail will be sent to validate your request." - ) - async def request_password_change(input: ForgottenPasswordRequest, background_tasks: BackgroundTasks): - return await ForgottenPasswordEndpoint().request(input, background_tasks) - - @self.api_router.post( - "/forgotten-password/reset", - summary="Set new password after successfull token validation" - ) - async def reset_password(input: ChangeForgottenPassword): - return await ForgottenPasswordEndpoint().change_password(input) - - @self.api_router.post("/contact", summary="Send email from contact form") + @api_router.post("/contact", summary="Send email from contact form") async def contact_form_message(input: ContactInput, background_tasks: BackgroundTasks): return await ContactEndpoint().on_post(input, background_tasks) - @self.api_router.get("/sitemap.xml") + @api_router.get("/sitemap.xml") async def sitemap(): return await SitemapEndpoint().on_get() - @self.api_router.get("/photo/editor-preview/{photo_id}", summary="Photo editor preview") - async def photo_editor_preview( - photo_id: int, - brightness: Optional[float] = None, - contrast: Optional[float] = None, - saturation: Optional[float] = None, - sharpness: Optional[float] = None, - rotate: Optional[float] = None, - crop_left: Optional[float] = None, - crop_top: Optional[float] = None, - crop_width: Optional[float] = None, - crop_height: Optional[float] = None, - ): - return await PhotoEditorEndpoint( - access_token=self.access_security, - refresh_token=self.refresh_security - ).show_preview( - photo_id=photo_id, - logged_user_id=0, - saturation=saturation, - brightness=brightness, - contrast=contrast, - sharpness=sharpness, - crop_top=crop_top, - crop_left=crop_left, - crop_height=crop_height, - crop_width=crop_width, - rotate=rotate, - ) - - setup_graphql_endpoint(app, self.access_security) # musi byt na konci - app.include_router(self.api_router) + app.include_router(auth.router) + app.include_router(forgotten_password.router) + app.include_router(photo_preview.router) + app.include_router(graphql.router) + app.include_router(api_router) diff --git a/src/routers/__init__.py b/src/routers/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/routers/auth.py b/src/routers/auth.py new file mode 100644 index 0000000..771b846 --- /dev/null +++ b/src/routers/auth.py @@ -0,0 +1,41 @@ +from fastapi import BackgroundTasks, Security, APIRouter +from fastapi_jwt import JwtAuthorizationCredentials +from starlette.responses import Response + +from endpoints.login import LoginInput, LoginEndpoint, RefreshEndpoint, LogoutEndpoint +from endpoints.registration import RegistrationInput, RegistrationEndpoint +from jwt import access_security, refresh_security + + +router = APIRouter() + +@router.post("/registration", status_code=201) +async def registration(user: RegistrationInput, background_tasks: BackgroundTasks): + return await RegistrationEndpoint().on_post(user, background_tasks) + + +@router.post("/login") +async def login(resp: Response, user: LoginInput): + return await LoginEndpoint( + access_token=access_security, + refresh_token=refresh_security + ).on_post(user, resp) + + +@router.post("/refresh", summary="Refresh access token") +async def refresh( + resp: Response, + credentials: JwtAuthorizationCredentials = Security(refresh_security) +): + return await RefreshEndpoint( + access_token=access_security, + refresh_token=refresh_security + ).on_post(resp, credentials) + + +@router.post("/logout") +async def logout(resp: Response): + return await LogoutEndpoint( + access_token=access_security, + refresh_token=refresh_security + ).on_post(resp) diff --git a/src/routers/forgotten_password.py b/src/routers/forgotten_password.py new file mode 100644 index 0000000..d5c284f --- /dev/null +++ b/src/routers/forgotten_password.py @@ -0,0 +1,25 @@ +from fastapi import BackgroundTasks, APIRouter +from endpoints.forgotten_password import ForgottenPasswordEndpoint, ForgottenPasswordRequest, ChangeForgottenPassword + +router = APIRouter() + +@router.get( + "/forgotten-password/token/{token}", + summary="Info about token used for resetting password" +) +async def token_info(token: str): + return await ForgottenPasswordEndpoint().token_info(token) + +@router.post( + "/forgotten-password/request", + summary="Request password change, e-mail will be sent to validate your request." +) +async def request_password_change(input: ForgottenPasswordRequest, background_tasks: BackgroundTasks): + return await ForgottenPasswordEndpoint().request(input, background_tasks) + +@router.post( + "/forgotten-password/reset", + summary="Set new password after successfull token validation" +) +async def reset_password(input: ChangeForgottenPassword): + return await ForgottenPasswordEndpoint().change_password(input) diff --git a/src/routers/graphql.py b/src/routers/graphql.py new file mode 100644 index 0000000..f1f16a4 --- /dev/null +++ b/src/routers/graphql.py @@ -0,0 +1,30 @@ +from datetime import timedelta +from fastapi import APIRouter +from starlette.responses import RedirectResponse +from strawberry.fastapi import GraphQLRouter +from config import GRAPHIQL, APP_DEBUG +from graphql_schema.context import setup_graphql_context +from graphql_schema.schema import schema +from jwt import access_security + +router = APIRouter() + +if APP_DEBUG: + @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 + +gql_router = GraphQLRouter( + schema, + graphiql=GRAPHIQL, + debug=APP_DEBUG, + context_getter=setup_graphql_context, + multipart_uploads_enabled=True, + prefix="/graphql" +) + +router.include_router(gql_router, tags=["login"]) diff --git a/src/routers/photo_preview.py b/src/routers/photo_preview.py new file mode 100644 index 0000000..3544d06 --- /dev/null +++ b/src/routers/photo_preview.py @@ -0,0 +1,36 @@ +from fastapi import APIRouter +from endpoints.photo_editor_preview import PhotoEditorEndpoint +from jwt import access_security, refresh_security + +router = APIRouter() + + +@router.get("/photo/editor-preview/{photo_id}", summary="Photo editor preview") +async def photo_editor_preview( + photo_id: int, + brightness: float | None = None, + contrast: float | None = None, + saturation: float | None = None, + sharpness: float | None = None, + rotate: float | None = None, + crop_left: float | None = None, + crop_top: float | None = None, + crop_width: float | None = None, + crop_height: float | None = None, +): + return await PhotoEditorEndpoint( + access_token=access_security, + refresh_token=refresh_security + ).show_preview( + photo_id=photo_id, + logged_user_id=0, + saturation=saturation, + brightness=brightness, + contrast=contrast, + sharpness=sharpness, + crop_top=crop_top, + crop_left=crop_left, + crop_height=crop_height, + crop_width=crop_width, + rotate=rotate, + ) diff --git a/src/scripts/add_gps_to_photos.py b/src/scripts/add_gps_to_photos.py new file mode 100644 index 0000000..e72cc5c --- /dev/null +++ b/src/scripts/add_gps_to_photos.py @@ -0,0 +1,69 @@ +import asyncio +import sys +from datetime import datetime +from itertools import groupby + +from sqlalchemy import select + +sys.path.insert(0, "/app/src") +from database import models +from database.transaction import get_session + + +def find_closest(needle: datetime, haystack, _best_difference: float = sys.maxsize): + if len(haystack) == 0: + return None + if len(haystack) == 1: + # nalezeno + return haystack[0] + + index = len(haystack) / 2 + diff = haystack[index].timestamp - needle + if diff < _best_difference: + _best_difference = diff + return find_closest(needle, haystack[:index], _best_difference) + else: + return find_closest(needle, haystack[index + 1:], _best_difference) + + +async def add_gps_to_photos(): + async with get_session() as db: + photos = (await db.execute( + select(models.Photo, models.Photo.flight) + .join(models.Photo.flight) + .filter(models.Photo.gps_latitude.is_(None)) + .filter(models.Photo.gps_longitude.is_(None)) + )).all() + + flight_ids = {photo.flight_id for photo, flight in photos} + + tracks = (await db.execute( + select(models.Flight.track_id, models.Flight.id) + .select_from(models.Flight) + .join(models.Flight.track) + .filter(models.Flight.id.in_(flight_ids)) + )).all() + + track_id_to_flight_id = {track_id: flight_id for track_id, flight_id in tracks} + + track_points_data = (await db.scalars( + select(models.TrackPoint) + .filter(models.TrackPoint.track_id.in_(track_id_to_flight_id.keys())) + .order_by(models.TrackPoint.timestamp) + )).all() + print(track_points_data) + grouped_points_by_track_id = groupby(track_points_data, key=lambda x: x.track_id) + + for photo, flight in photos: + best_track_point = find_closest(photo.exposed_at, grouped_points_by_track_id[flight.track_id]) + print(best_track_point) + break + + + print(grouped_points_by_track_id) + # tracks_by_flight_id = {track.flight_id: track_points for track, track_points in tracks} + + +if __name__ == "__main__": + loop = asyncio.get_event_loop() + loop.run_until_complete(add_gps_to_photos()) diff --git a/src/utils/image.py b/src/utils/image.py index c0aa2ea..994acd3 100644 --- a/src/utils/image.py +++ b/src/utils/image.py @@ -102,10 +102,10 @@ class PhotoEditor: def adjust( self, - brightness: Optional[float] = None, - contrast: Optional[float] = None, - saturation: Optional[float] = None, - sharpness: Optional[float] = None + brightness: float | None = None, + contrast: float | None = None, + saturation: float | None = None, + sharpness: float | None = None ): adjustments = [ (Brightness, brightness),