packages feed

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