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 +22/−4
- README.md +3/−1
- app/Main.hs +2/−2
- src/Network/Wai/Auth/AppRoot.hs +0/−1
- src/Network/Wai/Auth/ClientSession.hs +6/−3
- src/Network/Wai/Auth/Config.hs +1/−2
- src/Network/Wai/Auth/Internal.hs +123/−0
- src/Network/Wai/Middleware/Auth.hs +27/−4
- src/Network/Wai/Middleware/Auth/OAuth2.hs +59/−68
- src/Network/Wai/Middleware/Auth/OAuth2/Github.hs +8/−3
- src/Network/Wai/Middleware/Auth/OAuth2/Google.hs +14/−7
- src/Network/Wai/Middleware/Auth/OIDC.hs +293/−0
- src/Network/Wai/Middleware/Auth/Provider.hs +15/−2
- test/Main.hs +12/−48
- test/Network/Wai/Auth/Test.hs +157/−0
- test/Spec/Network/Wai/Auth/Internal.hs +55/−0
- test/Spec/Network/Wai/Middleware/Auth/OAuth2.hs +133/−0
- test/Spec/Network/Wai/Middleware/Auth/OIDC.hs +164/−0
- wai-middleware-auth.cabal +28/−4
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 +[](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