Create functions that mimic Postgres' format() command
This commit is contained in:
+3
-1
@@ -26,6 +26,7 @@ executable dbapi
|
|||||||
, containers, unordered-containers
|
, containers, unordered-containers
|
||||||
, optparse-applicative >= 0.9.1 && < 0.10
|
, optparse-applicative >= 0.9.1 && < 0.10
|
||||||
, regex-base, regex-tdfa
|
, regex-base, regex-tdfa
|
||||||
|
, regex-tdfa-text
|
||||||
, Ranged-sets
|
, Ranged-sets
|
||||||
, transformers
|
, transformers
|
||||||
, bcrypt, base64-string
|
, bcrypt, base64-string
|
||||||
@@ -45,7 +46,7 @@ Test-Suite spec
|
|||||||
ghc-options: -Wall -W -Werror
|
ghc-options: -Wall -W -Werror
|
||||||
Main-Is: Main.hs
|
Main-Is: Main.hs
|
||||||
Other-Modules: Dbapi, Spec, SpecHelper
|
Other-Modules: Dbapi, Spec, SpecHelper
|
||||||
Build-Depends: base, hspec2
|
Build-Depends: base, hspec2, QuickCheck
|
||||||
, hspec-wai >= 0.5.0, hspec-wai-json
|
, hspec-wai >= 0.5.0, hspec-wai-json
|
||||||
, HDBC, HDBC-postgresql
|
, HDBC, HDBC-postgresql
|
||||||
, warp, wai >= 3.0.1 && < 3.0.2
|
, warp, wai >= 3.0.1 && < 3.0.2
|
||||||
@@ -60,6 +61,7 @@ Test-Suite spec
|
|||||||
, regex-base
|
, regex-base
|
||||||
, string-conversions
|
, string-conversions
|
||||||
, http-media, regex-tdfa
|
, http-media, regex-tdfa
|
||||||
|
, regex-tdfa-text
|
||||||
, Ranged-sets
|
, Ranged-sets
|
||||||
, transformers
|
, transformers
|
||||||
, bcrypt
|
, bcrypt
|
||||||
|
|||||||
+27
-1
@@ -11,18 +11,23 @@ module PgQuery (
|
|||||||
, setRole
|
, setRole
|
||||||
, resetRole
|
, resetRole
|
||||||
, checkPass
|
, checkPass
|
||||||
|
, pgFormatIdentifier
|
||||||
|
, pgFormatLiteral
|
||||||
, RangedResult(..)
|
, RangedResult(..)
|
||||||
, LoginAttempt(..)
|
, LoginAttempt(..)
|
||||||
, DbRole
|
, DbRole
|
||||||
) where
|
) where
|
||||||
|
|
||||||
import Data.Text (Text, splitOn, intercalate)
|
import Data.Text (Text, splitOn, intercalate, replace, takeWhile)
|
||||||
import Data.String.Conversions (cs)
|
import Data.String.Conversions (cs)
|
||||||
import Data.Functor ( (<$>) )
|
import Data.Functor ( (<$>) )
|
||||||
import Data.Maybe (fromMaybe, mapMaybe)
|
import Data.Maybe (fromMaybe, mapMaybe)
|
||||||
import Data.Monoid ((<>), mconcat)
|
import Data.Monoid ((<>), mconcat)
|
||||||
import qualified Data.Map as M
|
import qualified Data.Map as M
|
||||||
|
|
||||||
|
import Text.Regex.TDFA ((=~))
|
||||||
|
import Text.Regex.TDFA.Text ()
|
||||||
|
|
||||||
import Control.Monad (join)
|
import Control.Monad (join)
|
||||||
|
|
||||||
import qualified RangeQuery as R
|
import qualified RangeQuery as R
|
||||||
@@ -238,6 +243,27 @@ populateSql conn sql = do
|
|||||||
ph :: [a] -> Text
|
ph :: [a] -> Text
|
||||||
ph = intercalate ", " . map (const "?::varchar")
|
ph = intercalate ", " . map (const "?::varchar")
|
||||||
|
|
||||||
|
pgFormatIdentifier :: Text -> Text
|
||||||
|
pgFormatIdentifier x =
|
||||||
|
let escaped = replace "\"" "\"\"" (trimNullChars x) in
|
||||||
|
if escaped =~ danger
|
||||||
|
then "\"" <> escaped <> "\""
|
||||||
|
else escaped
|
||||||
|
|
||||||
|
where danger = "^$|^[^a-z_]|[^a-z_0-9]" :: Text
|
||||||
|
|
||||||
|
pgFormatLiteral :: Text -> Text
|
||||||
|
pgFormatLiteral x =
|
||||||
|
let trimmed = trimNullChars x
|
||||||
|
escaped = "'" <> replace "'" "''" trimmed <> "'"
|
||||||
|
slashed = replace "\\" "\\\\" escaped in
|
||||||
|
if escaped =~ ("\\\\" :: Text)
|
||||||
|
then "E" <> slashed
|
||||||
|
else slashed
|
||||||
|
|
||||||
|
trimNullChars :: Text -> Text
|
||||||
|
trimNullChars = Data.Text.takeWhile (/= '\x0')
|
||||||
|
|
||||||
setRole :: Connection -> DbRole -> IO ()
|
setRole :: Connection -> DbRole -> IO ()
|
||||||
setRole conn role = runRaw conn $ "set role " <> cs role
|
setRole conn role = runRaw conn $ "set role " <> cs role
|
||||||
|
|
||||||
|
|||||||
@@ -3,11 +3,14 @@
|
|||||||
module Unit.PgQuerySpec where
|
module Unit.PgQuerySpec where
|
||||||
|
|
||||||
import Test.Hspec
|
import Test.Hspec
|
||||||
|
import Test.QuickCheck
|
||||||
|
import Test.QuickCheck.Monadic
|
||||||
|
|
||||||
import Database.HDBC (IConnection, SqlValue, toSql, prepare,
|
import Database.HDBC (IConnection, SqlValue, toSql, prepare,
|
||||||
quickQuery, fromSql, execute, seState, fetchAllRowsAL)
|
quickQuery, fromSql, execute, seState, fetchAllRowsAL)
|
||||||
|
|
||||||
import PgQuery (LoginAttempt(..), insert, addUser, signInRole, checkPass)
|
import PgQuery (LoginAttempt(..), insert, addUser, signInRole, checkPass
|
||||||
|
, pgFormatIdentifier, pgFormatLiteral)
|
||||||
import Types (SqlRow(SqlRow))
|
import Types (SqlRow(SqlRow))
|
||||||
import TestTypes (fromList, incStr, incNullableStr, incInsert, incId)
|
import TestTypes (fromList, incStr, incNullableStr, incInsert, incId)
|
||||||
import Data.Map (toList)
|
import Data.Map (toList)
|
||||||
@@ -83,3 +86,17 @@ spec = around dbWithSchema $ do
|
|||||||
it "returns nothing with bad creds" $ \conn -> do
|
it "returns nothing with bad creds" $ \conn -> do
|
||||||
signInRole "not-a-user" pass conn `shouldReturn` LoginFailed
|
signInRole "not-a-user" pass conn `shouldReturn` LoginFailed
|
||||||
signInRole user (pass <> "crap") conn `shouldReturn` LoginFailed
|
signInRole user (pass <> "crap") conn `shouldReturn` LoginFailed
|
||||||
|
|
||||||
|
describe "pgFormatIdentifier" $
|
||||||
|
it "Does what format %I would do" $ \conn ->
|
||||||
|
property $ monadicIO $ do
|
||||||
|
fuzz <- pick arbitrary
|
||||||
|
[[row]] <- run $ quickALQuery conn "select format('%I', ? :: varchar)" [toSql (fuzz :: String)]
|
||||||
|
assert $ fromSql (snd row) == pgFormatIdentifier (cs fuzz)
|
||||||
|
|
||||||
|
describe "pgFormatLiteral" $
|
||||||
|
it "Does what format %L would do" $ \conn ->
|
||||||
|
property $ monadicIO $ do
|
||||||
|
fuzz <- pick arbitrary
|
||||||
|
[[row]] <- run $ quickALQuery conn "select format('%L', ? :: varchar)" [toSql (fuzz :: String)]
|
||||||
|
assert $ fromSql (snd row) == pgFormatLiteral (cs fuzz)
|
||||||
|
|||||||
Reference in New Issue
Block a user