diff --git a/src/database/models.py b/src/database/models.py index 395496d..3ef7bb5 100644 --- a/src/database/models.py +++ b/src/database/models.py @@ -64,8 +64,8 @@ class Airport(BaseModel): id: Mapped[int] = mapped_column(primary_key=True) name: Mapped[str] = mapped_column(String(128), nullable=False) icao_code: Mapped[str] = mapped_column(String(4), nullable=False) - gps_latitude: Mapped[float] = mapped_column(Float, nullable=True) - gps_longitude: Mapped[float] = mapped_column(Float, nullable=True) + gps_latitude: Mapped[float] = mapped_column(Float, nullable=False) + gps_longitude: Mapped[float] = mapped_column(Float, nullable=False) elevation: Mapped[int] = mapped_column(Integer, nullable=True) created_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now()) deleted: Mapped[bool] = mapped_column(Boolean, nullable=False, server_default='0') diff --git a/src/dependencies/db.py b/src/dependencies/db.py index d8383f3..5db0448 100644 --- a/src/dependencies/db.py +++ b/src/dependencies/db.py @@ -6,3 +6,4 @@ async def db_session(): async with session.begin(): yield session await session.commit() + await session.flush() diff --git a/src/endpoints/login.py b/src/endpoints/login.py index 36a8519..47656b8 100644 --- a/src/endpoints/login.py +++ b/src/endpoints/login.py @@ -36,24 +36,52 @@ class LoginEndpoint(BaseEndpoint): access_token = self.access_security.create_access_token(subject=subject) refresh_token = self.refresh_security.create_refresh_token(subject=subject) - self.access_security.set_access_cookie(resp, access_token) - self.refresh_security.set_refresh_cookie(resp, refresh_token) + # self.access_security.set_access_cookie(resp, access_token) + self.refresh_security.set_refresh_cookie( + resp, refresh_token, + expires_delta=self.refresh_security.refresh_expires_delta + ) return { "user": logged_user.as_dict(), "access_token": access_token, - "refresh_token": refresh_token + "access_token_validity": self.access_security.access_expires_delta.total_seconds(), + } + +class RefreshEndpoint(BaseEndpoint): + def __init__(self, access_token: JwtAccess, refresh_token: JwtRefresh): + super().__init__(db=None) + self.access_security = access_token + self.refresh_security = refresh_token + + async def on_post(self, resp: Response, credentials: JwtAuthorizationCredentials): + access_token = self.access_security.create_access_token( + subject=credentials.subject, + expires_delta=self.access_security.access_expires_delta + ) + refresh_token = self.refresh_security.create_refresh_token( + subject=credentials.subject, + expires_delta=self.refresh_security.refresh_expires_delta + ) + + self.refresh_security.set_refresh_cookie( + resp, refresh_token, + expires_delta=self.refresh_security.refresh_expires_delta + ) + + return { + "access_token": access_token, + "access_token_validity": self.access_security.access_expires_delta.total_seconds(), } class MeEndpoint(BaseEndpoint): - def __init__(self, db, credentials: JwtAuthorizationCredentials): - super().__init__(db) - self.credentials = credentials + async def on_get(self, credentials) -> dict: + if not credentials: + raise HTTPException(status_code=401) - async def on_get(self) -> dict: - query = select(User).filter_by(id=self.credentials['id']) + query = select(User).filter_by(id=credentials['id']) user = (await self.db.scalars(query)).first() return user.as_dict() diff --git a/src/graphql_schema/dataloaders/flight.py b/src/graphql_schema/dataloaders/flight.py new file mode 100644 index 0000000..b7e58a1 --- /dev/null +++ b/src/graphql_schema/dataloaders/flight.py @@ -0,0 +1,27 @@ +from collections import defaultdict +from typing import List +from sqlalchemy import select +from strawberry.dataloader import DataLoader +from database import async_session +from database.models import Flight + + +class FlightsLoader: + def __init__(self, relationship_column: str): + self.relationship_column = relationship_column + + async def load(self, ids: List[int]): + async with async_session() as session: + query = ( + select(Flight) + .filter(getattr(Flight, self.relationship_column).in_(ids)) + ) + data = (await session.scalars(query)).all() + + result_data = defaultdict(list) + for poi in data: + result_data[getattr(poi, self.relationship_column)].append(poi) + + return [result_data[id_] for id_ in ids] + +flights_by_copilot_dataloader = DataLoader(load_fn=FlightsLoader("copilot_id").load, cache=False) \ No newline at end of file diff --git a/src/graphql_schema/entities/copilot.py b/src/graphql_schema/entities/copilot.py index 9bb1b9d..2776ef0 100644 --- a/src/graphql_schema/entities/copilot.py +++ b/src/graphql_schema/entities/copilot.py @@ -1,17 +1,41 @@ -from typing import List +from typing import List, Annotated, TYPE_CHECKING import strawberry from sqlalchemy import select -from database.models import Copilot +from database import models +from graphql_schema.dataloaders.flight import flights_by_copilot_dataloader from graphql_schema.sqlalchemy_to_strawberry_type import strawberry_sqlalchemy_type +if TYPE_CHECKING: + from .flight import Flight -@strawberry_sqlalchemy_type(Copilot) +@strawberry_sqlalchemy_type(models.Copilot) class CopilotType: - pass + async def load_flights(root): + return await flights_by_copilot_dataloader.load(root.id) + + flights: List[Annotated["Flight", strawberry.lazy('.flight')]] = strawberry.field(resolver=load_flights) + + +def get_base_query(user_id: int): + return ( + select(models.Copilot) + .filter(models.Copilot.created_by_id == user_id) + .filter(models.Copilot.deleted.is_(False)) + .order_by(models.Copilot.id.desc()) + ) @strawberry.type class CopilotQueries: @strawberry.field async def copilots(root, info) -> List[CopilotType]: - return (await info.context.db.scalars(select(Copilot))).all() + return (await info.context.db.scalars( + get_base_query(info.context.user_id) + )).all() + + @strawberry.field + async def copilot(root, info, id: int) -> CopilotType: + return (await info.context.db.scalars( + get_base_query(info.context.user_id) + .filter(models.Copilot.id == id) + )).one() diff --git a/src/graphql_schema/entities/flight.py b/src/graphql_schema/entities/flight.py index 9f709ce..adfa706 100644 --- a/src/graphql_schema/entities/flight.py +++ b/src/graphql_schema/entities/flight.py @@ -1,5 +1,5 @@ from datetime import timedelta -from typing import List, Optional +from typing import List, Optional, Annotated, TYPE_CHECKING import strawberry from sqlalchemy import select, delete from sqlalchemy.ext.asyncio import AsyncSession @@ -17,7 +17,8 @@ from graphql_schema.entities.poi import PointOfInterest from graphql_schema.sqlalchemy_to_strawberry_type import strawberry_sqlalchemy_type, strawberry_sqlalchemy_input -# Bude se hodit: https://strawberry.rocks/docs/types/lazy +if TYPE_CHECKING: + from .copilot import CopilotType @strawberry.input() @@ -74,7 +75,7 @@ class Flight: return 0 duration_min_calculated: int = strawberry.field(resolver=duration_min_calculated) - copilot: Optional[CopilotType] = strawberry.field(resolver=load_copilot) + copilot: Optional[Annotated["CopilotType", strawberry.lazy(".copilot")]] = strawberry.field(resolver=load_copilot) aircraft: Aircraft = strawberry.field(resolver=load_aircraft) takeoff_airport: Airport = strawberry.field(resolver=load_takeoff_airport) landing_airport: Airport = strawberry.field(resolver=load_landing_airport) @@ -186,6 +187,7 @@ class EditFlightMutation: @strawberry.mutation async def edit_flight(self, info, id: int, input: EditFlightInput) -> Flight: + # TODO: umoznit editovat jen vlastni lety! flight = await models.Flight.update(info.context.db, id=id, data=input.to_dict()) if input.track is not None: diff --git a/src/graphql_schema/entities/photo.py b/src/graphql_schema/entities/photo.py index 38e852c..f914232 100644 --- a/src/graphql_schema/entities/photo.py +++ b/src/graphql_schema/entities/photo.py @@ -39,7 +39,7 @@ class PhotoQueries: @strawberry.type class UploadPhotoMutation: - @strawberry_sqlalchemy_input(models.Photo, exclude_fields=["id", "filename"]) + @strawberry_sqlalchemy_input(models.Photo, exclude_fields=["id", "filename", "is_flight_cover"]) class UploadPhotoInput: photo: Upload @@ -59,6 +59,8 @@ class UploadPhotoMutation: "created_by_id": info.context.user_id, }, db_session=info.context.db) + await info.context.db.flush() + return created_photo diff --git a/src/graphql_schema/entities/poi.py b/src/graphql_schema/entities/poi.py index eb25205..ea4552d 100644 --- a/src/graphql_schema/entities/poi.py +++ b/src/graphql_schema/entities/poi.py @@ -82,7 +82,11 @@ class EditPointOfInterestMutation: # TODO: kontrola organizace # TODO: kontrola opravneni na akci - poi = (await get_base_query(info.context.user_id, only_my=True).filter(models.Photo.id == id)).one() + 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.to_dict()) @@ -91,6 +95,6 @@ class DeletePointOfInterestMutation: @strawberry.mutation async def delete_point_of_interest(self, info, id: int) -> PointOfInterest: - poi = (await get_base_query(info.context.user_id, only_my=True).filter(models.Photo.id == id)).one() + poi = (await 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/schema.py b/src/graphql_schema/schema.py index cd60c6d..0c5422e 100644 --- a/src/graphql_schema/schema.py +++ b/src/graphql_schema/schema.py @@ -1,7 +1,7 @@ import dataclasses import strawberry from fastapi_jwt import JwtAuthorizationCredentials -from fastapi_jwt.jwt import JwtAccessBearerCookie +from fastapi_jwt.jwt import JwtAccessBearer, JwtAccessBearerCookie from sqlalchemy.ext.asyncio import AsyncSession from strawberry.extensions import SchemaExtension from strawberry.fastapi import BaseContext diff --git a/src/main.py b/src/main.py index 512892e..24d9947 100644 --- a/src/main.py +++ b/src/main.py @@ -1,6 +1,6 @@ from datetime import timedelta from fastapi import FastAPI, APIRouter, Depends, Security, HTTPException -from fastapi_jwt import JwtAuthorizationCredentials, JwtAccessBearerCookie, JwtRefreshBearerCookie +from fastapi_jwt import JwtAuthorizationCredentials, JwtRefreshBearer, JwtAccessBearerCookie, JwtRefreshBearerCookie from sqlalchemy.ext.asyncio import AsyncSession from starlette.middleware.cors import CORSMiddleware from starlette.responses import RedirectResponse, Response @@ -9,7 +9,7 @@ from strawberry.fastapi import GraphQLRouter from config import APP_SECRET_KEY, GRAPHIQL, APP_DEBUG from dependencies.db import db_session from endpoints.init_data import InitDataEndpoint -from endpoints.login import LoginEndpoint, LoginInput, MeEndpoint +from endpoints.login import LoginEndpoint, LoginInput, MeEndpoint, RefreshEndpoint from endpoints.registration import RegistrationInput, RegistrationEndpoint from graphql_schema.schema import schema, GraphQLContext @@ -19,7 +19,7 @@ class App: access_security = JwtAccessBearerCookie( secret_key=APP_SECRET_KEY, auto_error=False, - access_expires_delta=timedelta(hours=1) + access_expires_delta=timedelta(seconds=30) ) refresh_security = JwtRefreshBearerCookie( secret_key=APP_SECRET_KEY, @@ -43,7 +43,7 @@ class App: def setup_middleware(app: FastAPI): app.add_middleware( CORSMiddleware, - allow_origins=["*"], + allow_origins=["http://localhost:9001"], allow_credentials=True, allow_methods=["*"], allow_headers=["*"], @@ -89,19 +89,15 @@ class App: db: AsyncSession = Depends(db_session), credentials: JwtAuthorizationCredentials = Security(self.access_security) ): - return await MeEndpoint(db, credentials).on_get() + return await MeEndpoint(db).on_get(credentials) + + @app.post("/refresh") + async def refresh( + resp: Response, + credentials: JwtAuthorizationCredentials = Security(self.refresh_security) + ): + return await RefreshEndpoint(self.access_security, self.refresh_security).on_post(resp, credentials) - # @app.post("/refresh") - # def refresh( - # credentials: JwtAuthorizationCredentials = Security(refresh_security) - # ): - # # Update access/refresh tokens pair - # # We can customize expires_delta when creating - # access_token = access_security.create_access_token(subject=credentials.subject) - # refresh_token = refresh_security.create_refresh_token(subject=credentials.subject, - # expires_delta=timedelta(days=2)) - # - # return {"access_token": access_token, "refresh_token": refresh_token} self.setup_graphql_endpoint(app) diff --git a/test_requests/me.sh b/test_requests/me.sh new file mode 100755 index 0000000..825ada6 --- /dev/null +++ b/test_requests/me.sh @@ -0,0 +1,11 @@ +TOKEN="eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJzdWJqZWN0Ijp7ImlkIjoxOCwiZW1haWwiOiIifSwidHlwZSI6ImFjY2VzcyIsImV4cCI6MTY5MzEzODU5MCwiaWF0IjoxNjkwNTQ2NTkwLCJqdGkiOiIwMjg3YjAzMy1jZGE3LTQ3MDItOGVmYi1kYTI4M2FlYzhlNGMifQ.89sQRRbGj5z0XEzQgdONrTFiGrNYtWf52oiHXGyEmUE" + +http POST localhost:8000/login email="" password="" + +#http GET http://localhost:8000/me "authorization: Bearer $TOKEN" + + +#http POST http://localhost:8000/me "authorization: Bearer $TOKEN" + +REFRESH_TOKEN="eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJzdWJqZWN0Ijp7ImlkIjoxOCwiZW1haWwiOiIifSwidHlwZSI6InJlZnJlc2giLCJleHAiOjE2OTA3MTk4MzMsImlhdCI6MTY5MDU0NzAzMywianRpIjoiMjVmMDIxYWEtNzZmOC00YjNlLWIxMmMtYzIyNjZiNTk5YzE5In0.k6XkzWx2BaeJw3dAfGdA9Qy0pkCO5gWAvKgO_0lrYwU" +http POST http://localhost:8000/refresh "Authorization: Bearer $REFRESH_TOKEN" \ No newline at end of file