Uklid v alembic migracich, error logging, upravy kvuli autorizaci
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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
@@ -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
@@ -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(
|
||||
|
||||
@@ -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,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)
|
||||
|
||||
@@ -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
@@ -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)):
|
||||
|
||||
Reference in New Issue
Block a user