Expose all claims via sql postgrest.claims
This commit is contained in:
+24
-20
@@ -18,12 +18,13 @@ module PostgREST.Auth (
|
||||
, tokenJWT
|
||||
) where
|
||||
|
||||
import Control.Monad (join)
|
||||
import Data.Aeson (Value (..), parseJSON)
|
||||
import Control.Lens
|
||||
import Data.Aeson (Value (..), parseJSON, toJSON)
|
||||
import Data.Aeson.Lens
|
||||
import Data.Aeson.Types (parseMaybe, emptyObject, emptyArray)
|
||||
import qualified Data.ByteString as BS
|
||||
import Data.Vector as V (null, head)
|
||||
import Data.Map as M (toList)
|
||||
import qualified Data.Vector as V
|
||||
import qualified Data.HashMap.Strict as M
|
||||
import Data.Maybe (fromMaybe)
|
||||
import Data.Monoid ((<>))
|
||||
import Data.String.Conversions (cs)
|
||||
@@ -39,12 +40,12 @@ import qualified Web.JWT as JWT
|
||||
this one is mapped to a SET ROLE statement.
|
||||
In case there is any problem decoding the JWT it returns Nothing.
|
||||
-}
|
||||
claimsToSQL :: JWT.ClaimsMap -> [BS.ByteString]
|
||||
claimsToSQL = map setVar . toList
|
||||
claimsToSQL :: M.HashMap Text Value -> [BS.ByteString]
|
||||
claimsToSQL = map setVar . M.toList
|
||||
where
|
||||
setVar ("role", String val) = setRole val
|
||||
setVar (k, val) = "set local postgrest.claims." <> cs (pgFmtIdent k) <>
|
||||
" = " <> cs (valueToVariable val) <> ";"
|
||||
setVar (k, val) = "set local " <> cs (pgFmtIdent $ "postgrest.claims." <> k)
|
||||
<> " = " <> cs (valueToVariable val) <> ";"
|
||||
valueToVariable = pgFmtLit . unquoted
|
||||
|
||||
{-|
|
||||
@@ -52,19 +53,22 @@ claimsToSQL = map setVar . toList
|
||||
returns a map of JWT claims
|
||||
In case there is any problem decoding the JWT it returns Nothing.
|
||||
-}
|
||||
jwtClaims :: JWT.Secret -> Text -> NominalDiffTime -> Maybe JWT.ClaimsMap
|
||||
|
||||
|
||||
jwtClaims :: JWT.Secret -> Text -> NominalDiffTime -> Either Text (M.HashMap Text Value)
|
||||
jwtClaims secret input time =
|
||||
case join $ claim JWT.exp of
|
||||
Just expires ->
|
||||
if JWT.secondsSinceEpoch expires > time
|
||||
then customClaims
|
||||
else Nothing
|
||||
_ -> customClaims
|
||||
where
|
||||
decoded = JWT.decodeAndVerifySignature secret input
|
||||
claim :: (JWT.JWTClaimsSet -> a) -> Maybe a
|
||||
claim prop = prop . JWT.claims <$> decoded
|
||||
customClaims = claim JWT.unregisteredClaims
|
||||
case mClaims of
|
||||
Nothing -> Right M.empty
|
||||
Just claims -> do
|
||||
let mExp = claims ^? key "exp" . _Integer
|
||||
expired = fromMaybe False $ (<= time) . fromInteger <$> mExp
|
||||
if expired
|
||||
then Left "JWT expired"
|
||||
else Right (value2map claims)
|
||||
where
|
||||
mClaims = toJSON . JWT.claims <$> JWT.decodeAndVerifySignature secret input
|
||||
value2map (Object o) = o
|
||||
value2map _ = M.empty
|
||||
|
||||
{-| Receives the name of a role and returns a SET ROLE statement -}
|
||||
setRole :: Text -> BS.ByteString
|
||||
|
||||
+16
-16
@@ -3,6 +3,7 @@
|
||||
|
||||
module PostgREST.Middleware where
|
||||
|
||||
import qualified Data.HashMap.Strict as M
|
||||
import Data.Maybe (fromMaybe)
|
||||
import Data.Text
|
||||
import Data.String.Conversions (cs)
|
||||
@@ -17,38 +18,37 @@ import Network.Wai.Middleware.Cors (cors)
|
||||
import Network.Wai.Middleware.Gzip (def, gzip)
|
||||
import Network.Wai.Middleware.Static (only, staticPolicy)
|
||||
|
||||
import PostgREST.ApiRequest (pickContentType)
|
||||
import PostgREST.ApiRequest (pickContentType)
|
||||
import PostgREST.Auth (setRole, jwtClaims, claimsToSQL)
|
||||
import PostgREST.Config (AppConfig (..), corsPolicy)
|
||||
import PostgREST.Error (errResponse)
|
||||
|
||||
import Prelude hiding(concat)
|
||||
|
||||
import qualified Data.Map.Lazy as M
|
||||
import Prelude hiding (concat, null)
|
||||
|
||||
runWithClaims :: AppConfig -> NominalDiffTime ->
|
||||
(Request -> H.Transaction Response) ->
|
||||
Request -> H.Transaction Response
|
||||
runWithClaims conf time app req = do
|
||||
H.sql setAnon
|
||||
case split (== ' ') (cs auth) of
|
||||
("Bearer" : tokenStr : _) ->
|
||||
case jwtClaims jwtSecret tokenStr time of
|
||||
Just claims ->
|
||||
if M.member "role" claims
|
||||
then do
|
||||
mapM_ H.sql $ claimsToSQL claims
|
||||
app req
|
||||
else invalidJWT
|
||||
_ -> invalidJWT
|
||||
_ -> app req
|
||||
let tokenStr = case split (== ' ') (cs auth) of
|
||||
("Bearer" : t : _) -> t
|
||||
_ -> ""
|
||||
eClaims = jwtClaims jwtSecret tokenStr time
|
||||
case eClaims of
|
||||
Left e -> clientErr e
|
||||
Right claims ->
|
||||
if M.null claims && not (null tokenStr)
|
||||
then clientErr "Invalid JWT"
|
||||
else do
|
||||
mapM_ H.sql $ claimsToSQL claims
|
||||
app req
|
||||
where
|
||||
hdrs = requestHeaders req
|
||||
jwtSecret = configJwtSecret conf
|
||||
auth = fromMaybe "" $ lookup hAuthorization hdrs
|
||||
anon = cs $ configAnonRole conf
|
||||
setAnon = setRole anon
|
||||
invalidJWT = return $ errResponse status400 "Invalid JWT"
|
||||
clientErr = return . errResponse status400
|
||||
|
||||
unsupportedAccept :: Application -> Application
|
||||
unsupportedAccept app req respond =
|
||||
|
||||
Reference in New Issue
Block a user