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) id: Mapped[int] = mapped_column(primary_key=True)
name: Mapped[str] = mapped_column(String(128), nullable=False) name: Mapped[str] = mapped_column(String(128), nullable=False)
icao_code: Mapped[str] = mapped_column(String(4), nullable=False) icao_code: Mapped[str] = mapped_column(String(4), nullable=False)
gps_latitude: Mapped[float] = mapped_column(Float, nullable=True) gps_latitude: Mapped[float] = mapped_column(Float, nullable=False)
gps_longitude: Mapped[float] = mapped_column(Float, nullable=True) gps_longitude: Mapped[float] = mapped_column(Float, nullable=False)
elevation: Mapped[int] = mapped_column(Integer, nullable=True) elevation: Mapped[int] = mapped_column(Integer, nullable=True)
created_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now()) created_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now())
deleted: Mapped[bool] = mapped_column(Boolean, nullable=False, server_default='0') 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(): async with session.begin():
yield session yield session
await session.commit() 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) access_token = self.access_security.create_access_token(subject=subject)
refresh_token = self.refresh_security.create_refresh_token(subject=subject) refresh_token = self.refresh_security.create_refresh_token(subject=subject)
self.access_security.set_access_cookie(resp, access_token) # self.access_security.set_access_cookie(resp, access_token)
self.refresh_security.set_refresh_cookie(resp, refresh_token) self.refresh_security.set_refresh_cookie(
resp, refresh_token,
expires_delta=self.refresh_security.refresh_expires_delta
)
return { return {
"user": logged_user.as_dict(), "user": logged_user.as_dict(),
"access_token": access_token, "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): class MeEndpoint(BaseEndpoint):
def __init__(self, db, credentials: JwtAuthorizationCredentials): async def on_get(self, credentials) -> dict:
super().__init__(db) if not credentials:
self.credentials = credentials raise HTTPException(status_code=401)
async def on_get(self) -> dict: query = select(User).filter_by(id=credentials['id'])
query = select(User).filter_by(id=self.credentials['id'])
user = (await self.db.scalars(query)).first() user = (await self.db.scalars(query)).first()
return user.as_dict() 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 import strawberry
from sqlalchemy import select 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 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: 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 @strawberry.type
class CopilotQueries: class CopilotQueries:
@strawberry.field @strawberry.field
async def copilots(root, info) -> List[CopilotType]: 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 datetime import timedelta
from typing import List, Optional from typing import List, Optional, Annotated, TYPE_CHECKING
import strawberry import strawberry
from sqlalchemy import select, delete from sqlalchemy import select, delete
from sqlalchemy.ext.asyncio import AsyncSession 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 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() @strawberry.input()
@@ -74,7 +75,7 @@ class Flight:
return 0 return 0
duration_min_calculated: int = strawberry.field(resolver=duration_min_calculated) 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) aircraft: Aircraft = strawberry.field(resolver=load_aircraft)
takeoff_airport: Airport = strawberry.field(resolver=load_takeoff_airport) takeoff_airport: Airport = strawberry.field(resolver=load_takeoff_airport)
landing_airport: Airport = strawberry.field(resolver=load_landing_airport) landing_airport: Airport = strawberry.field(resolver=load_landing_airport)
@@ -186,6 +187,7 @@ class EditFlightMutation:
@strawberry.mutation @strawberry.mutation
async def edit_flight(self, info, id: int, input: EditFlightInput) -> Flight: 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()) flight = await models.Flight.update(info.context.db, id=id, data=input.to_dict())
if input.track is not None: if input.track is not None:
+3 -1
View File
@@ -39,7 +39,7 @@ class PhotoQueries:
@strawberry.type @strawberry.type
class UploadPhotoMutation: 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: class UploadPhotoInput:
photo: Upload photo: Upload
@@ -59,6 +59,8 @@ class UploadPhotoMutation:
"created_by_id": info.context.user_id, "created_by_id": info.context.user_id,
}, db_session=info.context.db) }, db_session=info.context.db)
await info.context.db.flush()
return created_photo return created_photo
+6 -2
View File
@@ -82,7 +82,11 @@ class EditPointOfInterestMutation:
# TODO: kontrola organizace # TODO: kontrola organizace
# TODO: kontrola opravneni na akci # 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()) return await models.PointOfInterest.update(info.context.db, obj=poi, data=input.to_dict())
@@ -91,6 +95,6 @@ class DeletePointOfInterestMutation:
@strawberry.mutation @strawberry.mutation
async def delete_point_of_interest(self, info, id: int) -> PointOfInterest: 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)) 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 dataclasses
import strawberry import strawberry
from fastapi_jwt import JwtAuthorizationCredentials 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 sqlalchemy.ext.asyncio import AsyncSession
from strawberry.extensions import SchemaExtension from strawberry.extensions import SchemaExtension
from strawberry.fastapi import BaseContext from strawberry.fastapi import BaseContext
+12 -16
View File
@@ -1,6 +1,6 @@
from datetime import timedelta from datetime import timedelta
from fastapi import FastAPI, APIRouter, Depends, Security, HTTPException 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 sqlalchemy.ext.asyncio import AsyncSession
from starlette.middleware.cors import CORSMiddleware from starlette.middleware.cors import CORSMiddleware
from starlette.responses import RedirectResponse, Response 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 config import APP_SECRET_KEY, GRAPHIQL, APP_DEBUG
from dependencies.db import db_session from dependencies.db import db_session
from endpoints.init_data import InitDataEndpoint 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 endpoints.registration import RegistrationInput, RegistrationEndpoint
from graphql_schema.schema import schema, GraphQLContext from graphql_schema.schema import schema, GraphQLContext
@@ -19,7 +19,7 @@ class App:
access_security = JwtAccessBearerCookie( access_security = JwtAccessBearerCookie(
secret_key=APP_SECRET_KEY, secret_key=APP_SECRET_KEY,
auto_error=False, auto_error=False,
access_expires_delta=timedelta(hours=1) access_expires_delta=timedelta(seconds=30)
) )
refresh_security = JwtRefreshBearerCookie( refresh_security = JwtRefreshBearerCookie(
secret_key=APP_SECRET_KEY, secret_key=APP_SECRET_KEY,
@@ -43,7 +43,7 @@ class App:
def setup_middleware(app: FastAPI): def setup_middleware(app: FastAPI):
app.add_middleware( app.add_middleware(
CORSMiddleware, CORSMiddleware,
allow_origins=["*"], allow_origins=["http://localhost:9001"],
allow_credentials=True, allow_credentials=True,
allow_methods=["*"], allow_methods=["*"],
allow_headers=["*"], allow_headers=["*"],
@@ -89,19 +89,15 @@ class App:
db: AsyncSession = Depends(db_session), db: AsyncSession = Depends(db_session),
credentials: JwtAuthorizationCredentials = Security(self.access_security) 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) 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"