diff --git a/src/PostgREST/Auth.hs b/src/PostgREST/Auth.hs index e8b1e2f47..a608b684c 100644 --- a/src/PostgREST/Auth.hs +++ b/src/PostgREST/Auth.hs @@ -18,6 +18,7 @@ module PostgREST.Auth ( , tokenJWT ) where +import Control.Monad (join) import Data.Aeson (Value (..), Object) import Data.Aeson.Types (emptyObject, emptyArray) import Data.Vector as V (null, head) @@ -52,8 +53,8 @@ claimsToSQL = map setVar . toList -} jwtClaims :: Text -> Text -> NominalDiffTime -> Maybe JWT.ClaimsMap jwtClaims secret input time = - case claim JWT.exp of - Just (Just expires) -> + case join $ claim JWT.exp of + Just expires -> if JWT.secondsSinceEpoch expires > time then customClaims else Nothing diff --git a/src/PostgREST/Main.hs b/src/PostgREST/Main.hs index f51a936ae..1a49e9862 100644 --- a/src/PostgREST/Main.hs +++ b/src/PostgREST/Main.hs @@ -2,7 +2,6 @@ module Main where import PostgREST.App --- import PostgREST.QueryBuilder import PostgREST.Config (AppConfig (..), minimumPgVersion, prettyVersion, @@ -24,7 +23,6 @@ import qualified Hasql.Postgres as P import Network.Wai import Network.Wai.Handler.Warp hiding (Connection) import Network.Wai.Middleware.RequestLogger (logStdout) -import Data.Time.Clock.POSIX (getPOSIXTime) import System.IO (BufferMode (..), hSetBuffering, stderr, stdin, stdout) @@ -99,8 +97,7 @@ main = do -- print $ findRelation (fakeRels ++ allRels) "test" "pg_source" "clients" runSettings appSettings $ middle $ \ req respond -> do - time <- getPOSIXTime body <- strictRequestBody req 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 diff --git a/src/PostgREST/Middleware.hs b/src/PostgREST/Middleware.hs index 9d4ea5fc4..3444fe167 100644 --- a/src/PostgREST/Middleware.hs +++ b/src/PostgREST/Middleware.hs @@ -7,7 +7,7 @@ import Data.Maybe (fromMaybe, isNothing) import Data.Monoid import Data.Text import Data.String.Conversions (cs) -import Data.Time.Clock (NominalDiffTime) +import Data.Time.Clock.POSIX (getPOSIXTime) import qualified Hasql as H import qualified Hasql.Postgres as P @@ -27,17 +27,20 @@ import PostgREST.App (contentTypeForAccept) import PostgREST.Auth (setRole, jwtClaims, claimsToSQL) import PostgREST.Config (AppConfig (..), corsPolicy) +import System.IO.Unsafe (unsafePerformIO) + import Prelude hiding(concat) import qualified Data.Vector as V import qualified Hasql.Backend as B 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 -runWithClaims conf time app req = do +runWithClaims conf app req = do _ <- H.unitEx $ stmt setAnon + let time = unsafePerformIO getPOSIXTime case split (== ' ') (cs auth) of ("Bearer" : tokenStr : _) -> case jwtClaims jwtSecret tokenStr time of @@ -50,13 +53,13 @@ runWithClaims conf time app req = do _ -> invalidJWT _ -> app req where - stmt = (flip $ flip B.Stmt V.empty) True + stmt c = B.Stmt c V.empty True hdrs = requestHeaders req jwtSecret = (cs $ configJwtSecret conf) :: Text auth = fromMaybe "" $ lookup hAuthorization hdrs anon = cs $ configAnonRole conf setAnon = setRole anon - invalidJWT = return $ responseLBS status400 [] "Invalid JWT" + invalidJWT = return $ responseLBS status400 [("Content-Type","application/json")] "{\"message\":\"Invalid JWT\"}" redirectInsecure :: Application -> Application redirectInsecure app req respond = do diff --git a/test/SpecHelper.hs b/test/SpecHelper.hs index a7e97f812..62fb98694 100644 --- a/test/SpecHelper.hs +++ b/test/SpecHelper.hs @@ -11,7 +11,6 @@ import Hasql.Postgres as P import Data.String.Conversions (cs) import Data.Monoid import Data.Text hiding (map) -import Data.Time.Clock.POSIX (getPOSIXTime) import qualified Data.Vector as V import Control.Monad (void) import Control.Applicative @@ -74,10 +73,9 @@ withApp perform = do } perform $ middle $ \req resp -> do - time <- getPOSIXTime body <- strictRequestBody req 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 where middle = defaultMiddle False