feat: JWT cache implementation based on sieve algorithm (#4084)
Changes: 1. Refactoring and some cleanup of JWT handling code: * Instead of caching AuthResult cache decoded claims (which signature was verified). Validating claims and determining role is done after cache lookup * Cleaned up API so that usage of it is simplified: lookupJwtCache cache key >>= parseClaims configJwtAud time * Handling of JwtCacheState initialization and updates of configuration is encapsulated in Auth.JwtCache module 2. Generic high performance (hopefully) scalable, dynamically resizeable cache implementation based on stm, stm-hamt and sieve algorithm. It also integrates with PostgREST measurements infrastructure providing usage stats (ie. hit ratio, evictions count)
This commit is contained in:
@@ -57,7 +57,7 @@ import Data.IORef (IORef, atomicWriteIORef, newIORef,
|
||||
readIORef)
|
||||
import Data.Time.Clock (UTCTime, getCurrentTime)
|
||||
|
||||
import PostgREST.Auth.JwtCache (JwtCacheState)
|
||||
import PostgREST.Auth.JwtCache (JwtCacheState, update)
|
||||
import PostgREST.Config (AppConfig (..),
|
||||
addFallbackAppName,
|
||||
readAppConfig)
|
||||
@@ -127,14 +127,13 @@ init conf@AppConfig{configLogLevel, configDbPoolSize} = do
|
||||
|
||||
observer $ AppStartObs prettyVersion
|
||||
|
||||
jwtCacheState <- JwtCache.init
|
||||
pool <- initPool conf observer
|
||||
(sock, adminSock) <- initSockets conf
|
||||
state' <- initWithPool (sock, adminSock) pool conf jwtCacheState loggerState metricsState observer
|
||||
state' <- initWithPool (sock, adminSock) pool conf loggerState metricsState observer
|
||||
pure state' { stateSocketREST = sock, stateSocketAdmin = adminSock}
|
||||
|
||||
initWithPool :: AppSockets -> SQL.Pool -> AppConfig -> JwtCache.JwtCacheState -> Logger.LoggerState -> Metrics.MetricsState -> ObservationHandler -> IO AppState
|
||||
initWithPool (sock, adminSock) pool conf jwtCacheState loggerState metricsState observer = do
|
||||
initWithPool :: AppSockets -> SQL.Pool -> AppConfig -> Logger.LoggerState -> Metrics.MetricsState -> ObservationHandler -> IO AppState
|
||||
initWithPool (sock, adminSock) pool conf loggerState metricsState observer = do
|
||||
|
||||
appState <- AppState pool
|
||||
<$> newIORef minimumPgVersion -- assume we're in a supported version when starting, this will be corrected on a later step
|
||||
@@ -150,7 +149,7 @@ initWithPool (sock, adminSock) pool conf jwtCacheState loggerState metricsState
|
||||
<*> pure sock
|
||||
<*> pure adminSock
|
||||
<*> pure observer
|
||||
<*> pure jwtCacheState
|
||||
<*> JwtCache.init conf observer
|
||||
<*> pure loggerState
|
||||
<*> pure metricsState
|
||||
|
||||
@@ -471,10 +470,7 @@ readInDbConfig startingUp appState@AppState{stateObserver=observer} = do
|
||||
-- After the config has reloaded, jwt-secret might have changed, so
|
||||
-- if it has changed, it is important to invalidate the jwt cache
|
||||
-- entries, because they were cached using the old secret
|
||||
if configJwtSecret conf == configJwtSecret newConf then
|
||||
pass
|
||||
else
|
||||
JwtCache.emptyCache (getJwtCacheState appState) -- atomic O(1) operation
|
||||
update (getJwtCacheState appState) newConf
|
||||
|
||||
if startingUp then
|
||||
pass
|
||||
|
||||
+13
-31
@@ -30,50 +30,32 @@ import System.TimeIt (timeItT)
|
||||
|
||||
import PostgREST.AppState (AppState, getConfig, getJwtCacheState,
|
||||
getTime)
|
||||
import PostgREST.Auth.Jwt (parseClaims)
|
||||
import PostgREST.Auth.JwtCache (lookupJwtCache)
|
||||
import PostgREST.Auth.Types (AuthResult (..))
|
||||
import PostgREST.Config (AppConfig (..))
|
||||
import PostgREST.Error (Error (..), JwtError (..))
|
||||
import PostgREST.Error (Error (..))
|
||||
|
||||
import qualified Data.Aeson.KeyMap as KM
|
||||
import PostgREST.Auth.Jwt (parseAndDecodeClaims,
|
||||
parseClaims)
|
||||
import Protolude
|
||||
import Protolude
|
||||
|
||||
-- | Validate authorization header.
|
||||
-- | Validate authorization header
|
||||
-- Parse and store JWT claims for future use in the request.
|
||||
middleware :: AppState -> Wai.Middleware
|
||||
middleware appState app req respond = do
|
||||
cfg@AppConfig{..} <- getConfig appState
|
||||
conf@AppConfig{..} <- getConfig appState
|
||||
time <- getTime appState
|
||||
|
||||
let token = Wai.extractBearerAuth =<< lookup HTTP.hAuthorization (Wai.requestHeaders req)
|
||||
parseAuthToken = maybe (const $ throwError (JwtErr JwtSecretMissing)) parseAndDecodeClaims configJWKS
|
||||
parseJwt = runExceptT $ maybe (pure KM.empty) parseAuthToken token >>= parseClaims cfg time
|
||||
parseJwt = runExceptT $ lookupJwtCache jwtCacheState token >>= parseClaims conf time
|
||||
jwtCacheState = getJwtCacheState appState
|
||||
|
||||
-- If ServerTimingEnabled -> calculate JWT validation time
|
||||
-- If JwtCacheMaxLifetime -> cache JWT validation result
|
||||
req' <- case (configServerTimingEnabled, configJwtCacheMaxLifetime) of
|
||||
(True, 0) -> do
|
||||
(dur, authResult) <- timeItT parseJwt
|
||||
return $ req { Wai.vault = Wai.vault req & Vault.insert authResultKey authResult & Vault.insert jwtDurKey dur }
|
||||
|
||||
(True, maxLifetime) -> do
|
||||
(dur, authResult) <- timeItT $ case token of
|
||||
Just tkn -> lookupJwtCache jwtCacheState tkn maxLifetime parseJwt time
|
||||
Nothing -> parseJwt
|
||||
return $ req { Wai.vault = Wai.vault req & Vault.insert authResultKey authResult & Vault.insert jwtDurKey dur }
|
||||
|
||||
(False, 0) -> do
|
||||
authResult <- parseJwt
|
||||
return $ req { Wai.vault = Wai.vault req & Vault.insert authResultKey authResult }
|
||||
|
||||
(False, maxLifetime) -> do
|
||||
authResult <- case token of
|
||||
Just tkn -> lookupJwtCache jwtCacheState tkn maxLifetime parseJwt time
|
||||
Nothing -> parseJwt
|
||||
return $ req { Wai.vault = Wai.vault req & Vault.insert authResultKey authResult }
|
||||
-- If ServerTimingEnabled -> calculate JWT validation time
|
||||
req' <- if configServerTimingEnabled then do
|
||||
(dur, authResult) <- timeItT parseJwt
|
||||
pure $ req { Wai.vault = Wai.vault req & Vault.insert authResultKey authResult & Vault.insert jwtDurKey dur }
|
||||
else do
|
||||
authResult <- parseJwt
|
||||
pure $ req { Wai.vault = Wai.vault req & Vault.insert authResultKey authResult }
|
||||
|
||||
app req' respond
|
||||
|
||||
|
||||
@@ -1,99 +1,114 @@
|
||||
{-|
|
||||
Module : PostgREST.Auth.JwtCache
|
||||
Description : PostgREST Jwt Authentication Result Cache.
|
||||
Description : PostgREST JWT validation results Cache.
|
||||
|
||||
This module provides functions to deal with the JWT cache
|
||||
This module provides functions to deal with the JWT cache.
|
||||
-}
|
||||
{-# LANGUAGE NamedFieldPuns #-}
|
||||
{-# LANGUAGE ExistentialQuantification #-}
|
||||
{-# LANGUAGE FlexibleInstances #-}
|
||||
{-# LANGUAGE LambdaCase #-}
|
||||
{-# LANGUAGE MultiParamTypeClasses #-}
|
||||
{-# LANGUAGE NamedFieldPuns #-}
|
||||
{-# LANGUAGE StrictData #-}
|
||||
|
||||
module PostgREST.Auth.JwtCache
|
||||
( init
|
||||
, update
|
||||
, JwtCacheState
|
||||
, lookupJwtCache
|
||||
, emptyCache
|
||||
) where
|
||||
|
||||
import qualified Data.Aeson as JSON
|
||||
import qualified Data.Aeson.KeyMap as KM
|
||||
import qualified Data.Cache as C
|
||||
import qualified Data.Scientific as Sci
|
||||
|
||||
import Control.Debounce
|
||||
import PostgREST.Error (Error (..), JwtError (JwtSecretMissing))
|
||||
|
||||
import Data.Time.Clock (UTCTime, nominalDiffTimeToSeconds)
|
||||
import Data.Time.Clock.POSIX (utcTimeToPOSIXSeconds)
|
||||
import System.Clock (TimeSpec (..))
|
||||
import Control.Concurrent.STM (newTVarIO, readTVar,
|
||||
writeTVar)
|
||||
import Control.Concurrent.STM.TVar (TVar)
|
||||
import Control.Monad.Error.Class (liftEither)
|
||||
import Data.ByteString hiding (all, init)
|
||||
import Data.IORef (IORef, newIORef,
|
||||
readIORef, writeIORef)
|
||||
import Jose.Jwk (JwkSet)
|
||||
import PostgREST.Auth.Jwt (parseAndDecodeClaims)
|
||||
import PostgREST.Cache.Sieve (alwaysValid)
|
||||
import qualified PostgREST.Cache.Sieve as SC
|
||||
import PostgREST.Config (AppConfig (..))
|
||||
import PostgREST.Observation (Observation (JwtCacheEviction, JwtCacheLookup),
|
||||
ObservationHandler)
|
||||
import Protolude
|
||||
|
||||
import PostgREST.Auth.Types (AuthResult (..))
|
||||
import PostgREST.Error (Error (..))
|
||||
data JwtCacheState = JwtCacheState ObservationHandler (IORef JwtCache)
|
||||
|
||||
import Protolude
|
||||
class CacheVariant m v where
|
||||
cached :: SC.Cache m ByteString v -> ByteString -> ExceptT Error IO JSON.Object
|
||||
|
||||
-- | JWT Cache and IO action that triggers purging old entries from the cache
|
||||
data JwtCacheState = JwtCacheState
|
||||
{ jwtCache :: C.Cache ByteString AuthResult
|
||||
, purgeCache :: IO ()
|
||||
}
|
||||
{-|
|
||||
Jwt caching can have three different configurations:
|
||||
* missing JWT Key (no caching and throw error when JWT token present in the request)
|
||||
* JWT cache turned off
|
||||
* JWT cache turned on
|
||||
|
||||
All three options are represented by JwtCache data type.
|
||||
|
||||
Handling of reconfiguration is centralized in this module.
|
||||
-}
|
||||
data JwtCache =
|
||||
JwtNoJwks |
|
||||
JwtNoCache JwkSet |
|
||||
forall m v. CacheVariant m v => JwtCache JwkSet (TVar Int) (SC.Cache m ByteString v)
|
||||
|
||||
instance CacheVariant IO (Either Error JSON.Object) where
|
||||
cached c = lift . SC.cached c >=> liftEither
|
||||
|
||||
instance CacheVariant (ExceptT Error IO) JSON.Object where
|
||||
cached = SC.cached
|
||||
|
||||
decode :: JwtCache -> ByteString -> ExceptT Error IO JSON.Object
|
||||
decode JwtNoJwks = const $ throwError (JwtErr JwtSecretMissing)
|
||||
decode (JwtNoCache key) = parseAndDecodeClaims key
|
||||
decode (JwtCache _ _ c) = cached c
|
||||
|
||||
-- | Reconfigure JWT caching and update JwtCacheState accordingly
|
||||
update :: JwtCacheState -> AppConfig -> IO ()
|
||||
update (JwtCacheState observationHandler jwtCacheState) config@AppConfig{configJWKS, configJwtCacheMaxEntries} =
|
||||
let reinitialize =
|
||||
newJwtCache config observationHandler
|
||||
>>= writeIORef jwtCacheState
|
||||
in
|
||||
readIORef jwtCacheState >>= \case
|
||||
(JwtCache decodingKey maxSize _) ->
|
||||
if configJWKS /= Just decodingKey || configJwtCacheMaxEntries <= 0 then
|
||||
-- reinitialize if key changed or cache disabled
|
||||
reinitialize
|
||||
else
|
||||
-- max size changed - set it and let the cache shrink itself if necessary
|
||||
atomically $ writeTVar maxSize configJwtCacheMaxEntries
|
||||
|
||||
_ -> reinitialize
|
||||
|
||||
init :: AppConfig -> ObservationHandler -> IO JwtCacheState
|
||||
init config = fmap (<$>) JwtCacheState <*> (newJwtCache config >=> newIORef)
|
||||
|
||||
-- | Initialize JwtCacheState
|
||||
init :: IO JwtCacheState
|
||||
init = do
|
||||
cache <- C.newCache Nothing -- no default expiration
|
||||
-- purgeExpired has O(n^2) complexity
|
||||
-- so we wrap it in debounce to make sure it:
|
||||
-- 1) is executed asynchronously
|
||||
-- 2) only a single purge operation is running at a time
|
||||
debounce <- mkDebounce defaultDebounceSettings
|
||||
-- debounceFreq is set to default 1 second
|
||||
{ debounceAction = C.purgeExpired cache
|
||||
, debounceEdge = leadingEdge
|
||||
}
|
||||
pure $ JwtCacheState cache debounce
|
||||
newJwtCache :: AppConfig -> ObservationHandler -> IO JwtCache
|
||||
newJwtCache AppConfig{configJWKS, configJwtCacheMaxEntries} observationHandler = do
|
||||
maybe (pure JwtNoJwks) initCache configJWKS
|
||||
where
|
||||
initCache key = if configJwtCacheMaxEntries <= 0 then pure (JwtNoCache key) else createCache key configJwtCacheMaxEntries
|
||||
|
||||
-- | Used to retrieve and insert JWT to JWT Cache
|
||||
lookupJwtCache :: JwtCacheState -> ByteString -> Int -> IO (Either Error AuthResult) -> UTCTime -> IO (Either Error AuthResult)
|
||||
lookupJwtCache JwtCacheState{jwtCache, purgeCache} token maxLifetime parseJwt utc = do
|
||||
checkCache <- C.lookup jwtCache token
|
||||
authResult <- maybe parseJwt (pure . Right) checkCache
|
||||
createCache key maxSize = do
|
||||
maxSizeTVar <- newTVarIO maxSize
|
||||
JwtCache key maxSizeTVar <$>
|
||||
notCachingErrors (readTVar maxSizeTVar) key
|
||||
|
||||
case (authResult,checkCache) of
|
||||
-- From comment:
|
||||
-- https://github.com/PostgREST/postgrest/pull/3801#discussion_r1857987914
|
||||
--
|
||||
-- We purge expired cache entries on a cache miss
|
||||
-- The reasoning is that:
|
||||
--
|
||||
-- 1. We expect it to be rare (otherwise there is no point of the cache)
|
||||
-- 2. It makes sure the cache is not growing (as inserting new entries
|
||||
-- does garbage collection)
|
||||
-- 3. Since this is time expiration based cache there is no real risk of
|
||||
-- starvation - sooner or later we are going to have a cache miss.
|
||||
notCachingErrors :: STM Int -> JwkSet -> IO (SC.Cache (ExceptT Error IO) ByteString JSON.Object)
|
||||
notCachingErrors maxSize key = SC.cacheIO (SC.CacheConfig maxSize
|
||||
(parseAndDecodeClaims key)
|
||||
(lift . observationHandler . JwtCacheLookup) -- lookup metrics
|
||||
(const . const $ lift $ observationHandler JwtCacheEviction) -- evictions metrics
|
||||
alwaysValid) -- no invalidation for now
|
||||
|
||||
(Right res, Nothing) -> do -- cache miss
|
||||
|
||||
let timeSpec = getTimeSpec res maxLifetime utc
|
||||
|
||||
-- insert new cache entry
|
||||
C.insert' jwtCache (Just timeSpec) token res
|
||||
|
||||
-- Execute IO action to purge the cache
|
||||
-- It is assumed this action returns immidiately
|
||||
-- so that request processing is not blocked.
|
||||
purgeCache
|
||||
|
||||
_ -> pure ()
|
||||
|
||||
return authResult
|
||||
|
||||
-- Used to extract JWT exp claim and add to JWT Cache
|
||||
getTimeSpec :: AuthResult -> Int -> UTCTime -> TimeSpec
|
||||
getTimeSpec res maxLifetime utc = do
|
||||
let expireJSON = KM.lookup "exp" (authClaims res)
|
||||
utcToSecs = floor . nominalDiffTimeToSeconds . utcTimeToPOSIXSeconds
|
||||
sciToInt = fromMaybe 0 . Sci.toBoundedInteger
|
||||
case expireJSON of
|
||||
Just (JSON.Number seconds) -> TimeSpec (sciToInt seconds - utcToSecs utc) 0
|
||||
_ -> TimeSpec (fromIntegral maxLifetime :: Int64) 0
|
||||
|
||||
-- | Empty the cache (done when the config is reloaded)
|
||||
emptyCache :: JwtCacheState -> IO ()
|
||||
emptyCache JwtCacheState{jwtCache} = C.purge jwtCache
|
||||
lookupJwtCache :: JwtCacheState -> Maybe ByteString -> ExceptT Error IO JSON.Object
|
||||
lookupJwtCache (JwtCacheState _ cacheState) k = liftIO (readIORef cacheState) >>= flip (maybe (pure KM.empty)) k . decode
|
||||
|
||||
@@ -203,8 +203,8 @@ exampleConfigFile =
|
||||
|# jwt-secret = "secret_with_at_least_32_characters"
|
||||
|jwt-secret-is-base64 = false
|
||||
|
|
||||
|## Enables and set JWT Cache max lifetime, disables caching with 0
|
||||
|# jwt-cache-max-lifetime = 0
|
||||
|## Enables JWT Cache and sets its max size, disables caching with 0
|
||||
|# jwt-cache-max-entries = 0
|
||||
|
|
||||
|## Logging level, the admitted values are: crit, error, warn, info and debug.
|
||||
|log-level = "error"
|
||||
|
||||
@@ -0,0 +1,218 @@
|
||||
{-|
|
||||
Module : PostgREST.Cache.Sieve
|
||||
Description : PostgREST cache implementation based on Sieve algorithm.
|
||||
|
||||
This module provides implementation of a mutable cache on Sieve algorithm.
|
||||
-}
|
||||
{-# LANGUAGE DataKinds #-}
|
||||
{-# LANGUAGE GADTs #-}
|
||||
{-# LANGUAGE LambdaCase #-}
|
||||
{-# LANGUAGE NamedFieldPuns #-}
|
||||
{-# LANGUAGE PolyKinds #-}
|
||||
{-# LANGUAGE RecordWildCards #-}
|
||||
{-# LANGUAGE RecursiveDo #-}
|
||||
{-# LANGUAGE StrictData #-}
|
||||
{-# LANGUAGE TupleSections #-}
|
||||
|
||||
module PostgREST.Cache.Sieve (
|
||||
Cache
|
||||
, CacheConfig (..)
|
||||
, Discard (..)
|
||||
, alwaysValid
|
||||
, cache
|
||||
, cacheIO
|
||||
, cached
|
||||
)
|
||||
where
|
||||
|
||||
import Control.Concurrent.STM
|
||||
import Control.Monad.Extra (whileM)
|
||||
import Data.Some
|
||||
import qualified Focus as F
|
||||
import Protolude hiding (elem, head)
|
||||
import qualified StmHamt.SizedHamt as SH
|
||||
|
||||
data ListNode k v (b :: Bool) = ListNode {
|
||||
nextPtr :: NodePtr k v,
|
||||
prevNextPtrPtr :: NodePtrPtr k v,
|
||||
elem :: NodeElem k v b
|
||||
}
|
||||
|
||||
data NodeElem :: Type -> Type -> Bool -> Type where
|
||||
Head :: {
|
||||
entries :: SH.SizedHamt (HamtEntry k v),
|
||||
finger :: NodePtrPtr k v
|
||||
} -> NodeElem k v False
|
||||
Entry :: Hashable k => {
|
||||
visited :: TVar Bool,
|
||||
ekey :: k,
|
||||
entryValue :: v
|
||||
} -> NodeElem k v True
|
||||
|
||||
type HamtEntry k v = ListNode k v True
|
||||
type AnyNode k v = Some (ListNode k v)
|
||||
type NodePtr k v = TVar (AnyNode k v)
|
||||
type NodePtrPtr k v = TVar (NodePtr k v)
|
||||
|
||||
data Discard m v = Refresh (m ()) | Invalid (m v)
|
||||
|
||||
data Cache m k v = (MonadIO m, Hashable k) => Cache (ListNode k v False) (CacheConfig m k v)
|
||||
|
||||
data CacheConfig m k v = CacheConfig {
|
||||
maxSize :: STM Int,
|
||||
load :: k -> m v,
|
||||
requestListener :: Bool -> m (),
|
||||
evictionListener :: k -> v -> m (),
|
||||
validator :: m (k -> v -> Maybe (Discard m v))
|
||||
}
|
||||
|
||||
alwaysValid :: Applicative m => m (k -> v -> Maybe (Discard m v))
|
||||
alwaysValid = pure (const . const Nothing)
|
||||
|
||||
cacheIO :: (MonadIO m, Hashable k) => CacheConfig m k v -> IO (Cache m k v)
|
||||
cacheIO = atomically . cache
|
||||
|
||||
cache :: (MonadIO m, Hashable k) => CacheConfig m k v -> STM (Cache m k v)
|
||||
cache cacheConfig = mdo
|
||||
tail <- newTVar (Some head)
|
||||
entries <- SH.new
|
||||
finger <- newTVar tail
|
||||
head <- ListNode tail <$> newTVar tail <*> pure Head {..}
|
||||
pure $ Cache head cacheConfig
|
||||
|
||||
cached :: Cache m k v -> k -> m v
|
||||
cached (Cache head@ListNode{prevNextPtrPtr=neck, elem=Head{..}} CacheConfig{..}) k = do
|
||||
checkValid <- validator
|
||||
tryMaybe
|
||||
-- Fast path: lookup value, update stats and return the value if found and valid
|
||||
((liftIO . atomically) (lookup checkValid) >>= notify (requestListener . isJust) >>= validate)
|
||||
-- Slow path: load/calculate value and insert it (if still not found)
|
||||
(do
|
||||
value <- load k
|
||||
whileM (not <$> tryInsert value)
|
||||
pure value)
|
||||
where
|
||||
tryMaybe f notFound = f >>= maybe notFound pure
|
||||
|
||||
notify = ((<$) <*>)
|
||||
|
||||
validate = fmap join . traverse (\case
|
||||
-- valid value
|
||||
(Right v) -> pure $ Just v
|
||||
-- refresh value
|
||||
(Left (Refresh act)) -> act $> Nothing
|
||||
-- discard value and return alt result
|
||||
(Left (Invalid res)) -> Just <$> res)
|
||||
|
||||
lookup checkValid = SH.focus focus (ekey . elem) k entries
|
||||
where
|
||||
focus = F.Focus
|
||||
-- not found
|
||||
(pure (Nothing, F.Leave))
|
||||
-- found
|
||||
-- check entry validity
|
||||
(\e@ListNode{elem=Entry{visited, entryValue}} ->
|
||||
maybe
|
||||
-- entry valid
|
||||
(mark visited True $> (Just $ Right entryValue, F.Leave))
|
||||
-- entry invalid
|
||||
-- remove it
|
||||
((removeEntry e $>) . (, F.Remove) . Just . Left)
|
||||
(checkValid k entryValue)
|
||||
)
|
||||
|
||||
mark t b = whenM ((/= b) <$> readTVar t) (writeTVar t b)
|
||||
|
||||
-- perform a single entry eviction and possibly insertion atomically
|
||||
-- returning False if could not insert
|
||||
-- (either because entry currently pointed by the finger was visited
|
||||
-- or because after this entry eviction the cache is still full)
|
||||
-- so that other threads don't have to wait when visiting entries.
|
||||
-- First check if entry is still not in the cache - this time inside transaction.
|
||||
--
|
||||
-- Execute evictionListener if an entry was evicted
|
||||
tryInsert value = do
|
||||
(result, evicted) <- liftIO . atomically $ do
|
||||
-- Use SH.focus to performa a single lookup instead of 2
|
||||
-- we cannot modify Hamt from inside focus
|
||||
-- so if there is any entry to remove
|
||||
-- we need to delete it after
|
||||
(res, evictedKey) <- SH.focus focus (ekey . elem) k entries
|
||||
case evictedKey of
|
||||
(Just Entry{ekey=entryKey, entryValue}) -> do
|
||||
SH.focus F.delete (ekey . elem) entryKey entries
|
||||
pure (res, evictionListener entryKey entryValue)
|
||||
Nothing -> pure (res, pure ())
|
||||
|
||||
evicted $> result
|
||||
where
|
||||
focus = F.Focus (do
|
||||
(hasSpace, evictedKey) <- evictionStep
|
||||
if hasSpace then do
|
||||
entry <- newLinkedEntry value
|
||||
-- done, maybe evicted, insert entry
|
||||
pure ((True, evictedKey), F.Set entry)
|
||||
else
|
||||
-- not done, maybe evicted, don't modify entries
|
||||
pure ((False, evictedKey), F.Leave))
|
||||
-- Entry found case
|
||||
(\ListNode{elem=Entry{visited}} -> do
|
||||
-- mark as visited
|
||||
mark visited True
|
||||
-- done, no evictions, don't modify entries
|
||||
pure ((True, Nothing), F.Leave))
|
||||
|
||||
-- if the cache is full precoesses a single node
|
||||
-- removing it if it is marked as unvisited
|
||||
-- or clearing visited mark
|
||||
-- returns True if there is space in the cache
|
||||
-- puts evictionListener in state if an entry was evicted
|
||||
evictionStep = do
|
||||
currDiff <- liftA2 (-) (SH.size entries) (max 1 <$> maxSize)
|
||||
if currDiff >= 0 then do
|
||||
-- no space in the cache
|
||||
-- need to evict an entry
|
||||
(nextFinger, evictedKey) <- readTVar finger >>= evict
|
||||
writeTVar finger nextFinger
|
||||
-- return if enough space and evicted key if any
|
||||
pure (isJust evictedKey && currDiff == 0, evictedKey)
|
||||
else
|
||||
-- there is space in the cache
|
||||
pure (True, Nothing)
|
||||
|
||||
evict :: TVar (Some (ListNode k v)) -> STM (NodePtr k v, Maybe (NodeElem k v True))
|
||||
evict = readTVar >=> \case
|
||||
(Some e@ListNode{nextPtr, prevNextPtrPtr, elem=elem@Entry{visited}}) -> do
|
||||
ifM (readTVar visited)
|
||||
|
||||
(writeTVar visited False $> (nextPtr, Nothing))
|
||||
|
||||
(unlinkEntry e *> fmap (, Just elem) (readTVar prevNextPtrPtr))
|
||||
-- skip head
|
||||
(Some ListNode{nextPtr, elem=Head{}}) -> evict nextPtr
|
||||
|
||||
unlinkEntry :: HamtEntry k v -> STM ()
|
||||
unlinkEntry (ListNode{nextPtr, prevNextPtrPtr=currPrev}) = do
|
||||
nextEntry <- readTVar nextPtr
|
||||
withSome nextEntry $ \e -> do
|
||||
prevNextPtr <- readTVar currPrev
|
||||
writeTVar (prevNextPtrPtr e) prevNextPtr
|
||||
writeTVar prevNextPtr nextEntry
|
||||
|
||||
newLinkedEntry v = do
|
||||
oldNeckNextPtr <- readTVar neck
|
||||
newNeckNextPtr <- newTVar (Some head)
|
||||
newNeck <- ListNode newNeckNextPtr <$>
|
||||
newTVar oldNeckNextPtr <*>
|
||||
(Entry <$> newTVar False <*> pure k <*> pure v)
|
||||
-- update pointers
|
||||
writeTVar oldNeckNextPtr (Some newNeck)
|
||||
writeTVar neck newNeckNextPtr
|
||||
-- return HAMT entry
|
||||
pure newNeck
|
||||
|
||||
removeEntry = fmap (*>) unlinkEntry <*> adjustFinger
|
||||
|
||||
adjustFinger ListNode{nextPtr, prevNextPtrPtr} =
|
||||
whenM ((nextPtr ==) <$> readTVar finger) $
|
||||
readTVar prevNextPtrPtr >>= writeTVar finger
|
||||
@@ -97,7 +97,7 @@ data AppConfig = AppConfig
|
||||
, configJwtRoleClaimKey :: JSPath
|
||||
, configJwtSecret :: Maybe BS.ByteString
|
||||
, configJwtSecretIsBase64 :: Bool
|
||||
, configJwtCacheMaxLifetime :: Int
|
||||
, configJwtCacheMaxEntries :: Int
|
||||
, configLogLevel :: LogLevel
|
||||
, configLogQuery :: LogQuery
|
||||
, configOpenApiMode :: OpenAPIMode
|
||||
@@ -177,7 +177,7 @@ toText conf =
|
||||
,("jwt-role-claim-key", q . T.intercalate mempty . fmap dumpJSPath . configJwtRoleClaimKey)
|
||||
,("jwt-secret", q . T.decodeUtf8 . showJwtSecret)
|
||||
,("jwt-secret-is-base64", T.toLower . show . configJwtSecretIsBase64)
|
||||
,("jwt-cache-max-lifetime", show . configJwtCacheMaxLifetime)
|
||||
,("jwt-cache-max-entries", show . configJwtCacheMaxEntries)
|
||||
,("log-level", q . dumpLogLevel . configLogLevel)
|
||||
,("log-query", q . dumpLogQuery . configLogQuery)
|
||||
,("openapi-mode", q . dumpOpenApiMode . configOpenApiMode)
|
||||
@@ -287,7 +287,7 @@ parser optPath env dbSettings roleSettings roleIsolationLvl =
|
||||
<*> (fromMaybe False <$> optWithAlias
|
||||
(optBool "jwt-secret-is-base64")
|
||||
(optBool "secret-is-base64"))
|
||||
<*> (fromMaybe 0 <$> optInt "jwt-cache-max-lifetime")
|
||||
<*> (fromMaybe 1000 <$> optInt "jwt-cache-max-entries")
|
||||
<*> parseLogLevel "log-level"
|
||||
<*> parseLogQuery "log-query"
|
||||
<*> parseOpenAPIMode "openapi-mode"
|
||||
|
||||
@@ -100,6 +100,12 @@ observationLogger loggerState logLevel obs = case obs of
|
||||
o@PoolRequestFullfilled ->
|
||||
when (logLevel >= LogDebug) $ do
|
||||
logWithZTime loggerState $ observationMessage o
|
||||
o@JwtCacheEviction ->
|
||||
when (logLevel >= LogDebug) $ do
|
||||
logWithZTime loggerState $ observationMessage o
|
||||
o@(JwtCacheLookup _) ->
|
||||
when (logLevel >= LogDebug) $ do
|
||||
logWithZTime loggerState $ observationMessage o
|
||||
o ->
|
||||
logWithZTime loggerState $ observationMessage o
|
||||
|
||||
|
||||
@@ -26,7 +26,10 @@ data MetricsState =
|
||||
poolWaiting :: Gauge,
|
||||
poolMaxSize :: Gauge,
|
||||
schemaCacheLoads :: Vector Label1 Counter,
|
||||
schemaCacheQueryTime :: Gauge
|
||||
schemaCacheQueryTime :: Gauge,
|
||||
jwtCacheRequests :: Counter,
|
||||
jwtCacheHits :: Counter,
|
||||
jwtCacheEvictions :: Counter
|
||||
}
|
||||
|
||||
init :: Int -> IO MetricsState
|
||||
@@ -37,7 +40,10 @@ init configDbPoolSize = do
|
||||
register (gauge (Info "pgrst_db_pool_waiting" "Requests waiting to acquire a pool connection")) <*>
|
||||
register (gauge (Info "pgrst_db_pool_max" "Max pool connections")) <*>
|
||||
register (vector "status" $ counter (Info "pgrst_schema_cache_loads_total" "The total number of times the schema cache was loaded")) <*>
|
||||
register (gauge (Info "pgrst_schema_cache_query_time_seconds" "The query time in seconds of the last schema cache load"))
|
||||
register (gauge (Info "pgrst_schema_cache_query_time_seconds" "The query time in seconds of the last schema cache load")) <*>
|
||||
register (counter (Info "pgrst_jwt_cache_requests_total" "The total number of JWT cache lookups")) <*>
|
||||
register (counter (Info "pgrst_jwt_cache_hits_total" "The total number of JWT cache hits")) <*>
|
||||
register (counter (Info "pgrst_jwt_cache_evictions_total" "The total number of JWT cache evictions"))
|
||||
setGauge (poolMaxSize metricState) (fromIntegral configDbPoolSize)
|
||||
pure metricState
|
||||
|
||||
@@ -63,6 +69,9 @@ observationMetrics MetricsState{..} obs = case obs of
|
||||
setGauge schemaCacheQueryTime resTime
|
||||
SchemaCacheErrorObs{} -> do
|
||||
withLabel schemaCacheLoads "FAIL" incCounter
|
||||
JwtCacheLookup True -> incCounter jwtCacheRequests *> incCounter jwtCacheHits
|
||||
JwtCacheLookup False -> incCounter jwtCacheRequests
|
||||
JwtCacheEviction -> incCounter jwtCacheEvictions
|
||||
_ ->
|
||||
pure ()
|
||||
|
||||
|
||||
@@ -60,6 +60,8 @@ data Observation
|
||||
| HasqlPoolObs SQL.Observation
|
||||
| PoolRequest
|
||||
| PoolRequestFullfilled
|
||||
| JwtCacheLookup Bool
|
||||
| JwtCacheEviction
|
||||
|
||||
data ObsFatalError = ServerAuthError | ServerPgrstBug | ServerError42P05 | ServerError08P01
|
||||
|
||||
@@ -151,6 +153,10 @@ observationMessage = \case
|
||||
"Trying to borrow a connection from pool"
|
||||
PoolRequestFullfilled ->
|
||||
"Borrowed a connection from the pool"
|
||||
JwtCacheLookup _ ->
|
||||
"Looked up a JWT in JWT cache"
|
||||
JwtCacheEviction ->
|
||||
"Evicted entry from JWT cache"
|
||||
where
|
||||
showMillis :: Double -> Text
|
||||
showMillis x = toS $ showFFloat (Just 1) (x * 1000) ""
|
||||
|
||||
Reference in New Issue
Block a user