From 66e966d8641b877d1dbdb61799169266f51cf0da Mon Sep 17 00:00:00 2001 From: Taimoor Zaeem Date: Mon, 17 Feb 2025 12:45:26 +0500 Subject: [PATCH] refactor: move jwt caching logic to Auth/JwtCache.hs --- postgrest.cabal | 1 + src/PostgREST/AppState.hs | 23 +++++----- src/PostgREST/Auth.hs | 68 +++++------------------------ src/PostgREST/Auth/JwtCache.hs | 79 ++++++++++++++++++++++++++++++++++ test/spec/Main.hs | 10 +++-- 5 files changed, 109 insertions(+), 72 deletions(-) create mode 100644 src/PostgREST/Auth/JwtCache.hs diff --git a/postgrest.cabal b/postgrest.cabal index 2bf8e45e3..e5d81219a 100644 --- a/postgrest.cabal +++ b/postgrest.cabal @@ -49,6 +49,7 @@ library PostgREST.App PostgREST.AppState PostgREST.Auth + PostgREST.Auth.JwtCache PostgREST.Auth.Types PostgREST.CLI PostgREST.Config diff --git a/src/PostgREST/AppState.hs b/src/PostgREST/AppState.hs index 797400949..27131f25e 100644 --- a/src/PostgREST/AppState.hs +++ b/src/PostgREST/AppState.hs @@ -12,7 +12,7 @@ module PostgREST.AppState , getNextDelay , getNextListenerDelay , getTime - , getJwtCache + , getJwtCacheState , getSocketREST , getSocketAdmin , init @@ -31,7 +31,6 @@ module PostgREST.AppState ) where import qualified Data.ByteString.Char8 as BS -import qualified Data.Cache as C import Data.Either.Combinators (whenLeft) import qualified Data.Text as T (unpack) import qualified Hasql.Pool as SQL @@ -40,6 +39,7 @@ import qualified Hasql.Session as SQL import qualified Hasql.Transaction.Sessions as SQL import qualified Network.HTTP.Types.Status as HTTP import qualified Network.Socket as NS +import qualified PostgREST.Auth.JwtCache as JwtCache import qualified PostgREST.Error as Error import qualified PostgREST.Logger as Logger import qualified PostgREST.Metrics as Metrics @@ -57,7 +57,7 @@ import Data.IORef (IORef, atomicWriteIORef, newIORef, readIORef) import Data.Time.Clock (UTCTime, getCurrentTime) -import PostgREST.Auth.Types (AuthResult) +import PostgREST.Auth.JwtCache (JwtCacheState) import PostgREST.Config (AppConfig (..), addFallbackAppName, readAppConfig) @@ -99,14 +99,14 @@ data AppState = AppState , stateNextDelay :: IORef Int -- | Keeps track of the next delay for the listener , stateNextListenerDelay :: IORef Int - -- | JWT Cache - , jwtCache :: C.Cache ByteString AuthResult -- | Network socket for REST API , stateSocketREST :: NS.Socket -- | Network socket for the admin UI , stateSocketAdmin :: Maybe NS.Socket -- | Observation handler , stateObserver :: ObservationHandler + -- | JWT Cache + , stateJwtCache :: JwtCache.JwtCacheState , stateLogger :: Logger.LoggerState , stateMetrics :: Metrics.MetricsState } @@ -127,13 +127,14 @@ 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 loggerState metricsState observer + state' <- initWithPool (sock, adminSock) pool conf jwtCacheState loggerState metricsState observer pure state' { stateSocketREST = sock, stateSocketAdmin = adminSock} -initWithPool :: AppSockets -> SQL.Pool -> AppConfig -> Logger.LoggerState -> Metrics.MetricsState -> ObservationHandler -> IO AppState -initWithPool (sock, adminSock) pool conf loggerState metricsState observer = do +initWithPool :: AppSockets -> SQL.Pool -> AppConfig -> JwtCache.JwtCacheState -> Logger.LoggerState -> Metrics.MetricsState -> ObservationHandler -> IO AppState +initWithPool (sock, adminSock) pool conf jwtCacheState 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 @@ -146,10 +147,10 @@ initWithPool (sock, adminSock) pool conf loggerState metricsState observer = do <*> myThreadId <*> newIORef 0 <*> newIORef 1 - <*> C.newCache Nothing <*> pure sock <*> pure adminSock <*> pure observer + <*> pure jwtCacheState <*> pure loggerState <*> pure metricsState @@ -310,8 +311,8 @@ putConfig = atomicWriteIORef . stateConf getTime :: AppState -> IO UTCTime getTime = stateGetTime -getJwtCache :: AppState -> C.Cache ByteString AuthResult -getJwtCache = jwtCache +getJwtCacheState :: AppState -> JwtCacheState +getJwtCacheState = stateJwtCache getSocketREST :: AppState -> NS.Socket getSocketREST = stateSocketREST diff --git a/src/PostgREST/Auth.hs b/src/PostgREST/Auth.hs index 51a3d81af..a7fbc58d1 100644 --- a/src/PostgREST/Auth.hs +++ b/src/PostgREST/Auth.hs @@ -24,7 +24,6 @@ import qualified Data.Aeson.KeyMap as KM import qualified Data.Aeson.Types as JSON import qualified Data.ByteString as BS import qualified Data.ByteString.Lazy.Char8 as LBS -import qualified Data.Cache as C import qualified Data.Scientific as Sci import qualified Data.Text as T import qualified Data.Vault.Lazy as Vault @@ -40,20 +39,19 @@ import Data.Either.Combinators (mapLeft) import Data.List (lookup) import Data.Time.Clock (UTCTime, nominalDiffTimeToSeconds) import Data.Time.Clock.POSIX (utcTimeToPOSIXSeconds) -import System.Clock (TimeSpec (..)) import System.IO.Unsafe (unsafePerformIO) import System.TimeIt (timeItT) -import PostgREST.AppState (AppState, getConfig, getJwtCache, - getTime) -import PostgREST.Auth.Types (AuthResult (..)) -import PostgREST.Config (AppConfig (..), FilterExp (..), JSPath, - JSPathExp (..)) -import PostgREST.Error (Error (..)) +import PostgREST.AppState (AppState, getConfig, getJwtCacheState, + getTime) +import PostgREST.Auth.JwtCache (lookupJwtCache) +import PostgREST.Auth.Types (AuthResult (..)) +import PostgREST.Config (AppConfig (..), FilterExp (..), + JSPath, JSPathExp (..)) +import PostgREST.Error (Error (..)) import Protolude - -- | Receives the JWT secret and audience (from config) and a JWT and returns a -- JSON object of JWT claims. parseToken :: AppConfig -> ByteString -> UTCTime -> ExceptT Error IO JSON.Value @@ -152,8 +150,9 @@ middleware appState app req respond = do let token = fromMaybe "" $ Wai.extractBearerAuth =<< lookup HTTP.hAuthorization (Wai.requestHeaders req) parseJwt = runExceptT $ parseToken conf token time >>= parseClaims conf + jwtCacheState = getJwtCacheState appState --- If DbPlanEnabled -> calculate JWT validation time +-- If ServerTimingEnabled -> calculate JWT validation time -- If JwtCacheMaxLifetime -> cache JWT validation result req' <- case (configServerTimingEnabled conf, configJwtCacheMaxLifetime conf) of (True, 0) -> do @@ -161,7 +160,7 @@ middleware appState app req respond = do return $ req { Wai.vault = Wai.vault req & Vault.insert authResultKey authResult & Vault.insert jwtDurKey dur } (True, maxLifetime) -> do - (dur, authResult) <- timeItT $ getJWTFromCache appState token maxLifetime parseJwt time + (dur, authResult) <- timeItT $ lookupJwtCache jwtCacheState token maxLifetime parseJwt time return $ req { Wai.vault = Wai.vault req & Vault.insert authResultKey authResult & Vault.insert jwtDurKey dur } (False, 0) -> do @@ -169,56 +168,11 @@ middleware appState app req respond = do return $ req { Wai.vault = Wai.vault req & Vault.insert authResultKey authResult } (False, maxLifetime) -> do - authResult <- getJWTFromCache appState token maxLifetime parseJwt time + authResult <- lookupJwtCache jwtCacheState token maxLifetime parseJwt time return $ req { Wai.vault = Wai.vault req & Vault.insert authResultKey authResult } app req' respond --- | Used to retrieve and insert JWT to JWT Cache -getJWTFromCache :: AppState -> ByteString -> Int -> IO (Either Error AuthResult) -> UTCTime -> IO (Either Error AuthResult) -getJWTFromCache appState token maxLifetime parseJwt utc = do - checkCache <- C.lookup (getJwtCache appState) token - authResult <- maybe parseJwt (pure . Right) checkCache - - 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. - - (Right res, Nothing) -> do -- cache miss - - let timeSpec = getTimeSpec res maxLifetime utc - - -- purge expired cache entries - C.purgeExpired jwtCache - - -- insert new cache entry - C.insert' jwtCache timeSpec token res - - _ -> pure () - - return authResult - where - jwtCache = getJwtCache appState - --- Used to extract JWT exp claim and add to JWT Cache -getTimeSpec :: AuthResult -> Int -> UTCTime -> Maybe 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) -> Just $ TimeSpec (sciToInt seconds - utcToSecs utc) 0 - _ -> Just $ TimeSpec (fromIntegral maxLifetime :: Int64) 0 - authResultKey :: Vault.Key (Either Error AuthResult) authResultKey = unsafePerformIO Vault.newKey {-# NOINLINE authResultKey #-} diff --git a/src/PostgREST/Auth/JwtCache.hs b/src/PostgREST/Auth/JwtCache.hs new file mode 100644 index 000000000..e02193a9c --- /dev/null +++ b/src/PostgREST/Auth/JwtCache.hs @@ -0,0 +1,79 @@ +{-| +Module : PostgREST.Auth.JwtCache +Description : PostgREST Jwt Authentication Result Cache. + +This module provides functions to deal with the JWT cache +-} +{-# LANGUAGE NamedFieldPuns #-} +module PostgREST.Auth.JwtCache + ( init + , JwtCacheState + , lookupJwtCache + ) 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 Data.Time.Clock (UTCTime, nominalDiffTimeToSeconds) +import Data.Time.Clock.POSIX (utcTimeToPOSIXSeconds) +import System.Clock (TimeSpec (..)) + +import PostgREST.Auth.Types (AuthResult (..)) +import PostgREST.Error (Error (..)) + +import Protolude + +newtype JwtCacheState = JwtCacheState + { jwtCache :: C.Cache ByteString AuthResult + } + +-- | Initialize JwtCacheState +init :: IO JwtCacheState +init = do + cache <- C.newCache Nothing -- no default expiration + return $ JwtCacheState cache + +-- | 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} token maxLifetime parseJwt utc = do + checkCache <- C.lookup jwtCache token + authResult <- maybe parseJwt (pure . Right) checkCache + + 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. + + (Right res, Nothing) -> do -- cache miss + + let timeSpec = getTimeSpec res maxLifetime utc + + -- purge expired cache entries + C.purgeExpired jwtCache + + -- insert new cache entry + C.insert' jwtCache (Just timeSpec) token res + + _ -> 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 diff --git a/test/spec/Main.hs b/test/spec/Main.hs index 988ded030..09fa5ac77 100644 --- a/test/spec/Main.hs +++ b/test/spec/Main.hs @@ -15,9 +15,10 @@ import PostgREST.SchemaCache (querySchemaCache) import Protolude hiding (toList, toS) import SpecHelper -import qualified PostgREST.AppState as AppState -import qualified PostgREST.Logger as Logger -import qualified PostgREST.Metrics as Metrics +import qualified PostgREST.AppState as AppState +import qualified PostgREST.Auth.JwtCache as JwtCache +import qualified PostgREST.Logger as Logger +import qualified PostgREST.Metrics as Metrics import qualified Feature.Auth.AsymmetricJwtSpec import qualified Feature.Auth.AudienceJwtSecretSpec @@ -84,12 +85,13 @@ main = do -- cached schema cache so most tests run fast baseSchemaCache <- loadSCache pool testCfg sockets <- AppState.initSockets testCfg + jwtCacheState <- JwtCache.init loggerState <- Logger.init metricsState <- Metrics.init (configDbPoolSize testCfg) let initApp sCache config = do - appState <- AppState.initWithPool sockets pool config loggerState metricsState (const $ pure ()) + appState <- AppState.initWithPool sockets pool config jwtCacheState loggerState metricsState (const $ pure ()) AppState.putPgVersion appState actualPgVersion AppState.putSchemaCache appState (Just sCache) return ((), postgrest (configLogLevel config) appState (pure ()))