Refresh token, spolecne lety s copilotem, lazy typ

This commit is contained in:
Michal Kváček
2023-07-29 23:27:24 +02:00
parent 77dd5538d3
commit fcd80f20e8
11 changed files with 133 additions and 38 deletions
+2 -2
View File
@@ -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')
+1
View File
@@ -6,3 +6,4 @@ async def db_session():
async with session.begin():
yield session
await session.commit()
await session.flush()
+36 -8
View File
@@ -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()
+27
View File
@@ -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)
+29 -5
View File
@@ -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()
+5 -3
View File
@@ -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:
+3 -1
View File
@@ -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
+6 -2
View File
@@ -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))
+1 -1
View File
@@ -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
+12 -16
View File
@@ -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)
+11
View File
@@ -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"