{-# LANGUAGE DataKinds #-}
{-# LANGUAGE DoAndIfThenElse #-}
{-# LANGUAGE FlexibleContexts #-}
{-# LANGUAGE GADTs #-}
{-# LANGUAGE OverloadedStrings #-}
{-# LANGUAGE RankNTypes #-}
{-# LANGUAGE ScopedTypeVariables #-}

module Web.Spock.SessionActions
  ( SessionId,
    sessionRegenerateId,
    sessionDestroy,
    getSessionId,
    readSession,
    writeSession,
    modifySession,
    modifySession',
    modifyReadSession,
  )
where

import Web.Spock.Action
import Web.Spock.Internal.Monad ()
import Web.Spock.Internal.SessionManager
import Web.Spock.Internal.Types

-- | Regenerate the users sessionId. This preserves all stored data. Call this prior
-- to logging in a user to prevent session fixation attacks.
-- Server backends revoke the old ID; a client backend issues new ID/CSRF values
-- but cannot revoke copies of previously issued cookies.
sessionRegenerateId :: SpockActionCtx ctx conn sess st ()
sessionRegenerateId :: forall ctx conn sess st. SpockActionCtx ctx conn sess st ()
sessionRegenerateId =
  ()
-> ActionCtxT () (WebStateM conn sess st) ()
-> ActionCtxT ctx (WebStateM conn sess st) ()
forall (m :: * -> *) ctx' a ctx.
MonadIO m =>
ctx' -> ActionCtxT ctx' m a -> ActionCtxT ctx m a
runInContext () (ActionCtxT () (WebStateM conn sess st) ()
 -> ActionCtxT ctx (WebStateM conn sess st) ())
-> ActionCtxT () (WebStateM conn sess st) ()
-> ActionCtxT ctx (WebStateM conn sess st) ()
forall a b. (a -> b) -> a -> b
$
    ActionCtxT
  ()
  (WebStateM conn sess st)
  (SessionManager
     (ActionCtxT () (WebStateM conn sess st)) conn sess st)
ActionCtxT
  ()
  (WebStateM conn sess st)
  (SpockSessionManager
     (SpockConn (ActionCtxT () (WebStateM conn sess st)))
     (SpockSession (ActionCtxT () (WebStateM conn sess st)))
     (SpockState (ActionCtxT () (WebStateM conn sess st))))
forall (m :: * -> *).
HasSpock m =>
m (SpockSessionManager
     (SpockConn m) (SpockSession m) (SpockState m))
getSessMgr ActionCtxT
  ()
  (WebStateM conn sess st)
  (SessionManager
     (ActionCtxT () (WebStateM conn sess st)) conn sess st)
-> (SessionManager
      (ActionCtxT () (WebStateM conn sess st)) conn sess st
    -> ActionCtxT () (WebStateM conn sess st) ())
-> ActionCtxT () (WebStateM conn sess st) ()
forall a b.
ActionCtxT () (WebStateM conn sess st) a
-> (a -> ActionCtxT () (WebStateM conn sess st) b)
-> ActionCtxT () (WebStateM conn sess st) b
forall (m :: * -> *) a b. Monad m => m a -> (a -> m b) -> m b
>>= SessionManager
  (ActionCtxT () (WebStateM conn sess st)) conn sess st
-> ActionCtxT () (WebStateM conn sess st) ()
forall (m :: * -> *) conn sess st.
SessionManager m conn sess st -> m ()
sm_regenerateSessionId

-- | Expire this browser's cookie and discard the current request's session.
-- Server backends also revoke the stored session. Previously issued stateless
-- cookies remain replayable until expiry or key removal. No replacement is
-- created until another session action is used. Use a CSRF-protected logout.
sessionDestroy :: SpockActionCtx ctx conn sess st ()
sessionDestroy :: forall ctx conn sess st. SpockActionCtx ctx conn sess st ()
sessionDestroy = ()
-> ActionCtxT () (WebStateM conn sess st) ()
-> ActionCtxT ctx (WebStateM conn sess st) ()
forall (m :: * -> *) ctx' a ctx.
MonadIO m =>
ctx' -> ActionCtxT ctx' m a -> ActionCtxT ctx m a
runInContext () (ActionCtxT () (WebStateM conn sess st) ()
 -> ActionCtxT ctx (WebStateM conn sess st) ())
-> ActionCtxT () (WebStateM conn sess st) ()
-> ActionCtxT ctx (WebStateM conn sess st) ()
forall a b. (a -> b) -> a -> b
$ ActionCtxT
  ()
  (WebStateM conn sess st)
  (SessionManager
     (ActionCtxT () (WebStateM conn sess st)) conn sess st)
ActionCtxT
  ()
  (WebStateM conn sess st)
  (SpockSessionManager
     (SpockConn (ActionCtxT () (WebStateM conn sess st)))
     (SpockSession (ActionCtxT () (WebStateM conn sess st)))
     (SpockState (ActionCtxT () (WebStateM conn sess st))))
forall (m :: * -> *).
HasSpock m =>
m (SpockSessionManager
     (SpockConn m) (SpockSession m) (SpockState m))
getSessMgr ActionCtxT
  ()
  (WebStateM conn sess st)
  (SessionManager
     (ActionCtxT () (WebStateM conn sess st)) conn sess st)
-> (SessionManager
      (ActionCtxT () (WebStateM conn sess st)) conn sess st
    -> ActionCtxT () (WebStateM conn sess st) ())
-> ActionCtxT () (WebStateM conn sess st) ()
forall a b.
ActionCtxT () (WebStateM conn sess st) a
-> (a -> ActionCtxT () (WebStateM conn sess st) b)
-> ActionCtxT () (WebStateM conn sess st) b
forall (m :: * -> *) a b. Monad m => m a -> (a -> m b) -> m b
>>= SessionManager
  (ActionCtxT () (WebStateM conn sess st)) conn sess st
-> ActionCtxT () (WebStateM conn sess st) ()
forall (m :: * -> *) conn sess st.
SessionManager m conn sess st -> m ()
sm_destroySession

-- | Get the current users sessionId. Note that this ID should only be
-- shown to it's owner as otherwise sessions can be hijacked.
getSessionId :: SpockActionCtx ctx conn sess st SessionId
getSessionId :: forall ctx conn sess st. SpockActionCtx ctx conn sess st SessionId
getSessionId =
  ()
-> ActionCtxT () (WebStateM conn sess st) SessionId
-> ActionCtxT ctx (WebStateM conn sess st) SessionId
forall (m :: * -> *) ctx' a ctx.
MonadIO m =>
ctx' -> ActionCtxT ctx' m a -> ActionCtxT ctx m a
runInContext () (ActionCtxT () (WebStateM conn sess st) SessionId
 -> ActionCtxT ctx (WebStateM conn sess st) SessionId)
-> ActionCtxT () (WebStateM conn sess st) SessionId
-> ActionCtxT ctx (WebStateM conn sess st) SessionId
forall a b. (a -> b) -> a -> b
$
    ActionCtxT
  ()
  (WebStateM conn sess st)
  (SessionManager
     (ActionCtxT () (WebStateM conn sess st)) conn sess st)
ActionCtxT
  ()
  (WebStateM conn sess st)
  (SpockSessionManager
     (SpockConn (ActionCtxT () (WebStateM conn sess st)))
     (SpockSession (ActionCtxT () (WebStateM conn sess st)))
     (SpockState (ActionCtxT () (WebStateM conn sess st))))
forall (m :: * -> *).
HasSpock m =>
m (SpockSessionManager
     (SpockConn m) (SpockSession m) (SpockState m))
getSessMgr ActionCtxT
  ()
  (WebStateM conn sess st)
  (SessionManager
     (ActionCtxT () (WebStateM conn sess st)) conn sess st)
-> (SessionManager
      (ActionCtxT () (WebStateM conn sess st)) conn sess st
    -> ActionCtxT () (WebStateM conn sess st) SessionId)
-> ActionCtxT () (WebStateM conn sess st) SessionId
forall a b.
ActionCtxT () (WebStateM conn sess st) a
-> (a -> ActionCtxT () (WebStateM conn sess st) b)
-> ActionCtxT () (WebStateM conn sess st) b
forall (m :: * -> *) a b. Monad m => m a -> (a -> m b) -> m b
>>= SessionManager
  (ActionCtxT () (WebStateM conn sess st)) conn sess st
-> ActionCtxT () (WebStateM conn sess st) SessionId
forall (m :: * -> *) conn sess st.
SessionManager m conn sess st -> m SessionId
sm_getSessionId

-- | Write to the current session using the configured server or cookie backend.
writeSession :: forall sess ctx conn st. sess -> SpockActionCtx ctx conn sess st ()
writeSession :: forall sess ctx conn st. sess -> SpockActionCtx ctx conn sess st ()
writeSession sess
d =
  do
    mgr <- ActionCtxT
  ctx
  (WebStateM conn sess st)
  (SessionManager
     (ActionCtxT () (WebStateM conn sess st)) conn sess st)
ActionCtxT
  ctx
  (WebStateM conn sess st)
  (SpockSessionManager
     (SpockConn (ActionCtxT ctx (WebStateM conn sess st)))
     (SpockSession (ActionCtxT ctx (WebStateM conn sess st)))
     (SpockState (ActionCtxT ctx (WebStateM conn sess st))))
forall (m :: * -> *).
HasSpock m =>
m (SpockSessionManager
     (SpockConn m) (SpockSession m) (SpockState m))
getSessMgr
    runInContext () $ sm_writeSession mgr d

-- | Modify the current session. Server backends perform this atomically in the
-- store. Client backends modify this request's copy; simultaneous requests can
-- overwrite one another when the browser accepts their response cookies.
modifySession :: (sess -> sess) -> SpockActionCtx ctx conn sess st ()
modifySession :: forall sess ctx conn st.
(sess -> sess) -> SpockActionCtx ctx conn sess st ()
modifySession sess -> sess
f =
  (sess -> (sess, ())) -> SpockActionCtx ctx conn sess st ()
forall sess a ctx conn st.
(sess -> (sess, a)) -> SpockActionCtx ctx conn sess st a
modifySession' ((sess -> (sess, ())) -> SpockActionCtx ctx conn sess st ())
-> (sess -> (sess, ())) -> SpockActionCtx ctx conn sess st ()
forall a b. (a -> b) -> a -> b
$ \sess
sess -> (sess -> sess
f sess
sess, ())

-- | Modify the stored session and return a value
modifySession' :: (sess -> (sess, a)) -> SpockActionCtx ctx conn sess st a
modifySession' :: forall sess a ctx conn st.
(sess -> (sess, a)) -> SpockActionCtx ctx conn sess st a
modifySession' sess -> (sess, a)
f =
  do
    mgr <- ActionCtxT
  ctx
  (WebStateM conn sess st)
  (SessionManager
     (ActionCtxT () (WebStateM conn sess st)) conn sess st)
ActionCtxT
  ctx
  (WebStateM conn sess st)
  (SpockSessionManager
     (SpockConn (ActionCtxT ctx (WebStateM conn sess st)))
     (SpockSession (ActionCtxT ctx (WebStateM conn sess st)))
     (SpockState (ActionCtxT ctx (WebStateM conn sess st))))
forall (m :: * -> *).
HasSpock m =>
m (SpockSessionManager
     (SpockConn m) (SpockSession m) (SpockState m))
getSessMgr
    runInContext () $ sm_modifySession mgr f

-- | Modify the stored session and return the new value after modification
modifyReadSession :: (sess -> sess) -> SpockActionCtx ctx conn sess st sess
modifyReadSession :: forall sess ctx conn st.
(sess -> sess) -> SpockActionCtx ctx conn sess st sess
modifyReadSession sess -> sess
f =
  (sess -> (sess, sess)) -> SpockActionCtx ctx conn sess st sess
forall sess a ctx conn st.
(sess -> (sess, a)) -> SpockActionCtx ctx conn sess st a
modifySession' ((sess -> (sess, sess)) -> SpockActionCtx ctx conn sess st sess)
-> (sess -> (sess, sess)) -> SpockActionCtx ctx conn sess st sess
forall a b. (a -> b) -> a -> b
$ \sess
sess ->
    let x :: sess
x = sess -> sess
f sess
sess
     in (sess
x, sess
x)

-- | Read the stored session
readSession :: SpockActionCtx ctx conn sess st sess
readSession :: forall ctx conn sess st. SpockActionCtx ctx conn sess st sess
readSession =
  ()
-> ActionCtxT () (WebStateM conn sess st) sess
-> ActionCtxT ctx (WebStateM conn sess st) sess
forall (m :: * -> *) ctx' a ctx.
MonadIO m =>
ctx' -> ActionCtxT ctx' m a -> ActionCtxT ctx m a
runInContext () (ActionCtxT () (WebStateM conn sess st) sess
 -> ActionCtxT ctx (WebStateM conn sess st) sess)
-> ActionCtxT () (WebStateM conn sess st) sess
-> ActionCtxT ctx (WebStateM conn sess st) sess
forall a b. (a -> b) -> a -> b
$
    do
      mgr <- ActionCtxT
  ()
  (WebStateM conn sess st)
  (SessionManager
     (ActionCtxT () (WebStateM conn sess st)) conn sess st)
ActionCtxT
  ()
  (WebStateM conn sess st)
  (SpockSessionManager
     (SpockConn (ActionCtxT () (WebStateM conn sess st)))
     (SpockSession (ActionCtxT () (WebStateM conn sess st)))
     (SpockState (ActionCtxT () (WebStateM conn sess st))))
forall (m :: * -> *).
HasSpock m =>
m (SpockSessionManager
     (SpockConn m) (SpockSession m) (SpockState m))
getSessMgr
      sm_readSession mgr