diff --git a/src/endpoints/base.py b/src/endpoints/base.py index bacccfe..b5553e5 100644 --- a/src/endpoints/base.py +++ b/src/endpoints/base.py @@ -1,13 +1,13 @@ from fastapi_jwt.jwt import JwtAccess, JwtRefresh + class BaseEndpoint: - def __init__(self, *args, **kwargs): - self.db = kwargs.get("db") - super().__init__(*args, **kwargs) + def __init__(self, db): + self.db = db 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__() + self.access_security: JwtAccess = kwargs.pop("access_token") + self.refresh_security: JwtRefresh = kwargs.pop("refresh_token") + super().__init__(*args, **kwargs) diff --git a/src/endpoints/login.py b/src/endpoints/login.py index 4e2a87a..fd2ea5d 100644 --- a/src/endpoints/login.py +++ b/src/endpoints/login.py @@ -1,6 +1,5 @@ from fastapi import HTTPException from fastapi_jwt import JwtAuthorizationCredentials - from passlib.hash import bcrypt from sqlalchemy import select from starlette.responses import Response @@ -14,7 +13,7 @@ class LoginInput(BaseModel): password: str -class LoginEndpoint(BaseEndpoint, AuthEndpoint): +class LoginEndpoint(AuthEndpoint, BaseEndpoint): 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() diff --git a/src/endpoints/registration.py b/src/endpoints/registration.py index f91edbf..71f0d61 100644 --- a/src/endpoints/registration.py +++ b/src/endpoints/registration.py @@ -32,13 +32,13 @@ class RegistrationEndpoint(BaseEndpoint): if existing_user: raise HTTPException(status_code=422, detail="User already exists") - model = User( - name=user_data.name, - email=user_data.email, - password_hashed=bcrypt.hash(user_data.password) - ) + model = await User.create(self.db, { + "name": user_data.name, + "email": user_data.email, + "password_hashed": bcrypt.hash(user_data.password), + "description": "" + }) - self.db.add(model) await self.db.commit() return model.as_dict() diff --git a/src/graphql_schema/entities/aircraft.py b/src/graphql_schema/entities/aircraft.py index 415e0e9..25b56d5 100644 --- a/src/graphql_schema/entities/aircraft.py +++ b/src/graphql_schema/entities/aircraft.py @@ -3,6 +3,7 @@ import strawberry from strawberry.file_uploads import Upload from sqlalchemy import select from database import models +from decorators.endpoints import authenticated_user_only from graphql_schema.sqlalchemy_to_strawberry_type import strawberry_sqlalchemy_type, strawberry_sqlalchemy_input from upload_utils import handle_file_upload, delete_file, get_public_url from ..dataloaders.flight import flights_by_aircraft_dataloader @@ -37,6 +38,7 @@ def get_base_query(user_id: int): class AircraftQueries: @strawberry.field + @authenticated_user_only async def aircrafts(root, info) -> List[Aircraft]: query = ( get_base_query(info.context.user_id) @@ -46,6 +48,7 @@ class AircraftQueries: return (await info.context.db.scalars(query)).all() @strawberry.field + @authenticated_user_only async def aircraft(root, info, id: int) -> Aircraft: query = ( get_base_query(info.context.user_id) @@ -61,6 +64,7 @@ class CreateAircraftMutation: photo: Optional[Upload] @strawberry.mutation + @authenticated_user_only async def create_aircraft(root, info, input: CreateAircraftInput) -> Aircraft: # TODO: kontrola organizace @@ -84,6 +88,7 @@ class EditAircraftMutation: photo: Optional[Upload] @strawberry.mutation + @authenticated_user_only async def edit_aircraft(root, info, id: int, input: EditAircraftInput) -> Aircraft: # TODO: kontrola organizace # TODO: kontrola opravneni na akci @@ -103,6 +108,7 @@ class EditAircraftMutation: class DeleteAircraftMutation: @strawberry.mutation + @authenticated_user_only async def delete_aircraft(self, info, id: int) -> Aircraft: # TODO: kontrola opravneni na akci diff --git a/src/main.py b/src/main.py index 54c3f36..217a24c 100644 --- a/src/main.py +++ b/src/main.py @@ -110,11 +110,14 @@ class App: @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) + 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)): - return await RegistrationEndpoint(db).on_post(user) + return await RegistrationEndpoint(db=db).on_post(user) # musi byt na konci app.include_router(self.api_router)