Uklid v alembic migracich, error logging, upravy kvuli autorizaci

This commit is contained in:
Michal Kváček
2023-08-14 15:01:13 +02:00
parent bc260a2a8b
commit 97295c17c0
38 changed files with 203 additions and 957 deletions
+14 -2
View File
@@ -1,7 +1,19 @@
from functools import wraps
from fastapi import HTTPException
from starlette.status import HTTP_401_UNAUTHORIZED
def public_endpoint(func):
pass
def private_endpoint(func):
pass
def authenticated_user_only(func):
@wraps(func)
async def decorator(*args, **kwargs):
if 'info' in kwargs:
if not kwargs['info'].context.user_id:
raise HTTPException(HTTP_401_UNAUTHORIZED, "Not authorized")
return await func(*args, **kwargs)
return decorator
+3 -2
View File
@@ -1,5 +1,6 @@
from functools import wraps
from graphql import GraphQLError
from sqlalchemy.exc import NoResultFound
def error_logging(func):
@@ -7,7 +8,7 @@ def error_logging(func):
async def decorator(*args, **kwargs):
try:
return await func(*args, **kwargs)
except Exception as e:
raise GraphQLError(f"Not found")
except NoResultFound as e:
raise GraphQLError(f"Not found", original_error=e)
return decorator
+12 -2
View File
@@ -1,3 +1,13 @@
from fastapi_jwt.jwt import JwtAccess, JwtRefresh
class BaseEndpoint:
def __init__(self, db):
self.db = db
def __init__(self, *args, **kwargs):
self.db = kwargs.get("db")
super().__init__(*args, **kwargs)
class AuthEndpoint:
def __init__(self, *args, **kwargs):
self.access_security: JwtAccess = kwargs.get("access_token")
self.refresh_security: JwtRefresh = kwargs.get("refresh_token")
super().__init__()
+13 -14
View File
@@ -1,11 +1,11 @@
from fastapi import HTTPException
from fastapi_jwt import JwtAuthorizationCredentials
from fastapi_jwt.jwt import JwtAccess, JwtRefresh
from passlib.hash import bcrypt
from sqlalchemy import select
from starlette.responses import Response
from database.models import User
from endpoints.base import BaseEndpoint
from endpoints.base import BaseEndpoint, AuthEndpoint
from pydantic import BaseModel
@@ -14,13 +14,7 @@ class LoginInput(BaseModel):
password: str
class LoginEndpoint(BaseEndpoint):
def __init__(self, db, access_token: JwtAccess, refresh_token: JwtRefresh):
super().__init__(db)
self.access_security = access_token
self.refresh_security = refresh_token
class LoginEndpoint(BaseEndpoint, AuthEndpoint):
async def on_post(self, user_data: LoginInput, resp: Response) -> dict:
query = select(User).filter_by(email=user_data.email)
logged_user = (await self.db.scalars(query)).first()
@@ -48,11 +42,16 @@ class LoginEndpoint(BaseEndpoint):
}
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
class LogoutEndpoint(AuthEndpoint):
async def on_post(self, resp: Response):
self.refresh_security.unset_refresh_cookie(resp)
return {
"logged_out": True
}
class RefreshEndpoint(AuthEndpoint):
async def on_post(self, resp: Response, credentials: JwtAuthorizationCredentials):
access_token = self.access_security.create_access_token(
+18 -2
View File
@@ -1,9 +1,13 @@
from datetime import timedelta
from functools import wraps
from typing import List, Optional, Annotated, TYPE_CHECKING, Tuple
import strawberry
from fastapi import HTTPException
from sqlalchemy import select
from starlette.status import HTTP_401_UNAUTHORIZED
from strawberry.file_uploads import Upload
from database import models
from decorators.endpoints import authenticated_user_only
from decorators.error_logging import error_logging
from graphql_schema.dataloaders import copilots_dataloader
from graphql_schema.dataloaders.aircraft import aircraft_dataloader
@@ -119,6 +123,9 @@ class FlightQueries:
@strawberry.field
async def flights(root, info, username: Optional[str] = None) -> List[Flight]:
if not info.context.user_id and not username:
raise HTTPException(HTTP_401_UNAUTHORIZED)
query = (
get_base_query(user_id=info.context.user_id, username=username, is_auth=bool(info.context.user_id))
.order_by(models.Flight.id.desc())
@@ -128,6 +135,9 @@ class FlightQueries:
@strawberry.field
@error_logging
async def flight(root, info, id: int, username: Optional[str] = None) -> Flight:
if not info.context.user_id and not username:
raise HTTPException(HTTP_401_UNAUTHORIZED)
query = (
get_base_query(user_id=info.context.user_id, username=username, is_auth=bool(info.context.user_id))
.filter(models.Flight.id == id)
@@ -163,6 +173,7 @@ class CreateFlightMutation:
takeoff_airport: ComboboxInput
@strawberry.mutation
@authenticated_user_only
async def create_flight(self, info, input: CreateFlightInput) -> Flight:
db = info.context.db
aircraft_id = await handle_aircraft_save(db, info.context.user_id, input.aircraft)
@@ -187,6 +198,7 @@ class CreateFlightMutation:
})
@strawberry.type
class EditFlightMutation:
@strawberry_sqlalchemy_input(models.Flight, exclude_fields=[
@@ -202,6 +214,7 @@ class EditFlightMutation:
takeoff_airport: Optional[ComboboxInput] = None
@strawberry.mutation
@authenticated_user_only
async def edit_flight(self, info, id: int, input: EditFlightInput) -> Flight:
db = info.context.db
user_id = info.context.user_id
@@ -212,8 +225,11 @@ class EditFlightMutation:
)).one()
data = input.to_dict()
data['takeoff_datetime'] = data['takeoff_datetime'].astimezone()
data['landing_datetime'] = data['landing_datetime'].astimezone()
if input.takeoff_datetime:
data['takeoff_datetime'] = data['takeoff_datetime'].astimezone()
if input.landing_datetime:
data['landing_datetime'] = data['landing_datetime'].astimezone()
if input.gpx_track is not None:
# TODO: poresit validaci uploadovaneho souboru!
+8 -4
View File
@@ -8,6 +8,7 @@ from sqlalchemy.exc import NoResultFound
from strawberry.file_uploads import Upload
from database import models
from database.models import User
from decorators.endpoints import authenticated_user_only
from decorators.error_logging import error_logging
from graphql_schema.sqlalchemy_to_strawberry_type import strawberry_sqlalchemy_type, strawberry_sqlalchemy_input
from upload_utils import handle_file_upload, delete_file
@@ -45,14 +46,16 @@ class UserQueries:
)).one()
@strawberry.field
@authenticated_user_only
# @error_logging
async def logged_user(root, info) -> User:
if not info.context.user_id:
raise GraphQLError("Not authenticated")
return (await info.context.db.scalars(
user = (await info.context.db.scalars(
select(models.User).filter_by(id=info.context.user_id)
)).one()
print(user)
return user
@strawberry.type
class EditUserMutation:
@@ -68,6 +71,7 @@ class EditUserMutation:
title_image: Optional[Upload] = None
@strawberry.mutation
@authenticated_user_only
async def edit_logged_user(root, info, input: EditUserInput) -> User:
user = (await info.context.db.scalars(
select(models.User).filter_by(id=info.context.user_id)
-6
View File
@@ -1,17 +1,11 @@
import dataclasses
from typing import List, Optional
import strawberry
from fastapi_jwt import JwtAuthorizationCredentials
from fastapi_jwt.jwt import JwtAccessBearerCookie
from graphql import GraphQLError
from sqlalchemy.exc import NoResultFound
from sqlalchemy.ext.asyncio import AsyncSession
from starlette.background import BackgroundTasks
from strawberry.extensions import SchemaExtension
from strawberry.fastapi import BaseContext
from strawberry.types import ExecutionContext
from .mutation import Mutation
from .query import Query
+33 -25
View File
@@ -1,17 +1,15 @@
import time
from datetime import timedelta
from fastapi import FastAPI, APIRouter, Depends, Security, HTTPException
from fastapi import FastAPI, APIRouter, Depends, Security
from fastapi_jwt import JwtAuthorizationCredentials, JwtAccessBearerCookie, JwtRefreshBearerCookie
from sqlalchemy.ext.asyncio import AsyncSession
from starlette.background import BackgroundTasks
from starlette.middleware.cors import CORSMiddleware
from starlette.responses import RedirectResponse, Response
from starlette.staticfiles import StaticFiles
from starlette.status import HTTP_401_UNAUTHORIZED
from strawberry.fastapi import GraphQLRouter
from config import APP_SECRET_KEY, GRAPHIQL, APP_DEBUG
from dependencies.db import db_session
from endpoints.login import LoginEndpoint, LoginInput, RefreshEndpoint
from endpoints.login import LoginEndpoint, LoginInput, RefreshEndpoint, LogoutEndpoint
from endpoints.registration import RegistrationInput, RegistrationEndpoint
from graphql_schema.schema import schema, GraphQLContext
@@ -21,11 +19,12 @@ class App:
access_security = JwtAccessBearerCookie(
secret_key=APP_SECRET_KEY,
auto_error=False,
access_expires_delta=timedelta(minutes=20)
access_expires_delta=timedelta(seconds=20),
)
refresh_security = JwtRefreshBearerCookie(
secret_key=APP_SECRET_KEY,
auto_error=True
auto_error=True,
refresh_expires_delta=timedelta(days=30),
)
def create_app(self):
@@ -58,24 +57,14 @@ class App:
app.mount("/static", StaticFiles(directory="/app/static"), name="static")
def setup_graphql_endpoint(self, app: FastAPI):
if APP_DEBUG:
@self.api_router.get("/graphql/autologin")
async def autologin():
access_token = self.access_security.create_access_token(subject={"id": 18, "name": "Franta Vomacka"})
response = RedirectResponse(url="/graphql")
self.access_security.set_access_cookie(response, access_token)
return response
def setup_graphql_context(
credentials: JwtAuthorizationCredentials = Security(self.access_security),
db: AsyncSession = Depends(db_session)
):
return GraphQLContext(
jwt_auth_credentials=credentials,
user_id=credentials['id'] if credentials else None,
db=db,
user_id=credentials['id'] if credentials else None,
jwt_auth_credentials=credentials,
jwt=self.access_security,
background_tasks=Depends(BackgroundTasks)
)
@@ -89,20 +78,39 @@ class App:
app.include_router(graphql_app, prefix="/graphql")
def setup_routes(self, app: FastAPI):
# protected endpoints
@app.post("/refresh")
self.setup_graphql_endpoint(app)
if APP_DEBUG:
@self.api_router.get("/graphql/autologin")
async def autologin():
access_token = self.access_security.create_access_token(subject={"id": 18, "name": "Franta Vomacka"})
response = RedirectResponse(url="/graphql")
self.access_security.set_access_cookie(response, access_token, expires_delta=timedelta(days=14))
return response
@self.api_router.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)
return await RefreshEndpoint(
access_token=self.access_security,
refresh_token=self.refresh_security
).on_post(resp, credentials)
self.setup_graphql_endpoint(app)
# public endpoints
@self.api_router.post("/login")
async def login(resp: Response, user: LoginInput, db: AsyncSession = Depends(db_session)):
return await LoginEndpoint(db, self.access_security, self.refresh_security).on_post(user, resp)
return await LoginEndpoint(
db=db,
access_token=self.access_security,
refresh_token=self.refresh_security
).on_post(user, resp)
@self.api_router.post("/logout")
async def login(resp: Response):
return await LogoutEndpoint(access_token=self.access_security, refresh_token=self.refresh_security).on_post(resp)
@self.api_router.post("/registration", status_code=201)
async def registration(user: RegistrationInput, db: AsyncSession = Depends(db_session)):