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:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user