packages feed

wai-middleware-auth 0.2.1.0 → 0.2.3.0

raw patch · 19 files changed

+1122/−149 lines, 19 filesdep +josedep +microlensdep +mtldep ~aesondep ~basedep ~cookienew-uploaderPVP ok

version bump matches the API change (PVP)

Dependencies added: jose, microlens, mtl, tasty, tasty-hedgehog, tasty-hunit

Dependency ranges changed: aeson, base, cookie, hoauth2, wai, wai-extra

API changes (from Hackage documentation)

+ Network.Wai.Middleware.Auth: decodeKey :: ByteString -> Either String Key
+ Network.Wai.Middleware.Auth.OIDC: data OpenIDConnect
+ Network.Wai.Middleware.Auth.OIDC: discover :: Text -> IO OpenIDConnect
+ Network.Wai.Middleware.Auth.OIDC: getAccessToken :: Request -> Maybe OAuth2Token
+ Network.Wai.Middleware.Auth.OIDC: getIdToken :: Request -> Maybe ClaimsSet
+ Network.Wai.Middleware.Auth.OIDC: instance Data.Aeson.Types.FromJSON.FromJSON Network.Wai.Middleware.Auth.OIDC.OpenIDConnect
+ Network.Wai.Middleware.Auth.OIDC: instance Network.Wai.Middleware.Auth.Provider.AuthProvider Network.Wai.Middleware.Auth.OIDC.OpenIDConnect
+ Network.Wai.Middleware.Auth.OIDC: oidcAllowedSkew :: OpenIDConnect -> NominalDiffTime
+ Network.Wai.Middleware.Auth.OIDC: oidcClientId :: OpenIDConnect -> Text
+ Network.Wai.Middleware.Auth.OIDC: oidcClientSecret :: OpenIDConnect -> Text
+ Network.Wai.Middleware.Auth.OIDC: oidcManager :: OpenIDConnect -> Maybe Manager
+ Network.Wai.Middleware.Auth.OIDC: oidcProviderInfo :: OpenIDConnect -> ProviderInfo
+ Network.Wai.Middleware.Auth.OIDC: oidcScopes :: OpenIDConnect -> [Text]
+ Network.Wai.Middleware.Auth.Provider: instance GHC.Classes.Eq Network.Wai.Middleware.Auth.Provider.AuthUser
+ Network.Wai.Middleware.Auth.Provider: refreshLoginState :: AuthProvider ap => ap -> Request -> AuthUser -> IO (Maybe (Request, AuthUser))

Files

