Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 0 additions & 1 deletion app/core/constant.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,6 @@

class RedisKey(str, Enum):
UserSession = "user_session"
UserSessionByUser = "user_session:{user_id}"
INVALID_TOKEN_SET_KEY = "notifications:invalid_tokens"
MobileSessionCache = "session:{session_id}"

Expand Down
20 changes: 13 additions & 7 deletions app/router/mobile/auth.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
from fastapi import APIRouter, Depends, Request, UploadFile
from fastapi.responses import Response
from app.core.image_validation import build_image_payload
from app.core.exceptions import AppException
from uuid import UUID

from app.container import get_container, Container
Expand Down Expand Up @@ -123,21 +124,26 @@ async def revoke_device(
container: Container = Depends(get_container),
current_user: MobileUserSchema = Depends(get_current_mobile_user),
) -> dict[str, str]:
from app.core.constant import RedisKey

session = await container.session_service.session_querier.get_session_by_device(
device_id=device_id
device = await container.device_service.get_device_by_id(
device_id=device_id, user_id=current_user.user_id
)
if session:
await container.session_service.delete_session_cache(container.redis, session.id)
if device is None or device.user_id != current_user.user_id:
raise AppException.not_found("Device not found")

user_session_key = RedisKey.UserSessionByUser.value.format(user_id=current_user.user_id)
await container.redis.delete(user_session_key)
session = await container.session_service.session_querier.get_session_by_device_for_user(
device_id=device_id, user_id=current_user.user_id
)

await container.device_service.revoke_device(
device_id=device_id,
user_id=current_user.user_id,
)

if session:
await container.session_service.delete_session_cache(container.redis, session.id)


return {"message": "Device revoked successfully"}


Expand Down
4 changes: 2 additions & 2 deletions app/service/device.py
Original file line number Diff line number Diff line change
Expand Up @@ -86,7 +86,7 @@ async def inactivate_device(
user_id: uuid.UUID,
) -> None:
try:
device = await self.device_querier.get_device_by_id(id=device_id)
device = await self.device_querier.get_device_by_id(id=device_id, user_id=user_id)
if device is None or device.user_id != user_id:
raise AppException.not_found("Device not found")
await self.device_querier.deactivate_device(
Expand All @@ -112,7 +112,7 @@ async def get_device_by_id(
user_id: uuid.UUID,
) -> UserDevice:
try :
device = await self.device_querier.get_device_by_id(id=device_id)
device = await self.device_querier.get_device_by_id(id=device_id, user_id=user_id)
if device is None :
raise AppException.not_found("device not found ")
return device
Expand Down
109 changes: 1 addition & 108 deletions app/service/session.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,18 +3,9 @@
from db.generated import session as session_queries
import uuid
from db.generated.models import UserSession
from datetime import datetime, timedelta, timezone
from datetime import datetime
from app.infra.redis import RedisClient
from app.core.constant import RedisKey
from db.generated.session import UpsertSessionRow


class SessionRedis(BaseModel):
session_id: uuid.UUID
user_id: uuid.UUID
device_id: uuid.UUID
last_active: datetime
expires_at: datetime


class MobileSessionCache(BaseModel):
Expand Down Expand Up @@ -74,35 +65,6 @@ async def delete_session_cache(
key = RedisKey.MobileSessionCache.value.format(session_id=session_id)
await redis.delete(key)

@staticmethod
async def create_session(user_id: uuid.UUID, device_id: uuid.UUID) -> UpsertSessionRow:
try:
session = await SessionService.session_querier.upsert_session(
user_id=user_id,
device_id=device_id,
expires_at=datetime.now(timezone.utc) + timedelta(days=7),
)
if session is None:
raise AppException.internal_error("session creation failed ")

result = await SessionService.redis.set(
key=RedisKey.UserSessionByUser.format(user_id=user_id),
value=SessionRedis(
session_id=session.id,
user_id=session.user_id,
device_id=session.device_id,
last_active=session.last_active,
expires_at=session.expires_at,
).model_dump_json(),
expire=60 * 60 * 5,
nx=True,
)
if not result:
AppException.forbidden("You already logged in in another device")
return session
except Exception as e:
raise DBExceptionImpl.handle(e)

@staticmethod
async def get_session_by_id(session_id: uuid.UUID) -> UserSession:
try:
Expand All @@ -113,75 +75,6 @@ async def get_session_by_id(session_id: uuid.UUID) -> UserSession:
except Exception as e:
raise DBExceptionImpl.handle(e)

@staticmethod
async def check_session(
session_id: uuid.UUID,
user_id: uuid.UUID,
device_id: uuid.UUID,
) -> bool:
try:
session_in_redis = await SessionService.redis.get(
RedisKey.UserSessionByUser.format(user_id=user_id)
)

if session_in_redis is None:
return False

session_info = SessionRedis.model_validate_json(session_in_redis)

if session_info:
if session_info.device_id != device_id and session_info.session_id != session_id:
raise AppException.forbidden("You already logged in on another device")

await SessionService.redis.set(
key=RedisKey.UserSessionByUser.format(user_id=user_id),
value=SessionRedis(
session_id=session_info.session_id,
user_id=session_info.user_id,
device_id=session_info.device_id,
last_active=session_info.last_active,
expires_at=session_info.expires_at,
).model_dump_json(),
expire=60 * 60 * 5,
nx=False,
)

return True

session = await SessionService.session_querier.get_session_by_id(id=session_id)

if session is None:
raise AppException.forbidden("Session not found")

await SessionService.redis.set(
key=RedisKey.UserSessionByUser.format(user_id=user_id),
value=SessionRedis(
session_id=session.id,
user_id=session.user_id,
device_id=session.device_id,
last_active=session.last_active,
expires_at=session.expires_at,
).model_dump_json(),
expire=60 * 60 * 5,
nx=True,
)

return True

except Exception as e:
raise DBExceptionImpl.handle(e)

@staticmethod
async def delete_session(
session_id: uuid.UUID, user_id: uuid.UUID, device_id: uuid.UUID
) -> None:
try:
await SessionService.session_querier.delete_session_by_device(
user_id=user_id, device_id=device_id
)
except Exception as e:
raise DBExceptionImpl.handle(e)

@staticmethod
async def delete_expired_sessions() -> None:
try:
Expand Down
46 changes: 16 additions & 30 deletions app/service/users.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,6 @@
decode_refresh_mobile_token,
Get_expiry_time,
)
from app.core import constant
from app.core.config import settings
from app.infra.redis import RedisClient
from app.infra.minio import Bucket, IMAGES_BUCKET_NAME
Expand Down Expand Up @@ -63,7 +62,7 @@ async def _ensure_device_for_login(
user_id: uuid.UUID,
req: MobileAuthBaseRequest,
) -> UserDevice:
existing_device = await self.device_querier.get_device_by_id(id=req.device_id)
existing_device = await self.device_querier.get_device_by_id_any(id=req.device_id)

if existing_device:
if existing_device.user_id != user_id:
Expand Down Expand Up @@ -275,8 +274,6 @@ async def _create_mobile_session(
) -> MobileAuthResponse:
user_id: uuid.UUID = user.id

session_key = constant.RedisKey.UserSessionByUser.value.format(user_id=user_id)

session_count = await self.session_querier.count_user_sessions(user_id=user_id)
if session_count and session_count >= AuthService.SESSION_LIMIT:
logger.warning(
Expand All @@ -302,10 +299,6 @@ async def _create_mobile_session(
if not session:
raise AppException.internal_error("Failed to create session")

await redis.set(
session_key, str(session.id), expire=AuthService.REDIS_SESSION_TTL
)

access_token = create_acces_mobile_token(str(session.id))
refresh_token = create_refresh_mobile_token(str(session.id))
expiry = Get_expiry_time()
Expand Down Expand Up @@ -373,8 +366,10 @@ async def logout(
user_id: str,
session_id: str,
) -> dict[str, str]:
session_key = constant.RedisKey.UserSessionByUser.value.format(user_id=user_id)
await redis.delete(session_key)
sid = uuid.UUID(session_id)
await SessionService.delete_session_cache(redis, sid)
await self.session_querier.delete_session_by_id(id=sid, user_id=uuid.UUID(user_id))

return {"message": "Logged out successfully"}

async def add_embbed_user(
Expand Down Expand Up @@ -566,18 +561,14 @@ async def delete_user(self, *, redis: RedisClient, user_id: uuid.UUID) -> User:
existing = await self.user_querier.get_user_by_id(id=user_id)
if not existing:
raise AppException.not_found("User not found")

sessions = self.session_querier.list_sessions_by_user(user_id=user_id)
async for s in sessions:
await SessionService.delete_session_cache(redis=redis, session_id=s.id)
await self.session_querier.delete_all_user_sessions(user_id=user_id)

await self.user_querier.delete_user(id=user_id)
session_key = constant.RedisKey.UserSessionByUser.value.format(
user_id=user_id
)
raw_session_id = await redis.get(session_key)
if raw_session_id:
try:
session_id = uuid.UUID(raw_session_id)
await SessionService.delete_session_cache(redis=redis, session_id=session_id)
except (ValueError, Exception):
pass
await redis.delete(session_key)

return existing
except Exception as exc:
logger.error("Failed to delete user: %s", exc)
Expand All @@ -589,15 +580,10 @@ async def block_user(self, *, redis: RedisClient, user_id: uuid.UUID) -> User:
if not user:
raise AppException.not_found("User not found")

session_key = constant.RedisKey.UserSessionByUser.value.format(user_id=user_id)
raw_session_id = await redis.get(session_key)
if raw_session_id:
try:
session_id = uuid.UUID(raw_session_id)
await SessionService.delete_session_cache(redis=redis, session_id=session_id)
except (ValueError, Exception):
pass
await redis.delete(session_key)
sessions = self.session_querier.list_sessions_by_user(user_id=user_id)
async for s in sessions:
await SessionService.delete_session_cache(redis, s.id)
await self.session_querier.delete_all_user_sessions(user_id=user_id)

return user
except Exception as exc:
Expand Down
31 changes: 28 additions & 3 deletions db/generated/devices.py
Original file line number Diff line number Diff line change
Expand Up @@ -69,7 +69,14 @@ class CreateDeviceParams:

GET_DEVICE_BY_ID = """-- name: get_device_by_id \\:one
SELECT id, user_id, device_name, device_type, totp_secret, is_2fa_enabled, last_active, created_at, push_token, is_active, is_invalid_token from user_devices
WHERE id =:p1
WHERE id = :p1
AND user_id = :p2
"""


GET_DEVICE_BY_ID_ANY = """-- name: get_device_by_id_any \\:one
SELECT id, user_id, device_name, device_type, totp_secret, is_2fa_enabled, last_active, created_at, push_token, is_active, is_invalid_token from user_devices
WHERE id = :p1
"""


Expand Down Expand Up @@ -159,8 +166,26 @@ async def deactivate_device(self, *, id: uuid.UUID, user_id: uuid.UUID) -> None:
async def enable_device2_fa(self, *, id: uuid.UUID, user_id: uuid.UUID) -> None:
await self._conn.execute(sqlalchemy.text(ENABLE_DEVICE2_FA), {"p1": id, "p2": user_id})

async def get_device_by_id(self, *, id: uuid.UUID) -> Optional[models.UserDevice]:
row = (await self._conn.execute(sqlalchemy.text(GET_DEVICE_BY_ID), {"p1": id})).first()
async def get_device_by_id(self, *, id: uuid.UUID, user_id: uuid.UUID) -> Optional[models.UserDevice]:
row = (await self._conn.execute(sqlalchemy.text(GET_DEVICE_BY_ID), {"p1": id, "p2": user_id})).first()
if row is None:
return None
return models.UserDevice(
id=row[0],
user_id=row[1],
device_name=row[2],
device_type=row[3],
totp_secret=row[4],
is_2fa_enabled=row[5],
last_active=row[6],
created_at=row[7],
push_token=row[8],
is_active=row[9],
is_invalid_token=row[10],
)

async def get_device_by_id_any(self, *, id: uuid.UUID) -> Optional[models.UserDevice]:
row = (await self._conn.execute(sqlalchemy.text(GET_DEVICE_BY_ID_ANY), {"p1": id})).first()
if row is None:
return None
return models.UserDevice(
Expand Down
17 changes: 13 additions & 4 deletions db/generated/session.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,10 +37,16 @@
"""


GET_SESSION_BY_DEVICE = """-- name: get_session_by_device \\:one
DELETE_SESSION_BY_ID = """-- name: delete_session_by_id \\:exec
DELETE FROM user_sessions
WHERE id = :p1 AND user_id = :p2
"""


GET_SESSION_BY_DEVICE_FOR_USER = """-- name: get_session_by_device_for_user \\:one
SELECT id, user_id, device_id, created_at, last_active, expires_at
FROM user_sessions
WHERE device_id = :p1
WHERE device_id = :p1 AND user_id = :p2
"""


Expand Down Expand Up @@ -116,8 +122,11 @@ async def delete_expired_sessions(self) -> None:
async def delete_session_by_device(self, *, device_id: uuid.UUID, user_id: uuid.UUID) -> None:
await self._conn.execute(sqlalchemy.text(DELETE_SESSION_BY_DEVICE), {"p1": device_id, "p2": user_id})

async def get_session_by_device(self, *, device_id: uuid.UUID) -> Optional[models.UserSession]:
row = (await self._conn.execute(sqlalchemy.text(GET_SESSION_BY_DEVICE), {"p1": device_id})).first()
async def delete_session_by_id(self, *, id: uuid.UUID, user_id: uuid.UUID) -> None:
await self._conn.execute(sqlalchemy.text(DELETE_SESSION_BY_ID), {"p1": id, "p2": user_id})

async def get_session_by_device_for_user(self, *, device_id: uuid.UUID, user_id: uuid.UUID) -> Optional[models.UserSession]:
row = (await self._conn.execute(sqlalchemy.text(GET_SESSION_BY_DEVICE_FOR_USER), {"p1": device_id, "p2": user_id})).first()
if row is None:
return None
return models.UserSession(
Expand Down
7 changes: 6 additions & 1 deletion db/queries/devices.sql
Original file line number Diff line number Diff line change
Expand Up @@ -34,9 +34,14 @@ WHERE id = $1
AND user_id = $2
AND is_2fa_enabled = FALSE;

-- name: GetDeviceByIdAny :one
SELECT * from user_devices
WHERE id = $1;

-- name: GetDeviceById :one
SELECT * from user_devices
WHERE id =$1;
WHERE id = $1
AND user_id = $2;

-- name: CountUserDevices :one
SELECT COUNT(*)
Expand Down
Loading