diff --git a/dbapi.cabal b/dbapi.cabal index 1406c21ea..14c34d395 100644 --- a/dbapi.cabal +++ b/dbapi.cabal @@ -36,6 +36,7 @@ executable dbapi , PgStructure , PgQuery , RangeQuery + , Middleware hs-source-dirs: src Test-Suite spec diff --git a/src/Dbapi.hs b/src/Dbapi.hs index 4729dd6ef..7ce42bef5 100644 --- a/src/Dbapi.hs +++ b/src/Dbapi.hs @@ -5,7 +5,6 @@ module Dbapi where import Types (SqlRow, getRow) -import Control.Exception (try) import Control.Monad (join) import Control.Exception.Base (bracket_) import Control.Arrow ((***)) @@ -35,7 +34,6 @@ import Data.String.Conversions (cs) import qualified Data.CaseInsensitive as CI import Database.HDBC.PostgreSQL (Connection) -import Database.HDBC.Types (SqlError, seErrorMsg) import PgStructure (printTables, printColumns, primaryKeyColumns, columns, Column(colName)) @@ -94,7 +92,7 @@ app conn anonymous req respond = do MalformedAuth -> respond $ responseLBS status400 [] "Malformed basic auth header" LoginFailed -> - respond $ responseLBS status403 [] "Invalid username or password" + respond $ responseLBS status401 [] "Invalid username or password" LoginSuccess role -> bracket_ (pgSetRole conn role) (pgResetRole conn) $ appWithRole conn req respond NoCredentials -> @@ -102,73 +100,70 @@ app conn anonymous req respond = do appWithRole :: Connection -> Application -appWithRole conn req respond = do - r <- try $ - case (path, verb) of - ([], _) -> - responseLBS status200 [jsonContentType] <$> printTables ver conn +appWithRole conn req respond = + respond =<< case (path, verb) of + ([], _) -> + responseLBS status200 [jsonContentType] <$> printTables ver conn - ([table], "OPTIONS") -> - responseLBS status200 [jsonContentType, allOrigins] <$> - printColumns ver (cs table) conn + ([table], "OPTIONS") -> + responseLBS status200 [jsonContentType, allOrigins] <$> + printColumns ver (cs table) conn - ([table], "GET") -> - if range == Just emptyRange - then return $ responseLBS status416 [] "HTTP Range error" - else do - r <- respondWithRangedResult <$> getRows ver (cs table) qq range conn - let canonical = urlEncodeVars $ sort $ - map (join (***) cs) $ - parseSimpleQuery $ - rawQueryString req - return $ addHeaders [ - ("Content-Location", - "/" <> cs table <> "?" <> cs canonical - )] r + ([table], "GET") -> + if range == Just emptyRange + then return $ responseLBS status416 [] "HTTP Range error" + else do + r <- respondWithRangedResult <$> getRows ver (cs table) qq range conn + let canonical = urlEncodeVars $ sort $ + map (join (***) cs) $ + parseSimpleQuery $ + rawQueryString req + return $ addHeaders [ + ("Content-Location", + "/" <> cs table <> "?" <> cs canonical + )] r - ([table], "POST") -> - jsonBodyAction req (\row -> do - allvals <- insert ver table row conn - keys <- primaryKeyColumns ver (cs table) conn - let params = urlEncodeVars $ map (\t -> (fst t, "eq." <> convert (snd t) :: String)) $ toList $ filterByKeys allvals keys - return $ responseLBS status201 - [ jsonContentType - , (hLocation, "/" <> cs table <> "?" <> cs params) - ] "" - ) + ([table], "POST") -> + jsonBodyAction req (\row -> do + allvals <- insert ver table row conn + keys <- primaryKeyColumns ver (cs table) conn + let params = urlEncodeVars $ map (\t -> (fst t, "eq." <> convert (snd t) :: String)) $ toList $ filterByKeys allvals keys + return $ responseLBS status201 + [ jsonContentType + , (hLocation, "/" <> cs table <> "?" <> cs params) + ] "" + ) - ([table], "PUT") -> - jsonBodyAction req (\row -> do - keys <- primaryKeyColumns ver (cs table) conn - let specifiedKeys = map (cs . fst) qq - if S.fromList keys /= S.fromList specifiedKeys - then return $ responseLBS status405 [] - "You must speficy all and only primary keys as params" - else - if isJust cRange - then return $ responseLBS status400 [] - "Content-Range is not allowed in PUT request" - else do - cols <- columns ver (cs table) conn - let colNames = S.fromList $ map (cs . colName) cols - let specifiedCols = S.fromList $ map fst $ getRow row - if colNames == specifiedCols then do - allvals <- upsert ver table row qq conn - let params = urlEncodeVars $ map (\t -> (fst t, "eq." <> convert (snd t) :: String)) $ toList $ filterByKeys allvals keys - return $ responseLBS status201 - [ jsonContentType - , (hLocation, "/" <> cs table <> "?" <> cs params) - ] "" + ([table], "PUT") -> + jsonBodyAction req (\row -> do + keys <- primaryKeyColumns ver (cs table) conn + let specifiedKeys = map (cs . fst) qq + if S.fromList keys /= S.fromList specifiedKeys + then return $ responseLBS status405 [] + "You must speficy all and only primary keys as params" + else + if isJust cRange + then return $ responseLBS status400 [] + "Content-Range is not allowed in PUT request" + else do + cols <- columns ver (cs table) conn + let colNames = S.fromList $ map (cs . colName) cols + let specifiedCols = S.fromList $ map fst $ getRow row + if colNames == specifiedCols then do + allvals <- upsert ver table row qq conn + let params = urlEncodeVars $ map (\t -> (fst t, "eq." <> convert (snd t) :: String)) $ toList $ filterByKeys allvals keys + return $ responseLBS status201 + [ jsonContentType + , (hLocation, "/" <> cs table <> "?" <> cs params) + ] "" - else return $ if S.null colNames then responseLBS status404 [] "" - else responseLBS status400 [] - "You must specify all columns in PUT request" - ) + else return $ if S.null colNames then responseLBS status404 [] "" + else responseLBS status400 [] + "You must specify all columns in PUT request" + ) - (_, _) -> - return $ responseLBS status404 [] "" - - respond $ either sqlErrorHandler id r + (_, _) -> + return $ responseLBS status404 [] "" where path = pathInfo req @@ -232,9 +227,6 @@ requestedVersion hdrs = accept = cs <$> lookup hAccept hdrs :: Maybe String verStr = (=~ verRegex) <$> accept :: Maybe [[String]] -sqlErrorHandler :: SqlError -> Response -sqlErrorHandler e = - responseLBS status400 [] $ cs (seErrorMsg e) addHeaders :: ResponseHeaders -> Response -> Response addHeaders hdrs (ResponseFile s headers fp m) = diff --git a/src/Main.hs b/src/Main.hs index 91d892336..38f3598a4 100644 --- a/src/Main.hs +++ b/src/Main.hs @@ -4,6 +4,7 @@ module Main where import Dbapi +import Middleware (reportPgErrors) import Network.Wai.Handler.Warp hiding (Connection) import Database.HDBC.PostgreSQL (connectPostgreSQL') import Data.String.Conversions (cs) @@ -40,7 +41,7 @@ main = do Prelude.putStrLn $ "Listening on port " ++ (show $ configPort conf :: String) conn <- connectPostgreSQL' dburi - runTLS tls settings $ gzip def $ cors corsPolicy $ app conn (cs $ configAnonRole conf) + runTLS tls settings $ gzip def $ cors corsPolicy $ reportPgErrors $ app conn (cs $ configAnonRole conf) where describe = progDesc "create a REST API to an existing Postgres database" diff --git a/src/Middleware.hs b/src/Middleware.hs new file mode 100644 index 000000000..ce289346f --- /dev/null +++ b/src/Middleware.hs @@ -0,0 +1,33 @@ +{-# LANGUAGE OverloadedStrings #-} +{-# OPTIONS_GHC -fno-warn-orphans #-} + +module Middleware where + +import Data.Aeson + +import Network.HTTP.Types.Header (hContentType) +import Network.HTTP.Types.Status (status400) +import Database.HDBC.Types (SqlError(..)) +import Control.Exception (catchJust) +import Network.Wai + + +instance ToJSON SqlError where + toJSON t = object [ + "error" .= object [ + "code" .= seNativeError t + , "message" .= seErrorMsg t + , "state" .= seState t + ] + ] + +reportPgErrors :: Middleware +reportPgErrors app req respond = + catchJust isPgException (app req respond) ( + respond . responseLBS status400 [(hContentType, "application/json")] + . encode + ) + + where + isPgException :: SqlError -> Maybe SqlError + isPgException = Just diff --git a/test/Feature/AuthSpec.hs b/test/Feature/AuthSpec.hs index f7596a0f7..6be6e6933 100644 --- a/test/Feature/AuthSpec.hs +++ b/test/Feature/AuthSpec.hs @@ -13,12 +13,11 @@ spec :: Spec spec = around appWithFixture $ describe "authorization" $ do it "hides tables that anonymous does not own" $ - -- TODO: should be 404 - get "/authors_only" `shouldRespondWith` 400 + get "/authors_only" `shouldRespondWith` 400 -- TODO: should be 404 it "indicates login failure" $ do let auth = authHeader "dbapi_test_author_a" "fakefake" request methodGet "/authors_only" [auth] "" - `shouldRespondWith` 403 + `shouldRespondWith` 401 it "allows users with permissions to see their tables" $ do let auth = authHeader "dbapi_test_author_a" "" request methodGet "/authors_only" [auth] ""