Refresh token, spolecne lety s copilotem, lazy typ
This commit is contained in:
@@ -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')
|
||||||
|
|||||||
@@ -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
@@ -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()
|
||||||
|
|||||||
@@ -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)
|
||||||
@@ -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()
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -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,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
@@ -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)
|
||||||
|
|
||||||
|
|||||||
Executable
+11
@@ -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"
|
||||||
Reference in New Issue
Block a user