fix problems
This commit is contained in:
@@ -3,10 +3,10 @@ from pathlib import Path
|
||||
import uvicorn
|
||||
from fastapi import FastAPI
|
||||
|
||||
from src.web.protected_routes.auth_routes import router as protected_router
|
||||
from src.web.protected_routes.protected_user_action_routes import (
|
||||
router as protected_user_action_routes,
|
||||
)
|
||||
from src.web.protected_routes.routes import router as protected_router
|
||||
|
||||
app=FastAPI(root_path="/")
|
||||
app.include_router(router=protected_router)
|
||||
|
||||
@@ -1,10 +1,9 @@
|
||||
|
||||
from uuid import UUID
|
||||
|
||||
from sqlalchemy import and_, not_, select
|
||||
from sqlalchemy import and_, not_, select, update
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
from src.errors.http_errors.errors import Errors
|
||||
from src.models.database_models.model import RefreshTokens, engine
|
||||
from src.models.pydantic_models.model import RefreshTokensOut
|
||||
|
||||
@@ -12,75 +11,59 @@ from src.models.pydantic_models.model import RefreshTokensOut
|
||||
class JwtCrudActions:
|
||||
def __init__(self) -> None:
|
||||
self.Session=sessionmaker(bind=engine)
|
||||
self.error=Errors()
|
||||
|
||||
def get_token_by_user_id(self, user_id:UUID)->RefreshTokensOut|None:
|
||||
with self.Session() as session: # noqa: SIM117
|
||||
with session.begin():
|
||||
query=select(RefreshTokens).where(and_(RefreshTokens.user_id==user_id, not_(RefreshTokens.is_revoked)))
|
||||
response=session.scalars(query).one_or_none()
|
||||
if response is None:
|
||||
return None
|
||||
return RefreshTokensOut.model_validate(response)
|
||||
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()
|
||||
if response is None:
|
||||
return None
|
||||
return RefreshTokensOut.model_validate(response)
|
||||
|
||||
def get_token_by_id(self, id:UUID)->RefreshTokensOut|None:
|
||||
with self.Session() as session: # noqa: SIM117
|
||||
with session.begin():
|
||||
query=select(RefreshTokens).where(RefreshTokens.id==id)
|
||||
response=session.scalars(query).one_or_none()
|
||||
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():
|
||||
query=select(RefreshTokens).where(RefreshTokens.id==token_id)
|
||||
response=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: # noqa: SIM117
|
||||
with session.begin():
|
||||
new_token=RefreshTokens(**data)
|
||||
session.add(new_token)
|
||||
with self.Session() as session, session.begin():
|
||||
new_token=RefreshTokens(**data)
|
||||
session.add(new_token)
|
||||
|
||||
|
||||
def update_token(self, old_jti:UUID, new_jti:UUID)->bool:
|
||||
with self.Session() as session: # noqa: SIM117
|
||||
with session.begin():
|
||||
query=select(RefreshTokens).where(RefreshTokens.id==old_jti)
|
||||
response = session.scalars(query).one_or_none()
|
||||
if response is None:
|
||||
return False
|
||||
else:
|
||||
response.is_revoked=True
|
||||
response.replaced_by=new_jti
|
||||
return True
|
||||
|
||||
def create_and_update_token(self,data:dict, old_jti:UUID, new_jti:UUID)->bool:
|
||||
with self.Session() as session: #noqa:SIM117
|
||||
with session.begin():
|
||||
new_token=RefreshTokens(**data)
|
||||
query=select(RefreshTokens).where(RefreshTokens.id==old_jti)
|
||||
response = session.scalars(query).one_or_none()
|
||||
if response is None:
|
||||
return False
|
||||
else:
|
||||
session.add(new_token)
|
||||
response.is_revoked=True
|
||||
response.replaced_by=new_jti
|
||||
return True
|
||||
def create_and_update_token(self, data: dict, old_jti: UUID, new_jti: UUID) -> bool:
|
||||
with self.Session() as session, session.begin():
|
||||
new_token = RefreshTokens(**data)
|
||||
|
||||
query = (
|
||||
update(RefreshTokens)
|
||||
.where(RefreshTokens.id == old_jti, RefreshTokens.is_revoked.is_(False))
|
||||
.values(is_revoked=True, replaced_by=new_jti)
|
||||
.returning(RefreshTokens.id)
|
||||
)
|
||||
updated_id = session.execute(query).scalar_one_or_none()
|
||||
|
||||
if updated_id is None:
|
||||
return False
|
||||
|
||||
session.add(new_token)
|
||||
return True
|
||||
|
||||
|
||||
def revoke_all(self, user_id:UUID)->bool:
|
||||
with self.Session() as session: # noqa: SIM117
|
||||
with session.begin():
|
||||
query=select(RefreshTokens).where(RefreshTokens.user_id==user_id)
|
||||
response=session.scalars(query).all()
|
||||
for record in response:
|
||||
record.is_revoked=True
|
||||
return True
|
||||
|
||||
def logout(self,id:UUID)->bool:
|
||||
with self.Session() as session: # noqa: SIM117
|
||||
with session.begin():
|
||||
query=select(RefreshTokens).where(RefreshTokens.id == id)
|
||||
response=session.scalars(query).one_or_none()
|
||||
if response is None:
|
||||
return False
|
||||
else:
|
||||
response.is_revoked=True
|
||||
return True
|
||||
with self.Session() as session, session.begin():
|
||||
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():
|
||||
query=select(RefreshTokens).where(RefreshTokens.id == token_id)
|
||||
response=session.scalars(query).one_or_none()
|
||||
if response is None:
|
||||
return False
|
||||
else:
|
||||
response.is_revoked=True
|
||||
return True
|
||||
+12
-14
@@ -13,19 +13,17 @@ class UsersCrudActions:
|
||||
self.Session=sessionmaker(bind=engine)
|
||||
|
||||
def get_user_by_email(self, email:str)->UserOutDB|None:
|
||||
with self.Session() as session: # noqa: SIM117
|
||||
with session.begin():
|
||||
query=select(User).where(User.email==email)
|
||||
response=session.scalars(query).one_or_none()
|
||||
if response is None:
|
||||
return None
|
||||
return UserOutDB.model_validate(response)
|
||||
with self.Session() as session, session.begin():
|
||||
query=select(User).where(User.email==email)
|
||||
response=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: # noqa: SIM117
|
||||
with session.begin():
|
||||
query=select(User).where(User.id==id)
|
||||
response=session.scalars(query).one_or_none()
|
||||
if response is None:
|
||||
return None
|
||||
return UserOutDB.model_validate(response)
|
||||
with self.Session() as session, session.begin():
|
||||
query=select(User).where(User.id==id)
|
||||
response=session.scalars(query).one_or_none()
|
||||
if response is None:
|
||||
return None
|
||||
return UserOutDB.model_validate(response)
|
||||
+27
-22
@@ -36,6 +36,18 @@ class CurrentUserService:
|
||||
return user
|
||||
|
||||
|
||||
def _token_record_create(self, jti:UUID,user_id:UUID,token:str, request:Request)->RefreshTokensCreate:
|
||||
|
||||
return RefreshTokensCreate(
|
||||
id=jti,
|
||||
user_id=user_id,
|
||||
token_hash=self.hash.token_to_hash(token),
|
||||
device_info=request.headers.get("user-agent", "unknown"),
|
||||
ip_address=request.headers.get("x-forwarded-for", "").split(",")[0].strip() or (request.client.host if request.client else "unknown"),
|
||||
expires_at=datetime.now(UTC)+timedelta(days=env_settings.REFRESH_TOKEN_EXPIRE_DAYS)
|
||||
)
|
||||
|
||||
|
||||
def get_current_user(self, token:str)->UserOut:
|
||||
|
||||
payload=self.jwt_service.jwt_decode(token)
|
||||
@@ -46,6 +58,8 @@ class CurrentUserService:
|
||||
except (ValueError, TypeError) as e:
|
||||
raise self.error.credentials_error(detail="Jwt token is incorrect") from e
|
||||
|
||||
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)
|
||||
if user is None:
|
||||
@@ -74,14 +88,9 @@ 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=RefreshTokensCreate(
|
||||
id=jti,
|
||||
user_id=user_id,
|
||||
token_hash=self.hash.token_to_hash(token),
|
||||
device_info=request.headers.get("user-agent", "unknown"),
|
||||
ip_address=request.headers.get("x-forwarded-for", "").split(",")[0].strip() or (request.client.host if request.client else "unknown"),
|
||||
expires_at=datetime.now(UTC)+timedelta(days=env_settings.REFRESH_TOKEN_EXPIRE_DAYS)
|
||||
)
|
||||
token_record=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))
|
||||
|
||||
return token
|
||||
@@ -104,12 +113,14 @@ class CurrentUserService:
|
||||
|
||||
|
||||
'''old refresh token check'''
|
||||
|
||||
if (old_refresh_token.get("token_type")=="access"):
|
||||
raise self.error.credentials_error(detail="Jwt token type is incorrect")
|
||||
|
||||
|
||||
old_record=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")
|
||||
if old_record.is_revoked:
|
||||
self.jwt_db_actions.revoke_all(old_record.user_id)
|
||||
raise self.error.credentials_error(detail="Reuse token detected")
|
||||
|
||||
|
||||
'''sqlite constraints about timezone'''
|
||||
@@ -133,19 +144,12 @@ class CurrentUserService:
|
||||
|
||||
try:
|
||||
new_jti=UUID(new_jti)
|
||||
except (ValueError, TypeError):
|
||||
raise self.error.credentials_error(detail="Jwt token is incorrect")
|
||||
except (ValueError, TypeError) as e:
|
||||
raise self.error.credentials_error(detail="Jwt token is incorrect") from e
|
||||
|
||||
|
||||
'''create database record with the new token'''
|
||||
new_token_record=RefreshTokensCreate(
|
||||
id=new_jti,
|
||||
user_id=sub,
|
||||
token_hash=self.hash.token_to_hash(new_refresh_token),
|
||||
device_info=request.headers.get("user-agent", "unknown"),
|
||||
ip_address=request.headers.get("x-forwarded-for", "").split(",")[0].strip() or (request.client.host if request.client else "unknown"),
|
||||
expires_at=datetime.now(UTC)+timedelta(days=env_settings.REFRESH_TOKEN_EXPIRE_DAYS),
|
||||
)
|
||||
new_token_record=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)
|
||||
|
||||
@@ -188,4 +192,5 @@ class CurrentUserService:
|
||||
|
||||
return (access_token, refresh_token)
|
||||
|
||||
auth=CurrentUserService()
|
||||
def auth_service()->CurrentUserService:
|
||||
return CurrentUserService()
|
||||
+10
-4
@@ -30,12 +30,17 @@ class JwtService:
|
||||
def __init__(self) -> None:
|
||||
|
||||
self.error=Errors()
|
||||
|
||||
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:
|
||||
|
||||
user_info=data.copy()
|
||||
if not (user_info.get("sub")) or user_info.get("sub") == "":
|
||||
raise self.error.credentials_error(detail="Jwt token is incorrect")
|
||||
|
||||
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)
|
||||
@@ -45,8 +50,9 @@ class JwtService:
|
||||
|
||||
user_info=data.copy()
|
||||
jti=str(uuid4())
|
||||
if not (user_info.get("sub")) or user_info.get("sub") == "":
|
||||
raise self.error.credentials_error(detail="Jwt token is incorrect")
|
||||
|
||||
self._validate_sub(user_info)
|
||||
|
||||
user_info.update({"exp":datetime.now(UTC)+timedelta(days=env_settings.REFRESH_TOKEN_EXPIRE_DAYS),
|
||||
"token_type":"refresh",
|
||||
"jti":jti
|
||||
|
||||
@@ -3,13 +3,13 @@ from fastapi.security import OAuth2PasswordBearer, OAuth2PasswordRequestForm
|
||||
|
||||
from src.models.configs_read.env import env_settings
|
||||
from src.models.pydantic_models.model import UserOut
|
||||
from src.service.auth.auth import auth
|
||||
from src.service.auth.auth import CurrentUserService, auth_service
|
||||
|
||||
router=APIRouter(prefix="/protected")
|
||||
oauth2_schema=OAuth2PasswordBearer(tokenUrl="/protected/token", refreshUrl="/protected/refresh")
|
||||
|
||||
@router.post("/token")
|
||||
async def get_access_token(request: Request,response:Response, form_data:OAuth2PasswordRequestForm=Depends())->dict: # noqa: B008
|
||||
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)
|
||||
|
||||
@@ -25,7 +25,7 @@ async def get_access_token(request: Request,response:Response, form_data:OAuth2P
|
||||
|
||||
|
||||
@router.post("/refresh")
|
||||
async def get_refresh_token(request:Request,response:Response, refresh_token: str = Cookie())->dict:
|
||||
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)
|
||||
|
||||
@@ -41,12 +41,12 @@ async def get_refresh_token(request:Request,response:Response, refresh_token: st
|
||||
return {"access_token":access_token, "token_type": "bearer"}
|
||||
|
||||
|
||||
async def get_current_user(token:str = Depends(oauth2_schema)) -> UserOut:
|
||||
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))
|
||||
|
||||
|
||||
@router.get("/logout")
|
||||
async def logout(response:Response,refresh_token: str = Cookie(),current_user:UserOut=Depends(get_current_user))->bool: # noqa: B008
|
||||
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)
|
||||
|
||||
@@ -2,7 +2,7 @@ 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.web.protected_routes.routes import get_current_user
|
||||
from src.web.protected_routes.auth_routes import get_current_user
|
||||
|
||||
router=APIRouter(prefix="/user")
|
||||
|
||||
|
||||
@@ -184,20 +184,20 @@ class TestAuth:
|
||||
assert new_access_token!=token
|
||||
assert new_refresh_token!=token
|
||||
|
||||
@pytest.mark.parametrize("db_result_token, user_data_result_db, fake_token_data,expected_exception", [
|
||||
pytest.param(SimpleNamespace(is_revoked=True,user_id="123", expires_at=datetime.now(UTC)+timedelta(days=15)),SimpleNamespace(status=True),{"sub":str(uuid4()), "jti":str(uuid4()), "token_type":"refresh", "exp":datetime.now(UTC)+timedelta(days=45)}, HTTPException,id="false_revoke_status"),
|
||||
pytest.param(SimpleNamespace(is_revoked=False,user_id="123", expires_at=datetime.now(UTC)+timedelta(days=15)),SimpleNamespace(status=True),{"sub":str(uuid4()), "token_type":"refresh", "exp":datetime.now(UTC)+timedelta(days=45)}, HTTPException, id="jti_missing"),
|
||||
pytest.param(None,SimpleNamespace(status=True),{"sub":str(uuid4()), "jti":str(uuid4()),"token_type":"refresh", "exp":datetime.now(UTC)+timedelta(days=45)}, HTTPException, id="token_missing"),
|
||||
pytest.param(SimpleNamespace(is_revoked=False,user_id="123", expires_at=datetime.now(UTC)+timedelta(days=15)),SimpleNamespace(status=False),{"sub":str(uuid4()), "jti":str(uuid4()), "token_type":"refresh", "exp":datetime.now(UTC)+timedelta(days=45)}, HTTPException,id="false_user_status"),
|
||||
pytest.param(SimpleNamespace(is_revoked=False, user_id="123",expires_at=datetime.now(UTC)+timedelta(days=15)),None,{"sub":str(uuid4()), "jti":str(uuid4()), "token_type":"refresh", "exp":datetime.now(UTC)+timedelta(days=45)}, HTTPException,id="user_missing"),
|
||||
pytest.param(SimpleNamespace(is_revoked=False,user_id="123", expires_at=datetime.now(UTC)-timedelta(days=15)),SimpleNamespace(status=True),{"sub":str(uuid4()), "jti":str(uuid4()), "token_type":"refresh", "exp":datetime.now(UTC)+timedelta(days=45)}, HTTPException,id="wrong_exp")
|
||||
@pytest.mark.parametrize("db_result_token, user_data_result_db, update_result, fake_token_data,expected_exception", [
|
||||
pytest.param(SimpleNamespace(is_revoked=True,user_id="123", expires_at=datetime.now(UTC)+timedelta(days=15)),SimpleNamespace(status=True),False,{"sub":str(uuid4()), "jti":str(uuid4()), "token_type":"refresh", "exp":datetime.now(UTC)+timedelta(days=45)}, HTTPException,id="false_revoke_status"),
|
||||
pytest.param(SimpleNamespace(is_revoked=False,user_id="123", expires_at=datetime.now(UTC)+timedelta(days=15)),SimpleNamespace(status=True),True,{"sub":str(uuid4()), "token_type":"refresh", "exp":datetime.now(UTC)+timedelta(days=45)}, HTTPException, id="jti_missing"),
|
||||
pytest.param(None,SimpleNamespace(status=True),True,{"sub":str(uuid4()), "jti":str(uuid4()),"token_type":"refresh", "exp":datetime.now(UTC)+timedelta(days=45)}, HTTPException, id="token_missing"),
|
||||
pytest.param(SimpleNamespace(is_revoked=False,user_id="123", expires_at=datetime.now(UTC)+timedelta(days=15)),SimpleNamespace(status=False),True,{"sub":str(uuid4()), "jti":str(uuid4()), "token_type":"refresh", "exp":datetime.now(UTC)+timedelta(days=45)}, HTTPException,id="false_user_status"),
|
||||
pytest.param(SimpleNamespace(is_revoked=False, user_id="123",expires_at=datetime.now(UTC)+timedelta(days=15)),None,True,{"sub":str(uuid4()), "jti":str(uuid4()), "token_type":"refresh", "exp":datetime.now(UTC)+timedelta(days=45)}, HTTPException,id="user_missing"),
|
||||
pytest.param(SimpleNamespace(is_revoked=False,user_id="123", expires_at=datetime.now(UTC)-timedelta(days=15)),SimpleNamespace(status=True),True,{"sub":str(uuid4()), "jti":str(uuid4()), "token_type":"refresh", "exp":datetime.now(UTC)+timedelta(days=45)}, HTTPException,id="wrong_exp")
|
||||
])
|
||||
def test_refresh_token_negative(self, monkeypatch, current_user_service:CurrentUserService, db_result_token, requests, jwt_service:JwtService,user_data_result_db, expected_exception, fake_token_data)->None:
|
||||
def test_refresh_token_negative(self, monkeypatch, current_user_service:CurrentUserService, db_result_token, requests, jwt_service:JwtService,user_data_result_db, expected_exception, fake_token_data, update_result)->None:
|
||||
with allure.step("patching db call functions"):
|
||||
monkeypatch.setattr(current_user_service.jwt_db_actions,"get_token_by_id", lambda jti: db_result_token)
|
||||
monkeypatch.setattr(current_user_service.jwt_db_actions, "revoke_all", lambda user_id:True)
|
||||
monkeypatch.setattr(current_user_service.crud_db_actions, "get_user_by_id", lambda user:user_data_result_db)
|
||||
monkeypatch.setattr(current_user_service.jwt_db_actions, "create_and_update_token", lambda old_jti, new_jti, new_token_record:True)
|
||||
monkeypatch.setattr(current_user_service.jwt_db_actions, "create_and_update_token", lambda old_jti, new_jti, new_token_record:update_result)
|
||||
|
||||
fake_request = requests
|
||||
|
||||
|
||||
Reference in New Issue
Block a user