Cleanup JWT expires
This commit is contained in:
@@ -18,6 +18,7 @@ module PostgREST.Auth (
|
|||||||
, tokenJWT
|
, tokenJWT
|
||||||
) where
|
) where
|
||||||
|
|
||||||
|
import Control.Monad (join)
|
||||||
import Data.Aeson (Value (..), Object)
|
import Data.Aeson (Value (..), Object)
|
||||||
import Data.Aeson.Types (emptyObject, emptyArray)
|
import Data.Aeson.Types (emptyObject, emptyArray)
|
||||||
import Data.Vector as V (null, head)
|
import Data.Vector as V (null, head)
|
||||||
@@ -52,8 +53,8 @@ claimsToSQL = map setVar . toList
|
|||||||
-}
|
-}
|
||||||
jwtClaims :: Text -> Text -> NominalDiffTime -> Maybe JWT.ClaimsMap
|
jwtClaims :: Text -> Text -> NominalDiffTime -> Maybe JWT.ClaimsMap
|
||||||
jwtClaims secret input time =
|
jwtClaims secret input time =
|
||||||
case claim JWT.exp of
|
case join $ claim JWT.exp of
|
||||||
Just (Just expires) ->
|
Just expires ->
|
||||||
if JWT.secondsSinceEpoch expires > time
|
if JWT.secondsSinceEpoch expires > time
|
||||||
then customClaims
|
then customClaims
|
||||||
else Nothing
|
else Nothing
|
||||||
|
|||||||
@@ -2,7 +2,6 @@ module Main where
|
|||||||
|
|
||||||
|
|
||||||
import PostgREST.App
|
import PostgREST.App
|
||||||
-- import PostgREST.QueryBuilder
|
|
||||||
import PostgREST.Config (AppConfig (..),
|
import PostgREST.Config (AppConfig (..),
|
||||||
minimumPgVersion,
|
minimumPgVersion,
|
||||||
prettyVersion,
|
prettyVersion,
|
||||||
@@ -24,7 +23,6 @@ import qualified Hasql.Postgres as P
|
|||||||
import Network.Wai
|
import Network.Wai
|
||||||
import Network.Wai.Handler.Warp hiding (Connection)
|
import Network.Wai.Handler.Warp hiding (Connection)
|
||||||
import Network.Wai.Middleware.RequestLogger (logStdout)
|
import Network.Wai.Middleware.RequestLogger (logStdout)
|
||||||
import Data.Time.Clock.POSIX (getPOSIXTime)
|
|
||||||
import System.IO (BufferMode (..),
|
import System.IO (BufferMode (..),
|
||||||
hSetBuffering, stderr,
|
hSetBuffering, stderr,
|
||||||
stdin, stdout)
|
stdin, stdout)
|
||||||
@@ -99,8 +97,7 @@ main = do
|
|||||||
-- print $ findRelation (fakeRels ++ allRels) "test" "pg_source" "clients"
|
-- print $ findRelation (fakeRels ++ allRels) "test" "pg_source" "clients"
|
||||||
|
|
||||||
runSettings appSettings $ middle $ \ req respond -> do
|
runSettings appSettings $ middle $ \ req respond -> do
|
||||||
time <- getPOSIXTime
|
|
||||||
body <- strictRequestBody req
|
body <- strictRequestBody req
|
||||||
resOrError <- liftIO $ H.session pool $ H.tx txSettings $
|
resOrError <- liftIO $ H.session pool $ H.tx txSettings $
|
||||||
runWithClaims conf time (app dbstructure conf body) req
|
runWithClaims conf (app dbstructure conf body) req
|
||||||
either (respond . errResponse) respond resOrError
|
either (respond . errResponse) respond resOrError
|
||||||
|
|||||||
@@ -7,7 +7,7 @@ import Data.Maybe (fromMaybe, isNothing)
|
|||||||
import Data.Monoid
|
import Data.Monoid
|
||||||
import Data.Text
|
import Data.Text
|
||||||
import Data.String.Conversions (cs)
|
import Data.String.Conversions (cs)
|
||||||
import Data.Time.Clock (NominalDiffTime)
|
import Data.Time.Clock.POSIX (getPOSIXTime)
|
||||||
import qualified Hasql as H
|
import qualified Hasql as H
|
||||||
import qualified Hasql.Postgres as P
|
import qualified Hasql.Postgres as P
|
||||||
|
|
||||||
@@ -27,17 +27,20 @@ import PostgREST.App (contentTypeForAccept)
|
|||||||
import PostgREST.Auth (setRole, jwtClaims, claimsToSQL)
|
import PostgREST.Auth (setRole, jwtClaims, claimsToSQL)
|
||||||
import PostgREST.Config (AppConfig (..), corsPolicy)
|
import PostgREST.Config (AppConfig (..), corsPolicy)
|
||||||
|
|
||||||
|
import System.IO.Unsafe (unsafePerformIO)
|
||||||
|
|
||||||
import Prelude hiding(concat)
|
import Prelude hiding(concat)
|
||||||
|
|
||||||
import qualified Data.Vector as V
|
import qualified Data.Vector as V
|
||||||
import qualified Hasql.Backend as B
|
import qualified Hasql.Backend as B
|
||||||
import qualified Data.Map.Lazy as M
|
import qualified Data.Map.Lazy as M
|
||||||
|
|
||||||
runWithClaims :: forall s. AppConfig -> NominalDiffTime ->
|
runWithClaims :: forall s. AppConfig ->
|
||||||
(Request -> H.Tx P.Postgres s Response) ->
|
(Request -> H.Tx P.Postgres s Response) ->
|
||||||
Request -> H.Tx P.Postgres s Response
|
Request -> H.Tx P.Postgres s Response
|
||||||
runWithClaims conf time app req = do
|
runWithClaims conf app req = do
|
||||||
_ <- H.unitEx $ stmt setAnon
|
_ <- H.unitEx $ stmt setAnon
|
||||||
|
let time = unsafePerformIO getPOSIXTime
|
||||||
case split (== ' ') (cs auth) of
|
case split (== ' ') (cs auth) of
|
||||||
("Bearer" : tokenStr : _) ->
|
("Bearer" : tokenStr : _) ->
|
||||||
case jwtClaims jwtSecret tokenStr time of
|
case jwtClaims jwtSecret tokenStr time of
|
||||||
@@ -50,13 +53,13 @@ runWithClaims conf time app req = do
|
|||||||
_ -> invalidJWT
|
_ -> invalidJWT
|
||||||
_ -> app req
|
_ -> app req
|
||||||
where
|
where
|
||||||
stmt = (flip $ flip B.Stmt V.empty) True
|
stmt c = B.Stmt c V.empty True
|
||||||
hdrs = requestHeaders req
|
hdrs = requestHeaders req
|
||||||
jwtSecret = (cs $ configJwtSecret conf) :: Text
|
jwtSecret = (cs $ configJwtSecret conf) :: Text
|
||||||
auth = fromMaybe "" $ lookup hAuthorization hdrs
|
auth = fromMaybe "" $ lookup hAuthorization hdrs
|
||||||
anon = cs $ configAnonRole conf
|
anon = cs $ configAnonRole conf
|
||||||
setAnon = setRole anon
|
setAnon = setRole anon
|
||||||
invalidJWT = return $ responseLBS status400 [] "Invalid JWT"
|
invalidJWT = return $ responseLBS status400 [("Content-Type","application/json")] "{\"message\":\"Invalid JWT\"}"
|
||||||
|
|
||||||
redirectInsecure :: Application -> Application
|
redirectInsecure :: Application -> Application
|
||||||
redirectInsecure app req respond = do
|
redirectInsecure app req respond = do
|
||||||
|
|||||||
+1
-3
@@ -11,7 +11,6 @@ import Hasql.Postgres as P
|
|||||||
import Data.String.Conversions (cs)
|
import Data.String.Conversions (cs)
|
||||||
import Data.Monoid
|
import Data.Monoid
|
||||||
import Data.Text hiding (map)
|
import Data.Text hiding (map)
|
||||||
import Data.Time.Clock.POSIX (getPOSIXTime)
|
|
||||||
import qualified Data.Vector as V
|
import qualified Data.Vector as V
|
||||||
import Control.Monad (void)
|
import Control.Monad (void)
|
||||||
import Control.Applicative
|
import Control.Applicative
|
||||||
@@ -74,10 +73,9 @@ withApp perform = do
|
|||||||
}
|
}
|
||||||
|
|
||||||
perform $ middle $ \req resp -> do
|
perform $ middle $ \req resp -> do
|
||||||
time <- getPOSIXTime
|
|
||||||
body <- strictRequestBody req
|
body <- strictRequestBody req
|
||||||
result <- liftIO $ H.session pool $ H.tx txSettings
|
result <- liftIO $ H.session pool $ H.tx txSettings
|
||||||
$ runWithClaims cfg time (app dbstructure cfg body) req
|
$ runWithClaims cfg (app dbstructure cfg body) req
|
||||||
either (resp . errResponse) resp result
|
either (resp . errResponse) resp result
|
||||||
|
|
||||||
where middle = defaultMiddle False
|
where middle = defaultMiddle False
|
||||||
|
|||||||
Reference in New Issue
Block a user