diff --git a/dbapi.cabal b/dbapi.cabal index 7f6d02c68..1406c21ea 100644 --- a/dbapi.cabal +++ b/dbapi.cabal @@ -16,7 +16,7 @@ executable dbapi build-depends: base >=4.6 && <5 , HDBC, HDBC-postgresql , warp, wai >= 3.0.1 && < 3.0.2 - , wai-extra + , wai-extra, wai-cors , HTTP, convertible , case-insensitive , http-types, scientific, time @@ -51,7 +51,7 @@ Test-Suite spec , warp, wai >= 3.0.1 && < 3.0.2 , HTTP, convertible , case-insensitive - , wai-extra, containers + , wai-extra, wai-cors, containers , http-types, scientific, time , bytestring, aeson, network , text, optparse-applicative diff --git a/src/Dbapi.hs b/src/Dbapi.hs index 030aa0fed..a6eea6c2a 100644 --- a/src/Dbapi.hs +++ b/src/Dbapi.hs @@ -18,6 +18,7 @@ import Data.Map (intersection, fromList, toList, Map) import Data.List (sort) import qualified Data.Set as S import Data.Convertible.Base (convert) +import Data.Text (strip) import Network.HTTP.Types.Status import Network.HTTP.Types.Header @@ -27,9 +28,11 @@ import Network.HTTP.Base (urlEncodeVars) import Network.Wai import Network.Wai.Internal +import Network.Wai.Middleware.Cors (CorsResourcePolicy(..)) import qualified Data.ByteString.Char8 as BS import Data.String.Conversions (cs) +import qualified Data.CaseInsensitive as CI import Database.HDBC.PostgreSQL (Connection) import Database.HDBC.Types (SqlError, seErrorMsg) @@ -163,6 +166,25 @@ app conn anonymous req respond = do range = requestedRange hdrs cRange = requestedContentRange hdrs +defaultCorsPolicy :: CorsResourcePolicy +defaultCorsPolicy = CorsResourcePolicy Nothing + ["GET", "POST", "PUT", "PATCH", "DELETE", "OPTIONS"] ["Authorization"] Nothing + (Just $ 60*60*24) False False True + +corsPolicy :: Request -> Maybe CorsResourcePolicy +corsPolicy req = case lookup "origin" headers of + Just origin -> Just defaultCorsPolicy { + corsOrigins = Just ([origin], True), + corsRequestHeaders = "Authentication":accHeaders + } + Nothing -> Nothing + where + headers = requestHeaders req + accHeaders = case lookup "access-control-request-headers" headers of + Just hdrs -> map (CI.mk . cs . strip . cs) $ BS.split ',' hdrs + Nothing -> [] + + respondWithRangedResult :: RangedResult -> Response respondWithRangedResult rr = responseLBS status [ diff --git a/src/Main.hs b/src/Main.hs index 147639af6..91d892336 100644 --- a/src/Main.hs +++ b/src/Main.hs @@ -12,6 +12,7 @@ import Control.Applicative import Options.Applicative hiding (columns) import Network.Wai.Handler.WarpTLS (tlsSettings, runTLS) import Network.Wai.Middleware.Gzip (gzip, def) +import Network.Wai.Middleware.Cors (cors) -- }}} @@ -39,7 +40,7 @@ main = do Prelude.putStrLn $ "Listening on port " ++ (show $ configPort conf :: String) conn <- connectPostgreSQL' dburi - runTLS tls settings $ gzip def $ app conn (cs $ configAnonRole conf) + runTLS tls settings $ gzip def $ cors corsPolicy $ app conn (cs $ configAnonRole conf) where describe = progDesc "create a REST API to an existing Postgres database" diff --git a/test/Feature/CorsSpec.hs b/test/Feature/CorsSpec.hs new file mode 100644 index 000000000..52d64d950 --- /dev/null +++ b/test/Feature/CorsSpec.hs @@ -0,0 +1,42 @@ +{-# LANGUAGE OverloadedStrings #-} + +module Feature.CorsSpec where + +-- {{{ Imports +import Test.Hspec +import Test.Hspec.Wai +import Network.Wai.Test (SResponse(simpleHeaders)) + +import SpecHelper + +import Network.HTTP.Types +-- }}} + +spec :: Spec +spec = around appWithFixture $ + describe "CORS" $ + it "replies naively and permissively to preflight request" $ do + r <- request methodOptions "/" + [ + ("Accept", "*/*") + , ("Origin", "http://example.com") + , ("Access-Control-Request-Method", "POST") + , ("Access-Control-Request-Headers", "Foo,Bar") + ] "" + liftIO $ do + let respHeaders = simpleHeaders r + respHeaders `shouldSatisfy` matchHeader + "Access-Control-Allow-Origin" + "http://example.com" + respHeaders `shouldSatisfy` matchHeader + "Access-Control-Allow-Credentials" + "true" + respHeaders `shouldSatisfy` matchHeader + "Access-Control-Allow-Methods" + "GET, POST, PUT, PATCH, DELETE, OPTIONS, HEAD" + respHeaders `shouldSatisfy` matchHeader + "Access-Control-Allow-Headers" + "Authentication, Foo, Bar, Accept, Accept-Language, Content-Language" + respHeaders `shouldSatisfy` matchHeader + "Access-Control-Max-Age" + "86400" diff --git a/test/SpecHelper.hs b/test/SpecHelper.hs index 659a51e93..4f46d508a 100644 --- a/test/SpecHelper.hs +++ b/test/SpecHelper.hs @@ -14,8 +14,9 @@ import Data.CaseInsensitive (CI(..)) import Text.Regex.TDFA ((=~)) import qualified Data.HashMap.Strict as Hash import qualified Data.ByteString.Char8 as BS +import Network.Wai.Middleware.Cors (cors) -import Dbapi (app, AppConfig(..)) +import Dbapi (app, corsPolicy, AppConfig(..)) cfg :: AppConfig cfg = AppConfig "postgres://dbapi_test:@localhost:5432/dbapi_test" 9000 "test/test.crt" "test/test.key" "dbapi_anonymous" @@ -40,7 +41,7 @@ dbWithSchema action = withDatabaseConnection $ \c -> do appWithFixture :: ActionWith Application -> IO () appWithFixture action = withDatabaseConnection $ \c -> do runRaw c "begin;" - action $ app c "dbapi_anonymous" + action $ cors corsPolicy $ app c "dbapi_anonymous" rollback c rangeHdrs :: ByteRange -> [Header]