CHANGELOG.md view
@@ -1,10 +1,28 @@-# 0.2.1.0-========+0.2.3.0+======= +* Support `hoauth2-1.11.0`+* Drop support for `jose` versions < 0.8+* Expose `decodeKey`+* OAuth2 provider remove a session when an access token expires. It will use a+  refresh token if one is available to create a new session. If no refresh token+  is available it will redirect the user to re-authenticate.+* Providers can define logic for refreshing a session without user intervention.+* Add an OpenID Connect provider.++0.2.2.0+=======++* Add request logging to executable+* Newer multistage Docker build system++0.2.1.0+=======+ * Fix a bug in deserialization of `UserIdentity` -# 0.2.0.0-========+0.2.0.0+=======  * Drop compatiblity with hoauth2 versions <= 1.0.0. * Add a function for getting the oauth2 token from an authenticated request.
README.md view
@@ -1,5 +1,7 @@ # wai-middleware-auth +[![Build Status](https://dev.azure.com/fpco/wai-middleware-auth/_apis/build/status/fpco.wai-middleware-auth?branchName=master)](https://dev.azure.com/fpco/wai-middleware-auth/_build/latest?definitionId=4&branchName=master)+ Middleware that secures WAI application  ## Installation@@ -14,7 +16,7 @@  ## wai-auth -Along with middleware this package ships with an executbale `wai-auth`, which+Along with middleware this package ships with an executable `wai-auth`, which can function as a protected file server or a reverse proxy. Right from the box it supports OAuth2 authentication as well as it's custom implementations for Google and Github.
app/Main.hs view
@@ -3,7 +3,6 @@ {-# LANGUAGE TemplateHaskell   #-} module Main where import qualified Data.ByteString                           as S-import           Data.Monoid                               ((<>)) import           Data.Serialize                            (put, runPut) import           Network.Wai.Auth.Executable import           Network.Wai.Handler.Warp                  (run)@@ -11,6 +10,7 @@ import           Network.Wai.Middleware.Auth.OAuth2 import           Network.Wai.Middleware.Auth.OAuth2.Github import           Network.Wai.Middleware.Auth.OAuth2.Google+import           Network.Wai.Middleware.RequestLogger      (logStdout) import           Options.Applicative.Simple import           Web.ClientSession @@ -78,7 +78,7 @@       authConfig <- readAuthConfig configFile       mkMain authConfig [githubParser, googleParser, oAuth2Parser] $ \port app -> do         putStrLn $ "Listening on port " ++ show port-        run port app+        run port $ logStdout app     KeyFile (KeyOptions {..}) -> do       let key2str =             if keyBase64
src/Network/Wai/Auth/AppRoot.hs view
@@ -6,7 +6,6 @@ import           Data.ByteString          (ByteString) import           Data.CaseInsensitive     (CI, mk) import qualified Data.HashMap.Lazy        as HM-import           Data.Monoid              ((<>)) import qualified Data.Text                as T import           Data.Text.Encoding       (decodeUtf8With) import           Data.Text.Encoding.Error (lenientDecode)
src/Network/Wai/Auth/ClientSession.hs view
@@ -26,9 +26,10 @@ import           Web.ClientSession          (Key, decrypt, encryptIO,                                              getDefaultKey) import           Web.Cookie                 (def, parseCookies, renderSetCookie,-                                             setCookieExpires, setCookieHttpOnly,-                                             setCookieMaxAge, setCookieName,-                                             setCookiePath, setCookieValue)+                                             sameSiteLax, setCookieExpires,+                                             setCookieHttpOnly, setCookieMaxAge,+                                             setCookieName, setCookiePath,+                                             setCookieSameSite, setCookieValue)  data Wrapper value = Wrapper   { contained :: value@@ -81,6 +82,7 @@         , setCookiePath = Just "/"         , setCookieHttpOnly = True         , setCookieMaxAge = Just $ fromIntegral age+        , setCookieSameSite = Just sameSiteLax         })  deleteCookieValue@@ -96,4 +98,5 @@         , setCookiePath = Just "/"         , setCookieHttpOnly = True         , setCookieExpires = Just $ UTCTime (fromGregorian 1970 01 01) 0+        , setCookieSameSite = Just sameSiteLax         })
src/Network/Wai/Auth/Config.hs view
@@ -12,8 +12,7 @@   ) where  import           Data.Aeson-import           Data.Aeson.TH          (defaultOptions, deriveJSON,-                                         fieldLabelModifier)+import           Data.Aeson.TH          (deriveJSON) import qualified Data.Text              as T import           Data.Text.Encoding     (encodeUtf8) import           Network.Wai.Auth.Tools (decodeKey, encodeKey,
src/Network/Wai/Auth/Internal.hs view
@@ -1,15 +1,38 @@ {-# OPTIONS_HADDOCK hide, not-home #-}+{-# LANGUAGE DeriveGeneric     #-}+{-# LANGUAGE RecordWildCards   #-}+{-# LANGUAGE OverloadedStrings #-}+{-# LANGUAGE TupleSections     #-} module Network.Wai.Auth.Internal   ( OAuth2TokenBinary(..)+  , Metadata(..)   , encodeToken   , decodeToken+  , oauth2Login+  , refreshTokens   ) where +import qualified Data.Aeson                           as Aeson import           Data.Binary                          (Binary(get, put), encode,                                                       decodeOrFail) import qualified Data.ByteString                      as S+import qualified Data.ByteString.Char8                as S8 (pack) import qualified Data.ByteString.Lazy                 as SL+import qualified Data.Text                            as T+import           Data.Text.Encoding                   (encodeUtf8,+                                                       decodeUtf8With)+import           Data.Text.Encoding.Error             (lenientDecode)+import           GHC.Generics                         (Generic)+import           Network.HTTP.Client                  (Manager)+import           Network.HTTP.Types                   (Status, status303,+                                                       status403, status404,+                                                       status501) import qualified Network.OAuth.OAuth2                 as OA2+import           Network.Wai                          (Request, Response,+                                                       queryString, responseLBS)+import           Network.Wai.Middleware.Auth.Provider+import qualified URI.ByteString                       as U+import           URI.ByteString                       (URI)  decodeToken :: S.ByteString -> Either String OA2.OAuth2Token decodeToken bs =@@ -39,3 +62,103 @@     idToken <- fmap OA2.IdToken <$> get     pure $ OAuth2TokenBinary $       OA2.OAuth2Token accessToken refreshToken expiresIn tokenType idToken++oauth2Login+  :: OA2.OAuth2+  -> Manager+  -> Maybe [T.Text]+  -> T.Text+  -> Request +  -> [T.Text]+  -> (AuthLoginState -> IO Response)+  -> (Status -> S.ByteString -> IO Response)+  -> IO Response+oauth2Login oauth2 man oa2Scope providerName req suffix onSuccess onFailure = +  case suffix of+    [] -> do+      let scope = (encodeUtf8 . T.intercalate ",") <$> oa2Scope+      let redirectUrl =+            getRedirectURI $+            appendQueryParams+              (OA2.authorizationUrl oauth2)+              (maybe [] ((: []) . ("scope", )) scope)+      return $+        responseLBS+          status303+          [("Location", redirectUrl)]+          "Redirect to OAuth2 Authentication server"+    ["complete"] ->+      let params = queryString req+      in case lookup "code" params of+            Just (Just code) -> do+              eRes <- OA2.fetchAccessToken man oauth2 $ getExchangeToken code+              case eRes of+                Left err    -> onFailure status501 $ S8.pack $ show err+                Right token -> onSuccess $ encodeToken token+            _ ->+              case lookup "error" params of+                (Just (Just "access_denied")) ->+                  onFailure+                    status403+                    "User rejected access to the application."+                (Just (Just error_code)) ->+                  onFailure status501 $ "Received an error: " <> error_code+                (Just Nothing) ->+                  onFailure status501 $+                  "Unknown error connecting to " <>+                  encodeUtf8 providerName+                Nothing ->+                  onFailure+                    status404+                    "Page not found. Please continue with login."+    _ -> onFailure status404 "Page not found. Please continue with login."++refreshTokens :: OA2.OAuth2Token -> Manager -> OA2.OAuth2 -> IO (Maybe OA2.OAuth2Token)+refreshTokens tokens manager oauth2 = +  case OA2.refreshToken tokens of+    Nothing -> pure Nothing+    Just refreshToken -> do+      res <- OA2.refreshAccessToken manager oauth2 refreshToken+      case res of+        Left _ -> pure Nothing+        Right newTokens -> pure (Just newTokens)++getExchangeToken :: S.ByteString -> OA2.ExchangeToken+getExchangeToken = OA2.ExchangeToken . decodeUtf8With lenientDecode++appendQueryParams :: URI -> [(S.ByteString, S.ByteString)] -> URI+appendQueryParams uri params =+  OA2.appendQueryParams params uri++getRedirectURI :: U.URIRef a -> S.ByteString+getRedirectURI = U.serializeURIRef'++data Metadata+  = Metadata+      { issuer :: T.Text+      , authorizationEndpoint :: U.URI+      , tokenEndpoint :: U.URI+      , userinfoEndpoint :: Maybe T.Text+      , revocationEndpoint :: Maybe T.Text+      , jwksUri :: T.Text+      , responseTypesSupported :: [T.Text]+      , subjectTypesSupported :: [T.Text]+      , idTokenSigningAlgValuesSupported :: [T.Text]+      , scopesSupported :: Maybe [T.Text]+      , tokenEndpointAuthMethodsSupported :: Maybe [T.Text]+      , claimsSupported :: Maybe [T.Text]+      }+  deriving (Generic)++instance Aeson.FromJSON Metadata where+  parseJSON = Aeson.genericParseJSON metadataAesonOptions++instance Aeson.ToJSON Metadata where++  toJSON = Aeson.genericToJSON metadataAesonOptions++  toEncoding = Aeson.genericToEncoding metadataAesonOptions++metadataAesonOptions :: Aeson.Options+metadataAesonOptions =+  Aeson.defaultOptions {Aeson.fieldLabelModifier = Aeson.camelTo2 '_'}
src/Network/Wai/Middleware/Auth.hs view
@@ -22,6 +22,7 @@     , waiMiddlewareAuthVersion     , getAuthUser     , getDeleteSessionHeader+    , decodeKey     ) where  import           Blaze.ByteString.Builder             (fromByteString)@@ -29,7 +30,6 @@ import qualified Data.ByteString                      as S import           Data.ByteString.Builder              (Builder) import qualified Data.HashMap.Strict                  as HM-import           Data.Monoid                          ((<>)) import qualified Data.Text                            as T import           Data.Text.Encoding                   (decodeUtf8With,                                                        encodeUtf8)@@ -41,7 +41,8 @@ import           Network.HTTP.Types                   (Header, status200,                                                        status303, status404,                                                        status501)-import           Network.Wai                          (Middleware, Request,+import           Network.Wai                          (mapResponseHeaders,+                                                       Middleware, Request,                                                        pathInfo, rawPathInfo,                                                        rawQueryString,                                                        responseBuilder,@@ -49,6 +50,7 @@ import           Network.Wai.Auth.AppRoot import           Network.Wai.Auth.ClientSession import           Network.Wai.Middleware.Auth.Provider+import           Network.Wai.Auth.Tools               (decodeKey) import qualified Paths_wai_middleware_auth            as Paths import           System.IO.Unsafe                     (unsafePerformIO) import           System.PosixCompat.Time              (epochTime)@@ -272,8 +274,29 @@     authState <- loadCookieValue secretKey asStateKey req     case authState of       Just (AuthLoggedIn user) ->-        let req' = req {vault = Vault.insert userKey user $ vault req}-        in app req' respond+        let providerName = decodeUtf8With lenientDecode (authProviderName user)+        in case HM.lookup providerName asProviders of        +          Nothing ->+            -- We can no longer find the provider the user originally+            -- authenticated with, and as a result have no way to check if the+            -- session is still valid. For backwards compatibility with older+            -- versions of this library we'll assume the session remains valid.+            let req' = req {vault = Vault.insert userKey user $ vault req}+            in app req' respond+          Just provider -> do+            refreshResult <- refreshLoginState provider req user+            case refreshResult of+              Nothing ->+                -- The session has expired, the user needs to re-authenticate.+                enforceLogin "/" req respond+              Just (req', user') ->+                let req'' = req' {vault = Vault.insert userKey user' $ vault req'}+                    respond' response +                      | user' == user = respond response+                      | otherwise = do+                          cookieHeader <- saveAuthState (AuthLoggedIn user')+                          respond $ mapResponseHeaders (cookieHeader :) response +                in app req'' respond'       Just (AuthNeedRedirect url) -> enforceLogin url req respond       Nothing -> enforceLogin "/" req respond 
src/Network/Wai/Middleware/Auth/OAuth2.hs view
@@ -14,26 +14,23 @@ import           Data.Aeson.TH                        (defaultOptions,                                                        deriveJSON,                                                        fieldLabelModifier)-import qualified Data.ByteString                      as S-import qualified Data.ByteString.Char8                as S8 (pack)-import           Data.Monoid                          ((<>))+import           Data.Functor                         ((<&>))+import           Data.Int                             (Int64) import           Data.Proxy                           (Proxy (..)) import qualified Data.Text                            as T-import           Data.Text.Encoding                   (encodeUtf8,-                                                       decodeUtf8With)-import           Data.Text.Encoding.Error             (lenientDecode)+import           Data.Text.Encoding                   (encodeUtf8)+import           Foreign.C.Types                      (CTime (..)) import           Network.HTTP.Client.TLS              (getGlobalManager)-import           Network.HTTP.Types                   (status303, status403,-                                                       status404, status501) import qualified Network.OAuth.OAuth2                 as OA2-import           Network.Wai                          (Request, queryString,-                                                       responseLBS)-import           Network.Wai.Auth.Internal            (encodeToken, decodeToken)+import           Network.Wai                          (Request)+import           Network.Wai.Auth.Internal            (decodeToken, encodeToken,+                                                       oauth2Login,+                                                       refreshTokens) import           Network.Wai.Auth.Tools               (toLowerUnderscore) import qualified Network.Wai.Middleware.Auth          as MA import           Network.Wai.Middleware.Auth.Provider+import           System.PosixCompat.Time              (epochTime) import qualified URI.ByteString                       as U-import           URI.ByteString                       (URI)  -- | General OAuth2 authentication `Provider`. data OAuth2 = OAuth2@@ -63,25 +60,12 @@     Left err  -> throwM $ URIParseException err     Right url -> return url -parseAbsoluteURI' :: MonadThrow m => T.Text -> m U.URI-parseAbsoluteURI' = parseAbsoluteURI--getExchangeToken :: S.ByteString -> OA2.ExchangeToken-getExchangeToken = OA2.ExchangeToken . decodeUtf8With lenientDecode--appendQueryParams :: URI -> [(S.ByteString, S.ByteString)] -> URI-appendQueryParams uri params =-  OA2.appendQueryParams params uri- getClientId :: T.Text -> T.Text getClientId = id  getClientSecret :: T.Text -> T.Text getClientSecret = id -getRedirectURI :: U.URIRef a -> S.ByteString-getRedirectURI = U.serializeURIRef'- -- | Aeson parser for `OAuth2` provider. -- -- @since 0.1.0@@ -93,57 +77,64 @@   getProviderName _ = "oauth2"   getProviderInfo = oa2ProviderInfo   handleLogin oa2@OAuth2 {..} req suffix renderUrl onSuccess onFailure = do-    authEndpointURI <- parseAbsoluteURI' oa2AuthorizeEndpoint-    accessTokenEndpointURI <- parseAbsoluteURI' oa2AccessTokenEndpoint-    callbackURI <- parseAbsoluteURI' $ renderUrl (ProviderUrl ["complete"]) []+    authEndpointURI <- parseAbsoluteURI oa2AuthorizeEndpoint+    accessTokenEndpointURI <- parseAbsoluteURI oa2AccessTokenEndpoint+    callbackURI <- parseAbsoluteURI $ renderUrl (ProviderUrl ["complete"]) []     let oauth2 =           OA2.OAuth2           { oauthClientId = getClientId oa2ClientId-          , oauthClientSecret = getClientSecret oa2ClientSecret+          , oauthClientSecret = Just $ getClientSecret oa2ClientSecret           , oauthOAuthorizeEndpoint = authEndpointURI           , oauthAccessTokenEndpoint = accessTokenEndpointURI           , oauthCallback = Just callbackURI           }-    case suffix of-      [] -> do-        let scope = (encodeUtf8 . T.intercalate ",") <$> oa2Scope-        let redirectUrl =-              getRedirectURI $-              appendQueryParams-                (OA2.authorizationUrl oauth2)-                (maybe [] ((: []) . ("scope", )) scope)-        return $-          responseLBS-            status303-            [("Location", redirectUrl)]-            "Redirect to OAuth2 Authentication server"-      ["complete"] ->-        let params = queryString req-        in case lookup "code" params of-             Just (Just code) -> do-               man <- getGlobalManager-               eRes <- OA2.fetchAccessToken man oauth2 $ getExchangeToken code-               case eRes of-                 Left err    -> onFailure status501 $ S8.pack $ show err-                 Right token -> onSuccess $ encodeToken token-             _ ->-               case lookup "error" params of-                 (Just (Just "access_denied")) ->-                   onFailure-                     status403-                     "User rejected access to the application."-                 (Just (Just error_code)) ->-                   onFailure status501 $ "Received an error: " <> error_code-                 (Just Nothing) ->-                   onFailure status501 $-                   "Unknown error connecting to " <>-                   encodeUtf8 (getProviderName oa2)-                 Nothing ->-                   onFailure-                     status404-                     "Page not found. Please continue with login."-      _ -> onFailure status404 "Page not found. Please continue with login."+    man <- getGlobalManager+    oauth2Login+      oauth2+      man+      oa2Scope+      (getProviderName oa2)+      req+      suffix+      onSuccess+      onFailure+  refreshLoginState OAuth2 {..} req user = do+    authEndpointURI <- parseAbsoluteURI oa2AuthorizeEndpoint+    accessTokenEndpointURI <- parseAbsoluteURI oa2AccessTokenEndpoint+    let loginState = authLoginState user+    case decodeToken loginState of+      Left _ -> pure Nothing+      Right tokens -> do+        CTime now <- epochTime+        if tokenExpired user now tokens then do+          let oauth2 =+                OA2.OAuth2+                { oauthClientId = getClientId oa2ClientId+                , oauthClientSecret = Just (getClientSecret oa2ClientSecret)+                , oauthOAuthorizeEndpoint = authEndpointURI+                , oauthAccessTokenEndpoint = accessTokenEndpointURI+                -- Setting callback endpoint to `Nothing` below is a lie.+                -- We do have a callback endpoint but in this context+                -- don't have access to the function that can render it.+                -- We get away with this because the callback endpoint is+                -- not needed for obtaining a refresh token, the only+                -- way we use the config here constructed.+                , oauthCallback = Nothing+                }+          man <- getGlobalManager+          rRes <- refreshTokens tokens man oauth2+          pure (rRes <&> \newTokens -> (req, user {+                 authLoginState = encodeToken newTokens,+                 authLoginTime = fromIntegral now+               }))+        else+          pure (Just (req, user)) +tokenExpired :: AuthUser -> Int64 -> OA2.OAuth2Token -> Bool+tokenExpired user now tokens =+  case OA2.expiresIn tokens of+    Nothing -> False+    Just expiresIn -> authLoginTime user + (fromIntegral expiresIn) < now  $(deriveJSON defaultOptions { fieldLabelModifier = toLowerUnderscore . drop 3} ''OAuth2) 
src/Network/Wai/Middleware/Auth/OAuth2/Github.hs view
@@ -9,7 +9,6 @@ import           Data.Maybe                           (fromMaybe) import           Data.Aeson import qualified Data.ByteString                      as S-import           Data.Monoid                          ((<>)) import           Data.Proxy                           (Proxy (..)) import qualified Data.Text                            as T import           Data.Text.Encoding                   (encodeUtf8)@@ -17,6 +16,8 @@                                                        httpJSON, parseRequest,                                                        setRequestHeaders) import           Network.HTTP.Types+import qualified Network.OAuth.OAuth2                 as OA2+import           Network.Wai.Auth.Internal            (decodeToken) import           Network.Wai.Auth.Tools               (getValidEmail) import           Network.Wai.Middleware.Auth.OAuth2 import           Network.Wai.Middleware.Auth.Provider@@ -114,9 +115,13 @@   getProviderName _ = "github"   getProviderInfo = getProviderInfo . githubOAuth2   handleLogin Github {..} req suffix renderUrl onSuccess onFailure = do-    let onOAuth2Success accessToken = do+    let onOAuth2Success oauth2Tokens = do           catchAny-            (do emails <-+            (do accessToken <-+                  case decodeToken oauth2Tokens of+                    Left err -> fail err+                    Right tokens -> pure $ encodeUtf8 $ OA2.atoken $ OA2.accessToken tokens+                emails <-                   map githubEmail <$>                   retrieveEmails                     githubAppName
src/Network/Wai/Middleware/Auth/OAuth2/Google.hs view
@@ -10,17 +10,19 @@ import           Data.Aeson import qualified Data.ByteString                      as S import           Data.Maybe                           (fromMaybe)-import           Data.Monoid                          ((<>)) import           Data.Proxy                           (Proxy (..)) import qualified Data.Text                            as T import           Data.Text.Encoding                   (encodeUtf8) import           Network.HTTP.Simple                  (getResponseBody,-                                                       httpJSON, parseRequest,+                                                       httpJSON, parseRequestThrow,                                                        setRequestHeaders) import           Network.HTTP.Types+import qualified Network.OAuth.OAuth2                 as OA2+import           Network.Wai.Auth.Internal            (decodeToken) import           Network.Wai.Auth.Tools               (getValidEmail) import           Network.Wai.Middleware.Auth.OAuth2 import           Network.Wai.Middleware.Auth.Provider+import           System.IO                            (hPutStrLn, stderr)   -- | Create a google authentication provider@@ -96,7 +98,7 @@ -- | Makes a call to google API and retrieves user's main email. retrieveEmail :: T.Text -> S.ByteString -> IO GoogleEmail retrieveEmail emailApiEndpoint accessToken = do-  req <- parseRequest (T.unpack emailApiEndpoint)+  req <- parseRequestThrow (T.unpack emailApiEndpoint)   resp <- httpJSON $ setRequestHeaders headers req   return $ getResponseBody resp   where@@ -107,9 +109,13 @@   getProviderName _ = "google"   getProviderInfo = getProviderInfo . googleOAuth2   handleLogin Google {..} req suffix renderUrl onSuccess onFailure = do-    let onOAuth2Success accessToken = do+    let onOAuth2Success oauth2Tokens = do           catchAny-            (do email <-+            (do accessToken <-+                  case decodeToken oauth2Tokens of+                    Left err -> fail err+                    Right tokens -> pure $ encodeUtf8 $ OA2.atoken $ OA2.accessToken tokens+                email <-                   googleEmail <$>                   retrieveEmail googleAPIEmailEndpoint accessToken                 let mEmail = getValidEmail googleEmailWhitelist [email]@@ -118,6 +124,7 @@                   Nothing ->                     onFailure                       status403-                      "No valid email with permission to access was found.") $ \_err ->-            onFailure status501 "Issue communicating with google."+                      "No valid email with permission to access was found.") $ \err -> do+            hPutStrLn stderr $ "Issue communicating with Google: " ++ show err+            onFailure status501 "Issue communicating with Google."     handleLogin googleOAuth2 req suffix renderUrl onOAuth2Success onFailure
+ src/Network/Wai/Middleware/Auth/OIDC.hs view
@@ -0,0 +1,293 @@+{-# LANGUAGE FlexibleInstances   #-}     +{-# LANGUAGE RecordWildCards   #-}     +{-# LANGUAGE OverloadedStrings #-}+-- | An OpenID connect provider.+--+-- OpenID Connect is a simple identity layer on top of the OAuth2 protocol.+-- Learn more about it here: <https://openid.net/connect/>+--+-- @since 0.2.3.0+module Network.Wai.Middleware.Auth.OIDC+  ( -- * Creating a provider+    OpenIDConnect+  , discover+  -- * Customizing a provider+  , oidcClientId+  , oidcClientSecret+  , oidcProviderInfo+  , oidcManager+  , oidcScopes+  , oidcAllowedSkew+  -- * Accessing session data+  , getAccessToken+  , getIdToken+  ) where++import           Control.Applicative                  ((<|>))+import qualified Crypto.JOSE                          as JOSE+import qualified Crypto.JWT                           as JWT+import           Control.Monad.Except                 (runExceptT)+import           Data.Aeson                           (FromJSON(parseJSON),+                                                       withObject, (.:), (.!=))+import qualified Data.ByteString.Char8                as S8+import           Data.Function                        ((&))+import qualified Data.Time.Clock                      as Clock+import           Data.Traversable                     (for)+import qualified Data.Text                            as T+import qualified Data.Text.Lazy                       as TL+import qualified Data.Text.Lazy.Encoding              as TLE+import qualified Data.Vault.Lazy                      as Vault+import           Foreign.C.Types                      (CTime (..))+import qualified Lens.Micro                           as Lens+import qualified Lens.Micro.Extras                    as Lens.Extras+import           Network.HTTP.Simple                  (httpJSON,+                                                       getResponseBody,+                                                       parseRequestThrow)+import           Network.Wai.Middleware.Auth.OAuth2   (parseAbsoluteURI,+                                                      getAccessToken)+import qualified Network.OAuth.OAuth2                 as OA2+import           Network.HTTP.Client                  (Manager)+import           Network.HTTP.Client.TLS              (getGlobalManager)+import           Network.Wai                          (Request, vault)+import           Network.Wai.Auth.Internal            (Metadata(..),+                                                       decodeToken, encodeToken,+                                                       oauth2Login,+                                                       refreshTokens)+import           Network.Wai.Middleware.Auth.Provider+import           System.IO.Unsafe                     (unsafePerformIO)+import           System.PosixCompat.Time              (epochTime)+import qualified Text.Hamlet+import qualified URI.ByteString                       as U++-- | An Open ID Connect provider.+--+-- To create a value use `discover` to download configuration for an existing+-- provider, then use various setter functions to customize it.+--+-- @since 0.2.3.0+data OpenIDConnect+  = OpenIDConnect+      { oidcMetadata :: Metadata+      , oidcJwkSet :: JOSE.JWKSet+      -- | The client id this application is registered with at the Open ID+      -- Connect provider. The default is an empty string, you will need to+      -- overwrite this.+      --+      -- @since 0.2.3.0+      , oidcClientId :: T.Text+      -- | The client secret of this application. The default is an empty+      -- string, you will need to overwrite this.+      --+      -- @since 0.2.3.0+      , oidcClientSecret :: T.Text+      -- | The information for this provider. The default contains some+      -- placeholder texts. If you're using the provider screen you'll want to+      -- overwrite this.+      --+      -- @since 0.2.3.0+      , oidcProviderInfo :: ProviderInfo+      -- | The HTTP manager to use. Defaults to the global manager when not set.+      --+      -- @since 0.2.3.0+      , oidcManager :: Maybe Manager+      -- | The scopes to set. Defaults to only the "openid" scope.+      --+      -- @since 0.2.3.0+      , oidcScopes :: [T.Text]+      -- | The amount of clock skew to allow when validating id tokens. Defaults+      -- to 0.+      --+      -- @since 0.2.3.0+      , oidcAllowedSkew :: Clock.NominalDiffTime+      }++instance FromJSON OpenIDConnect where+  parseJSON =+    withObject "OpenIDConnect Object" $ \obj -> do+      metadata <- obj .: "metadata"+      jwkSet <- obj .: "jwk_set"+      clientId <- obj .: "client_id"+      clientSecret <- obj .: "client_secret"+      providerInfo <- obj .: "provider_info" .!= defProviderInfo+      scopes <- obj .: "scopes" .!= ["openid"]+      allowedSkew <- obj .: "allowed_skew" .!= 0+      pure OpenIDConnect {+        oidcMetadata = metadata,+        oidcJwkSet = jwkSet,+        oidcClientId = clientId,+        oidcClientSecret = clientSecret,+        oidcProviderInfo = providerInfo,+        oidcManager = Nothing,+        oidcScopes = scopes,+        oidcAllowedSkew = allowedSkew+      }++instance AuthProvider OpenIDConnect where+  getProviderName _ = "oidc"+  getProviderInfo = oidcProviderInfo+  handleLogin oidc@OpenIDConnect {.. } req suffix renderUrl onSuccess onFailure = do+    oauth2 <- mkOauth2 oidc (Just renderUrl)+    manager <- maybe getGlobalManager pure oidcManager+    oauth2Login+      oauth2+      manager+      (Just oidcScopes)+      (getProviderName oidc)+      req+      suffix+      onSuccess+      onFailure+  refreshLoginState oidc req user =+    let loginState = authLoginState user+    in case decodeToken loginState of+      Left _ -> pure Nothing+      Right tokens -> do+        vRes <- validateIdToken' oidc tokens+        case vRes of+          Nothing -> do+            oauth2 <- mkOauth2 oidc Nothing+            manager <- maybe getGlobalManager pure (oidcManager oidc)+            rRes <- refreshTokens tokens manager oauth2+            case rRes of+              Nothing -> pure Nothing+              Just newTokens -> do+                v2Res <- validateIdToken' oidc newTokens+                case v2Res of+                  Nothing -> pure Nothing+                  Just claims -> do+                    CTime now <- epochTime+                    let newUser =+                          user {+                            authLoginState = encodeToken newTokens,+                            authLoginTime = fromIntegral now+                          }+                    pure (Just (storeClaims claims req, newUser))+          Just claims -> +            pure (Just (storeClaims claims req, user))++-- | Fetch configuration for a provider from its discovery endpoint.+--+-- @since 0.2.3.0+discover :: T.Text -> IO OpenIDConnect+discover urlText = do+  base <- parseAbsoluteURI urlText+  let uri = base { U.uriPath = "/.well-known/openid-configuration" }+  metadata <- fetchMetadata uri+  jwkset <- fetchJWKSet (jwksUri metadata)+  pure OpenIDConnect +    { oidcClientId = ""+    , oidcClientSecret = ""+    , oidcMetadata = metadata+    , oidcJwkSet = jwkset+    , oidcProviderInfo = defProviderInfo+    , oidcManager = Nothing+    , oidcScopes = ["openid"]+    , oidcAllowedSkew = 0+    }++defProviderInfo :: ProviderInfo+defProviderInfo = ProviderInfo "OpenID Connect Provider" "" ""++fetchMetadata :: U.URI -> IO Metadata+fetchMetadata metadataEndpoint = do+  req <- parseRequestThrow (S8.unpack $ U.serializeURIRef' metadataEndpoint) +  getResponseBody <$> httpJSON req++fetchJWKSet :: T.Text -> IO JOSE.JWKSet+fetchJWKSet jwkSetEndpoint = do+  req <- parseRequestThrow (T.unpack jwkSetEndpoint) +  getResponseBody <$> httpJSON req++mkOauth2 :: OpenIDConnect -> Maybe (Text.Hamlet.Render ProviderUrl) -> IO OA2.OAuth2+mkOauth2 OpenIDConnect {..} renderUrl = do+  callbackURI <- for renderUrl $ \render -> parseAbsoluteURI $ render (ProviderUrl ["complete"]) []+  pure OA2.OAuth2+        { oauthClientId = oidcClientId+        , oauthClientSecret = Just oidcClientSecret+        , oauthOAuthorizeEndpoint = authorizationEndpoint oidcMetadata+        , oauthAccessTokenEndpoint = tokenEndpoint oidcMetadata+        , oauthCallback = callbackURI+        }++validateIdToken :: OpenIDConnect -> OA2.IdToken -> IO (Either JWT.JWTError JWT.ClaimsSet)+validateIdToken oidc (OA2.IdToken idToken) = runExceptT $ do+  signedJwt <- JOSE.decodeCompact (TLE.encodeUtf8 $ TL.fromStrict idToken)+  JWT.verifyClaims (validationSettings oidc) (oidcJwkSet oidc) signedJwt++validateIdToken' :: OpenIDConnect -> OA2.OAuth2Token -> IO (Maybe JWT.ClaimsSet)+validateIdToken' oidc tokens = +  case OA2.idToken tokens of+    Nothing -> pure Nothing+    Just idToken ->+      either (const Nothing) Just <$> validateIdToken oidc idToken++-- The validation of the ID token below is stricter then specified in the OIDC+-- spec, to make the job of validating tokens easier. If this is too limiting+-- for your user case please open an issue.+--+-- Full spec for ID token validation:+-- https://openid.net/specs/openid-connect-core-1_0.html#IDTokenValidation+--+-- Ways in which the validation below is stricter then the spec requires:+-- - We don't allow the `aud` claim to contain any audiences beyond ourselves.+validationSettings :: OpenIDConnect -> JWT.JWTValidationSettings+validationSettings oidc =+  -- The Client MUST validate that the aud (audience) Claim contains its+  -- client_id value registered at the Issuer identified by the iss (issuer)+  -- Claim as an audience. The aud (audience) Claim MAY contain an array with+  -- more than one element. The ID Token MUST be rejected if the ID Token does+  -- not list the Client as a valid audience, or if it contains additional+  -- audiences not trusted by the Client.+  validateAudience oidc+    -- If the ID Token is encrypted, decrypt it using the keys and algorithms+    -- that the Client specified during Registration that the OP was to use to+    -- encrypt the ID Token. If encryption was negotiated with the OP at+    -- Registration time and the ID Token is not encrypted, the RP SHOULD+    -- reject it.+    & JWT.defaultJWTValidationSettings+    -- The current time MUST be before the time represented by the exp Claim.+    & Lens.set JWT.jwtValidationSettingsCheckIssuedAt True+    -- The Issuer Identifier for the OpenID Provider (which is typically+    -- obtained during Discovery) MUST exactly match the value of the iss+    -- (issuer) Claim.+    & Lens.set JWT.jwtValidationSettingsIssuerPredicate (validateIssuer oidc)+    & Lens.set JWT.jwtValidationSettingsAllowedSkew (oidcAllowedSkew oidc)++validateAudience :: OpenIDConnect -> JWT.StringOrURI -> Bool+validateAudience oidc audClaim =+  audienceFromJWT == Just correctClientId+  where+    correctClientId = oidcClientId oidc+    audienceFromJWT = fromStringOrURI audClaim++validateIssuer :: OpenIDConnect -> JWT.StringOrURI -> Bool+validateIssuer oidc issClaim =+  issuerFromJWT == Just correctIssuer+  where+    correctIssuer = issuer (oidcMetadata oidc)+    issuerFromJWT = fromStringOrURI issClaim++fromStringOrURI :: JWT.StringOrURI -> Maybe T.Text+fromStringOrURI stringOrURI =+  Lens.Extras.preview JWT.string stringOrURI+   <|> fmap (T.pack . show) (Lens.Extras.preview JWT.uri stringOrURI)++storeClaims :: JWT.ClaimsSet -> Request -> Request+storeClaims claims req =+  req { vault = Vault.insert idTokenKey claims (vault req) }++-- | Get the @IdToken@ for the current user.+--+-- If called on a @Request@ behind the middleware, should always return a+-- @Just@ value.+--+-- The token returned was validated when the request was processed by the+-- middleware.+--+-- @since 0.2.3.0+getIdToken :: Request -> Maybe JWT.ClaimsSet+getIdToken req = Vault.lookup idTokenKey (vault req)++idTokenKey :: Vault.Key JWT.ClaimsSet+idTokenKey = unsafePerformIO Vault.newKey+{-# NOINLINE idTokenKey #-}
src/Network/Wai/Middleware/Auth/Provider.hs view
@@ -40,7 +40,6 @@ import qualified Data.HashMap.Strict           as HM import           Data.Int import           Data.Maybe                    (fromMaybe)-import           Data.Monoid                   ((<>)) import           Data.Proxy                    (Proxy) import qualified Data.Text                     as T import           Data.Text.Encoding            (decodeUtf8With)@@ -106,6 +105,19 @@     -> (Status -> S.ByteString -> IO Response)     -> IO Response +  -- | Check if the login state in a session is still valid, and have the+  -- opportunity to update it. Return `Nothing` to indicate a session has+  -- expired, and the user will be directed to re-authenticate. +  --+  -- The default implementation never invalidates a session once set.+  --+  -- @since 0.2.3.0+  refreshLoginState +    :: ap+    -> Request+    -> AuthUser+    -> IO (Maybe (Request, AuthUser))+  refreshLoginState _ req loginState = pure (Just (req, loginState))  -- | Generic authentication provider wrapper. data Provider where@@ -120,6 +132,7 @@    handleLogin (Provider p) = handleLogin p +  refreshLoginState (Provider p) = refreshLoginState p   -- | Collection of supported providers. type Providers = HM.HashMap T.Text Provider@@ -154,7 +167,7 @@   { authLoginState   :: !UserIdentity   , authProviderName :: !S.ByteString   , authLoginTime    :: !Int64-  } deriving (Generic, Show)+  } deriving (Eq, Generic, Show)  instance Binary AuthUser 
test/Main.hs view
@@ -2,53 +2,17 @@ {-# OPTIONS_GHC -fno-warn-orphans #-} module Main (main) where -import           Data.Binary                   (encode, decodeOrFail)-import qualified Data.ByteString.Lazy.Char8    as BSL8-import qualified Data.Text                     as T-import           Hedgehog-import           Hedgehog.Gen                  as Gen-import           Hedgehog.Range                as Range-import           Network.Wai.Auth.Internal-import qualified Network.OAuth.OAuth2.Internal as OA2--main :: IO Bool-main =-  checkParallel $ Group "Main" [-    ("oAuth2TokenBinaryDuality", oAuth2TokenBinaryDuality)-  ]--oAuth2TokenBinaryDuality :: Property-oAuth2TokenBinaryDuality = property $ do-  token <- forAll oauth2TokenBinary-  let checkUnconsumed ("", _, roundTripToken) = roundTripToken-      checkUnconsumed (unconsumed, _, _) =-        error $ "Unexpected unconsumed in bytes: " <> BSL8.unpack unconsumed-  tripping token encode (fmap checkUnconsumed . decodeOrFail)-  tripping token (encodeToken . unOAuth2TokenBinary) (fmap OAuth2TokenBinary . decodeToken)--oauth2TokenBinary :: Gen OAuth2TokenBinary-oauth2TokenBinary = do-  accessToken <- OA2.AccessToken <$> anyText-  refreshToken <- Gen.maybe $ OA2.RefreshToken <$> anyText-  expiresIn <- Gen.maybe $ Gen.int (Range.linear 0 1000)-  tokenType <- Gen.maybe anyText-  idToken <- Gen.maybe $ OA2.IdToken <$> anyText-  pure $-    OAuth2TokenBinary $-    OA2.OAuth2Token accessToken refreshToken expiresIn tokenType idToken+import           Test.Tasty+import qualified Spec.Network.Wai.Auth.Internal+import qualified Spec.Network.Wai.Middleware.Auth.OAuth2+import qualified Spec.Network.Wai.Middleware.Auth.OIDC -anyText :: Gen T.Text-anyText = Gen.text (Range.linear 0 100) Gen.unicodeAll+main :: IO ()+main = defaultMain tests --- The `OAuth2Token` type from the `hoauth2` library does not have a `Eq`--- instance, and it's constituent parts don't have a `Generic` instance. Hence--- this orphan instance here.-instance Eq OAuth2TokenBinary where-  (OAuth2TokenBinary t1) == (OAuth2TokenBinary t2) =-    and-      [ OA2.atoken (OA2.accessToken t1) == OA2.atoken (OA2.accessToken t2)-      , (OA2.rtoken <$> OA2.refreshToken t1) == (OA2.rtoken <$> OA2.refreshToken t2)-      , OA2.expiresIn t1 == OA2.expiresIn t2-      , OA2.tokenType t1 == OA2.tokenType t2-      , (OA2.idtoken <$> OA2.idToken t1) == (OA2.idtoken <$> OA2.idToken t2)-      ]+tests :: TestTree+tests = testGroup "wai-middleware-auth"+  [ Spec.Network.Wai.Auth.Internal.tests+  , Spec.Network.Wai.Middleware.Auth.OAuth2.tests+  , Spec.Network.Wai.Middleware.Auth.OIDC.tests+  ]
+ test/Network/Wai/Auth/Test.hs view
@@ -0,0 +1,157 @@+{-# LANGUAGE OverloadedStrings   #-}+{-# LANGUAGE ScopedTypeVariables #-}++module Network.Wai.Auth.Test+  (ChangeProvider+  , FakeProviderConf(..)+  , fakeProvider+  , const200+  , get+  ) where++import           Control.Monad.IO.Class                 (liftIO)+import           Data.ByteString                        (ByteString)+import qualified Data.IORef                             as IORef+import qualified Crypto.JOSE                            as JOSE+import qualified Crypto.JWT                             as JWT+import qualified Control.Monad.Except+import qualified Data.Aeson                             as Aeson+import           Data.Function                          ((&))+import qualified Data.Text                              as T+import qualified Data.Text.Encoding                     as TE+import qualified Data.Text.Lazy                         as TL+import qualified Data.Text.Lazy.Encoding                as TLE+import qualified Data.Time.Clock                        as Clock+import           GHC.Exts                               (fromString)+import qualified Network.HTTP.Types.Status              as Status+import qualified Network.OAuth.OAuth2                   as OA2+import qualified Network.Wai as Wai+import           Network.Wai.Auth.Internal              (Metadata(..))+import           Network.Wai.Test                       (Session, SResponse,+                                                         defaultRequest,+                                                         request, setPath)+import qualified Lens.Micro                             as Lens+import qualified URI.ByteString                         as U++get :: ByteString -> Session SResponse+get = request . setPath defaultRequest ++const200 :: Wai.Application+const200 _ respond = respond $ Wai.responseLBS Status.ok200 [] ""++data FakeProviderConf+  = FakeProviderConf+      { jwtExpiresIn :: Clock.NominalDiffTime,+        jwtAudience :: JWT.StringOrURI,+        jwtIssuer :: T.Text,+        jwtJWK :: JOSE.JWK,+        jwtSub :: String,+        accessTokenExpiresIn :: Int,+        returnIdToken :: Bool,+        returnRefreshToken :: Bool+      }++defaultConfig :: IO FakeProviderConf+defaultConfig = do+  jwk <- JOSE.genJWK (JOSE.RSAGenParam 256)+  pure+    FakeProviderConf+      { jwtExpiresIn = 600,+        jwtAudience = "client-id",+        jwtIssuer = "test-oidc-provider",+        jwtJWK = jwk,+        jwtSub = "1234",+        accessTokenExpiresIn = 600,+        returnIdToken = True,+        returnRefreshToken = True+      }++type ChangeProvider = (FakeProviderConf -> FakeProviderConf) -> Session ()++fakeProvider :: IO (Wai.Application, ChangeProvider)+fakeProvider = do+  config <- defaultConfig+  configRef <- IORef.newIORef config+  let changeProvider = IORef.modifyIORef configRef+  pure (fakeProvider' configRef, liftIO . changeProvider)++fakeProvider' :: IORef.IORef FakeProviderConf -> Wai.Application+fakeProvider' configRef req respond = do+  config <- IORef.readIORef configRef+  case Wai.pathInfo req of+    [".well-known", "openid-configuration"] ->+      case TE.decodeUtf8 <$> Wai.requestHeaderHost req of+        Nothing ->+          Wai.responseLBS Status.badRequest400 [] ""+            & respond+        Just host ->+          Metadata+            { issuer = jwtIssuer config,+              authorizationEndpoint = parseURI ("http://" <> host <> "/authorize"),+              tokenEndpoint = parseURI ("http://" <> host <> "/token"),+              userinfoEndpoint = Nothing,+              revocationEndpoint = Nothing,+              jwksUri = "http://" <> host <> "/jwks",+              responseTypesSupported = ["code"],+              subjectTypesSupported = ["public"],+              idTokenSigningAlgValuesSupported = ["RS256"],+              scopesSupported = Just ["openid"],+              tokenEndpointAuthMethodsSupported = Just ["client_secret_basic"],+              claimsSupported = Just ["iss", "sub", "aud", "exp", "iat"]+            }+            & Aeson.encode+            & Wai.responseLBS Status.ok200 [("Content-Type", "application/json")]+            & respond+    ["jwks"] ->+      JOSE.JWKSet [jwtJWK config]+        & Aeson.encode+        & Wai.responseLBS Status.ok200 [("Content-Type", "application/json")]+        & respond+    ["token"] -> do+      now <- Clock.getCurrentTime+      let claims =+            JWT.emptyClaimsSet+              & Lens.set JWT.claimIss (Just (fromString (T.unpack (jwtIssuer config))))+              & Lens.set JWT.claimAud (Just (JWT.Audience [jwtAudience config]))+              & Lens.set JWT.claimIat (Just (JWT.NumericDate now))+              & Lens.set JWT.claimExp (Just (JWT.NumericDate (Clock.addUTCTime (jwtExpiresIn config) now)))+              & Lens.set JWT.claimSub (Just (fromString (jwtSub config)))+      idToken <- doJwtSign (jwtJWK config) claims+      OA2.OAuth2Token+        { OA2.accessToken = OA2.AccessToken "access-granted",+          OA2.refreshToken =+            if returnRefreshToken config+              then Just (OA2.RefreshToken "refresh-token")+              else Nothing,+          OA2.expiresIn = Just (accessTokenExpiresIn config),+          OA2.tokenType = Nothing,+          OA2.idToken =+            if returnIdToken config+              then Just (OA2.IdToken idToken)+              else Nothing+        }+        & Aeson.encode+        & Wai.responseLBS Status.ok200 [("Content-Type", "application/json")]+        & respond+    _ ->+      Wai.responseLBS Status.notFound404 [] ""+        & respond++doJwtSign :: JOSE.JWK -> JWT.ClaimsSet -> IO T.Text+doJwtSign jwk claims = do+  result <- Control.Monad.Except.runExceptT $ do+    alg <- JOSE.bestJWSAlg jwk+    JWT.signClaims jwk (JOSE.newJWSHeader ((), alg)) claims+  case result of+    Left (err :: JOSE.Error) -> fail (show err)+    Right bytestring ->+      JOSE.encodeCompact bytestring+        & TLE.decodeUtf8+        & TL.toStrict+        & pure++parseURI :: T.Text -> U.URIRef U.Absolute+parseURI uri =+  TE.encodeUtf8 uri+    & U.parseURI U.laxURIParserOptions+    & either (error . show) id
+ test/Spec/Network/Wai/Auth/Internal.hs view
@@ -0,0 +1,55 @@+{-# LANGUAGE OverloadedStrings #-}+{-# OPTIONS_GHC -fno-warn-orphans #-}+module Spec.Network.Wai.Auth.Internal (tests) where++import           Data.Binary                   (encode, decodeOrFail)+import qualified Data.ByteString.Lazy.Char8    as BSL8+import qualified Data.Text                     as T+import           Test.Tasty                    (TestTree, testGroup)+import           Test.Tasty.Hedgehog           (testProperty)+import           Hedgehog+import           Hedgehog.Gen                  as Gen+import           Hedgehog.Range                as Range+import           Network.Wai.Auth.Internal+import qualified Network.OAuth.OAuth2.Internal as OA2++tests :: TestTree+tests = testGroup "Network.Wai.Auth.Internal"+  [ testProperty "oAuth2TokenBinaryDuality" oAuth2TokenBinaryDuality+  ]+  +oAuth2TokenBinaryDuality :: Property+oAuth2TokenBinaryDuality = property $ do+  token <- forAll oauth2TokenBinary+  let checkUnconsumed ("", _, roundTripToken) = roundTripToken+      checkUnconsumed (unconsumed, _, _) =+        error $ "Unexpected unconsumed in bytes: " <> BSL8.unpack unconsumed+  tripping token encode (fmap checkUnconsumed . decodeOrFail)+  tripping token (encodeToken . unOAuth2TokenBinary) (fmap OAuth2TokenBinary . decodeToken)++oauth2TokenBinary :: Gen OAuth2TokenBinary+oauth2TokenBinary = do+  accessToken <- OA2.AccessToken <$> anyText+  refreshToken <- Gen.maybe $ OA2.RefreshToken <$> anyText+  expiresIn <- Gen.maybe $ Gen.int (Range.linear 0 1000)+  tokenType <- Gen.maybe anyText+  idToken <- Gen.maybe $ OA2.IdToken <$> anyText+  pure $+    OAuth2TokenBinary $+    OA2.OAuth2Token accessToken refreshToken expiresIn tokenType idToken++anyText :: Gen T.Text+anyText = Gen.text (Range.linear 0 100) Gen.unicodeAll++-- The `OAuth2Token` type from the `hoauth2` library does not have a `Eq`+-- instance, and it's constituent parts don't have a `Generic` instance. Hence+-- this orphan instance here.+instance Eq OAuth2TokenBinary where+  (OAuth2TokenBinary t1) == (OAuth2TokenBinary t2) =+    and+      [ OA2.atoken (OA2.accessToken t1) == OA2.atoken (OA2.accessToken t2)+      , (OA2.rtoken <$> OA2.refreshToken t1) == (OA2.rtoken <$> OA2.refreshToken t2)+      , OA2.expiresIn t1 == OA2.expiresIn t2+      , OA2.tokenType t1 == OA2.tokenType t2+      , (OA2.idtoken <$> OA2.idToken t1) == (OA2.idtoken <$> OA2.idToken t2)+      ]
+ test/Spec/Network/Wai/Middleware/Auth/OAuth2.hs view
@@ -0,0 +1,133 @@+{-# LANGUAGE OverloadedStrings #-}++module Spec.Network.Wai.Middleware.Auth.OAuth2 (tests) where++import           Control.Monad                          (void)+import           Data.Function                          ((&))+import qualified Data.Text                              as T+import qualified Data.Text.Encoding                     as TE+import           GHC.Exts                               (fromList)+import qualified Network.HTTP.Types.Status              as Status+import qualified Network.Wai                            as Wai+import           Network.Wai.Auth.Test                  (ChangeProvider,+                                                         FakeProviderConf(..),+                                                         fakeProvider,+                                                         const200, get)+import qualified Network.Wai.Handler.Warp               as Warp+import qualified Network.Wai.Middleware.Auth            as Auth+import           Network.Wai.Middleware.Auth.OAuth2     (OAuth2(..),+                                                         getAccessToken)+import           Network.Wai.Middleware.Auth.Provider   (Provider(..),+                                                         ProviderInfo(..))+import           Network.Wai.Test                       (Session, assertHeader,+                                                         assertStatus,+                                                         runSession,+                                                         setClientCookie)+import           Test.Tasty                             (TestTree, testGroup)+import           Test.Tasty.HUnit                       (testCase)+import qualified Web.Cookie                             as Cookie+import qualified Web.ClientSession++tests :: TestTree+tests = testGroup "Network.Wai.Auth.OAuth2"+  [ testCase "when a request without a session is made then redirect to re-authorize" $+      runSessionWithProvider const200 $ \host _ -> do+        redirect1 <- get "/hi"+        assertStatus 303 redirect1+        assertHeader "Location" "/prefix" redirect1+        redirect2 <- get "/prefix"+        assertStatus 303 redirect2+        assertHeader "location" "/prefix/oauth2" redirect2+        redirect3 <- get "/prefix/oauth2"+        assertStatus 303 redirect3+        assertHeader+          "location"+          (TE.encodeUtf8 host <> "/authorize?scope=scope1%2Cscope2&client_id=client-id&response_type=code&redirect_uri=http%3A%2F%2Flocalhost%2Fprefix%2Foauth2%2Fcomplete")+          redirect3++  , testCase "when a request is made with a valid session then pass the request through" $+      runSessionWithProvider const200 $ \_ _ -> do+        createSession+        response <- get "/some/endpoint"+        assertStatus 200 response++  , testCase "when an access token expired and no refresh token is available then redirect to re-authorize" $+      runSessionWithProvider const200 $ \_ changeProvider -> do+        changeProvider (\c -> c { accessTokenExpiresIn = -600, returnRefreshToken = False })+        createSession+        response <- get "/some/endpoint"+        assertStatus 303 response++  , testCase "when an access token expired then use a refresh token" $+      runSessionWithProvider const200 $ \_ changeProvider -> do+        changeProvider (\c -> c { accessTokenExpiresIn = -600 })+        createSession+        response <- get "/some/endpoint"+        assertStatus 200 response++  , testCase "when a request is made with an invalid session redirect to re-authorize" $+      runSessionWithProvider const200 $ \_ _ -> do+        -- First create a known valid session, so we can see that it's the act+        -- of corrupting it that makes the test fail.+        createSession+        setClientCookie+          Cookie.defaultSetCookie+            { Cookie.setCookieName = "auth-cookie"+            , Cookie.setCookieValue = "garbage"+            }+        response <- get "/some/endpoint"+        assertStatus 303 response++  , testCase "when a request is made to the complete endpoint then create a session" $+      runSessionWithProvider const200 $ \_ _ -> do+        response <- get "/prefix/oauth2/complete?code=1234"+        assertStatus 303 response+        assertHeader "location" "/" response++  , testCase "when a request with a valid session is made then the app can access the session" $+      let app req respond = +            case getAccessToken req of+              Nothing -> respond $ Wai.responseLBS Status.badRequest400 [] ""+              Just _ -> respond $ Wai.responseLBS Status.ok200 [] ""+      in runSessionWithProvider app $ \_ _ -> do+          createSession+          response <- get "/some/endpoint"+          assertStatus 200 response+  ]++createSession :: Session ()+createSession = void $ get "/prefix/oauth2/complete?code=1234"++authSettings :: T.Text -> Auth.AuthSettings+authSettings host =+  Auth.defaultAuthSettings+    & Auth.setAuthProviders (fromList [("oauth2", provider host)])+    & Auth.setAuthPrefix "prefix"+    & Auth.setAuthCookieName "auth-cookie"+    & Auth.setAuthKey (snd <$> Web.ClientSession.randomKey)++provider :: T.Text -> Provider+provider host =+  Provider+    OAuth2+      { oa2ClientId            = "client-id"+      , oa2ClientSecret        = "client-secret"+      , oa2AuthorizeEndpoint   = host <> "/authorize"+      , oa2AccessTokenEndpoint = host <> "/token"+      , oa2Scope               = Just ["scope1", "scope2"]+      , oa2ProviderInfo        = +          ProviderInfo+            { providerTitle    = ""+            , providerLogoUrl  = ""+            , providerDescr    = ""+            }+      }++runSessionWithProvider :: Wai.Application -> (T.Text -> ChangeProvider -> Session a) -> IO a+runSessionWithProvider app session = do+  (p, changeProvider) <- fakeProvider+  Warp.testWithApplication (pure p) $ \port -> do+    let host = "http://localhost:" <> T.pack (show port)+    middleware <- Auth.mkAuthMiddleware $ authSettings host+    let app' = middleware app+    runSession (session host changeProvider) app'
+ test/Spec/Network/Wai/Middleware/Auth/OIDC.hs view
@@ -0,0 +1,164 @@+{-# LANGUAGE OverloadedStrings   #-}+{-# LANGUAGE ScopedTypeVariables #-}++module Spec.Network.Wai.Middleware.Auth.OIDC (tests) where++import           Control.Monad                          (void)+import           Control.Monad.IO.Class                 (liftIO)+import qualified Crypto.JOSE                            as JOSE+import           Data.Function                          ((&))+import qualified Data.Text                              as T+import qualified Data.Text.Encoding                     as TE+import           GHC.Exts                               (fromList, fromString)+import qualified Network.HTTP.Types.Status              as Status+import qualified Network.Wai                            as Wai+import           Network.Wai.Auth.Test                  (ChangeProvider,+                                                         FakeProviderConf(..),+                                                         fakeProvider,+                                                         const200, get)+import qualified Network.Wai.Handler.Warp               as Warp+import qualified Network.Wai.Middleware.Auth            as Auth+import           Network.Wai.Middleware.Auth.OIDC +import           Network.Wai.Middleware.Auth.Provider   (Provider(..))+import           Network.Wai.Test                       (Session, assertHeader,+                                                         assertStatus,+                                                         runSession,+                                                         setClientCookie)+import           Test.Tasty                             (TestTree, testGroup)+import           Test.Tasty.HUnit                       (testCase)+import qualified Web.Cookie                             as Cookie+import qualified Web.ClientSession++tests :: TestTree+tests = testGroup "Network.Wai.Auth.OIDC"+  [ testCase "when a request without a session is made then redirect to re-authorize" $+      runSessionWithProvider const200 $ \host _ -> do+        redirect1 <- get "/hi"+        assertStatus 303 redirect1+        assertHeader "Location" "/prefix" redirect1+        redirect2 <- get "/prefix"+        assertStatus 303 redirect2+        assertHeader "location" "/prefix/oidc" redirect2+        redirect3 <- get "/prefix/oidc"+        assertStatus 303 redirect3+        assertHeader+          "location"+          (TE.encodeUtf8 host <> "/authorize?scope=openid%2Cscope1&client_id=client-id&response_type=code&redirect_uri=http%3A%2F%2Flocalhost%2Fprefix%2Foidc%2Fcomplete")+          redirect3++  , testCase "when a request is made with a valid session then pass the request through" $+      runSessionWithProvider const200 $ \_ _ -> do+        createSession+        response <- get "/some/endpoint"+        assertStatus 200 response++  , testCase "when an ID token expired and no refresh token is available then redirect to re-authorize" $+      runSessionWithProvider const200 $ \_ changeProvider -> do+        changeProvider (\c -> c { jwtExpiresIn = -600, returnRefreshToken = False })+        createSession+        response <- get "/some/endpoint"+        assertStatus 303 response++  , testCase "when an ID token expired then use a refresh token" $+      runSessionWithProvider const200 $ \_ changeProvider -> do+        changeProvider (\c -> c { jwtExpiresIn = -600 })+        createSession+        changeProvider (\c -> c { jwtExpiresIn = 600 })+        response <- get "/some/endpoint"+        assertStatus 200 response++  , testCase "when a request is made with an invalid session redirect to re-authorize" $ +      runSessionWithProvider const200 $ \_ _ -> do+        -- First create a known valid session, so we can see that it's the act+        -- of corrupting it that makes the test fail.+        createSession+        setClientCookie+          Cookie.defaultSetCookie+            { Cookie.setCookieName = "auth-cookie"+            , Cookie.setCookieValue = "garbage"+            }+        response <- get "/some/endpoint"+        assertStatus 303 response++  , testCase "when a request is made to the complete endpoint then create a session" $+      runSessionWithProvider const200 $ \_ _ -> do+        response <- get "/prefix/oidc/complete?code=1234"+        assertStatus 303 response+        assertHeader "location" "/" response++  , testCase "when a request with a valid session is made then the app can access the access token" $+      let app req respond = +            case getAccessToken req of+              Nothing -> respond $ Wai.responseLBS Status.badRequest400 [] ""+              Just _ -> respond $ Wai.responseLBS Status.ok200 [] ""+      in runSessionWithProvider app $ \_ _ -> do+          createSession+          response <- get "/some/endpoint"+          assertStatus 200 response++  , testCase "when a request with a valid session is made then the app can access the id token" $+      let app req respond = +            case getIdToken req of+              Nothing -> respond $ Wai.responseLBS Status.badRequest400 [] ""+              Just _ -> respond $ Wai.responseLBS Status.ok200 [] ""+      in runSessionWithProvider app $ \_ _ -> do+          createSession+          response <- get "/some/endpoint"+          assertStatus 200 response++  , testCase "when an ID token has an invalid audience then redirect to re-authorize" $+      runSessionWithProvider const200 $ \_ changeProvider -> do+        changeProvider (\c -> c { jwtAudience = fromString "wrong-audience" })+        createSession+        response <- get "/some/endpoint"+        assertStatus 303 response++  , testCase "when an ID token has an invalid issuer then redirect to re-authorize" $+      runSessionWithProvider const200 $ \_ changeProvider -> do+        changeProvider (\c -> c { jwtIssuer = "wrong-issuer" })+        createSession+        response <- get "/some/endpoint"+        assertStatus 303 response++  , testCase "when a session does not contain an ID token then redirect to re-authorize" $+      runSessionWithProvider const200 $ \_ changeProvider -> do+        changeProvider (\c -> c { returnIdToken = False })+        createSession+        response <- get "/some/endpoint"+        assertStatus 303 response++  , testCase "when an ID token has an invalid signature then redirect to re-authorize" $+      runSessionWithProvider const200 $ \_ changeProvider -> do+        newJWK <- liftIO $ JOSE.genJWK (JOSE.RSAGenParam 256)+        changeProvider (\c -> c { jwtJWK = newJWK })+        createSession+        response <- get "/some/endpoint"+        assertStatus 303 response+  ]++createSession :: Session ()+createSession = void $ get "/prefix/oidc/complete?code=1234"++runSessionWithProvider :: Wai.Application -> (T.Text -> ChangeProvider -> Session a) -> IO a+runSessionWithProvider app session = do+  (provider, changeProvider) <- fakeProvider+  Warp.testWithApplication (pure provider) $ \port -> do+    let host = "http://localhost:" <> T.pack (show port)+    middleware <- Auth.mkAuthMiddleware =<< authSettings host+    let app' = middleware app+    runSession (session host changeProvider) app'++authSettings :: T.Text -> IO Auth.AuthSettings+authSettings host = do+  oidc' <- discover host+  let oidc =+        oidc'+          { oidcClientId = "client-id"+          , oidcClientSecret = "client-secret"+          , oidcScopes = ["openid", "scope1"]+          }+  pure $ Auth.defaultAuthSettings+    & Auth.setAuthProviders (fromList [("oidc", Provider oidc)])+    & Auth.setAuthPrefix "prefix"+    & Auth.setAuthCookieName "auth-cookie"+    & Auth.setAuthKey (snd <$> Web.ClientSession.randomKey)
wai-middleware-auth.cabal view
@@ -1,6 +1,6 @@ cabal-version:       1.18 name:                wai-middleware-auth-version:             0.2.1.0+version:             0.2.3.0 synopsis:            Authentication middleware that secures WAI application description:         Please see the README and Haddocks at <https://www.stackage.org/package/wai-middleware-auth> license:             MIT@@ -16,6 +16,7 @@                        Network.Wai.Middleware.Auth.OAuth2                        Network.Wai.Middleware.Auth.OAuth2.Github                        Network.Wai.Middleware.Auth.OAuth2.Google+                       Network.Wai.Middleware.Auth.OIDC                        Network.Wai.Middleware.Auth.Provider                        Network.Wai.Auth.Executable                        Network.Wai.Auth.Internal@@ -25,7 +26,7 @@                        Network.Wai.Auth.ClientSession                        Network.Wai.Auth.Tools   build-depends:       aeson-                     , base                 >= 4.7 && < 5+                     , base                 >= 4.12 && < 5                      , base64-bytestring                      , binary                      , blaze-builder@@ -34,14 +35,17 @@                      , case-insensitive                      , cereal                      , clientsession-                     , cookie+                     , cookie               >= 0.4.2                      , exceptions-                     , hoauth2              >= 1.0+                     , hoauth2              >= 1.11                      , http-client                      , http-client-tls                      , http-conduit                      , http-reverse-proxy                      , http-types+                     , jose                 >= 0.8.0+                     , microlens+                     , mtl                      , regex-posix                      , safe-exceptions                      , shakespeare@@ -68,6 +72,7 @@                      , cereal                      , clientsession                      , optparse-simple+                     , wai-extra                      , wai-middleware-auth                      , warp   ghc-options:         -Wall -threaded -rtsopts -with-rtsopts=-N@@ -77,13 +82,32 @@   type:                exitcode-stdio-1.0   main-is:             Main.hs   hs-source-dirs:      test+  other-modules:       Network.Wai.Auth.Test+                     , Spec.Network.Wai.Auth.Internal+                     , Spec.Network.Wai.Middleware.Auth.OAuth2+                     , Spec.Network.Wai.Middleware.Auth.OIDC   build-depends:       base+                     , aeson                      , binary                      , bytestring+                     , clientsession+                     , cookie                      , hedgehog                      , hoauth2+                     , http-types+                     , jose+                     , microlens+                     , mtl+                     , tasty+                     , tasty-hedgehog+                     , tasty-hunit                      , text+                     , time+                     , uri-bytestring+                     , wai+                     , wai-extra                      , wai-middleware-auth+                     , warp   ghc-options:         -Wall -threaded -rtsopts -with-rtsopts=-N  source-repository head