asyncio style refactoring
This commit is contained in:
@@ -2,6 +2,7 @@
|
||||
from uuid import UUID
|
||||
|
||||
from sqlalchemy import and_, not_, select, update
|
||||
from sqlalchemy.ext.asyncio import async_sessionmaker
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
from src.models.database_models.model import RefreshTokens, engine
|
||||
@@ -10,32 +11,32 @@ from src.models.pydantic_models.model import RefreshTokensOut
|
||||
|
||||
class JwtCrudActions:
|
||||
def __init__(self) -> None:
|
||||
self.Session=sessionmaker(bind=engine)
|
||||
self.Session=async_sessionmaker(bind=engine)
|
||||
|
||||
def get_token_by_user_id(self, user_id:UUID)->RefreshTokensOut|None:
|
||||
with self.Session() as session, session.begin():
|
||||
async def get_token_by_user_id(self, user_id:UUID)->RefreshTokensOut|None:
|
||||
async with self.Session() as session, session.begin():
|
||||
query=select(RefreshTokens).where(and_(RefreshTokens.user_id==user_id, not_(RefreshTokens.is_revoked)))
|
||||
response=session.scalars(query).first()
|
||||
response= (await session.scalars(query)).first()
|
||||
if response is None:
|
||||
return None
|
||||
return RefreshTokensOut.model_validate(response)
|
||||
|
||||
def get_token_by_id(self, token_id:UUID)->RefreshTokensOut|None:
|
||||
with self.Session() as session, session.begin():
|
||||
async def get_token_by_id(self, token_id:UUID)->RefreshTokensOut|None:
|
||||
async with self.Session() as session, session.begin():
|
||||
query=select(RefreshTokens).where(RefreshTokens.id==token_id)
|
||||
response=session.scalars(query).one_or_none()
|
||||
response= (await session.scalars(query)).one_or_none()
|
||||
if response is None:
|
||||
return None
|
||||
return RefreshTokensOut.model_validate(response)
|
||||
|
||||
def create_token(self, data:dict)->None:
|
||||
with self.Session() as session, session.begin():
|
||||
async def create_token(self, data:dict)->None:
|
||||
async with self.Session() as session, session.begin():
|
||||
new_token=RefreshTokens(**data)
|
||||
session.add(new_token)
|
||||
|
||||
|
||||
def create_and_update_token(self, data: dict, old_jti: UUID, new_jti: UUID) -> bool:
|
||||
with self.Session() as session, session.begin():
|
||||
async def create_and_update_token(self, data: dict, old_jti: UUID, new_jti: UUID) -> bool:
|
||||
async with self.Session() as session, session.begin():
|
||||
new_token = RefreshTokens(**data)
|
||||
|
||||
query = (
|
||||
@@ -44,7 +45,7 @@ class JwtCrudActions:
|
||||
.values(is_revoked=True, replaced_by=new_jti)
|
||||
.returning(RefreshTokens.id)
|
||||
)
|
||||
updated_id = session.execute(query).scalar_one_or_none()
|
||||
updated_id = (await session.execute(query)).scalar_one_or_none()
|
||||
|
||||
if updated_id is None:
|
||||
return False
|
||||
@@ -53,15 +54,15 @@ class JwtCrudActions:
|
||||
return True
|
||||
|
||||
|
||||
def revoke_all(self, user_id:UUID)->bool:
|
||||
with self.Session() as session, session.begin():
|
||||
session.execute(update(RefreshTokens).where(RefreshTokens.user_id==user_id).values(is_revoked=True)) #bulk update
|
||||
async def revoke_all(self, user_id:UUID)->bool:
|
||||
async with self.Session() as session, session.begin():
|
||||
await session.execute(update(RefreshTokens).where(RefreshTokens.user_id==user_id).values(is_revoked=True)) #bulk update
|
||||
return True
|
||||
|
||||
def logout(self,token_id:UUID)->bool:
|
||||
with self.Session() as session, session.begin():
|
||||
async def logout(self,token_id:UUID)->bool:
|
||||
async with self.Session() as session, session.begin():
|
||||
query=select(RefreshTokens).where(RefreshTokens.id == token_id)
|
||||
response=session.scalars(query).one_or_none()
|
||||
response= (await session.scalars(query)).one_or_none()
|
||||
if response is None:
|
||||
return False
|
||||
else:
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
from uuid import UUID
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
from sqlalchemy.ext.asyncio import async_sessionmaker
|
||||
|
||||
from src.models.database_models.model import User, engine
|
||||
from src.models.pydantic_models.model import UserOutDB
|
||||
@@ -10,20 +10,20 @@ from src.models.pydantic_models.model import UserOutDB
|
||||
class UsersCrudActions:
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.Session=sessionmaker(bind=engine)
|
||||
self.Session=async_sessionmaker(bind=engine)
|
||||
|
||||
def get_user_by_email(self, email:str)->UserOutDB|None:
|
||||
with self.Session() as session, session.begin():
|
||||
async def get_user_by_email(self, email:str)->UserOutDB|None:
|
||||
async with self.Session() as session, session.begin():
|
||||
query=select(User).where(User.email==email)
|
||||
response=session.scalars(query).one_or_none()
|
||||
response=(await session.scalars(query)).one_or_none()
|
||||
if response is None:
|
||||
return None
|
||||
return UserOutDB.model_validate(response)
|
||||
|
||||
def get_user_by_id(self, id:UUID)->UserOutDB|None:
|
||||
with self.Session() as session, session.begin():
|
||||
async def get_user_by_id(self, id:UUID)->UserOutDB|None:
|
||||
async with self.Session() as session, session.begin():
|
||||
query=select(User).where(User.id==id)
|
||||
response=session.scalars(query).one_or_none()
|
||||
response=(await session.scalars(query)).one_or_none()
|
||||
if response is None:
|
||||
return None
|
||||
return UserOutDB.model_validate(response)
|
||||
@@ -10,12 +10,12 @@ from sqlalchemy import (
|
||||
String,
|
||||
Table,
|
||||
Uuid,
|
||||
create_engine,
|
||||
func,
|
||||
)
|
||||
from sqlalchemy.ext.asyncio import create_async_engine
|
||||
from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column, relationship
|
||||
|
||||
engine = create_engine("sqlite:///DB/database.db", echo=True)
|
||||
engine = create_async_engine("sqlite+aiosqlite:///DB/database.db", echo=True)
|
||||
|
||||
'''remember as a boilerplate, or just cp/pst'''
|
||||
class Model(DeclarativeBase):
|
||||
|
||||
+31
-30
@@ -1,3 +1,4 @@
|
||||
import asyncio
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from uuid import UUID
|
||||
|
||||
@@ -21,14 +22,14 @@ class CurrentUserService:
|
||||
self.jwt_db_actions=JwtCrudActions()
|
||||
self.error=Errors()
|
||||
|
||||
def _check(self, form_data_email:str, form_data_password:str,):
|
||||
async def _check(self, form_data_email:str, form_data_password:str,):
|
||||
'''check user by email'''
|
||||
user=self.crud_db_actions.get_user_by_email(form_data_email)
|
||||
user=await self.crud_db_actions.get_user_by_email(form_data_email)
|
||||
|
||||
if user is None:
|
||||
raise self.error.credentials_error(detail="Wrong credentials")
|
||||
|
||||
if not self.hash.verify_password(plain_password=form_data_password, hashed_password=user.hashed_password):
|
||||
if not await asyncio.to_thread(self.hash.verify_password, plain_password=form_data_password, hashed_password=user.hashed_password):
|
||||
raise self.error.credentials_error(detail="Wrong credentials")
|
||||
|
||||
if user.status is False:
|
||||
@@ -36,7 +37,7 @@ class CurrentUserService:
|
||||
return user
|
||||
|
||||
|
||||
def _token_record_create(self, jti:UUID,user_id:UUID,token:str, request:Request)->RefreshTokensCreate:
|
||||
async def _token_record_create(self, jti:UUID,user_id:UUID,token:str, request:Request)->RefreshTokensCreate:
|
||||
|
||||
return RefreshTokensCreate(
|
||||
id=jti,
|
||||
@@ -48,9 +49,9 @@ class CurrentUserService:
|
||||
)
|
||||
|
||||
|
||||
def get_current_user(self, token:str)->UserOut:
|
||||
async def get_current_user(self, token:str)->UserOut:
|
||||
|
||||
payload=self.jwt_service.jwt_decode(token)
|
||||
payload= await self.jwt_service.jwt_decode(token)
|
||||
sub=payload.get("sub")
|
||||
|
||||
try:
|
||||
@@ -61,7 +62,7 @@ class CurrentUserService:
|
||||
if not (payload.get("token_type")=="access"):
|
||||
raise self.error.credentials_error(detail="Jwt token type is incorrect")
|
||||
|
||||
user=self.crud_db_actions.get_user_by_id(sub)
|
||||
user=await self.crud_db_actions.get_user_by_id(sub)
|
||||
if user is None:
|
||||
raise self.error.not_found_error(detail="User with this email address not found")
|
||||
|
||||
@@ -72,15 +73,15 @@ class CurrentUserService:
|
||||
|
||||
|
||||
|
||||
def create_access_token(self, user_id:UUID)->str:
|
||||
async def create_access_token(self, user_id:UUID)->str:
|
||||
'''create new access token if all the checks are successful'''
|
||||
return self.jwt_service.create_access_token({"sub":str(user_id)})
|
||||
return await self.jwt_service.create_access_token({"sub":str(user_id)})
|
||||
|
||||
|
||||
|
||||
def create_refresh_token(self,user_id:UUID, request:Request)->str:
|
||||
async def create_refresh_token(self,user_id:UUID, request:Request)->str:
|
||||
|
||||
token, jti=self.jwt_service.create_refresh_token({"sub":str(user_id)})
|
||||
token, jti= await self.jwt_service.create_refresh_token({"sub":str(user_id)})
|
||||
|
||||
try:
|
||||
jti=UUID(jti)
|
||||
@@ -88,18 +89,18 @@ class CurrentUserService:
|
||||
raise self.error.credentials_error(detail="Jwt token is incorrect") from e
|
||||
|
||||
'''create new refresh token if all the checks are successful'''
|
||||
token_record=self._token_record_create(jti=jti, user_id=user_id, token=token, request=request)
|
||||
token_record=await self._token_record_create(jti=jti, user_id=user_id, token=token, request=request)
|
||||
|
||||
|
||||
self.jwt_db_actions.create_token(RefreshTokensCreate.model_dump(token_record))
|
||||
await self.jwt_db_actions.create_token(RefreshTokensCreate.model_dump(token_record))
|
||||
|
||||
return token
|
||||
|
||||
|
||||
def refresh_token(self, refresh_token:str, request:Request)->tuple[str, str]:
|
||||
async def refresh_token(self, refresh_token:str, request:Request)->tuple[str, str]:
|
||||
|
||||
'''decode old refresh token'''
|
||||
old_refresh_token=self.jwt_service.jwt_decode(refresh_token)
|
||||
old_refresh_token= await self.jwt_service.jwt_decode(refresh_token)
|
||||
sub=old_refresh_token.get("sub")
|
||||
|
||||
if (old_jti:=old_refresh_token.get("jti")) is None:
|
||||
@@ -118,7 +119,7 @@ class CurrentUserService:
|
||||
raise self.error.credentials_error(detail="Jwt token type is incorrect")
|
||||
|
||||
|
||||
old_record=self.jwt_db_actions.get_token_by_id(old_jti)
|
||||
old_record=await self.jwt_db_actions.get_token_by_id(old_jti)
|
||||
if old_record is None:
|
||||
raise self.error.not_found_error(detail="Token not found")
|
||||
|
||||
@@ -131,7 +132,7 @@ class CurrentUserService:
|
||||
raise self.error.credentials_error(detail="Token expired")
|
||||
|
||||
'''user check'''
|
||||
user = self.crud_db_actions.get_user_by_id(sub)
|
||||
user = await self.crud_db_actions.get_user_by_id(sub)
|
||||
if user is None:
|
||||
raise self.error.not_found_error(detail="User not found")
|
||||
if user.status is False:
|
||||
@@ -139,8 +140,8 @@ class CurrentUserService:
|
||||
|
||||
|
||||
'''create new refresh token if all the checks are successful'''
|
||||
new_refresh_token, new_jti=self.jwt_service.create_refresh_token({"sub":str(sub)})
|
||||
new_access_token=self.create_access_token(user_id=sub)
|
||||
new_refresh_token, new_jti= await self.jwt_service.create_refresh_token({"sub":str(sub)})
|
||||
new_access_token=await self.create_access_token(user_id=sub)
|
||||
|
||||
try:
|
||||
new_jti=UUID(new_jti)
|
||||
@@ -149,9 +150,9 @@ class CurrentUserService:
|
||||
|
||||
|
||||
'''create database record with the new token'''
|
||||
new_token_record=self._token_record_create(jti=new_jti, user_id=sub, token=new_refresh_token, request=request)
|
||||
new_token_record=await self._token_record_create(jti=new_jti, user_id=sub, token=new_refresh_token, request=request)
|
||||
|
||||
success = self.jwt_db_actions.create_and_update_token(RefreshTokensCreate.model_dump(new_token_record), old_jti, new_jti)
|
||||
success = await self.jwt_db_actions.create_and_update_token(RefreshTokensCreate.model_dump(new_token_record), old_jti, new_jti)
|
||||
|
||||
if not success:
|
||||
raise self.error.not_found_error(detail="Token not found")
|
||||
@@ -160,10 +161,10 @@ class CurrentUserService:
|
||||
|
||||
|
||||
|
||||
def logout(self, refresh_token:str)->bool:
|
||||
async def logout(self, refresh_token:str)->bool:
|
||||
|
||||
'''decode current refresh token'''
|
||||
payload=self.jwt_service.jwt_decode(refresh_token)
|
||||
payload=await self.jwt_service.jwt_decode(refresh_token)
|
||||
|
||||
if (jti:=payload.get("jti")) is None:
|
||||
raise self.error.credentials_error(detail="Jwt token is incorrect")
|
||||
@@ -174,23 +175,23 @@ class CurrentUserService:
|
||||
raise self.error.credentials_error(detail="Jwt token is incorrect") from e
|
||||
|
||||
'''logout by assigning revoked flag'''
|
||||
if self.jwt_db_actions.logout(jti):
|
||||
if await self.jwt_db_actions.logout(jti):
|
||||
return True
|
||||
else:
|
||||
raise self.error.not_found_error(detail="Refresh Token Not Found")
|
||||
|
||||
|
||||
|
||||
def login(self, form_data_email:str, form_data_password:str, request:Request)->tuple[str, str]:
|
||||
async def login(self, form_data_email:str, form_data_password:str, request:Request)->tuple[str, str]:
|
||||
'''revoke all the old refresh tokens'''
|
||||
user = self._check(form_data_email, form_data_password)
|
||||
self.jwt_db_actions.revoke_all(user_id=user.id)
|
||||
user = await self._check(form_data_email, form_data_password)
|
||||
await self.jwt_db_actions.revoke_all(user_id=user.id)
|
||||
|
||||
'''create access and refresh tokens'''
|
||||
access_token=self.create_access_token(user_id=user.id)
|
||||
refresh_token=self.create_refresh_token(user_id=user.id,request=request)
|
||||
access_token=await self.create_access_token(user_id=user.id)
|
||||
refresh_token=await self.create_refresh_token(user_id=user.id,request=request)
|
||||
|
||||
return (access_token, refresh_token)
|
||||
|
||||
def auth_service()->CurrentUserService:
|
||||
async def auth_service()->CurrentUserService:
|
||||
return CurrentUserService()
|
||||
@@ -31,27 +31,27 @@ class JwtService:
|
||||
|
||||
self.error=Errors()
|
||||
|
||||
def _validate_sub(self,data:dict)->None:
|
||||
async def _validate_sub(self,data:dict)->None:
|
||||
if not (data.get("sub")) or data.get("sub") == "":
|
||||
raise self.error.credentials_error(detail="Jwt token is incorrect")
|
||||
|
||||
def create_access_token(self, data:dict)->str:
|
||||
async def create_access_token(self, data:dict)->str:
|
||||
|
||||
user_info=data.copy()
|
||||
|
||||
self._validate_sub(user_info)
|
||||
await self._validate_sub(user_info)
|
||||
|
||||
user_info.update({"exp": datetime.now(UTC)+timedelta(minutes=env_settings.ACCESS_TOKEN_EXPIRE_MINUTES),
|
||||
"token_type":"access"})
|
||||
return jwt.encode(user_info, env_settings.SECRET_KEY, env_settings.ALGORITHM)
|
||||
|
||||
|
||||
def create_refresh_token(self, data:dict)->tuple[str, str]:
|
||||
async def create_refresh_token(self, data:dict)->tuple[str, str]:
|
||||
|
||||
user_info=data.copy()
|
||||
jti=str(uuid4())
|
||||
|
||||
self._validate_sub(user_info)
|
||||
await self._validate_sub(user_info)
|
||||
|
||||
user_info.update({"exp":datetime.now(UTC)+timedelta(days=env_settings.REFRESH_TOKEN_EXPIRE_DAYS),
|
||||
"token_type":"refresh",
|
||||
@@ -62,7 +62,7 @@ class JwtService:
|
||||
|
||||
|
||||
|
||||
def jwt_decode(self, token:str)->dict:
|
||||
async def jwt_decode(self, token:str)->dict:
|
||||
|
||||
try:
|
||||
payload=jwt.decode(token, env_settings.SECRET_KEY, algorithms=[env_settings.ALGORITHM], options={"require_exp": True} )
|
||||
|
||||
@@ -9,10 +9,11 @@ class CrudService:
|
||||
self.errors=Errors()
|
||||
self.crud_db_actions=UsersCrudActions()
|
||||
|
||||
def get_user_by_email(self, email:str)->UserOut:
|
||||
user_entity=self.crud_db_actions.get_user_by_email(email)
|
||||
async def get_user_by_email(self, email:str)->UserOut:
|
||||
user_entity=await self.crud_db_actions.get_user_by_email(email)
|
||||
if not user_entity:
|
||||
raise self.errors.not_found_error(detail="User wasn't found")
|
||||
return UserOut.model_validate(user_entity)
|
||||
|
||||
crud_service=CrudService()
|
||||
async def crud_service()->CrudService:
|
||||
return CrudService()
|
||||
@@ -11,7 +11,7 @@ oauth2_schema=OAuth2PasswordBearer(tokenUrl="/protected/token", refreshUrl="/pro
|
||||
@router.post("/token")
|
||||
async def get_access_token(request: Request,response:Response,auth:CurrentUserService=Depends(auth_service), form_data:OAuth2PasswordRequestForm=Depends())->dict: # noqa: B008
|
||||
|
||||
access_token, refresh_token=auth.login(form_data_email=form_data.username, form_data_password=form_data.password, request=request)
|
||||
access_token, refresh_token=await auth.login(form_data_email=form_data.username, form_data_password=form_data.password, request=request)
|
||||
|
||||
response.set_cookie(
|
||||
key="refresh_token",
|
||||
@@ -27,7 +27,7 @@ async def get_access_token(request: Request,response:Response,auth:CurrentUserSe
|
||||
@router.post("/refresh")
|
||||
async def get_refresh_token(request:Request,response:Response, refresh_token: str = Cookie(), auth:CurrentUserService=Depends(auth_service))->dict: # noqa: B008
|
||||
|
||||
access_token, refresh_token= auth.refresh_token(refresh_token=refresh_token,request=request)
|
||||
access_token, refresh_token= await auth.refresh_token(refresh_token=refresh_token,request=request)
|
||||
|
||||
response.set_cookie(
|
||||
key="refresh_token",
|
||||
@@ -42,13 +42,13 @@ async def get_refresh_token(request:Request,response:Response, refresh_token: st
|
||||
|
||||
|
||||
async def get_current_user(token:str = Depends(oauth2_schema), auth:CurrentUserService=Depends(auth_service)) -> UserOut: # noqa: B008
|
||||
return UserOut.model_validate(auth.get_current_user(token))
|
||||
return UserOut.model_validate(await auth.get_current_user(token))
|
||||
|
||||
|
||||
@router.get("/logout")
|
||||
async def logout(response:Response,refresh_token: str = Cookie(),auth:CurrentUserService=Depends(auth_service),current_user:UserOut=Depends(get_current_user))->bool: # noqa: B008
|
||||
response.delete_cookie("refresh_token")
|
||||
return auth.logout(refresh_token)
|
||||
return await auth.logout(refresh_token)
|
||||
|
||||
@router.get("")
|
||||
async def protected(current_user:UserOut=Depends(get_current_user))->dict: # noqa: B008
|
||||
|
||||
@@ -1,12 +1,11 @@
|
||||
from fastapi import APIRouter, Depends
|
||||
|
||||
from src.models.pydantic_models.model import UserOut
|
||||
from src.service.users_crud.users_crud import crud_service
|
||||
from src.service.users_crud.users_crud import CrudService, crud_service
|
||||
from src.web.protected_routes.auth_routes import get_current_user
|
||||
|
||||
router=APIRouter(prefix="/user")
|
||||
|
||||
|
||||
@router.get("/get_by_email")
|
||||
async def get_current_user_by_email(email:str, current_user=Depends(get_current_user))->UserOut: # noqa: B008
|
||||
return crud_service.get_user_by_email(email)
|
||||
async def get_current_user_by_email(email:str, crud:CrudService=Depends(crud_service), current_user=Depends(get_current_user))->UserOut: # noqa: B008
|
||||
return await crud.get_user_by_email(email)
|
||||
Reference in New Issue
Block a user