{-# LANGUAGE CPP #-}
{-# LANGUAGE DoAndIfThenElse #-}
{-# LANGUAGE FlexibleContexts #-}
{-# LANGUAGE GADTs #-}
{-# LANGUAGE OverloadedStrings #-}
{-# LANGUAGE RankNTypes #-}

module Web.Spock.Internal.SessionManager
  ( createSessionManager,
    withSessionManager,
    SessionId,
    Session (..),
    SessionManager (..),
    ServerSessionManager (..),
    SessionIf (..),
  )
where

import Control.Concurrent
import Control.Exception
import Control.Monad
import Control.Monad.Trans
import qualified Data.ByteString as BS
import qualified Data.HashMap.Strict as HM
import Data.IORef
import qualified Data.Text as T
import Data.Time
import qualified Data.Traversable as T
import qualified Data.Vault.Lazy as V
import qualified Network.Wai as Wai
import Web.Spock.Core
import Web.Spock.Internal.Cookies
import Web.Spock.Internal.Types
import Web.Spock.Internal.Util
import Web.Spock.Internal.SessionCommon
import Web.Spock.Internal.ClientSession (createClientSessionManager)

withSessionManager ::
  MonadIO m => SessionCfg conn sess st -> SessionIf m -> (SessionManager m conn sess st -> IO a) -> IO a
withSessionManager :: forall (m :: * -> *) conn sess st a.
MonadIO m =>
SessionCfg conn sess st
-> SessionIf m -> (SessionManager m conn sess st -> IO a) -> IO a
withSessionManager SessionCfg conn sess st
sessCfg SessionIf m
sif =
  IO (SessionManager m conn sess st)
-> (SessionManager m conn sess st -> IO ())
-> (SessionManager m conn sess st -> IO a)
-> IO a
forall a b c. IO a -> (a -> IO b) -> (a -> IO c) -> IO c
bracket (SessionCfg conn sess st
-> SessionIf m -> IO (SessionManager m conn sess st)
forall (m :: * -> *) conn sess st.
MonadIO m =>
SessionCfg conn sess st
-> SessionIf m -> IO (SessionManager m conn sess st)
createSessionManager SessionCfg conn sess st
sessCfg SessionIf m
sif) SessionManager m conn sess st -> IO ()
forall (m :: * -> *) conn sess st.
SessionManager m conn sess st -> IO ()
sm_closeSessionManager

createSessionManager ::
  MonadIO m => SessionCfg conn sess st -> SessionIf m -> IO (SessionManager m conn sess st)
createSessionManager :: forall (m :: * -> *) conn sess st.
MonadIO m =>
SessionCfg conn sess st
-> SessionIf m -> IO (SessionManager m conn sess st)
createSessionManager SessionCfg conn sess st
cfg SessionIf m
sif = case SessionCfg conn sess st -> SessionBackend conn sess st
forall conn a st. SessionCfg conn a st -> SessionBackend conn a st
sc_backend SessionCfg conn sess st
cfg of
  ServerSessions ServerSessionCfg conn sess st
server -> SessionCfg conn sess st
-> ServerSessionCfg conn sess st
-> SessionIf m
-> IO (SessionManager m conn sess st)
forall (m :: * -> *) conn sess st.
MonadIO m =>
SessionCfg conn sess st
-> ServerSessionCfg conn sess st
-> SessionIf m
-> IO (SessionManager m conn sess st)
createServerSessionManager SessionCfg conn sess st
cfg ServerSessionCfg conn sess st
server SessionIf m
sif
  ClientSessions ClientSessionCfg sess
client -> SessionCfg conn sess st
-> ClientSessionCfg sess
-> SessionIf m
-> IO (SessionManager m conn sess st)
forall (m :: * -> *) conn sess st.
MonadIO m =>
SessionCfg conn sess st
-> ClientSessionCfg sess
-> SessionIf m
-> IO (SessionManager m conn sess st)
createClientSessionManager SessionCfg conn sess st
cfg ClientSessionCfg sess
client SessionIf m
sif

createServerSessionManager :: MonadIO m => SessionCfg conn sess st -> ServerSessionCfg conn sess st -> SessionIf m -> IO (SessionManager m conn sess st)
createServerSessionManager :: forall (m :: * -> *) conn sess st.
MonadIO m =>
SessionCfg conn sess st
-> ServerSessionCfg conn sess st
-> SessionIf m
-> IO (SessionManager m conn sess st)
createServerSessionManager SessionCfg conn sess st
cfg ServerSessionCfg conn sess st
server SessionIf m
originalIf =
  do
    cookieKey <- IO (Key (IORef (Maybe ByteString)))
forall a. IO (Key a)
V.newKey
    -- Share the pending cookie with the response middleware. The last session
    -- action wins, even on the first request or when a handler falls through.
    let sif = SessionIf m
originalIf
          { si_setRawMultiHeader = \MultiHeader
header ByteString
value -> do
              pending <- SessionIf m -> forall a. Key a -> m (Maybe a)
forall (m :: * -> *). SessionIf m -> forall a. Key a -> m (Maybe a)
si_queryVault SessionIf m
originalIf Key (IORef (Maybe ByteString))
cookieKey
              case pending of
                Maybe (IORef (Maybe ByteString))
Nothing -> SessionIf m -> MultiHeader -> ByteString -> m ()
forall (m :: * -> *).
SessionIf m -> MultiHeader -> ByteString -> m ()
si_setRawMultiHeader SessionIf m
originalIf MultiHeader
header ByteString
value
                Just IORef (Maybe ByteString)
ref -> IO () -> m ()
forall a. IO a -> m a
forall (m :: * -> *) a. MonadIO m => IO a -> m a
liftIO (IO () -> m ()) -> IO () -> m ()
forall a b. (a -> b) -> a -> b
$ IORef (Maybe ByteString) -> Maybe ByteString -> IO ()
forall a. IORef a -> a -> IO ()
writeIORef IORef (Maybe ByteString)
ref (ByteString -> Maybe ByteString
forall a. a -> Maybe a
Just ByteString
value)
          }
    vaultKey <- si_vaultKey sif
    housekeepThread <-
      if sc_sessionMode cfg == SessionsDisabled
        then pure Nothing
        else Just <$> forkIO (forever (housekeepSessions server))
    return
      SessionManager
        { sm_getSessionId = enabled $ getSessionIdImpl vaultKey store cfg sif,
          sm_getCsrfToken = enabled $ getCsrfTokenImpl vaultKey store cfg sif,
          sm_regenerateSessionId = enabled $ regenerateSessionIdImpl vaultKey store cfg sif,
          sm_destroySession = enabled $ destroySessionImpl vaultKey store cfg sif,
          sm_readSession = enabled $ readSessionImpl vaultKey store cfg sif,
          sm_writeSession = \sess
value -> m () -> m ()
forall (n :: * -> *) a. MonadIO n => n a -> n a
enabled (m () -> m ()) -> m () -> m ()
forall a b. (a -> b) -> a -> b
$ Key SessionId
-> SessionStoreInstance (Session conn sess st)
-> SessionCfg conn sess st
-> SessionIf m
-> sess
-> m ()
forall (m :: * -> *) conn sess st.
MonadIO m =>
Key SessionId
-> SessionStoreInstance (Session conn sess st)
-> SessionCfg conn sess st
-> SessionIf m
-> sess
-> m ()
writeSessionImpl Key SessionId
vaultKey SessionStoreInstance (Session conn sess st)
store SessionCfg conn sess st
cfg SessionIf m
sif sess
value,
          sm_modifySession = \sess -> (sess, a)
f -> m a -> m a
forall (n :: * -> *) a. MonadIO n => n a -> n a
enabled (m a -> m a) -> m a -> m a
forall a b. (a -> b) -> a -> b
$ Key SessionId
-> SessionStoreInstance (Session conn sess st)
-> SessionCfg conn sess st
-> SessionIf m
-> (sess -> (sess, a))
-> m a
forall (m :: * -> *) conn sess st a.
MonadIO m =>
Key SessionId
-> SessionStoreInstance (Session conn sess st)
-> SessionCfg conn sess st
-> SessionIf m
-> (sess -> (sess, a))
-> m a
modifySessionImpl Key SessionId
vaultKey SessionStoreInstance (Session conn sess st)
store SessionCfg conn sess st
cfg SessionIf m
sif sess -> (sess, a)
f,
          sm_serverSessions = if sc_sessionMode cfg == SessionsDisabled then Nothing else
            Just $ ServerSessionManager (mapAllSessionsImpl store) (clearAllSessionsImpl store),
          sm_middleware = sessionMiddleware cfg store vaultKey cookieKey,
          sm_closeSessionManager = mapM_ killThread housekeepThread
        }
  where
    store :: SessionStoreInstance (Session conn sess st)
store = ServerSessionCfg conn sess st
-> SessionStoreInstance (Session conn sess st)
forall conn sess st.
ServerSessionCfg conn sess st
-> SessionStoreInstance (Session conn sess st)
ssc_store ServerSessionCfg conn sess st
server
    enabled :: MonadIO n => n a -> n a
    enabled :: forall (n :: * -> *) a. MonadIO n => n a -> n a
enabled n a
action
      | SessionCfg conn sess st -> SessionMode
forall conn a st. SessionCfg conn a st -> SessionMode
sc_sessionMode SessionCfg conn sess st
cfg SessionMode -> SessionMode -> Bool
forall a. Eq a => a -> a -> Bool
== SessionMode
SessionsDisabled = IO a -> n a
forall a. IO a -> n a
forall (m :: * -> *) a. MonadIO m => IO a -> m a
liftIO (IO a -> n a) -> IO a -> n a
forall a b. (a -> b) -> a -> b
$ SessionError -> IO a
forall e a. (HasCallStack, Exception e) => e -> IO a
throwIO SessionError
SessionUseWhenDisabled
      | Bool
otherwise = n a
action

regenerateSessionIdImpl ::
  MonadIO m =>
  V.Key SessionId ->
  SessionStoreInstance (Session conn sess st) ->
  SessionCfg conn sess st ->
  SessionIf m ->
  m ()
regenerateSessionIdImpl :: forall (m :: * -> *) conn sess st.
MonadIO m =>
Key SessionId
-> SessionStoreInstance (Session conn sess st)
-> SessionCfg conn sess st
-> SessionIf m
-> m ()
regenerateSessionIdImpl Key SessionId
vK SessionStoreInstance (Session conn sess st)
sessionRef SessionCfg conn sess st
cfg SessionIf m
sif =
  do
    sid <- SessionIf m -> forall a. Key a -> m (Maybe a)
forall (m :: * -> *). SessionIf m -> forall a. Key a -> m (Maybe a)
si_queryVault SessionIf m
sif Key SessionId
vK
    fresh <- liftIO $ createSession cfg (sc_emptySession cfg)
    now <- liftIO getCurrentTime
    newSession <- liftIO $ case sessionRef of
      SessionStoreInstance SessionStore (Session conn sess st) tx
store -> SessionStore (Session conn sess st) tx -> forall a. tx a -> IO a
forall sess (tx :: * -> *).
SessionStore sess tx -> forall a. tx a -> IO a
ss_runTx SessionStore (Session conn sess st) tx
store (tx (Session conn sess st) -> IO (Session conn sess st))
-> tx (Session conn sess st) -> IO (Session conn sess st)
forall a b. (a -> b) -> a -> b
$ do
        previous <- tx (Maybe (Session conn sess st))
-> (SessionId -> tx (Maybe (Session conn sess st)))
-> Maybe SessionId
-> tx (Maybe (Session conn sess st))
forall b a. b -> (a -> b) -> Maybe a -> b
maybe (Maybe (Session conn sess st) -> tx (Maybe (Session conn sess st))
forall a. a -> tx a
forall (f :: * -> *) a. Applicative f => a -> f a
pure Maybe (Session conn sess st)
forall a. Maybe a
Nothing) (\SessionId
key -> SessionCfg conn sess st
-> SessionStore (Session conn sess st) tx
-> SessionId
-> UTCTime
-> tx (Maybe (Session conn sess st))
forall (tx :: * -> *) conn sess st.
Monad tx =>
SessionCfg conn sess st
-> SessionStore (Session conn sess st) tx
-> SessionId
-> UTCTime
-> tx (Maybe (Session conn sess st))
loadSessionTx SessionCfg conn sess st
cfg SessionStore (Session conn sess st) tx
store SessionId
key UTCTime
now) Maybe SessionId
sid
        let replacement = Session conn sess st
fresh { sess_data = maybe (sc_emptySession cfg) sess_data previous }
        mapM_ (ss_deleteSession store) sid
        ss_storeSession store replacement
        pure replacement
    si_setRawMultiHeader sif MultiHeaderSetCookie (makeSessionIdCookie cfg newSession now)
    si_modifyVault sif $ V.insert vK (sess_id newSession)

destroySessionImpl :: MonadIO m => V.Key SessionId -> SessionStoreInstance (Session conn sess st) -> SessionCfg conn sess st -> SessionIf m -> m ()
destroySessionImpl :: forall (m :: * -> *) conn sess st.
MonadIO m =>
Key SessionId
-> SessionStoreInstance (Session conn sess st)
-> SessionCfg conn sess st
-> SessionIf m
-> m ()
destroySessionImpl Key SessionId
vK SessionStoreInstance (Session conn sess st)
store SessionCfg conn sess st
cfg SessionIf m
sif = do
  sid <- SessionIf m -> forall a. Key a -> m (Maybe a)
forall (m :: * -> *). SessionIf m -> forall a. Key a -> m (Maybe a)
si_queryVault SessionIf m
sif Key SessionId
vK
  liftIO $ mapM_ (deleteSessionImpl store) sid
  si_modifyVault sif $ V.delete vK
  now <- liftIO getCurrentTime
  let settings = (SessionCfg conn sess st -> CookieSettings
forall conn a st. SessionCfg conn a st -> CookieSettings
sc_cookieSettings SessionCfg conn sess st
cfg) { cs_EOL = CookieValidFor 0 }
  si_setRawMultiHeader sif MultiHeaderSetCookie $
    generateCookieHeaderString (sc_cookieName cfg) "" settings now

getSessionIdImpl ::
  MonadIO m =>
  V.Key SessionId ->
  SessionStoreInstance (Session conn sess st) ->
  SessionCfg conn sess st ->
  SessionIf m ->
  m SessionId
getSessionIdImpl :: forall (m :: * -> *) conn sess st.
MonadIO m =>
Key SessionId
-> SessionStoreInstance (Session conn sess st)
-> SessionCfg conn sess st
-> SessionIf m
-> m SessionId
getSessionIdImpl Key SessionId
vK SessionStoreInstance (Session conn sess st)
store SessionCfg conn sess st
cfg SessionIf m
sif =
  do
    sess <- Key SessionId
-> SessionStoreInstance (Session conn sess st)
-> SessionCfg conn sess st
-> SessionIf m
-> m (Session conn sess st)
forall (m :: * -> *) conn sess st.
MonadIO m =>
Key SessionId
-> SessionStoreInstance (Session conn sess st)
-> SessionCfg conn sess st
-> SessionIf m
-> m (Session conn sess st)
readSessionBase Key SessionId
vK SessionStoreInstance (Session conn sess st)
store SessionCfg conn sess st
cfg SessionIf m
sif
    return $ sess_id sess

getCsrfTokenImpl ::
  (MonadIO m) =>
  V.Key SessionId ->
  SessionStoreInstance (Session conn sess st) ->
  SessionCfg conn sess st ->
  SessionIf m ->
  m T.Text
getCsrfTokenImpl :: forall (m :: * -> *) conn sess st.
MonadIO m =>
Key SessionId
-> SessionStoreInstance (Session conn sess st)
-> SessionCfg conn sess st
-> SessionIf m
-> m SessionId
getCsrfTokenImpl Key SessionId
vK SessionStoreInstance (Session conn sess st)
store SessionCfg conn sess st
cfg SessionIf m
sif =
  do
    sess <- Key SessionId
-> SessionStoreInstance (Session conn sess st)
-> SessionCfg conn sess st
-> SessionIf m
-> m (Session conn sess st)
forall (m :: * -> *) conn sess st.
MonadIO m =>
Key SessionId
-> SessionStoreInstance (Session conn sess st)
-> SessionCfg conn sess st
-> SessionIf m
-> m (Session conn sess st)
readSessionBase Key SessionId
vK SessionStoreInstance (Session conn sess st)
store SessionCfg conn sess st
cfg SessionIf m
sif
    return $ sess_csrfToken sess

modifySessionBase ::
  MonadIO m =>
  V.Key SessionId ->
  SessionStoreInstance (Session conn sess st) ->
  SessionCfg conn sess st ->
  SessionIf m ->
  (Session conn sess st -> (Session conn sess st, a)) ->
  m a
modifySessionBase :: forall (m :: * -> *) conn sess st a.
MonadIO m =>
Key SessionId
-> SessionStoreInstance (Session conn sess st)
-> SessionCfg conn sess st
-> SessionIf m
-> (Session conn sess st -> (Session conn sess st, a))
-> m a
modifySessionBase Key SessionId
vK (SessionStoreInstance SessionStore (Session conn sess st) tx
sessionRef) SessionCfg conn sess st
cfg SessionIf m
sif Session conn sess st -> (Session conn sess st, a)
modFun =
  do
    mValue <- SessionIf m -> forall a. Key a -> m (Maybe a)
forall (m :: * -> *). SessionIf m -> forall a. Key a -> m (Maybe a)
si_queryVault SessionIf m
sif Key SessionId
vK
    now <- liftIO getCurrentTime
    mResult <-
      liftIO $ ss_runTx sessionRef $
        do
          mSession <- maybe (pure Nothing) (\SessionId
sid -> SessionCfg conn sess st
-> SessionStore (Session conn sess st) tx
-> SessionId
-> UTCTime
-> tx (Maybe (Session conn sess st))
forall (tx :: * -> *) conn sess st.
Monad tx =>
SessionCfg conn sess st
-> SessionStore (Session conn sess st) tx
-> SessionId
-> UTCTime
-> tx (Maybe (Session conn sess st))
loadSessionTx SessionCfg conn sess st
cfg SessionStore (Session conn sess st) tx
sessionRef SessionId
sid UTCTime
now) mValue
          forM mSession $ \Session conn sess st
session ->
            do
              let (Session conn sess st
sessionNew, a
result) = Session conn sess st -> (Session conn sess st, a)
modFun Session conn sess st
session
              SessionStore (Session conn sess st) tx
-> Session conn sess st -> tx ()
forall sess (tx :: * -> *). SessionStore sess tx -> sess -> tx ()
ss_storeSession SessionStore (Session conn sess st) tx
sessionRef Session conn sess st
sessionNew
              a -> tx a
forall a. a -> tx a
forall (m :: * -> *) a. Monad m => a -> m a
return a
result
    case mResult of
      Just a
result -> a -> m a
forall a. a -> m a
forall (m :: * -> *) a. Monad m => a -> m a
return a
result
      Maybe a
Nothing ->
        do
          session <- IO (Session conn sess st) -> m (Session conn sess st)
forall a. IO a -> m a
forall (m :: * -> *) a. MonadIO m => IO a -> m a
liftIO (IO (Session conn sess st) -> m (Session conn sess st))
-> IO (Session conn sess st) -> m (Session conn sess st)
forall a b. (a -> b) -> a -> b
$ SessionCfg conn sess st -> sess -> IO (Session conn sess st)
forall conn sess st.
SessionCfg conn sess st -> sess -> IO (Session conn sess st)
createSession SessionCfg conn sess st
cfg (SessionCfg conn sess st -> sess
forall conn a st. SessionCfg conn a st -> a
sc_emptySession SessionCfg conn sess st
cfg)
          let (sessionNew, result) = modFun session
          liftIO $ ss_runTx sessionRef $ ss_storeSession sessionRef sessionNew
          cookieTime <- liftIO getCurrentTime
          si_setRawMultiHeader sif MultiHeaderSetCookie (makeSessionIdCookie cfg sessionNew cookieTime)
          si_modifyVault sif $ V.insert vK (sess_id sessionNew)
          return result

readSessionBase ::
  MonadIO m =>
  V.Key SessionId ->
  SessionStoreInstance (Session conn sess st) ->
  SessionCfg conn sess st ->
  SessionIf m ->
  m (Session conn sess st)
readSessionBase :: forall (m :: * -> *) conn sess st.
MonadIO m =>
Key SessionId
-> SessionStoreInstance (Session conn sess st)
-> SessionCfg conn sess st
-> SessionIf m
-> m (Session conn sess st)
readSessionBase Key SessionId
vK SessionStoreInstance (Session conn sess st)
store SessionCfg conn sess st
cfg SessionIf m
sif =
  do
    mValue <- SessionIf m -> forall a. Key a -> m (Maybe a)
forall (m :: * -> *). SessionIf m -> forall a. Key a -> m (Maybe a)
si_queryVault SessionIf m
sif Key SessionId
vK
    readOrNewSession cfg store vK sif mValue

readSessionImpl ::
  MonadIO m =>
  V.Key SessionId ->
  SessionStoreInstance (Session conn sess st) ->
  SessionCfg conn sess st ->
  SessionIf m ->
  m sess
readSessionImpl :: forall (m :: * -> *) conn sess st.
MonadIO m =>
Key SessionId
-> SessionStoreInstance (Session conn sess st)
-> SessionCfg conn sess st
-> SessionIf m
-> m sess
readSessionImpl Key SessionId
vK SessionStoreInstance (Session conn sess st)
store SessionCfg conn sess st
cfg SessionIf m
sif =
  do
    base <- Key SessionId
-> SessionStoreInstance (Session conn sess st)
-> SessionCfg conn sess st
-> SessionIf m
-> m (Session conn sess st)
forall (m :: * -> *) conn sess st.
MonadIO m =>
Key SessionId
-> SessionStoreInstance (Session conn sess st)
-> SessionCfg conn sess st
-> SessionIf m
-> m (Session conn sess st)
readSessionBase Key SessionId
vK SessionStoreInstance (Session conn sess st)
store SessionCfg conn sess st
cfg SessionIf m
sif
    return (sess_data base)

writeSessionImpl ::
  MonadIO m =>
  V.Key SessionId ->
  SessionStoreInstance (Session conn sess st) ->
  SessionCfg conn sess st ->
  SessionIf m ->
  sess ->
  m ()
writeSessionImpl :: forall (m :: * -> *) conn sess st.
MonadIO m =>
Key SessionId
-> SessionStoreInstance (Session conn sess st)
-> SessionCfg conn sess st
-> SessionIf m
-> sess
-> m ()
writeSessionImpl Key SessionId
vK SessionStoreInstance (Session conn sess st)
sessionRef SessionCfg conn sess st
cfg SessionIf m
sif sess
value =
  Key SessionId
-> SessionStoreInstance (Session conn sess st)
-> SessionCfg conn sess st
-> SessionIf m
-> (sess -> (sess, ()))
-> m ()
forall (m :: * -> *) conn sess st a.
MonadIO m =>
Key SessionId
-> SessionStoreInstance (Session conn sess st)
-> SessionCfg conn sess st
-> SessionIf m
-> (sess -> (sess, a))
-> m a
modifySessionImpl Key SessionId
vK SessionStoreInstance (Session conn sess st)
sessionRef SessionCfg conn sess st
cfg SessionIf m
sif ((sess, ()) -> sess -> (sess, ())
forall a b. a -> b -> a
const (sess
value, ()))

modifySessionImpl ::
  MonadIO m =>
  V.Key SessionId ->
  SessionStoreInstance (Session conn sess st) ->
  SessionCfg conn sess st ->
  SessionIf m ->
  (sess -> (sess, a)) ->
  m a
modifySessionImpl :: forall (m :: * -> *) conn sess st a.
MonadIO m =>
Key SessionId
-> SessionStoreInstance (Session conn sess st)
-> SessionCfg conn sess st
-> SessionIf m
-> (sess -> (sess, a))
-> m a
modifySessionImpl Key SessionId
vK SessionStoreInstance (Session conn sess st)
sessionRef SessionCfg conn sess st
cfg SessionIf m
sif sess -> (sess, a)
f =
  do
    let modFun :: Session conn sess st -> (Session conn sess st, a)
modFun Session conn sess st
session =
          let (sess
sessData', a
out) = sess -> (sess, a)
f (Session conn sess st -> sess
forall conn sess st. Session conn sess st -> sess
sess_data Session conn sess st
session)
           in (Session conn sess st
session {sess_data = sessData'}, a
out)
    Key SessionId
-> SessionStoreInstance (Session conn sess st)
-> SessionCfg conn sess st
-> SessionIf m
-> (Session conn sess st -> (Session conn sess st, a))
-> m a
forall (m :: * -> *) conn sess st a.
MonadIO m =>
Key SessionId
-> SessionStoreInstance (Session conn sess st)
-> SessionCfg conn sess st
-> SessionIf m
-> (Session conn sess st -> (Session conn sess st, a))
-> m a
modifySessionBase Key SessionId
vK SessionStoreInstance (Session conn sess st)
sessionRef SessionCfg conn sess st
cfg SessionIf m
sif Session conn sess st -> (Session conn sess st, a)
modFun

makeSessionIdCookie :: SessionCfg conn sess st -> Session conn sess st -> UTCTime -> BS.ByteString
makeSessionIdCookie :: forall conn sess st.
SessionCfg conn sess st
-> Session conn sess st -> UTCTime -> ByteString
makeSessionIdCookie SessionCfg conn sess st
cfg Session conn sess st
sess = SessionId -> SessionId -> CookieSettings -> UTCTime -> ByteString
generateCookieHeaderString SessionId
name SessionId
value CookieSettings
settings
  where
    name :: SessionId
name = SessionCfg conn sess st -> SessionId
forall conn a st. SessionCfg conn a st -> SessionId
sc_cookieName SessionCfg conn sess st
cfg
    value :: SessionId
value = Session conn sess st -> SessionId
forall conn sess st. Session conn sess st -> SessionId
sess_id Session conn sess st
sess
    settings :: CookieSettings
settings = SessionCfg conn sess st -> CookieSettings
forall conn a st. SessionCfg conn a st -> CookieSettings
sc_cookieSettings SessionCfg conn sess st
cfg

readOrNewSession ::
  MonadIO m =>
  SessionCfg conn sess st ->
  SessionStoreInstance (Session conn sess st) ->
  V.Key SessionId ->
  SessionIf m ->
  Maybe SessionId ->
  m (Session conn sess st)
readOrNewSession :: forall (m :: * -> *) conn sess st.
MonadIO m =>
SessionCfg conn sess st
-> SessionStoreInstance (Session conn sess st)
-> Key SessionId
-> SessionIf m
-> Maybe SessionId
-> m (Session conn sess st)
readOrNewSession SessionCfg conn sess st
cfg SessionStoreInstance (Session conn sess st)
store Key SessionId
vK SessionIf m
sif Maybe SessionId
mSid =
  do
    (sess, write) <- SessionCfg conn sess st
-> SessionStoreInstance (Session conn sess st)
-> Maybe SessionId
-> m (Session conn sess st, Bool)
forall (m :: * -> *) conn sess st.
MonadIO m =>
SessionCfg conn sess st
-> SessionStoreInstance (Session conn sess st)
-> Maybe SessionId
-> m (Session conn sess st, Bool)
loadOrSpanSession SessionCfg conn sess st
cfg SessionStoreInstance (Session conn sess st)
store Maybe SessionId
mSid
    when write $
      do
        now <- liftIO getCurrentTime
        si_setRawMultiHeader sif MultiHeaderSetCookie (makeSessionIdCookie cfg sess now)
        si_modifyVault sif $ V.insert vK (sess_id sess)
    return sess

loadOrSpanSession ::
  MonadIO m =>
  SessionCfg conn sess st ->
  SessionStoreInstance (Session conn sess st) ->
  Maybe SessionId ->
  m (Session conn sess st, Bool)
loadOrSpanSession :: forall (m :: * -> *) conn sess st.
MonadIO m =>
SessionCfg conn sess st
-> SessionStoreInstance (Session conn sess st)
-> Maybe SessionId
-> m (Session conn sess st, Bool)
loadOrSpanSession SessionCfg conn sess st
cfg SessionStoreInstance (Session conn sess st)
sessionRef Maybe SessionId
mSid =
  do
    mSess <-
      IO (Maybe (Session conn sess st))
-> m (Maybe (Session conn sess st))
forall a. IO a -> m a
forall (m :: * -> *) a. MonadIO m => IO a -> m a
liftIO (IO (Maybe (Session conn sess st))
 -> m (Maybe (Session conn sess st)))
-> IO (Maybe (Session conn sess st))
-> m (Maybe (Session conn sess st))
forall a b. (a -> b) -> a -> b
$
        Maybe (Maybe (Session conn sess st))
-> Maybe (Session conn sess st)
forall (m :: * -> *) a. Monad m => m (m a) -> m a
join (Maybe (Maybe (Session conn sess st))
 -> Maybe (Session conn sess st))
-> IO (Maybe (Maybe (Session conn sess st)))
-> IO (Maybe (Session conn sess st))
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> (SessionId -> IO (Maybe (Session conn sess st)))
-> Maybe SessionId -> IO (Maybe (Maybe (Session conn sess st)))
forall (t :: * -> *) (m :: * -> *) a b.
(Traversable t, Monad m) =>
(a -> m b) -> t a -> m (t b)
forall (m :: * -> *) a b.
Monad m =>
(a -> m b) -> Maybe a -> m (Maybe b)
T.mapM (SessionCfg conn sess st
-> SessionStoreInstance (Session conn sess st)
-> SessionId
-> IO (Maybe (Session conn sess st))
forall conn sess st.
SessionCfg conn sess st
-> SessionStoreInstance (Session conn sess st)
-> SessionId
-> IO (Maybe (Session conn sess st))
loadSessionImpl SessionCfg conn sess st
cfg SessionStoreInstance (Session conn sess st)
sessionRef) Maybe SessionId
mSid
    case mSess of
      Maybe (Session conn sess st)
Nothing ->
        do
          newSess <-
            IO (Session conn sess st) -> m (Session conn sess st)
forall a. IO a -> m a
forall (m :: * -> *) a. MonadIO m => IO a -> m a
liftIO (IO (Session conn sess st) -> m (Session conn sess st))
-> IO (Session conn sess st) -> m (Session conn sess st)
forall a b. (a -> b) -> a -> b
$
              SessionCfg conn sess st
-> SessionStoreInstance (Session conn sess st)
-> sess
-> IO (Session conn sess st)
forall conn sess st.
SessionCfg conn sess st
-> SessionStoreInstance (Session conn sess st)
-> sess
-> IO (Session conn sess st)
newSessionImpl SessionCfg conn sess st
cfg SessionStoreInstance (Session conn sess st)
sessionRef (SessionCfg conn sess st -> sess
forall conn a st. SessionCfg conn a st -> a
sc_emptySession SessionCfg conn sess st
cfg)
          return (newSess, True)
      Just Session conn sess st
s -> (Session conn sess st, Bool) -> m (Session conn sess st, Bool)
forall a. a -> m a
forall (m :: * -> *) a. Monad m => a -> m a
return (Session conn sess st
s, Bool
False)
sessionMiddleware ::
  SessionCfg conn sess st ->
  SessionStoreInstance (Session conn sess st) ->
  V.Key SessionId ->
  V.Key (IORef (Maybe BS.ByteString)) ->
  Wai.Middleware
sessionMiddleware :: forall conn sess st.
SessionCfg conn sess st
-> SessionStoreInstance (Session conn sess st)
-> Key SessionId
-> Key (IORef (Maybe ByteString))
-> Middleware
sessionMiddleware SessionCfg conn sess st
cfg SessionStoreInstance (Session conn sess st)
store Key SessionId
vK Key (IORef (Maybe ByteString))
cookieKey Application
app Request
req Response -> IO ResponseReceived
respond
  | SessionCfg conn sess st -> SessionMode
forall conn a st. SessionCfg conn a st -> SessionMode
sc_sessionMode SessionCfg conn sess st
cfg SessionMode -> SessionMode -> Bool
forall a. Eq a => a -> a -> Bool
== SessionMode
SessionsDisabled = Application
app Request
req Response -> IO ResponseReceived
respond
  | Bool
otherwise = do
      (sid, pendingCookie) <- case SessionCfg conn sess st -> SessionMode
forall conn a st. SessionCfg conn a st -> SessionMode
sc_sessionMode SessionCfg conn sess st
cfg of
        SessionMode
SessionsAlways -> do
          (sess, writeCookie) <- SessionCfg conn sess st
-> SessionStoreInstance (Session conn sess st)
-> Maybe SessionId
-> IO (Session conn sess st, Bool)
forall (m :: * -> *) conn sess st.
MonadIO m =>
SessionCfg conn sess st
-> SessionStoreInstance (Session conn sess st)
-> Maybe SessionId
-> m (Session conn sess st, Bool)
loadOrSpanSession SessionCfg conn sess st
cfg SessionStoreInstance (Session conn sess st)
store Maybe SessionId
cookieId
          now <- getCurrentTime
          pure (Just $ sess_id sess, if writeCookie then Just (makeSessionIdCookie cfg sess now) else Nothing)
        SessionMode
_ -> (Maybe SessionId, Maybe ByteString)
-> IO (Maybe SessionId, Maybe ByteString)
forall a. a -> IO a
forall (f :: * -> *) a. Applicative f => a -> f a
pure (Maybe SessionId
cookieId, Maybe ByteString
forall a. Maybe a
Nothing)
      pending <- newIORef pendingCookie
      let requestVault = Key (IORef (Maybe ByteString))
-> IORef (Maybe ByteString) -> Vault -> Vault
forall a. Key a -> a -> Vault -> Vault
V.insert Key (IORef (Maybe ByteString))
cookieKey IORef (Maybe ByteString)
pending (Vault -> Vault) -> Vault -> Vault
forall a b. (a -> b) -> a -> b
$ Vault -> (SessionId -> Vault) -> Maybe SessionId -> Vault
forall b a. b -> (a -> b) -> Maybe a -> b
maybe Vault
v (\SessionId
key -> Key SessionId -> SessionId -> Vault -> Vault
forall a. Key a -> a -> Vault -> Vault
V.insert Key SessionId
vK SessionId
key Vault
v) Maybe SessionId
sid
      app (req { Wai.vault = requestVault }) $ \Response
response -> do
        cookie <- IORef (Maybe ByteString) -> IO (Maybe ByteString)
forall a. IORef a -> IO a
readIORef IORef (Maybe ByteString)
pending
        respond $ maybe response (\ByteString
value -> (ResponseHeaders -> ResponseHeaders) -> Response -> Response
mapReqHeaders ((HeaderName
"Set-Cookie", ByteString
value) Header -> ResponseHeaders -> ResponseHeaders
forall a. a -> [a] -> [a]
:) Response
response) cookie
  where
    cookieId :: Maybe SessionId
cookieId = SessionId -> Maybe SessionId
getCookieFromReq (SessionCfg conn sess st -> SessionId
forall conn a st. SessionCfg conn a st -> SessionId
sc_cookieName SessionCfg conn sess st
cfg)
    getCookieFromReq :: SessionId -> Maybe SessionId
getCookieFromReq SessionId
name =
      HeaderName -> ResponseHeaders -> Maybe ByteString
forall a b. Eq a => a -> [(a, b)] -> Maybe b
lookup HeaderName
"cookie" (Request -> ResponseHeaders
Wai.requestHeaders Request
req) Maybe ByteString
-> (ByteString -> Maybe SessionId) -> Maybe SessionId
forall a b. Maybe a -> (a -> Maybe b) -> Maybe b
forall (m :: * -> *) a b. Monad m => m a -> (a -> m b) -> m b
>>= SessionId -> [(SessionId, SessionId)] -> Maybe SessionId
forall a b. Eq a => a -> [(a, b)] -> Maybe b
lookup SessionId
name ([(SessionId, SessionId)] -> Maybe SessionId)
-> (ByteString -> [(SessionId, SessionId)])
-> ByteString
-> Maybe SessionId
forall b c a. (b -> c) -> (a -> b) -> a -> c
. ByteString -> [(SessionId, SessionId)]
parseCookies
    v :: Vault
v = Request -> Vault
Wai.vault Request
req

newSessionImpl ::
  SessionCfg conn sess st ->
  SessionStoreInstance (Session conn sess st) ->
  sess ->
  IO (Session conn sess st)
newSessionImpl :: forall conn sess st.
SessionCfg conn sess st
-> SessionStoreInstance (Session conn sess st)
-> sess
-> IO (Session conn sess st)
newSessionImpl SessionCfg conn sess st
sessCfg (SessionStoreInstance SessionStore (Session conn sess st) tx
sessionRef) sess
content =
  do
    sess <- SessionCfg conn sess st -> sess -> IO (Session conn sess st)
forall conn sess st.
SessionCfg conn sess st -> sess -> IO (Session conn sess st)
createSession SessionCfg conn sess st
sessCfg sess
content
    ss_runTx sessionRef $ ss_storeSession sessionRef sess
    return $! sess

loadSessionImpl ::
  SessionCfg conn sess st ->
  SessionStoreInstance (Session conn sess st) ->
  SessionId ->
  IO (Maybe (Session conn sess st))
loadSessionImpl :: forall conn sess st.
SessionCfg conn sess st
-> SessionStoreInstance (Session conn sess st)
-> SessionId
-> IO (Maybe (Session conn sess st))
loadSessionImpl SessionCfg conn sess st
sessCfg (SessionStoreInstance SessionStore (Session conn sess st) tx
store) SessionId
sid =
  do
    now <- IO UTCTime
getCurrentTime
    ss_runTx store $
      do
        mSess <- loadSessionTx sessCfg store sid now
        when (sc_sessionExpandTTL sessCfg) $
          forM_ mSess (ss_storeSession store)
        return mSess

-- The caller must store a renewed or modified session in this same transaction.
loadSessionTx ::
  Monad tx =>
  SessionCfg conn sess st ->
  SessionStore (Session conn sess st) tx ->
  SessionId ->
  UTCTime ->
  tx (Maybe (Session conn sess st))
loadSessionTx :: forall (tx :: * -> *) conn sess st.
Monad tx =>
SessionCfg conn sess st
-> SessionStore (Session conn sess st) tx
-> SessionId
-> UTCTime
-> tx (Maybe (Session conn sess st))
loadSessionTx SessionCfg conn sess st
sessCfg SessionStore (Session conn sess st) tx
store SessionId
sid UTCTime
now =
  do
    mSess <- SessionStore (Session conn sess st) tx
-> SessionId -> tx (Maybe (Session conn sess st))
forall sess (tx :: * -> *).
SessionStore sess tx -> SessionId -> tx (Maybe sess)
ss_loadSession SessionStore (Session conn sess st) tx
store SessionId
sid
    case mSess of
      Just Session conn sess st
sess
        | Session conn sess st -> UTCTime
forall conn sess st. Session conn sess st -> UTCTime
sess_validUntil Session conn sess st
sess UTCTime -> UTCTime -> Bool
forall a. Ord a => a -> a -> Bool
<= UTCTime
now ->
            do
              SessionStore (Session conn sess st) tx -> SessionId -> tx ()
forall sess (tx :: * -> *).
SessionStore sess tx -> SessionId -> tx ()
ss_deleteSession SessionStore (Session conn sess st) tx
store SessionId
sid
              Maybe (Session conn sess st) -> tx (Maybe (Session conn sess st))
forall a. a -> tx a
forall (m :: * -> *) a. Monad m => a -> m a
return Maybe (Session conn sess st)
forall a. Maybe a
Nothing
        | SessionCfg conn sess st -> Bool
forall conn a st. SessionCfg conn a st -> Bool
sc_sessionExpandTTL SessionCfg conn sess st
sessCfg ->
            Maybe (Session conn sess st) -> tx (Maybe (Session conn sess st))
forall a. a -> tx a
forall (m :: * -> *) a. Monad m => a -> m a
return (Maybe (Session conn sess st) -> tx (Maybe (Session conn sess st)))
-> Maybe (Session conn sess st)
-> tx (Maybe (Session conn sess st))
forall a b. (a -> b) -> a -> b
$ Session conn sess st -> Maybe (Session conn sess st)
forall a. a -> Maybe a
Just (Session conn sess st -> Maybe (Session conn sess st))
-> Session conn sess st -> Maybe (Session conn sess st)
forall a b. (a -> b) -> a -> b
$
              Session conn sess st
sess
                { sess_validUntil =
                    max (sess_validUntil sess) (addUTCTime (sc_sessionTTL sessCfg) now)
                }
        | Bool
otherwise -> Maybe (Session conn sess st) -> tx (Maybe (Session conn sess st))
forall a. a -> tx a
forall (m :: * -> *) a. Monad m => a -> m a
return (Session conn sess st -> Maybe (Session conn sess st)
forall a. a -> Maybe a
Just Session conn sess st
sess)
      Maybe (Session conn sess st)
Nothing -> Maybe (Session conn sess st) -> tx (Maybe (Session conn sess st))
forall a. a -> tx a
forall (m :: * -> *) a. Monad m => a -> m a
return Maybe (Session conn sess st)
forall a. Maybe a
Nothing

deleteSessionImpl ::
  SessionStoreInstance (Session conn sess st) ->
  SessionId ->
  IO ()
deleteSessionImpl :: forall conn sess st.
SessionStoreInstance (Session conn sess st) -> SessionId -> IO ()
deleteSessionImpl (SessionStoreInstance SessionStore (Session conn sess st) tx
sessionRef) SessionId
sid =
  SessionStore (Session conn sess st) tx -> forall a. tx a -> IO a
forall sess (tx :: * -> *).
SessionStore sess tx -> forall a. tx a -> IO a
ss_runTx SessionStore (Session conn sess st) tx
sessionRef (tx () -> IO ()) -> tx () -> IO ()
forall a b. (a -> b) -> a -> b
$ SessionStore (Session conn sess st) tx -> SessionId -> tx ()
forall sess (tx :: * -> *).
SessionStore sess tx -> SessionId -> tx ()
ss_deleteSession SessionStore (Session conn sess st) tx
sessionRef SessionId
sid

clearAllSessionsImpl ::
  MonadIO m =>
  SessionStoreInstance (Session conn sess st) ->
  m ()
clearAllSessionsImpl :: forall (m :: * -> *) conn sess st.
MonadIO m =>
SessionStoreInstance (Session conn sess st) -> m ()
clearAllSessionsImpl (SessionStoreInstance SessionStore (Session conn sess st) tx
sessionRef) =
  IO () -> m ()
forall a. IO a -> m a
forall (m :: * -> *) a. MonadIO m => IO a -> m a
liftIO (IO () -> m ()) -> IO () -> m ()
forall a b. (a -> b) -> a -> b
$ SessionStore (Session conn sess st) tx -> forall a. tx a -> IO a
forall sess (tx :: * -> *).
SessionStore sess tx -> forall a. tx a -> IO a
ss_runTx SessionStore (Session conn sess st) tx
sessionRef (tx () -> IO ()) -> tx () -> IO ()
forall a b. (a -> b) -> a -> b
$ SessionStore (Session conn sess st) tx
-> (Session conn sess st -> Bool) -> tx ()
forall sess (tx :: * -> *).
SessionStore sess tx -> (sess -> Bool) -> tx ()
ss_filterSessions SessionStore (Session conn sess st) tx
sessionRef (Bool -> Session conn sess st -> Bool
forall a b. a -> b -> a
const Bool
False)

mapAllSessionsImpl ::
  MonadIO m =>
  SessionStoreInstance (Session conn sess st) ->
  (forall n. Monad n => sess -> n sess) ->
  m ()
mapAllSessionsImpl :: forall (m :: * -> *) conn sess st.
MonadIO m =>
SessionStoreInstance (Session conn sess st)
-> (forall (n :: * -> *). Monad n => sess -> n sess) -> m ()
mapAllSessionsImpl (SessionStoreInstance SessionStore (Session conn sess st) tx
sessionRef) forall (n :: * -> *). Monad n => sess -> n sess
f =
  IO () -> m ()
forall a. IO a -> m a
forall (m :: * -> *) a. MonadIO m => IO a -> m a
liftIO (IO () -> m ()) -> IO () -> m ()
forall a b. (a -> b) -> a -> b
$
    SessionStore (Session conn sess st) tx -> forall a. tx a -> IO a
forall sess (tx :: * -> *).
SessionStore sess tx -> forall a. tx a -> IO a
ss_runTx SessionStore (Session conn sess st) tx
sessionRef (tx () -> IO ()) -> tx () -> IO ()
forall a b. (a -> b) -> a -> b
$
      SessionStore (Session conn sess st) tx
-> (Session conn sess st -> tx (Session conn sess st)) -> tx ()
forall sess (tx :: * -> *).
SessionStore sess tx -> (sess -> tx sess) -> tx ()
ss_mapSessions SessionStore (Session conn sess st) tx
sessionRef ((Session conn sess st -> tx (Session conn sess st)) -> tx ())
-> (Session conn sess st -> tx (Session conn sess st)) -> tx ()
forall a b. (a -> b) -> a -> b
$ \Session conn sess st
sess ->
        do
          newData <- sess -> tx sess
forall (n :: * -> *). Monad n => sess -> n sess
f (Session conn sess st -> sess
forall conn sess st. Session conn sess st -> sess
sess_data Session conn sess st
sess)
          return $ sess {sess_data = newData}

housekeepSessions :: ServerSessionCfg conn sess st -> IO ()
housekeepSessions :: forall conn sess st. ServerSessionCfg conn sess st -> IO ()
housekeepSessions ServerSessionCfg conn sess st
cfg =
  case ServerSessionCfg conn sess st
-> SessionStoreInstance (Session conn sess st)
forall conn sess st.
ServerSessionCfg conn sess st
-> SessionStoreInstance (Session conn sess st)
ssc_store ServerSessionCfg conn sess st
cfg of
    SessionStoreInstance SessionStore (Session conn sess st) tx
store ->
      do
        now <- IO UTCTime
getCurrentTime
        (newStatus, oldStatus) <-
          ss_runTx store $
            do
              oldSt <- ss_toList store
              ss_filterSessions store (\Session conn sess st
sess -> Session conn sess st -> UTCTime
forall conn sess st. Session conn sess st -> UTCTime
sess_validUntil Session conn sess st
sess UTCTime -> UTCTime -> Bool
forall a. Ord a => a -> a -> Bool
> UTCTime
now)
              (,) <$> ss_toList store <*> pure oldSt
        let packSessionHm = [(SessionId, Session conn sess st)]
-> HashMap SessionId (Session conn sess st)
forall k v. Hashable k => [(k, v)] -> HashMap k v
HM.fromList ([(SessionId, Session conn sess st)]
 -> HashMap SessionId (Session conn sess st))
-> ([Session conn sess st] -> [(SessionId, Session conn sess st)])
-> [Session conn sess st]
-> HashMap SessionId (Session conn sess st)
forall b c a. (b -> c) -> (a -> b) -> a -> c
. (Session conn sess st -> (SessionId, Session conn sess st))
-> [Session conn sess st] -> [(SessionId, Session conn sess st)]
forall a b. (a -> b) -> [a] -> [b]
map (\Session conn sess st
v -> (Session conn sess st -> SessionId
forall conn sess st. Session conn sess st -> SessionId
sess_id Session conn sess st
v, Session conn sess st
v))
            oldHm = [Session conn sess st] -> HashMap SessionId (Session conn sess st)
forall {conn} {sess} {st}.
[Session conn sess st] -> HashMap SessionId (Session conn sess st)
packSessionHm [Session conn sess st]
oldStatus
            newHm = [Session conn sess st] -> HashMap SessionId (Session conn sess st)
forall {conn} {sess} {st}.
[Session conn sess st] -> HashMap SessionId (Session conn sess st)
packSessionHm [Session conn sess st]
newStatus
        sh_removed (ssc_hooks cfg) (HM.map sess_data $ oldHm `HM.difference` newHm)
        threadDelay (1000 * 1000 * (round $ ssc_housekeepingInterval cfg))

createSession :: SessionCfg conn sess st -> sess -> IO (Session conn sess st)
createSession :: forall conn sess st.
SessionCfg conn sess st -> sess -> IO (Session conn sess st)
createSession SessionCfg conn sess st
sessCfg sess
content =
  IO UTCTime
getCurrentTime IO UTCTime
-> (UTCTime -> IO (Session conn sess st))
-> IO (Session conn sess st)
forall a b. IO a -> (a -> IO b) -> IO b
forall (m :: * -> *) a b. Monad m => m a -> (a -> m b) -> m b
>>= \UTCTime
now -> SessionCfg conn sess st
-> UTCTime -> sess -> IO (Session conn sess st)
forall conn sess st.
SessionCfg conn sess st
-> UTCTime -> sess -> IO (Session conn sess st)
createSessionAt SessionCfg conn sess st
sessCfg UTCTime
now sess
content