Support asymmetric JWK (#919)
This commit is contained in:
+9
-12
@@ -11,7 +11,6 @@ import qualified Data.ByteString.Char8 as BS
|
||||
import Data.Maybe
|
||||
import Data.IORef (IORef, readIORef)
|
||||
import Data.Text (intercalate)
|
||||
import Data.Time.Clock.POSIX (POSIXTime)
|
||||
|
||||
import qualified Hasql.Pool as P
|
||||
import qualified Hasql.Transaction as HT
|
||||
@@ -22,7 +21,6 @@ import Network.HTTP.Types.Status
|
||||
import Network.HTTP.Types.URI (renderSimpleQuery)
|
||||
import Network.Wai
|
||||
import Network.Wai.Middleware.RequestLogger (logStdout)
|
||||
import Web.JWT (binarySecret)
|
||||
|
||||
import qualified Data.Vector as V
|
||||
import qualified Hasql.Transaction as H
|
||||
@@ -35,7 +33,7 @@ import PostgREST.ApiRequest ( ApiRequest(..), ContentType(..)
|
||||
, mutuallyAgreeable
|
||||
, userApiRequest
|
||||
)
|
||||
import PostgREST.Auth (jwtClaims, containsRole)
|
||||
import PostgREST.Auth (jwtClaims, containsRole, parseJWK)
|
||||
import PostgREST.Config (AppConfig (..))
|
||||
import PostgREST.DbStructure
|
||||
import PostgREST.DbRequestBuilder( readRequest
|
||||
@@ -63,13 +61,12 @@ import Data.Function (id)
|
||||
import Protolude hiding (intercalate, Proxy)
|
||||
import Safe (headMay)
|
||||
|
||||
postgrest :: AppConfig -> IORef (Maybe DbStructure) -> P.Pool -> IO POSIXTime ->
|
||||
IO () -> Application
|
||||
postgrest conf refDbStructure pool getTime worker =
|
||||
let middle = (if configQuiet conf then id else logStdout) . defaultMiddle in
|
||||
postgrest :: AppConfig -> IORef (Maybe DbStructure) -> P.Pool -> IO () -> Application
|
||||
postgrest conf refDbStructure pool worker =
|
||||
let middle = (if configQuiet conf then id else logStdout) . defaultMiddle
|
||||
jwtSecret = parseJWK <$> configJwtSecret conf in
|
||||
|
||||
middle $ \ req respond -> do
|
||||
time <- getTime
|
||||
body <- strictRequestBody req
|
||||
maybeDbStructure <- readIORef refDbStructure
|
||||
case maybeDbStructure of
|
||||
@@ -78,9 +75,9 @@ postgrest conf refDbStructure pool getTime worker =
|
||||
response <- case userApiRequest (configSchema conf) req body of
|
||||
Left err -> return $ apiRequestError err
|
||||
Right apiRequest -> do
|
||||
let jwtSecret = binarySecret <$> configJwtSecret conf
|
||||
eClaims = jwtClaims jwtSecret (iJWT apiRequest) time
|
||||
authed = containsRole eClaims
|
||||
eClaims <- jwtClaims jwtSecret (toS $ iJWT apiRequest)
|
||||
|
||||
let authed = containsRole eClaims
|
||||
handleReq = runWithClaims conf eClaims (app dbStructure conf) apiRequest
|
||||
txMode = transactionMode dbStructure
|
||||
(iTarget apiRequest) (iAction apiRequest)
|
||||
@@ -324,7 +321,7 @@ responseContentTypeOrError accepts action = serves contentTypesForRequest accept
|
||||
case mutuallyAgreeable sProduces cAccepts of
|
||||
Nothing -> do
|
||||
let failed = intercalate ", " $ map (toS . toMime) cAccepts
|
||||
Left $ simpleError status415 $
|
||||
Left $ simpleError status415 [] $
|
||||
"None of these Content-Types are available: " <> failed
|
||||
Just ct -> Right ct
|
||||
|
||||
|
||||
+51
-39
@@ -14,63 +14,48 @@ very simple authentication system inside the PostgreSQL database.
|
||||
module PostgREST.Auth (
|
||||
containsRole
|
||||
, jwtClaims
|
||||
, tokenJWT
|
||||
, JWTAttempt(..)
|
||||
, parseJWK
|
||||
) where
|
||||
|
||||
import Protolude
|
||||
import Protolude hiding ((&))
|
||||
import Control.Lens
|
||||
import Data.Aeson (Value (..), parseJSON, toJSON)
|
||||
import Data.Aeson.Lens
|
||||
import Data.Aeson.Types (parseMaybe, emptyObject, emptyArray)
|
||||
import qualified Data.Vector as V
|
||||
import Data.Aeson (Value (..), decode, toJSON)
|
||||
import qualified Data.ByteString.Lazy as BL
|
||||
import qualified Data.HashMap.Strict as M
|
||||
import Data.Maybe (fromJust)
|
||||
import Data.Time.Clock (NominalDiffTime)
|
||||
import qualified Web.JWT as JWT
|
||||
|
||||
import Crypto.JOSE.Compact
|
||||
import Crypto.JOSE.JWK
|
||||
import Crypto.JOSE.JWS
|
||||
import Crypto.JOSE.Types
|
||||
import Crypto.JWT
|
||||
|
||||
{-|
|
||||
Possible situations encountered with client JWTs
|
||||
-}
|
||||
data JWTAttempt = JWTExpired
|
||||
| JWTInvalid
|
||||
data JWTAttempt = JWTInvalid JWTError
|
||||
| JWTMissingSecret
|
||||
| JWTClaims (M.HashMap Text Value)
|
||||
deriving Eq
|
||||
deriving (Eq, Show)
|
||||
|
||||
{-|
|
||||
Receives the JWT secret (from config) and a JWT and returns a map
|
||||
of JWT claims.
|
||||
-}
|
||||
jwtClaims :: Maybe JWT.Secret -> Text -> NominalDiffTime -> JWTAttempt
|
||||
jwtClaims _ "" _ = JWTClaims M.empty
|
||||
jwtClaims secret jwt time =
|
||||
jwtClaims :: Maybe JWK -> BL.ByteString -> IO JWTAttempt
|
||||
jwtClaims _ "" = return $ JWTClaims M.empty
|
||||
jwtClaims secret payload =
|
||||
case secret of
|
||||
Nothing -> JWTMissingSecret
|
||||
Just s ->
|
||||
let mClaims = toJSON . JWT.claims <$> JWT.decodeAndVerifySignature s jwt in
|
||||
case isExpired <$> mClaims of
|
||||
Just True -> JWTExpired
|
||||
Nothing -> JWTInvalid
|
||||
Just False -> JWTClaims $ value2map $ fromJust mClaims
|
||||
where
|
||||
isExpired claims =
|
||||
let mExp = claims ^? key "exp" . _Integer
|
||||
in fromMaybe False $ (<= time) . fromInteger <$> mExp
|
||||
value2map (Object o) = o
|
||||
value2map _ = M.empty
|
||||
|
||||
{-|
|
||||
Receives the JWT secret (from config) and a JWT and a JSON value
|
||||
and returns a signed JWT.
|
||||
-}
|
||||
tokenJWT :: JWT.Secret -> Value -> Text
|
||||
tokenJWT secret (Array arr) =
|
||||
let obj = if V.null arr then emptyObject else V.head arr
|
||||
jcs = parseMaybe parseJSON obj :: Maybe JWT.JWTClaimsSet in
|
||||
JWT.encodeSigned JWT.HS256 secret $ fromMaybe JWT.def jcs
|
||||
tokenJWT secret _ = tokenJWT secret emptyArray
|
||||
Nothing -> return JWTMissingSecret
|
||||
Just jwk -> do
|
||||
let validation = defaultJWTValidationSettings
|
||||
eJwt <- runExceptT $ do
|
||||
jwt <- decodeCompact payload
|
||||
validateJWSJWT validation jwk jwt
|
||||
return jwt
|
||||
return $ case eJwt of
|
||||
Left e -> JWTInvalid e
|
||||
Right jwt -> JWTClaims . claims2map . jwtClaimsSet $ jwt
|
||||
|
||||
{-|
|
||||
Whether a response from jwtClaims contains a role claim
|
||||
@@ -78,3 +63,30 @@ tokenJWT secret _ = tokenJWT secret emptyArray
|
||||
containsRole :: JWTAttempt -> Bool
|
||||
containsRole (JWTClaims claims) = M.member "role" claims
|
||||
containsRole _ = False
|
||||
|
||||
{-|
|
||||
Internal helper used to turn JWT ClaimSet into something
|
||||
easier to work with
|
||||
-}
|
||||
claims2map :: ClaimsSet -> M.HashMap Text Value
|
||||
claims2map = val2map . toJSON
|
||||
where
|
||||
val2map (Object o) = o
|
||||
val2map _ = M.empty
|
||||
|
||||
parseJWK :: ByteString -> JWK
|
||||
parseJWK str =
|
||||
fromMaybe (hs256jwk str) (decode (toS str) :: Maybe JWK)
|
||||
|
||||
{-|
|
||||
Internal helper to generate HMAC-SHA256. When the jwt key in the
|
||||
config file is a simple string rather than a JWK object, we'll
|
||||
apply this function to it.
|
||||
-}
|
||||
hs256jwk :: ByteString -> JWK
|
||||
hs256jwk key =
|
||||
fromKeyMaterial km
|
||||
& jwkUse .~ Just Sig
|
||||
& jwkAlg .~ (Just $ JWSAlg HS256)
|
||||
where
|
||||
km = OctKeyMaterial (OctKeyParameters Oct (Base64Octets key))
|
||||
|
||||
+13
-9
@@ -18,12 +18,15 @@ import qualified Data.Aeson as JSON
|
||||
import Data.Text (unwords)
|
||||
import qualified Hasql.Pool as P
|
||||
import qualified Hasql.Session as H
|
||||
import Network.HTTP.Types.Header
|
||||
import qualified Network.HTTP.Types.Status as HT
|
||||
import Network.Wai (Response, responseLBS)
|
||||
import PostgREST.Types
|
||||
|
||||
apiRequestError :: ApiRequestError -> Response
|
||||
apiRequestError err = errorResponse status err
|
||||
apiRequestError err =
|
||||
errorResponse status
|
||||
[toHeader CTApplicationJSON] err
|
||||
where
|
||||
status =
|
||||
case err of
|
||||
@@ -35,13 +38,14 @@ apiRequestError err = errorResponse status err
|
||||
InvalidRange -> HT.status416
|
||||
UnknownRelation -> HT.status404
|
||||
|
||||
simpleError :: HT.Status -> Text -> Response
|
||||
simpleError status message =
|
||||
errorResponse status $ JSON.object ["message" .= message]
|
||||
simpleError :: HT.Status -> [Header] -> Text -> Response
|
||||
simpleError status hdrs message =
|
||||
errorResponse status (toHeader CTApplicationJSON : hdrs) $
|
||||
JSON.object ["message" .= message]
|
||||
|
||||
errorResponse :: JSON.ToJSON a => HT.Status -> a -> Response
|
||||
errorResponse status e =
|
||||
responseLBS status [toHeader CTApplicationJSON] $ encodeError e
|
||||
errorResponse :: JSON.ToJSON a => HT.Status -> [Header] -> a -> Response
|
||||
errorResponse status hdrs e =
|
||||
responseLBS status hdrs $ encodeError e
|
||||
|
||||
pgError :: Bool -> P.UsageError -> Response
|
||||
pgError authed e =
|
||||
@@ -71,12 +75,12 @@ singularityError numRows =
|
||||
|
||||
binaryFieldError :: Response
|
||||
binaryFieldError =
|
||||
simpleError HT.status406 (toS (toMime CTOctetStream) <>
|
||||
simpleError HT.status406 [] (toS (toMime CTOctetStream) <>
|
||||
" requested but a single column was not selected")
|
||||
|
||||
connectionLostError :: Response
|
||||
connectionLostError =
|
||||
simpleError HT.status503 "Database connection lost, retrying the connection."
|
||||
simpleError HT.status503 [] "Database connection lost, retrying the connection."
|
||||
|
||||
encodeError :: JSON.ToJSON a => a -> LByteString
|
||||
encodeError = JSON.encode
|
||||
|
||||
+11
-13
@@ -4,13 +4,13 @@
|
||||
|
||||
module PostgREST.Middleware where
|
||||
|
||||
import Crypto.JWT
|
||||
import Data.Aeson (Value (..))
|
||||
import qualified Data.HashMap.Strict as M
|
||||
import qualified Hasql.Transaction as H
|
||||
|
||||
import Network.HTTP.Types.Status (unauthorized401, status500)
|
||||
import Network.Wai (Application, Response,
|
||||
responseLBS)
|
||||
import Network.Wai (Application, Response)
|
||||
import Network.Wai.Middleware.Cors (cors)
|
||||
import Network.Wai.Middleware.Gzip (def, gzip)
|
||||
import Network.Wai.Middleware.Static (only, staticPolicy)
|
||||
@@ -19,7 +19,6 @@ import PostgREST.ApiRequest (ApiRequest(..))
|
||||
import PostgREST.Auth (JWTAttempt(..))
|
||||
import PostgREST.Config (AppConfig (..), corsPolicy)
|
||||
import PostgREST.Error (simpleError)
|
||||
import PostgREST.Types (ContentType (..), toHeader)
|
||||
import PostgREST.QueryBuilder (pgFmtLit, unquoted, pgFmtEnvVar)
|
||||
|
||||
import Protolude hiding (concat, null)
|
||||
@@ -29,9 +28,9 @@ runWithClaims :: AppConfig -> JWTAttempt ->
|
||||
ApiRequest -> H.Transaction Response
|
||||
runWithClaims conf eClaims app req =
|
||||
case eClaims of
|
||||
JWTExpired -> return $ unauthed "JWT expired"
|
||||
JWTInvalid -> return $ unauthed "JWT invalid"
|
||||
JWTMissingSecret -> return $ simpleError status500 "Server lacks JWT secret"
|
||||
JWTInvalid JWTExpired -> return $ unauthed "JWT expired"
|
||||
JWTInvalid e -> return $ unauthed $ show e
|
||||
JWTMissingSecret -> return $ simpleError status500 [] "Server lacks JWT secret"
|
||||
JWTClaims claims -> do
|
||||
H.sql $ toS.mconcat $ setRoleSql ++ claimsSql ++ headersSql ++ cookiesSql
|
||||
mapM_ H.sql customReqCheck
|
||||
@@ -47,14 +46,13 @@ runWithClaims conf eClaims app req =
|
||||
anon = String . toS $ configAnonRole conf
|
||||
customReqCheck = (\f -> "select " <> toS f <> "();") <$> configReqCheck conf
|
||||
where
|
||||
unauthed message = responseLBS unauthorized401
|
||||
[ toHeader CTApplicationJSON
|
||||
, ( "WWW-Authenticate"
|
||||
unauthed message = simpleError
|
||||
unauthorized401
|
||||
[( "WWW-Authenticate"
|
||||
, "Bearer error=\"invalid_token\", " <>
|
||||
"error_description=\"" <> message <> "\""
|
||||
)
|
||||
]
|
||||
(toS $ "{\"message\":\""<>message<>"\"}")
|
||||
"error_description=" <> show message
|
||||
)]
|
||||
message
|
||||
|
||||
defaultMiddle :: Application -> Application
|
||||
defaultMiddle =
|
||||
|
||||
Reference in New Issue
Block a user