diff --git a/src/Middleware.hs b/src/Middleware.hs index b9c730d8e..3a0604b7c 100644 --- a/src/Middleware.hs +++ b/src/Middleware.hs @@ -3,15 +3,24 @@ module Middleware where -import Data.Aeson +import Data.Aeson ((.=), toJSON, ToJSON, object, encode) +import Data.Maybe (fromMaybe) import Database.HDBC (runRaw) import Database.HDBC.PostgreSQL (Connection) -import Network.HTTP.Types.Header (hContentType) -import Network.HTTP.Types.Status (status400) import Database.HDBC.Types (SqlError(..)) -import Network.Wai (Application, responseLBS) -import Control.Exception (finally, throw, catchJust, catch, SomeException) + +import Data.String.Conversions(cs) +import qualified Data.ByteString.Char8 as BS +import Control.Exception (finally, throw, catchJust, catch, SomeException, + bracket_) + +import Network.HTTP.Types.Header (RequestHeaders, hContentType, hAuthorization) +import Network.HTTP.Types.Status (status400, status401) +import Network.Wai (Application, requestHeaders, responseLBS) + +import PgQuery(LoginAttempt(..), signInRole, setRole, resetRole) +import Codec.Binary.Base64.String (decode) inTransaction :: (Connection -> Application) -> (Connection -> Application) @@ -24,6 +33,32 @@ withSavepoint app conn req respond = do catch (app conn req respond) (\e -> let _ = (e::SomeException) in runRaw conn "rollback to savepoint req_sp" >> throw e) +authenticated :: BS.ByteString -> (Connection -> Application) + -> Connection -> Application +authenticated anon app conn req respond = do + attempt <- httpRequesterRole (requestHeaders req) + case attempt of + MalformedAuth -> + respond $ responseLBS status400 [] "Malformed basic auth header" + LoginFailed -> + respond $ responseLBS status401 [] "Invalid username or password" + LoginSuccess role -> + bracket_ (setRole conn role) (resetRole conn) $ app conn req respond + NoCredentials -> + bracket_ (setRole conn anon) (resetRole conn) $ app conn req respond + + where + httpRequesterRole :: RequestHeaders -> IO LoginAttempt + httpRequesterRole hdrs = do + let auth = fromMaybe "" $ lookup hAuthorization hdrs + case BS.split ' ' (cs auth) of + ("Basic" : b64 : _) -> + case BS.split ':' $ cs (decode $ cs b64) of + (u:p:_) -> signInRole u p conn + _ -> return MalformedAuth + _ -> return NoCredentials + + instance ToJSON SqlError where toJSON t = object [ "error" .= object [