diff --git a/main/Main.hs b/main/Main.hs index f6030ab75..ab201d08f 100644 --- a/main/Main.hs +++ b/main/Main.hs @@ -16,24 +16,13 @@ import Data.Either.Combinators (whenLeft) import Data.IORef (IORef, atomicWriteIORef, newIORef, readIORef) import Data.String (IsString (..)) -import Data.Text (pack, replace, strip, stripPrefix, - unpack) +import Data.Text (pack, replace, strip, stripPrefix) import Data.Text.Encoding (decodeUtf8, encodeUtf8) import Data.Text.IO (hPutStrLn, readFile) import Data.Time.Clock (getCurrentTime) -import Network.Socket (Family (AF_UNIX), - SockAddr (SockAddrUnix), Socket, - SocketType (Stream), bind, close, - defaultProtocol, listen, - maxListenQueue, socket) import Network.Wai.Handler.Warp (defaultSettings, runSettings, - runSettingsSocket, setHost, setPort, - setServerName) -import System.Directory (removeFile) + setHost, setPort, setServerName) import System.IO (BufferMode (..), hSetBuffering) -import System.IO.Error (isDoesNotExistError) -import System.Posix.Files (setFileMode) -import System.Posix.Types (FileMode) import PostgREST.App (postgrest) import PostgREST.Config (AppConfig (..), configPoolTimeout', @@ -50,8 +39,10 @@ import Protolude hiding (hPutStrLn, head, replace) #ifndef mingw32_HOST_OS import System.Posix.Signals +import UnixSocket #endif + {-| The purpose of this worker is to fill the refDbStructure created in 'main' with the 'DbStructure' returned from calling 'getDbStructure'. This method @@ -188,7 +179,6 @@ main = do whenLeft roleClaimKey $ panic $ show roleClaimKey - -- -- create connection pool with the provided settings, returns either -- a 'Connection' or a 'ConnectionError'. Does not throw. pool <- P.acquire (configPool conf, configPoolTimeout' conf, pgSettings) @@ -251,19 +241,17 @@ main = do schemas refDbStructure refIsWorkerOn) - in case maybeSocketAddr of - Nothing -> do - -- run the postgrest application - putStrLn $ ("Listening on port " :: Text) <> show (configPort conf) - runSettings appSettings postgrestApplication - Just socketAddr -> do - -- run postgrest application with user defined socket - sock <- createAndBindSocket (unpack socketAddr) (rightToMaybe socketFileMode) - listen sock maxListenQueue - putStrLn $ ("Listening on unix socket " :: Text) <> show socketAddr - runSettingsSocket appSettings sock postgrestApplication - -- clean socket up when done - close sock + + -- run the postgrest application with user defined socket. Only for UNIX systems. +#ifndef mingw32_HOST_OS + whenJust maybeSocketAddr $ + runAppInSocket appSettings postgrestApplication socketFileMode +#endif + + -- run the postgrest application + whenNothing maybeSocketAddr $ do + putStrLn $ ("Listening on port " :: Text) <> show (configPort conf) + runSettings appSettings postgrestApplication {-| The purpose of this function is to load the JWT secret from a file if @@ -332,15 +320,11 @@ loadDbUriFile conf = extractDbUri mDbUri Just filename -> strip <$> readFile (toS filename) setDbUri dbUri = conf {configDatabase = dbUri} -createAndBindSocket :: FilePath -> Maybe FileMode -> IO Socket -createAndBindSocket socketFilePath maybeSocketFileMode = do - deleteSocketFileIfExist socketFilePath - sock <- socket AF_UNIX Stream defaultProtocol - bind sock $ SockAddrUnix socketFilePath - mapM_ (setFileMode socketFilePath) maybeSocketFileMode - return sock - where - deleteSocketFileIfExist path = removeFile path `catch` handleDoesNotExist - handleDoesNotExist e - | isDoesNotExistError e = return () - | otherwise = throwIO e +-- Utilitarian functions. +whenJust :: Applicative f => Maybe a -> (a -> f ()) -> f () +whenJust (Just x) f = f x +whenJust Nothing _ = pass + +whenNothing :: Applicative f => Maybe a -> f () -> f () +whenNothing Nothing f = f +whenNothing _ _ = pass diff --git a/main/UnixSocket.hs b/main/UnixSocket.hs new file mode 100644 index 000000000..772f080fd --- /dev/null +++ b/main/UnixSocket.hs @@ -0,0 +1,40 @@ +module UnixSocket ( + runAppInSocket +)where + +import Network.Socket (Family (AF_UNIX), + SockAddr (SockAddrUnix), Socket, + SocketType (Stream), bind, close, + defaultProtocol, listen, + maxListenQueue, socket) +import Network.Wai (Application) +import Network.Wai.Handler.Warp +import System.Directory (removeFile) +import System.IO.Error (isDoesNotExistError) +import System.Posix.Files (setFileMode) +import System.Posix.Types (FileMode) + +import Protolude + +createAndBindSocket :: FilePath -> Maybe FileMode -> IO Socket +createAndBindSocket socketFilePath maybeSocketFileMode = do + deleteSocketFileIfExist socketFilePath + sock <- socket AF_UNIX Stream defaultProtocol + bind sock $ SockAddrUnix socketFilePath + mapM_ (setFileMode socketFilePath) maybeSocketFileMode + return sock + where + deleteSocketFileIfExist path = removeFile path `catch` handleDoesNotExist + handleDoesNotExist e + | isDoesNotExistError e = return () + | otherwise = throwIO e + +-- run the postgrest application with user defined socket. +runAppInSocket :: Settings -> Application -> Either Text FileMode -> FilePath -> IO () +runAppInSocket settings app socketFileMode sockPath = do + sock <- createAndBindSocket sockPath (rightToMaybe socketFileMode) + putStrLn $ ("Listening on unix socket " :: Text) <> show sockPath + listen sock maxListenQueue + runSettingsSocket settings sock app + -- clean socket up when done + close sock diff --git a/postgrest.cabal b/postgrest.cabal index 24b6bfebf..f1036bff1 100644 --- a/postgrest.cabal +++ b/postgrest.cabal @@ -107,6 +107,7 @@ executable postgrest , retry >= 0.7.4 && < 0.9 , text >= 1.2.2 && < 1.3 , time >= 1.6 && < 1.10 + , wai >= 3.2.1 && < 3.3 , warp >= 3.2.12 && < 3.4 default-language: Haskell2010 default-extensions: OverloadedStrings @@ -116,6 +117,7 @@ executable postgrest if !os(windows) build-depends: unix + other-modules: UnixSocket test-suite spec type: exitcode-stdio-1.0 diff --git a/src/PostgREST/Config.hs b/src/PostgREST/Config.hs index 77302a5b5..5b6993ae4 100644 --- a/src/PostgREST/Config.hs +++ b/src/PostgREST/Config.hs @@ -75,7 +75,7 @@ data AppConfig = AppConfig { , configSchemas :: NonEmpty Text , configHost :: Text , configPort :: Int - , configSocket :: Maybe Text + , configSocket :: Maybe FilePath , configSocketMode :: Either Text FileMode , configJwtSecret :: Maybe B.ByteString @@ -159,7 +159,7 @@ readOptions = do <*> (fromList . splitOnCommas <$> reqValue "db-schema") <*> (fromMaybe "!4" <$> optString "server-host") <*> (fromMaybe 3000 <$> optInt "server-port") - <*> optString "server-unix-socket" + <*> (fmap unpack <$> optString "server-unix-socket") <*> parseSocketFileMode "server-unix-socket-mode" <*> (fmap encodeUtf8 <$> optString "jwt-secret") <*> (fromMaybe False <$> optBool "secret-is-base64")