Main compiles. Still bugs though

This commit is contained in:
Joe Nelson
2014-12-06 17:42:18 -08:00
parent 0f849e9bf1
commit 1b8f0f2829
2 changed files with 25 additions and 22 deletions
+10 -5
View File
@@ -5,19 +5,24 @@ import Paths_dbapi (version)
import App import App
import Middleware (inTransaction, authenticated, withSavepoint, clientErrors, import Middleware (inTransaction, authenticated, withSavepoint, clientErrors,
redirectInsecure, withDBConnection, Environment(..)) redirectInsecure, withDBConnection, Environment(..))
import Network.Wai.Handler.Warp hiding (Connection)
import Data.String.Conversions (cs) import Data.String.Conversions (cs)
import qualified Data.CaseInsensitive as CI
import qualified Data.ByteString.Char8 as BS
import Control.Monad (unless) import Control.Monad (unless)
import Control.Applicative import Control.Applicative
import Control.Exception(bracket) import Control.Exception(bracket)
import Options.Applicative hiding (columns) import Options.Applicative hiding (columns)
import Network.Wai
import Network.Wai.Handler.Warp hiding (Connection)
import Network.Wai.Middleware.Gzip (gzip, def) import Network.Wai.Middleware.Gzip (gzip, def)
import Network.Wai.Middleware.Cors (cors) import Network.Wai.Middleware.Cors (cors, CorsResourcePolicy(..))
import Network.Wai.Middleware.Static (staticPolicy, only) import Network.Wai.Middleware.Static (staticPolicy, only)
import Data.Pool(createPool, destroyAllResources) import Data.Pool(createPool, destroyAllResources)
import Data.List (intercalate) import Data.List (intercalate)
import Data.Version (versionBranch) import Data.Version (versionBranch)
import Data.Text (strip)
import Database.PostgreSQL.Simple
data AppConfig = AppConfig { data AppConfig = AppConfig {
configDbUri :: String configDbUri :: String
@@ -44,8 +49,8 @@ main :: IO ()
main = do main = do
conf <- execParser (info (helper <*> argParser) describe) conf <- execParser (info (helper <*> argParser) describe)
bracket bracket
(createPool (connectPostgreSQL' (configDbUri conf)) (createPool (connectPostgreSQL $ cs (configDbUri conf))
disconnect 1 600 (configPool conf)) close 1 600 (configPool conf))
destroyAllResources destroyAllResources
(\pool -> do (\pool -> do
let port = configPort conf let port = configPort conf
@@ -61,7 +66,7 @@ main = do
. gzip def . cors corsPolicy . clientErrors . gzip def . cors corsPolicy . clientErrors
. staticPolicy (only [("favicon.ico", "static/favicon.ico")]) . staticPolicy (only [("favicon.ico", "static/favicon.ico")])
. withDBConnection pool . inTransaction Production . withDBConnection pool . inTransaction Production
. authenticated (cs $ configAnonRole conf) . withSavepoint Production $ app . authenticated (cs $ configAnonRole conf) . Middleware.withSavepoint Production $ app
) )
where where
describe = progDesc "create a REST API to an existing Postgres database" describe = progDesc "create a REST API to an existing Postgres database"
+15 -17
View File
@@ -7,10 +7,10 @@ import Data.Maybe (fromMaybe)
import Data.Monoid (mconcat) import Data.Monoid (mconcat)
import Data.Pool(withResource, Pool) import Data.Pool(withResource, Pool)
import Database.PostgreSQL.Simple
import Data.String.Conversions(cs) import Data.String.Conversions(cs)
import qualified Data.ByteString.Char8 as BS import qualified Data.ByteString.Char8 as BS
import Control.Exception (finally, throw, catchJust, catch, SomeException, import Control.Exception (catchJust, bracket_)
bracket_)
import Network.HTTP.Types.Header (RequestHeaders, hContentType, hAuthorization, import Network.HTTP.Types.Header (RequestHeaders, hContentType, hAuthorization,
hLocation) hLocation)
@@ -19,7 +19,7 @@ import Network.Wai (Application, requestHeaders, responseLBS, rawPathInfo,
rawQueryString, isSecure, requestMethod, Request) rawQueryString, isSecure, requestMethod, Request)
import Network.URI (URI(..), parseURI) import Network.URI (URI(..), parseURI)
import PgQuery(LoginAttempt(..), signInRole, setRole, resetRole) import Auth (LoginAttempt(..), signInRole, setRole, resetRole)
import Codec.Binary.Base64.String (decode) import Codec.Binary.Base64.String (decode)
import Debug.Trace import Debug.Trace
@@ -37,20 +37,17 @@ inTransaction :: Environment -> (Connection -> Application) ->
Connection -> Application Connection -> Application
inTransaction env app conn req respond = inTransaction env app conn req respond =
if env == Production && safeAction req if env == Production && safeAction req
then then go
app conn req respond else withTransaction conn go
else where go = app conn req respond
finally (runRaw conn "begin" >> app conn req respond) (runRaw conn "commit")
withSavepoint :: Environment -> (Connection -> Application) -> withSavepoint :: Environment -> (Connection -> Application) ->
Connection -> Application Connection -> Application
withSavepoint env app conn req respond = withSavepoint env app conn req respond =
if env == Production && safeAction req if env == Production && safeAction req
then app conn req respond then go
else do else Database.PostgreSQL.Simple.withSavepoint conn go
runRaw conn "savepoint req_sp" where go = app conn req respond
catch (app conn req respond) (\e -> let _ = (e::SomeException) in
runRaw conn "rollback to savepoint req_sp" >> throw e)
authenticated :: BS.ByteString -> (Connection -> Application) -> authenticated :: BS.ByteString -> (Connection -> Application) ->
Connection -> Application Connection -> Application
@@ -73,23 +70,24 @@ authenticated anon app conn req respond = do
case BS.split ' ' (cs auth) of case BS.split ' ' (cs auth) of
("Basic" : b64 : _) -> ("Basic" : b64 : _) ->
case BS.split ':' $ cs (decode $ cs b64) of case BS.split ':' $ cs (decode $ cs b64) of
(u:p:_) -> signInRole u p conn (u:p:_) -> signInRole conn u p
_ -> return MalformedAuth _ -> return MalformedAuth
_ -> return NoCredentials _ -> return NoCredentials
instance ToJSON SqlError where instance ToJSON SqlError where
toJSON t = object [ toJSON t = object [
"error" .= object [ "error" .= object [
"code" .= seNativeError t "message" .= (cs $ sqlErrorMsg t :: String)
, "message" .= seErrorMsg t , "detail" .= (cs $ sqlErrorDetail t :: String)
, "state" .= seState t , "state" .= (cs $ sqlState t :: String)
, "hint" .= (cs $ sqlErrorHint t :: String)
] ]
] ]
clientErrors :: Application -> Application clientErrors :: Application -> Application
clientErrors app req respond = clientErrors app req respond =
catchJust isPgException (app req respond) $ \err -> catchJust isPgException (app req respond) $ \err ->
respond $ if seState err == "42P01" respond $ if sqlState err == "42P01"
then responseLBS status404 [] "" then responseLBS status404 [] ""
else responseLBS status400 [(hContentType, "application/json")] (encode err) else responseLBS status400 [(hContentType, "application/json")] (encode err)