80 lines
2.2 KiB
Python
80 lines
2.2 KiB
Python
from datetime import UTC, datetime
|
|
from hashlib import sha256
|
|
from secrets import token_urlsafe
|
|
|
|
from sqlalchemy import select
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
|
|
from src.models.refresh_token import RefreshToken
|
|
|
|
|
|
def hash_refresh_token(raw_token: str) -> str:
|
|
return sha256(raw_token.encode("utf-8")).hexdigest()
|
|
|
|
|
|
async def create_refresh_token(
|
|
*,
|
|
session: AsyncSession,
|
|
user_id: object,
|
|
expires_at: datetime,
|
|
user_agent: str | None,
|
|
ip_address: str | None,
|
|
) -> tuple[str, RefreshToken]:
|
|
raw_token = token_urlsafe(48)
|
|
record = RefreshToken(
|
|
user_id=user_id,
|
|
token_hash=hash_refresh_token(raw_token),
|
|
expires_at=expires_at,
|
|
created_at=datetime.now(UTC),
|
|
user_agent=user_agent,
|
|
ip_address=ip_address,
|
|
)
|
|
session.add(record)
|
|
await session.commit()
|
|
await session.refresh(record)
|
|
return raw_token, record
|
|
|
|
|
|
async def rotate_refresh_token(
|
|
*,
|
|
session: AsyncSession,
|
|
raw_token: str,
|
|
user_agent: str | None,
|
|
ip_address: str | None,
|
|
) -> tuple[str, RefreshToken]:
|
|
existing_hash = hash_refresh_token(raw_token)
|
|
existing = await session.scalar(
|
|
select(RefreshToken).where(
|
|
RefreshToken.token_hash == existing_hash,
|
|
RefreshToken.revoked_at.is_(None),
|
|
)
|
|
)
|
|
if existing is None:
|
|
raise ValueError("refresh token not found")
|
|
if existing.expires_at <= datetime.now(UTC):
|
|
raise ValueError("refresh token expired")
|
|
|
|
existing.revoked_at = datetime.now(UTC)
|
|
await session.flush()
|
|
|
|
return await create_refresh_token(
|
|
session=session,
|
|
user_id=existing.user_id,
|
|
expires_at=existing.expires_at,
|
|
user_agent=user_agent,
|
|
ip_address=ip_address,
|
|
)
|
|
|
|
|
|
async def revoke_refresh_token(*, session: AsyncSession, raw_token: str) -> bool:
|
|
token_hash = hash_refresh_token(raw_token)
|
|
existing = await session.scalar(select(RefreshToken).where(RefreshToken.token_hash == token_hash))
|
|
if existing is None:
|
|
return False
|
|
if existing.revoked_at is not None:
|
|
return True
|
|
|
|
existing.revoked_at = datetime.now(UTC)
|
|
await session.commit()
|
|
return True
|