From 2da58c748353ad23d90974d458698f3c0978c44c Mon Sep 17 00:00:00 2001 From: "MH.Dmitrii" Date: Wed, 5 Aug 2026 12:22:45 +0300 Subject: [PATCH] fix problems --- main.py | 2 +- src/database/auth/refresh_tokens.py | 113 ++++++++---------- src/database/users/crud.py | 26 ++-- src/service/auth/auth.py | 49 ++++---- src/service/auth/jwt.py | 14 ++- .../{routes.py => auth_routes.py} | 10 +- .../protected_user_action_routes.py | 2 +- tests/integrated/test_auth.py | 18 +-- 8 files changed, 113 insertions(+), 121 deletions(-) rename src/web/protected_routes/{routes.py => auth_routes.py} (74%) diff --git a/main.py b/main.py index 7638a9a..6130758 100644 --- a/main.py +++ b/main.py @@ -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) diff --git a/src/database/auth/refresh_tokens.py b/src/database/auth/refresh_tokens.py index ccd19c5..95adfe8 100644 --- a/src/database/auth/refresh_tokens.py +++ b/src/database/auth/refresh_tokens.py @@ -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 \ No newline at end of file + 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 \ No newline at end of file diff --git a/src/database/users/crud.py b/src/database/users/crud.py index 9db4a8f..d23c79d 100644 --- a/src/database/users/crud.py +++ b/src/database/users/crud.py @@ -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) \ No newline at end of file + 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) \ No newline at end of file diff --git a/src/service/auth/auth.py b/src/service/auth/auth.py index 3eb9c78..f216314 100644 --- a/src/service/auth/auth.py +++ b/src/service/auth/auth.py @@ -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() \ No newline at end of file +def auth_service()->CurrentUserService: + return CurrentUserService() \ No newline at end of file diff --git a/src/service/auth/jwt.py b/src/service/auth/jwt.py index fdecdfd..4767275 100644 --- a/src/service/auth/jwt.py +++ b/src/service/auth/jwt.py @@ -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 diff --git a/src/web/protected_routes/routes.py b/src/web/protected_routes/auth_routes.py similarity index 74% rename from src/web/protected_routes/routes.py rename to src/web/protected_routes/auth_routes.py index e2c7789..57299a9 100644 --- a/src/web/protected_routes/routes.py +++ b/src/web/protected_routes/auth_routes.py @@ -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) diff --git a/src/web/protected_routes/protected_user_action_routes.py b/src/web/protected_routes/protected_user_action_routes.py index ee1ddb2..5d315a6 100644 --- a/src/web/protected_routes/protected_user_action_routes.py +++ b/src/web/protected_routes/protected_user_action_routes.py @@ -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") diff --git a/tests/integrated/test_auth.py b/tests/integrated/test_auth.py index fd75c30..6c3c576 100644 --- a/tests/integrated/test_auth.py +++ b/tests/integrated/test_auth.py @@ -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