🔒 Made session-tokens more secure
This commit is contained in:
+6
-6
@@ -43,12 +43,12 @@ class OAuth2PasswordBearerWithCookie(OAuth2):
|
|||||||
super().__init__(flows=flows, scheme_name=scheme_name, auto_error=auto_error)
|
super().__init__(flows=flows, scheme_name=scheme_name, auto_error=auto_error)
|
||||||
|
|
||||||
async def __call__(self, request: Request) -> Optional[str]:
|
async def __call__(self, request: Request) -> Optional[str]:
|
||||||
authorization: str = request.cookies.get("access_token") # changed to accept access token from httpOnly Cookie
|
try:
|
||||||
if authorization is None:
|
authorization = request.state.access_token
|
||||||
try:
|
except AttributeError:
|
||||||
authorization = request.state.access_token
|
authorization: str = request.cookies.get(
|
||||||
except AttributeError:
|
"access_token"
|
||||||
pass
|
) # changed to accept access token from httpOnly Cookie
|
||||||
scheme, param = get_authorization_scheme_param(authorization)
|
scheme, param = get_authorization_scheme_param(authorization)
|
||||||
if not authorization or scheme.lower() != "bearer":
|
if not authorization or scheme.lower() != "bearer":
|
||||||
if self.auto_error:
|
if self.auto_error:
|
||||||
|
|||||||
@@ -4,11 +4,15 @@
|
|||||||
from datetime import timedelta, datetime
|
from datetime import timedelta, datetime
|
||||||
|
|
||||||
from fastapi import APIRouter, Request, Response
|
from fastapi import APIRouter, Request, Response
|
||||||
|
from fastapi.security.utils import get_authorization_scheme_param
|
||||||
|
from jose import jws, jwt, JWTError, JWSError
|
||||||
|
|
||||||
from classquiz.auth import ACCESS_TOKEN_EXPIRE_MINUTES, create_access_token
|
from classquiz.auth import ACCESS_TOKEN_EXPIRE_MINUTES, create_access_token
|
||||||
from classquiz.db.models import UserSession
|
from classquiz.db.models import UserSession
|
||||||
from classquiz.oauth import google, github
|
from classquiz.oauth import google, github
|
||||||
|
from classquiz.config import settings
|
||||||
|
|
||||||
|
settings = settings()
|
||||||
|
|
||||||
router = APIRouter()
|
router = APIRouter()
|
||||||
router.include_router(google.router, prefix="/google")
|
router.include_router(google.router, prefix="/google")
|
||||||
@@ -18,7 +22,45 @@ router.include_router(github.router, prefix="/github")
|
|||||||
async def rememberme_middleware(request: Request, call_next):
|
async def rememberme_middleware(request: Request, call_next):
|
||||||
rememberme_cookie = request.cookies.get("rememberme_token")
|
rememberme_cookie = request.cookies.get("rememberme_token")
|
||||||
bearer_token = request.cookies.get("access_token")
|
bearer_token = request.cookies.get("access_token")
|
||||||
if rememberme_cookie is not None and bearer_token is None:
|
conditions_to_handle_met = True
|
||||||
|
# print(bearer_token)
|
||||||
|
# if bearer_token is not None:
|
||||||
|
# bearer_token = bearer_token.replace("Bearer ", "")
|
||||||
|
# test = jws.verify(bearer_token, settings.secret_key, algorithms=["HS256"])
|
||||||
|
# try:
|
||||||
|
# jwt.decode(bearer_token, settings.secret_key, algorithms=["HS256"])
|
||||||
|
# print("jwt ok")
|
||||||
|
# except JWTError as e:
|
||||||
|
# print("jwt failed")
|
||||||
|
# print(test)
|
||||||
|
|
||||||
|
scheme, param = get_authorization_scheme_param(bearer_token)
|
||||||
|
|
||||||
|
# if bearer token is none, we can just do the request, since you can't be signed in
|
||||||
|
if scheme is None or param is None:
|
||||||
|
conditions_to_handle_met = False
|
||||||
|
# if rememberme token is none, we can just do the request, since you can't be signed in
|
||||||
|
if rememberme_cookie is None:
|
||||||
|
conditions_to_handle_met = False
|
||||||
|
|
||||||
|
if scheme.lower() != "bearer":
|
||||||
|
conditions_to_handle_met = False
|
||||||
|
|
||||||
|
# Verifying the bearer
|
||||||
|
try:
|
||||||
|
jwt.decode(
|
||||||
|
param, settings.secret_key, algorithms=["HS256"]
|
||||||
|
) # checking if the token is valid, throws error if not
|
||||||
|
conditions_to_handle_met = False
|
||||||
|
except JWTError:
|
||||||
|
try:
|
||||||
|
jws.verify(
|
||||||
|
param, settings.secret_key, algorithms=["HS256"]
|
||||||
|
) # Verifying only the signature of the jwt, throws error if signature is invalid
|
||||||
|
except JWSError:
|
||||||
|
conditions_to_handle_met = False
|
||||||
|
|
||||||
|
if conditions_to_handle_met:
|
||||||
user_session: UserSession | None = (
|
user_session: UserSession | None = (
|
||||||
await UserSession.objects.filter(session_key=rememberme_cookie)
|
await UserSession.objects.filter(session_key=rememberme_cookie)
|
||||||
.select_related(UserSession.user)
|
.select_related(UserSession.user)
|
||||||
@@ -27,17 +69,19 @@ async def rememberme_middleware(request: Request, call_next):
|
|||||||
if (user_session is None) or (user_session.user is None):
|
if (user_session is None) or (user_session.user is None):
|
||||||
response: Response = await call_next(request)
|
response: Response = await call_next(request)
|
||||||
return response
|
return response
|
||||||
access_token_expires = timedelta(minutes=ACCESS_TOKEN_EXPIRE_MINUTES * 60)
|
access_token_expires = timedelta(minutes=ACCESS_TOKEN_EXPIRE_MINUTES)
|
||||||
|
# access_token_expires = timedelta(seconds=1)
|
||||||
access_token = create_access_token(data={"sub": user_session.user.email}, expires_delta=access_token_expires)
|
access_token = create_access_token(data={"sub": user_session.user.email}, expires_delta=access_token_expires)
|
||||||
await user_session.update(last_seen=datetime.now())
|
await user_session.update(last_seen=datetime.now())
|
||||||
request.state.access_token = f"Bearer {access_token}"
|
request.state.access_token = f"Bearer {access_token}"
|
||||||
|
request.cookies.pop("access_token")
|
||||||
response: Response = await call_next(request)
|
response: Response = await call_next(request)
|
||||||
response.set_cookie(
|
response.set_cookie(
|
||||||
key="access_token",
|
key="access_token",
|
||||||
value=f"Bearer {access_token}",
|
value=f"Bearer {access_token}",
|
||||||
httponly=True,
|
httponly=True,
|
||||||
samesite="lax",
|
samesite="lax",
|
||||||
max_age=ACCESS_TOKEN_EXPIRE_MINUTES * 60,
|
max_age=60 * 60 * 24 * 365,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
response: Response = await call_next(request)
|
response: Response = await call_next(request)
|
||||||
|
|||||||
@@ -43,7 +43,7 @@ async def log_user_in(user: User, request: Request, response: Response):
|
|||||||
value=f"Bearer {access_token}",
|
value=f"Bearer {access_token}",
|
||||||
httponly=True,
|
httponly=True,
|
||||||
samesite="lax",
|
samesite="lax",
|
||||||
max_age=settings.access_token_expire_minutes * 60,
|
max_age=60 * 60 * 24 * 365,
|
||||||
)
|
)
|
||||||
response.set_cookie(
|
response.set_cookie(
|
||||||
key="rememberme_token", value=session_key, httponly=True, samesite="lax", max_age=60 * 60 * 24 * 365
|
key="rememberme_token", value=session_key, httponly=True, samesite="lax", max_age=60 * 60 * 24 * 365
|
||||||
@@ -64,7 +64,7 @@ async def rememberme_check(rememberme_token: str, response: Response):
|
|||||||
value=f"Bearer {access_token}",
|
value=f"Bearer {access_token}",
|
||||||
httponly=True,
|
httponly=True,
|
||||||
samesite="lax",
|
samesite="lax",
|
||||||
max_age=settings.access_token_expire_minutes * 60,
|
max_age=60 * 60 * 24 * 365,
|
||||||
)
|
)
|
||||||
response.set_cookie(key="expiry", value="", max_age=settings.access_token_expire_minutes * 60)
|
response.set_cookie(key="expiry", value="", max_age=settings.access_token_expire_minutes * 60)
|
||||||
response.status_code = 200
|
response.status_code = 200
|
||||||
|
|||||||
Reference in New Issue
Block a user