refresh tokens 0.1.3

This commit is contained in:
2026-07-23 21:44:35 +03:00
parent eab78b6679
commit be8eb0c485
7 changed files with 78 additions and 28 deletions

View File

@@ -1,9 +1,10 @@
from uuid import UUID from uuid import UUID
from src.models.database_models.model import Model, engine, RefreshTokens from src.models.database_models.model import engine, RefreshTokens
from sqlalchemy import and_, not_, select from sqlalchemy import and_, not_, select
from sqlalchemy.orm import sessionmaker from sqlalchemy.orm import sessionmaker
from src.models.pydantic_models.model import RefreshTokensCreate, RefreshTokensOut from src.models.pydantic_models.model import RefreshTokensOut
class JwtCrudActions: class JwtCrudActions:
def __init__(self) -> None: def __init__(self) -> None:
self.Session=sessionmaker(bind=engine) self.Session=sessionmaker(bind=engine)

View File

@@ -0,0 +1,32 @@
"""empty message
Revision ID: 74814eb1b7f8
Revises: 8c136ff14180
Create Date: 2026-07-23 21:06:37.254211
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
# revision identifiers, used by Alembic.
revision: str = '74814eb1b7f8'
down_revision: Union[str, Sequence[str], None] = '8c136ff14180'
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 ###

View File

@@ -11,9 +11,9 @@ class Stored(Model):
id:Mapped[int]=mapped_column(primary_key=True, index=True) id:Mapped[int]=mapped_column(primary_key=True, index=True)
doc_id:Mapped[UUID]=mapped_column(default=uuid1, unique=True) doc_id:Mapped[UUID]=mapped_column(default=uuid1, unique=True)
filename:Mapped[str]=mapped_column(String(255), index=True) filename:Mapped[str]=mapped_column(String(255), index=True)
created_at:Mapped[datetime]=mapped_column(TIMESTAMP, server_default=func.now()) created_at:Mapped[datetime]=mapped_column(TIMESTAMP(timezone=True), server_default=func.now())
uploaded_at:Mapped[datetime]=mapped_column(TIMESTAMP,onupdate=func.now(), nullable=True) uploaded_at:Mapped[datetime]=mapped_column(TIMESTAMP(timezone=True),onupdate=func.now(), nullable=True)
doc_date:Mapped[datetime]=mapped_column(TIMESTAMP) doc_date:Mapped[datetime]=mapped_column(TIMESTAMP(timezone=True))
status:Mapped[str]=mapped_column(String(64)) status:Mapped[str]=mapped_column(String(64))
user_id:Mapped[UUID]=mapped_column(ForeignKey("users.id", ondelete="CASCADE"), index=True) user_id:Mapped[UUID]=mapped_column(ForeignKey("users.id", ondelete="CASCADE"), index=True)
market_id:Mapped[int]=mapped_column(ForeignKey("markets.id", ondelete="CASCADE"),index=True) market_id:Mapped[int]=mapped_column(ForeignKey("markets.id", ondelete="CASCADE"),index=True)

View File

@@ -92,8 +92,8 @@ class RefreshTokens(Model):
ip_address:Mapped[str]=mapped_column(String(45)) ip_address:Mapped[str]=mapped_column(String(45))
is_revoked:Mapped[bool]=mapped_column(Boolean, default=False) is_revoked:Mapped[bool]=mapped_column(Boolean, default=False)
expires_at:Mapped[datetime]=mapped_column(TIMESTAMP) expires_at:Mapped[datetime]=mapped_column(TIMESTAMP(timezone=True))
created_at:Mapped[datetime]=mapped_column(TIMESTAMP, server_default=func.now()) created_at:Mapped[datetime]=mapped_column(TIMESTAMP(timezone=True), server_default=func.now())
replaced_by:Mapped[UUID|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)

View File

@@ -8,6 +8,8 @@ from src.database.auth.refresh_tokens import JwtCrudActions
from src.errors.http_errors.errors import Errors from src.errors.http_errors.errors import Errors
from src.models.pydantic_models.model import RefreshTokensCreate, UserOut from src.models.pydantic_models.model import RefreshTokensCreate, UserOut
from src.models.configs_read.env import env_settings from src.models.configs_read.env import env_settings
class CurrentUser: class CurrentUser:
def __init__(self) -> None: def __init__(self) -> None:
@@ -52,19 +54,14 @@ class CurrentUser:
def create_access_token(self, form_data_email:str, form_data_password:str)->str: def create_access_token(self, user_id:UUID)->str:
'''check user info'''
user = self._check(form_data_email, form_data_password)
'''create new access token if all the checks are successful''' '''create new access token if all the checks are successful'''
return self.jwt_service.create_access_token({"sub":str(user.id)}) 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: def create_refresh_token(self,user_id:UUID, request:Request)->str:
'''check user info''' token, jti=self.jwt_service.create_refresh_token({"sub":str(user_id)})
user=self._check(form_data_email, form_data_password)
token, jti=self.jwt_service.create_refresh_token({"sub":str(user.id)})
try: try:
jti=UUID(jti) jti=UUID(jti)
@@ -74,7 +71,7 @@ class CurrentUser:
'''create new refresh token if all the checks are successful''' '''create new refresh token if all the checks are successful'''
token_record=RefreshTokensCreate( token_record=RefreshTokensCreate(
id=jti, id=jti,
user_id=user.id, user_id=user_id,
token_hash=self.hash.token_to_hash(token), token_hash=self.hash.token_to_hash(token),
device_info=request.headers.get("user-agent", "unknown"), 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"), ip_address=request.headers.get("x-forwarded-for", "").split(",")[0].strip() or (request.client.host if request.client else "unknown"),
@@ -85,7 +82,7 @@ class CurrentUser:
return token return token
def refresh_token(self, refresh_token:str, expires_delta:timedelta, request:Request)->str: def refresh_token(self, refresh_token:str, request:Request)->tuple[str, str]:
'''decode old refresh token''' '''decode old refresh token'''
old_refresh_token=self.jwt_service.jwt_decode(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: if (sub:=old_refresh_token.get("sub")) is None or (old_jti:=old_refresh_token.get("jti")) is None:
@@ -105,7 +102,13 @@ class CurrentUser:
if old_record.is_revoked: if old_record.is_revoked:
self.jwt_db_actions.revoke_all(old_record.user_id) self.jwt_db_actions.revoke_all(old_record.user_id)
raise self.error.credentials_error(detail="Reuse token detected") raise self.error.credentials_error(detail="Reuse token detected")
if old_record.expires_at<datetime.now(timezone.utc):
'''sqlite constraints about timezone'''
expires_at=old_record.expires_at
if expires_at.tzinfo is None:
expires_at = expires_at.replace(tzinfo=timezone.utc)
if expires_at<datetime.now(timezone.utc):
raise self.error.credentials_error(detail="Token expired") raise self.error.credentials_error(detail="Token expired")
'''user check''' '''user check'''
@@ -118,7 +121,7 @@ class CurrentUser:
'''create new refresh token if all the checks are successful''' '''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_refresh_token, new_jti=self.jwt_service.create_refresh_token({"sub":str(sub)})
new_access_token=self.create_access_token(user_id=sub)
try: try:
new_jti=UUID(new_jti) new_jti=UUID(new_jti)
except (ValueError, TypeError): except (ValueError, TypeError):
@@ -132,14 +135,14 @@ class CurrentUser:
token_hash=self.hash.token_to_hash(new_refresh_token), token_hash=self.hash.token_to_hash(new_refresh_token),
device_info=request.headers.get("user-agent", "unknown"), 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"), 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)+expires_delta, expires_at=datetime.now(timezone.utc)+timedelta(days=env_settings.REFRESH_TOKEN_EXPIRE_DAYS),
) )
self.jwt_db_actions.create_token(RefreshTokensCreate.model_dump(new_token_record)) self.jwt_db_actions.create_token(RefreshTokensCreate.model_dump(new_token_record))
'''update old token to deactivate it and assign replaced_by''' '''update old token to deactivate it and assign replaced_by'''
self.jwt_db_actions.update_token(old_jti, new_jti) self.jwt_db_actions.update_token(old_jti, new_jti)
return new_refresh_token return (new_access_token,new_refresh_token)
@@ -165,10 +168,13 @@ class CurrentUser:
def login(self, form_data_email:str, form_data_password:str, request:Request)->tuple[str, str]: 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)
'''create access and refresh tokens''' '''create access and refresh tokens'''
access_token=self.create_access_token(form_data_email=form_data_email, form_data_password=form_data_password) access_token=self.create_access_token(user_id=user.id)
refresh_token=self.create_refresh_token(form_data_password=form_data_password, form_data_email=form_data_email,request=request) refresh_token=self.create_refresh_token(user_id=user.id,request=request)
return (access_token, refresh_token) return (access_token, refresh_token)

View File

@@ -34,7 +34,7 @@ class Jwt:
user_info=data.copy() 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"}) "token_type":"access"})
print(f"DEBUG: expires at {datetime.now(timezone.utc)+timedelta(minutes=env_settings.ACCESS_TOKEN_EXPIRE_MINUTES)}, minutes={env_settings.ACCESS_TOKEN_EXPIRE_MINUTES}")
return jwt.encode(user_info, env_settings.SECRET_KEY, env_settings.ALGORITHM) return jwt.encode(user_info, env_settings.SECRET_KEY, env_settings.ALGORITHM)

View File

@@ -1,4 +1,3 @@
from datetime import timedelta
from fastapi import APIRouter, Depends, Request, Response, Cookie from fastapi import APIRouter, Depends, Request, Response, Cookie
from fastapi.security import OAuth2PasswordRequestForm, OAuth2PasswordBearer from fastapi.security import OAuth2PasswordRequestForm, OAuth2PasswordBearer
from src.models.configs_read.env import env_settings from src.models.configs_read.env import env_settings
@@ -26,8 +25,20 @@ async def get_access_token(request: Request,response:Response, form_data:OAuth2P
@router.post("/refresh") @router.post("/refresh")
async def get_refresh_token(request:Request, refresh_token: str = Cookie())->dict: async def get_refresh_token(request:Request,response:Response, 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"}
access_token, refresh_token= auth.refresh_token(refresh_token=refresh_token,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"}
async def get_current_user(token:str = Depends(oauth2_schema)) -> UserOut: async def get_current_user(token:str = Depends(oauth2_schema)) -> UserOut: