shomei-postgres-0.2.0.0: src/Shomei/Session/RefreshToken/Postgres.hs
-- | PostgreSQL interpreter for the 'RefreshTokenStore' port, including the recursive-CTE
-- family revocation used by standalone OAuth revocation.
module Shomei.Session.RefreshToken.Postgres
( runRefreshTokenStorePostgres,
-- * Statements shared with the unit-of-work interpreter
-- | Exported so @Shomei.Session.UnitOfWork.Postgres@ can lift them into a transaction with
-- @Hasql.Transaction.statement@ instead of restating the SQL. 'markUsedStmt' in
-- particular is a compare-and-swap whose exact shape is owned by
-- @docs/plans/28-enforce-absolute-session-expiry-and-atomic-token-state-transitions.md@;
-- lift it, never retype it.
RefreshTokenRow,
insertRefreshTokenStmt,
markUsedStmt,
mkPersisted,
refreshTokenHashText,
revokeSessionTokensStmt,
revokeUserTokensStmt,
)
where
import Contravariant.Extras (contrazip2, contrazip9)
import Data.UUID (UUID)
import Effectful (Eff, IOE, (:>))
import Effectful.Dispatch.Dynamic (interpret_)
import Effectful.Error.Static (Error, throwError)
import Hasql.Decoders qualified as D
import Hasql.Encoders qualified as E
import Hasql.Session qualified as Session
import Hasql.Statement (Statement, preparable)
import Shomei.Error (AuthError (..))
import Shomei.Id
( RefreshTokenId,
genRefreshTokenId,
refreshTokenIdFromUUID,
refreshTokenIdToUUID,
sessionIdFromUUID,
sessionIdToUUID,
userIdToUUID,
)
import Shomei.Persistence.Codec.Postgres (refreshTokenStatusFromText, refreshTokenStatusToText)
import Shomei.Persistence.Database.Postgres (Database, postgresUnavailable, runSession)
import Shomei.Prelude
import Shomei.Session.RefreshToken.Domain
( NewRefreshToken (..),
PersistedRefreshToken (..),
RefreshTokenHash (..),
RefreshTokenStatus (RefreshTokenActive),
)
import Shomei.Session.RefreshToken.Store (RefreshTokenStore (..))
type RefreshTokenRow =
(UUID, UUID, Text, Maybe UUID, Text, UTCTime, UTCTime, Maybe UTCTime, Maybe UTCTime)
runRefreshTokenStorePostgres ::
(Database :> es, IOE :> es, Error AuthError :> es) =>
Eff (RefreshTokenStore : es) a ->
Eff es a
runRefreshTokenStorePostgres = interpret_ \case
CreateRefreshToken nrt -> do
rid <- genRefreshTokenId
let persisted = mkPersisted rid nrt
row =
( refreshTokenIdToUUID rid,
sessionIdToUUID nrt.sessionId,
refreshTokenHashText nrt.tokenHash,
fmap refreshTokenIdToUUID nrt.parentTokenId,
refreshTokenStatusToText RefreshTokenActive,
nrt.createdAt,
nrt.expiresAt,
Nothing,
Nothing
)
res <- runSession (Session.statement row insertRefreshTokenStmt)
either dbFail (const (pure persisted)) res
FindRefreshTokenByHash h -> do
res <- runSession (Session.statement (refreshTokenHashText h) findByHashStmt)
row <- either dbFail pure res
traverse rebuild row
MarkRefreshTokenUsed rid t -> do
res <- runSession (Session.statement (refreshTokenIdToUUID rid, t) markUsedStmt)
either dbFail (pure . isJust) res
RevokeRefreshTokenFamily rid t -> do
res <- runSession (Session.statement (refreshTokenIdToUUID rid, t) revokeFamilyStmt)
either dbFail (const (pure ())) res
RevokeSessionRefreshTokens sid t -> do
res <- runSession (Session.statement (sessionIdToUUID sid, t) revokeSessionTokensStmt)
either dbFail (const (pure ())) res
RevokeAllUserRefreshTokens uid t -> do
res <- runSession (Session.statement (userIdToUUID uid, t) revokeUserTokensStmt)
either dbFail (const (pure ())) res
where
dbFail = throwError . postgresUnavailable
rebuild r = either (throwError . InternalAuthError) pure (rebuildToken r)
refreshTokenHashText :: RefreshTokenHash -> Text
refreshTokenHashText (RefreshTokenHash t) = t
mkPersisted :: RefreshTokenId -> NewRefreshToken -> PersistedRefreshToken
mkPersisted rid nrt =
PersistedRefreshToken
{ refreshTokenId = rid,
sessionId = nrt.sessionId,
tokenHash = nrt.tokenHash,
parentTokenId = nrt.parentTokenId,
status = RefreshTokenActive,
createdAt = nrt.createdAt,
expiresAt = nrt.expiresAt,
usedAt = Nothing,
revokedAt = Nothing
}
rebuildToken :: RefreshTokenRow -> Either Text PersistedRefreshToken
rebuildToken (rid, sid, h, parent, st, c, e, used, revoked) = do
status <- refreshTokenStatusFromText st
pure
PersistedRefreshToken
{ refreshTokenId = refreshTokenIdFromUUID rid,
sessionId = sessionIdFromUUID sid,
tokenHash = RefreshTokenHash h,
parentTokenId = fmap refreshTokenIdFromUUID parent,
status = status,
createdAt = c,
expiresAt = e,
usedAt = used,
revokedAt = revoked
}
tokenRowDecoder :: D.Row RefreshTokenRow
tokenRowDecoder =
(,,,,,,,,)
<$> D.column (D.nonNullable D.uuid)
<*> D.column (D.nonNullable D.uuid)
<*> D.column (D.nonNullable D.text)
<*> D.column (D.nullable D.uuid)
<*> D.column (D.nonNullable D.text)
<*> D.column (D.nonNullable D.timestamptz)
<*> D.column (D.nonNullable D.timestamptz)
<*> D.column (D.nullable D.timestamptz)
<*> D.column (D.nullable D.timestamptz)
insertRefreshTokenStmt :: Statement RefreshTokenRow ()
insertRefreshTokenStmt =
preparable
"""
INSERT INTO shomei.shomei_refresh_tokens
(refresh_token_id, session_id, token_hash, parent_token_id, status,
created_at, expires_at, used_at, revoked_at)
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9)
"""
( contrazip9
(E.param (E.nonNullable E.uuid))
(E.param (E.nonNullable E.uuid))
(E.param (E.nonNullable E.text))
(E.param (E.nullable E.uuid))
(E.param (E.nonNullable E.text))
(E.param (E.nonNullable E.timestamptz))
(E.param (E.nonNullable E.timestamptz))
(E.param (E.nullable E.timestamptz))
(E.param (E.nullable E.timestamptz))
)
D.noResult
findByHashStmt :: Statement Text (Maybe RefreshTokenRow)
findByHashStmt =
preparable
"""
SELECT refresh_token_id, session_id, token_hash, parent_token_id, status,
created_at, expires_at, used_at, revoked_at
FROM shomei.shomei_refresh_tokens
WHERE token_hash = $1
"""
(E.param (E.nonNullable E.text))
(D.rowMaybe tokenRowDecoder)
-- | Compare-and-swap: the @status = 'active'@ guard and the write are one statement, so two
-- concurrent presentations of the same refresh token cannot both transition it. Under READ
-- COMMITTED the second UPDATE blocks on the first's row lock, re-evaluates the guard against
-- the committed row (now @used@), matches nothing, and returns no row.
markUsedStmt :: Statement (UUID, UTCTime) (Maybe UUID)
markUsedStmt =
preparable
"""
UPDATE shomei.shomei_refresh_tokens
SET status = 'used', used_at = $2
WHERE refresh_token_id = $1
AND status = 'active'
RETURNING refresh_token_id
"""
(contrazip2 (E.param (E.nonNullable E.uuid)) (E.param (E.nonNullable E.timestamptz)))
(D.rowMaybe (D.column (D.nonNullable D.uuid)))
-- Walk up from the presented token to the family root (the ancestor with no parent),
-- then walk down from that root to collect every descendant, and revoke the whole family.
revokeFamilyStmt :: Statement (UUID, UTCTime) ()
revokeFamilyStmt =
preparable
"""
WITH RECURSIVE ancestors AS (
SELECT refresh_token_id, parent_token_id
FROM shomei.shomei_refresh_tokens
WHERE refresh_token_id = $1
UNION
SELECT t.refresh_token_id, t.parent_token_id
FROM shomei.shomei_refresh_tokens t
JOIN ancestors a ON t.refresh_token_id = a.parent_token_id
),
root AS (
SELECT refresh_token_id FROM ancestors WHERE parent_token_id IS NULL LIMIT 1
),
family AS (
SELECT refresh_token_id FROM root
UNION
SELECT t.refresh_token_id
FROM shomei.shomei_refresh_tokens t
JOIN family f ON t.parent_token_id = f.refresh_token_id
)
UPDATE shomei.shomei_refresh_tokens
SET status = 'revoked', revoked_at = $2
WHERE refresh_token_id IN (SELECT refresh_token_id FROM family)
"""
(contrazip2 (E.param (E.nonNullable E.uuid)) (E.param (E.nonNullable E.timestamptz)))
D.noResult
revokeSessionTokensStmt :: Statement (UUID, UTCTime) ()
revokeSessionTokensStmt =
preparable
"""
UPDATE shomei.shomei_refresh_tokens
SET status = 'revoked', revoked_at = $2
WHERE session_id = $1
"""
(contrazip2 (E.param (E.nonNullable E.uuid)) (E.param (E.nonNullable E.timestamptz)))
D.noResult
revokeUserTokensStmt :: Statement (UUID, UTCTime) ()
revokeUserTokensStmt =
preparable
"""
UPDATE shomei.shomei_refresh_tokens rt
SET status = 'revoked', revoked_at = $2
FROM shomei.shomei_sessions s
WHERE rt.session_id = s.session_id
AND s.user_id = $1
"""
(contrazip2 (E.param (E.nonNullable E.uuid)) (E.param (E.nonNullable E.timestamptz)))
D.noResult