From eab78b66798403bc49fb156340540cce315d896f Mon Sep 17 00:00:00 2001 From: "MH.Dmitrii" Date: Thu, 23 Jul 2026 20:51:26 +0300 Subject: [PATCH] refresh tokens 0.1.2 --- src/database/auth/refresh_tokens.py | 63 ++++++++ src/migrations/versions/385d4efec15f_.py | 50 +++++++ src/migrations/versions/8c136ff14180_.py | 32 ++++ src/models/database_models/model.py | 5 +- src/models/pydantic_models/model.py | 13 +- src/service/auth/auth.py | 182 +++++++++++++++++++---- src/service/auth/jwt.py | 34 ++++- src/web/protected_routes/routes.py | 36 ++++- 8 files changed, 369 insertions(+), 46 deletions(-) create mode 100644 src/database/auth/refresh_tokens.py create mode 100644 src/migrations/versions/385d4efec15f_.py create mode 100644 src/migrations/versions/8c136ff14180_.py diff --git a/src/database/auth/refresh_tokens.py b/src/database/auth/refresh_tokens.py new file mode 100644 index 0000000..330b0e0 --- /dev/null +++ b/src/database/auth/refresh_tokens.py @@ -0,0 +1,63 @@ + +from uuid import UUID +from src.models.database_models.model import Model, engine, RefreshTokens +from sqlalchemy import and_, not_, select +from sqlalchemy.orm import sessionmaker +from src.models.pydantic_models.model import RefreshTokensCreate, RefreshTokensOut +class JwtCrudActions: + def __init__(self) -> None: + self.Session=sessionmaker(bind=engine) + + def get_token_by_user_id(self, user_id:UUID)->RefreshTokensOut|None: + with self.Session() as session: + 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) + + def get_token_by_id(self, id:UUID)->RefreshTokensOut|None: + with self.Session() as session: + 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 create_token(self, data:dict)->None: + with self.Session() as session: + with session.begin(): + new_token=RefreshTokens(**data) + response=session.add(new_token) + return response + + def update_token(self, old_jti:UUID, new_jti:UUID)->bool: + with self.Session() as session: + with session.begin(): + query=select(RefreshTokens).where(RefreshTokens.id==old_jti) + response=session.scalars(query).one() + response.is_revoked=True + response.replaced_by=new_jti + return True + + def revoke_all(self, user_id:UUID)->bool: + with self.Session() as session: + 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: + 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 diff --git a/src/migrations/versions/385d4efec15f_.py b/src/migrations/versions/385d4efec15f_.py new file mode 100644 index 0000000..cd11529 --- /dev/null +++ b/src/migrations/versions/385d4efec15f_.py @@ -0,0 +1,50 @@ +"""empty message + +Revision ID: 385d4efec15f +Revises: 23dd6d3efe4b +Create Date: 2026-07-23 15:03:33.582896 + +""" +from typing import Sequence, Union + +from alembic import op +import sqlalchemy as sa + + +# revision identifiers, used by Alembic. +revision: str = '385d4efec15f' +down_revision: Union[str, Sequence[str], None] = '23dd6d3efe4b' +branch_labels: Union[str, Sequence[str], None] = None +depends_on: Union[str, Sequence[str], None] = None + + +def upgrade() -> None: + """Upgrade schema.""" + # ### commands auto generated by Alembic - please adjust! ### + with op.batch_alter_table('refresh_tokens', schema=None) as batch_op: + batch_op.alter_column('id', + existing_type=sa.INTEGER(), + type_=sa.Uuid(), + existing_nullable=False) + batch_op.alter_column('replaced_by', + existing_type=sa.INTEGER(), + type_=sa.Uuid(), + existing_nullable=True) + + # ### end Alembic commands ### + + +def downgrade() -> None: + """Downgrade schema.""" + # ### commands auto generated by Alembic - please adjust! ### + with op.batch_alter_table('refresh_tokens', schema=None) as batch_op: + batch_op.alter_column('replaced_by', + existing_type=sa.Uuid(), + type_=sa.INTEGER(), + existing_nullable=True) + batch_op.alter_column('id', + existing_type=sa.Uuid(), + type_=sa.INTEGER(), + existing_nullable=False) + + # ### end Alembic commands ### diff --git a/src/migrations/versions/8c136ff14180_.py b/src/migrations/versions/8c136ff14180_.py new file mode 100644 index 0000000..3154207 --- /dev/null +++ b/src/migrations/versions/8c136ff14180_.py @@ -0,0 +1,32 @@ +"""empty message + +Revision ID: 8c136ff14180 +Revises: 385d4efec15f +Create Date: 2026-07-23 17:56:26.345954 + +""" +from typing import Sequence, Union + +from alembic import op +import sqlalchemy as sa + + +# revision identifiers, used by Alembic. +revision: str = '8c136ff14180' +down_revision: Union[str, Sequence[str], None] = '385d4efec15f' +branch_labels: Union[str, Sequence[str], None] = None +depends_on: Union[str, Sequence[str], None] = None + + +def upgrade() -> None: + """Upgrade schema.""" + # ### commands auto generated by Alembic - please adjust! ### + pass + # ### end Alembic commands ### + + +def downgrade() -> None: + """Downgrade schema.""" + # ### commands auto generated by Alembic - please adjust! ### + pass + # ### end Alembic commands ### diff --git a/src/models/database_models/model.py b/src/models/database_models/model.py index ffa5337..85e4eb1 100644 --- a/src/models/database_models/model.py +++ b/src/models/database_models/model.py @@ -31,6 +31,7 @@ class User(Model): refresh_token:Mapped[list['RefreshTokens']]=relationship(back_populates="user") report:Mapped[list["Stored"]]=relationship(back_populates="user") + def __repr__(self) -> str: return f"ID: {self.id}, Name: {self.first_name}, Status: {self.status}" @@ -84,7 +85,7 @@ class RefreshTokens(Model): __tablename__= "refresh_tokens" - id:Mapped[int]=mapped_column(primary_key=True, index=True) + id:Mapped[UUID]=mapped_column(Uuid(as_uuid=True),primary_key=True, index=True) user_id:Mapped[UUID]=mapped_column(ForeignKey("users.id", ondelete="CASCADE"), index=True) token_hash:Mapped[str]=mapped_column(String(255), unique=True) device_info:Mapped[str]=mapped_column(String(255)) @@ -93,7 +94,7 @@ class RefreshTokens(Model): expires_at:Mapped[datetime]=mapped_column(TIMESTAMP) created_at:Mapped[datetime]=mapped_column(TIMESTAMP, server_default=func.now()) - replaced_by:Mapped[int|None]=mapped_column(ForeignKey("refresh_tokens.id"), nullable=True, default=None) + replaced_by:Mapped[UUID|None]=mapped_column(ForeignKey("refresh_tokens.id"), nullable=True, default=None) user:Mapped["User"]=relationship(back_populates="refresh_token") diff --git a/src/models/pydantic_models/model.py b/src/models/pydantic_models/model.py index 34bf301..057d55b 100644 --- a/src/models/pydantic_models/model.py +++ b/src/models/pydantic_models/model.py @@ -8,6 +8,7 @@ class Base(BaseModel): model_config = {"from_attributes": True} + class PermissionsCreate(Base): permission:Annotated[str, Field(..., max_length=30, description="permission name")] @@ -67,14 +68,18 @@ class UserUpdate(Base): class RefreshTokensCreate(Base): - + + id:Annotated[UUID, Field(..., description="jti")] user_id:Annotated[UUID, Field(..., description="foreign key for the user")] token_hash:Annotated[str, Field(...,max_length=255, description="token hash")] device_info:Annotated[str, Field(...,max_length=255, description="User device info")] ip_address:Annotated[str, Field(...,max_length=45, description="ip v4/v6 of the user")] - is_revoked:Annotated[bool|None, Field(None, description="revoke token if logout was made")] expires_at:Annotated[datetime, Field(..., description="when token is going to be expired")] +class RefreshTokensUpdate(Base): + is_revoked:Annotated[bool, Field(..., description="revoke token if logout was made")] + replaced_by:Annotated[UUID, Field(...,description="old_jti")] + class RefreshTokensOut(Base): user_id:Annotated[UUID, Field(..., description="foreign key for the user")] @@ -83,3 +88,7 @@ class RefreshTokensOut(Base): ip_address:Annotated[str, Field(...,max_length=45, description="ip v4/v6 of the user")] is_revoked:Annotated[bool|None, Field(None, description="revoke token if logout was made")] expires_at:Annotated[datetime, Field(..., description="when token is going to be expired")] + replaced_by:Annotated[UUID|None, Field(..., description="Old refresh token")] + +class RefreshRequest(Base): + refresh_token:str diff --git a/src/service/auth/auth.py b/src/service/auth/auth.py index cad1560..b57fb91 100644 --- a/src/service/auth/auth.py +++ b/src/service/auth/auth.py @@ -1,42 +1,25 @@ +from datetime import datetime, timedelta, timezone from uuid import UUID + +from fastapi import Request from .jwt import Jwt, Hashes from src.database.users.crud import UsersCrudActions +from src.database.auth.refresh_tokens import JwtCrudActions from src.errors.http_errors.errors import Errors - -from src.models.pydantic_models.model import UserOut +from src.models.pydantic_models.model import RefreshTokensCreate, UserOut +from src.models.configs_read.env import env_settings class CurrentUser: def __init__(self) -> None: self.jwt_service=Jwt() self.hash=Hashes() - self.crud_actions=UsersCrudActions() + self.crud_db_actions=UsersCrudActions() + self.jwt_db_actions=JwtCrudActions() self.error=Errors() - - def get_current_user(self, token:str)->UserOut: - - payload=self.jwt_service.jwt_decode(token) - sub=payload.get("sub") - print(sub) - if sub is None: - raise self.error.credentials_error(detail="Jwt token is incorrect") - try: - UUID(sub) - except (ValueError, TypeError): - raise self.error.credentials_error(detail="Jwt token is incorrect") - - user=self.crud_actions.get_user_by_id(UUID(sub)) - - if user is None: - raise self.error.not_found_error(detail="User with this email address not found") - - return UserOut.model_validate(user) - - - def create_token(self, form_data_email:str, form_data_password:str)->dict: - - user=self.crud_actions.get_user_by_email(form_data_email) - + 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) if user is None: raise self.error.credentials_error(detail="Wrong credentials") @@ -45,7 +28,148 @@ class CurrentUser: if user.status is False: raise self.error.credentials_error(detail="This user is deactivated") + return user + + + def get_current_user(self, token:str)->UserOut: - return {"access_token":self.jwt_service.create_token({"sub":str(user.id)}), "token_type":"bearer"} + payload=self.jwt_service.jwt_decode(token) + + if (sub:=payload.get("sub")) is None: + raise self.error.credentials_error(detail="Jwt token is incorrect") + + try: + sub=UUID(sub) + except (ValueError, TypeError): + raise self.error.credentials_error(detail="Jwt token is incorrect") + + + user=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") + + return UserOut.model_validate(user) + + + + def create_access_token(self, form_data_email:str, form_data_password:str)->str: + '''check user info''' + user = self._check(form_data_email, form_data_password) + + '''create new access token if all the checks are successful''' + return self.jwt_service.create_access_token({"sub":str(user.id)}) + + + + def create_refresh_token(self,form_data_email:str, form_data_password:str, request:Request)->str: + '''check user info''' + user=self._check(form_data_email, form_data_password) + token, jti=self.jwt_service.create_refresh_token({"sub":str(user.id)}) + + try: + jti=UUID(jti) + except (ValueError, TypeError): + raise self.error.credentials_error(detail="Jwt token is incorrect") + + '''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(timezone.utc)+timedelta(days=env_settings.REFRESH_TOKEN_EXPIRE_DAYS) + ) + self.jwt_db_actions.create_token(RefreshTokensCreate.model_dump(token_record)) + + return token + + + def refresh_token(self, refresh_token:str, expires_delta:timedelta, request:Request)->str: + '''decode old refresh token''' + old_refresh_token=self.jwt_service.jwt_decode(refresh_token) + if (sub:=old_refresh_token.get("sub")) is None or (old_jti:=old_refresh_token.get("jti")) is None: + raise self.error.credentials_error(detail="Jwt token is incorrect") + + try: + old_jti=UUID(old_jti) + sub=UUID(sub) + except (ValueError, TypeError): + raise self.error.credentials_error(detail="Jwt token is incorrect") + + + '''old refresh token check''' + 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") + if old_record.expires_atbool: + + '''decode current refresh token''' + payload=self.jwt_service.jwt_decode(refresh_token) + + if (jti:=payload.get("jti")) is None: + raise self.error.credentials_error(detail="Invalid Refresh Token") + try: + jti=UUID(jti) + except (ValueError, TypeError): + raise self.error.credentials_error(detail="Jwt token is incorrect") + + current_token = self.jwt_db_actions.get_token_by_id(jti) + if current_token is None: + raise self.error.not_found_error(detail="Refresh Token Not Found") + + '''logout by assigning revoked flag''' + return self.jwt_db_actions.logout(jti) + + + + def login(self, form_data_email:str, form_data_password:str, request:Request)->tuple[str, str]: + + '''create access and refresh tokens''' + access_token=self.create_access_token(form_data_email=form_data_email, form_data_password=form_data_password) + refresh_token=self.create_refresh_token(form_data_password=form_data_password, form_data_email=form_data_email,request=request) + + return (access_token, refresh_token) auth=CurrentUser() \ No newline at end of file diff --git a/src/service/auth/jwt.py b/src/service/auth/jwt.py index 850bc61..b91fbbb 100644 --- a/src/service/auth/jwt.py +++ b/src/service/auth/jwt.py @@ -3,6 +3,9 @@ import bcrypt from src.errors.http_errors.errors import Errors from datetime import datetime, timedelta, timezone from src.models.configs_read.env import env_settings +from uuid import uuid4 +import hashlib + '''Hash/Check hash''' class Hashes: @@ -15,32 +18,49 @@ class Hashes: def verify_password(self, plain_password:str, hashed_password:str)->bool: return bcrypt.checkpw(plain_password.encode("utf-8"), hashed_password.encode("utf-8")) + def token_to_hash(self, token:str)->str: + return hashlib.sha256(token.encode("utf-8")).hexdigest() + '''jwt''' class Jwt: - def __init__(self) -> None: - self.error=Errors() - def create_token(self, data:dict)->str: + def __init__(self) -> None: + + self.error=Errors() + + def create_access_token(self, data:dict)->str: user_info=data.copy() - user_info.update({"exp": datetime.now(timezone.utc)+timedelta(minutes=env_settings.ACCESS_TOKEN_EXPIRE_MINUTES)}) + user_info.update({"exp": datetime.now(timezone.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]: + user_info=data.copy() + jti=str(uuid4()) + user_info.update({"exp":datetime.now(timezone.utc)+timedelta(days=env_settings.REFRESH_TOKEN_EXPIRE_DAYS), + "token_type":"refresh", + "jti":jti + }) + return jwt.encode(user_info, env_settings.SECRET_KEY, env_settings.ALGORITHM), jti + + + def jwt_decode(self, token:str)->dict: try: - payload=jwt.decode(token, env_settings.SECRET_KEY, algorithms=[env_settings.ALGORITHM]) if (payload.get("sub")) is None: raise self.error.credentials_error(detail="Sub block is missing") except JWTError as e: - raise self.error.credentials_error(detail="JWTerror") from e return payload - + diff --git a/src/web/protected_routes/routes.py b/src/web/protected_routes/routes.py index 9720495..1989c70 100644 --- a/src/web/protected_routes/routes.py +++ b/src/web/protected_routes/routes.py @@ -1,20 +1,44 @@ -from fastapi import APIRouter, Depends +from datetime import timedelta +from fastapi import APIRouter, Depends, Request, Response, Cookie from fastapi.security import OAuth2PasswordRequestForm, OAuth2PasswordBearer - +from src.models.configs_read.env import env_settings from src.models.pydantic_models.model import UserOut from src.service.auth.auth import auth router=APIRouter(prefix="/protected") -oauth2_schema=OAuth2PasswordBearer(tokenUrl="/protected/token") +oauth2_schema=OAuth2PasswordBearer(tokenUrl="/protected/token", refreshUrl="/protected/refresh") @router.post("/token") -async def get_access_token(form_data:OAuth2PasswordRequestForm=Depends())->dict: - return auth.create_token(form_data_email=form_data.username, form_data_password=form_data.password) +async def get_access_token(request: Request,response:Response, form_data:OAuth2PasswordRequestForm=Depends())->dict: + + access_token, refresh_token=auth.login(form_data_email=form_data.username, form_data_password=form_data.password, request=request) + + response.set_cookie( + key="refresh_token", + value=refresh_token, + httponly=True, + secure=True, + samesite="strict", + max_age=env_settings.REFRESH_TOKEN_EXPIRE_DAYS * 24 * 60 * 60 + ) + return {"access_token": access_token, "token_type": "bearer"} -def get_current_user(token:str = Depends(oauth2_schema)) -> UserOut: + +@router.post("/refresh") +async def get_refresh_token(request:Request, refresh_token: str = Cookie())->dict: + return {"access_token":auth.refresh_token(refresh_token=refresh_token,request=request, expires_delta=timedelta(minutes=env_settings.ACCESS_TOKEN_EXPIRE_MINUTES)), "token_type": "bearer"} + + +async def get_current_user(token:str = Depends(oauth2_schema)) -> UserOut: 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: + response.delete_cookie("refresh_token") + return auth.logout(refresh_token) + @router.get("") async def protected(current_user:UserOut=Depends(get_current_user))->dict: return {"protected router": "Hello, this is a protected router"}