WIP: cleaner PgQuery functions

This commit is contained in:
Joe Nelson
2014-12-06 17:42:16 -08:00
parent 94ab57941d
commit cb467b62c1
3 changed files with 116 additions and 275 deletions
+7 -2
View File
@@ -14,8 +14,9 @@ executable dbapi
ghc-options: -Wall -W -Werror -O2 ghc-options: -Wall -W -Werror -O2
default-language: Haskell2010 default-language: Haskell2010
default-extensions: OverloadedStrings default-extensions: OverloadedStrings
other-extensions: QuasiQuotes
build-depends: base >=4.6 && <5 build-depends: base >=4.6 && <5
, HDBC, HDBC-postgresql , postgresql-simple >= 0.4.7.0
, warp >= 3.0.2, wai >= 3.0.1 , warp >= 3.0.2, wai >= 3.0.1
, wai-extra, wai-cors , wai-extra, wai-cors
, wai-middleware-static >= 0.6.0 , wai-middleware-static >= 0.6.0
@@ -24,6 +25,7 @@ executable dbapi
, scientific, time , scientific, time
, aeson, network >= 2.6 , aeson, network >= 2.6
, bytestring, text, split, string-conversions , bytestring, text, split, string-conversions
, stringsearch
, 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
@@ -33,6 +35,7 @@ executable dbapi
, bcrypt, base64-string , bcrypt, base64-string
, network-uri >= 2.6 , network-uri >= 2.6
, resource-pool, process , resource-pool, process
, blaze-builder
Other-Modules: Dbapi Other-Modules: Dbapi
, PgStructure , PgStructure
, PgQuery , PgQuery
@@ -51,7 +54,7 @@ Test-Suite spec
Other-Modules: Dbapi, Spec, SpecHelper Other-Modules: Dbapi, Spec, SpecHelper
Build-Depends: base, hspec2, QuickCheck Build-Depends: base, hspec2, QuickCheck
, hspec-wai >= 0.5.0, hspec-wai-json , hspec-wai >= 0.5.0, hspec-wai-json
, HDBC, HDBC-postgresql , postgresql-simple >= 0.4.7.0
, warp >= 3.0.2, wai >= 3.0.1 , warp >= 3.0.2, wai >= 3.0.1
, HTTP, convertible , HTTP, convertible
, case-insensitive , case-insensitive
@@ -60,6 +63,7 @@ Test-Suite spec
, http-types, scientific, time , http-types, scientific, time
, bytestring, aeson, network >= 2.6 , bytestring, aeson, network >= 2.6
, text, optparse-applicative , text, optparse-applicative
, stringsearch
, unordered-containers , unordered-containers
, regex-base , regex-base
, string-conversions , string-conversions
@@ -72,3 +76,4 @@ Test-Suite spec
, split , split
, network-uri >= 2.6 , network-uri >= 2.6
, resource-pool , resource-pool
, blaze-builder
+73 -242
View File
@@ -1,262 +1,93 @@
-- {{{ Imports module PgQuery where
module PgQuery (
getRows
, insert
, update
, upsert
, addUser
, signInRole
, setRole
, resetRole
, checkPass
, pgFmtIdent
, pgFmtLit
, RangedResult(..)
, LoginAttempt(..)
, DbRole
) where
import Data.Text (Text, splitOn, intercalate, replace, takeWhile) import RangeQuery
import Data.String.Conversions (cs) import Database.PostgreSQL.Simple
import Data.Functor ( (<$>) ) import Database.PostgreSQL.Simple.ToField
import Data.Maybe (fromMaybe, mapMaybe)
import Data.Monoid ((<>), mconcat)
import qualified Data.Map as M
import Text.Regex.TDFA ((=~))
import Text.Regex.TDFA.Text ()
import Control.Monad (join)
import qualified RangeQuery as R
import qualified Data.ByteString.Char8 as BS import qualified Data.ByteString.Char8 as BS
import qualified Data.ByteString.Lazy as BL import Data.ByteString.Search (split)
import qualified Data.List as L
import Database.HDBC hiding (colType, colNullable)
import Database.HDBC.PostgreSQL
import qualified Network.HTTP.Types.URI as Net import qualified Network.HTTP.Types.URI as Net
import Blaze.ByteString.Builder.ByteString (fromByteString)
import Types (SqlRow(..), getRow, sqlRowColumns, sqlRowValues) import Data.Text hiding (map, intersperse, split)
import Crypto.BCrypt (hashPasswordUsingPolicy, fastBcryptHashingPolicy, validatePassword) import Data.Monoid
import Data.Maybe (fromMaybe)
-- }}} import Data.Functor ( (<$>) )
import Data.String.Conversions (cs)
import qualified Data.List as L
data RangedResult = RangedResult { data RangedResult = RangedResult {
rrFrom :: Int rrFrom :: Int
, rrTo :: Int , rrTo :: Int
, rrTotal :: Int , rrTotal :: Int
, rrBody :: BL.ByteString , rrBody :: BS.ByteString
} deriving (Show) } deriving (Show)
type Schema = Text type CompleteQuery = (Query, [Action])
type DbRole = BS.ByteString type CompleteQueryT = CompleteQuery -> CompleteQuery
type JsonQuery = CompleteQuery
data LoginAttempt = data QualifiedTable = QualifiedTable {
NoCredentials qtSchema :: Text
| MalformedAuth , qtName :: Text
| LoginFailed } deriving (Show)
| LoginSuccess DbRole
deriving (Eq, Show)
getRows :: Schema -> Text -> Net.Query -> Maybe R.NonnegRange -> Connection -> IO RangedResult
getRows schema table qq range conn = do
r <- quickQuery conn (cs query) []
return $ case r of
[[total, _, SqlNull]] -> RangedResult offset 0 (fromSql total) "[]"
[[total, limited_total, json]] ->
RangedResult offset (offset + fromSql limited_total - 1)
(fromSql total) (fromSql json)
_ -> RangedResult 0 0 0 "[]"
where
offset = fromMaybe 0 $ R.offset <$> range
query = globalAndLimitedCounts schema table qq <> jsonArrayRows (
selectStarClause schema table
<> whereClause qq
<> orderClause qq
<> limitClause range)
whereClause :: Net.Query -> Text
whereClause qs =
if null qs then "" else " where " <> conjunction
where
cols = [ col | col <- qs, fst col `notElem` ["order"] ]
conjunction = mconcat $ L.intersperse " and " (map wherePred cols)
orderClause :: Net.Query -> Text
orderClause qs = do
let order = fromMaybe "" $ join $ lookup "order" qs
terms = mapMaybe parseOrderTerm $ splitOn "," $ cs order
termPred = mconcat $ L.intersperse ", " (map orderTermSql terms)
if null terms
then ""
else " order by " <> termPred
where
parseOrderTerm :: Text -> Maybe OrderTerm
parseOrderTerm s =
case splitOn "." s of
[d,c] ->
if d `elem` ["asc", "desc"]
then Just $ OrderTerm d c
else Nothing
_ -> Nothing
orderTermSql :: OrderTerm -> Text
orderTermSql t = pgFmtIdent (otColumn t) <> " " <> otDirection t
data OrderTerm = OrderTerm { data OrderTerm = OrderTerm {
otDirection :: Text otTerm :: BS.ByteString
, otColumn :: Text , otDirection :: BS.ByteString
} }
limitT :: Maybe NonnegRange -> CompleteQueryT
limitT r q =
q <> (" LIMIT ? OFFSET ? ", [toField limit, toField offset])
where
limit = fromMaybe "ALL" $ show . rangeLimit <$> r
offset = fromMaybe 0 $ rangeOffset <$> r
wherePred :: Net.QueryItem -> Text whereT :: Net.Query -> CompleteQueryT
wherePred (column, predicate) = whereT params q =
pgFmtIdent (cs column) <> " " <> op <> " " <> pgFmtLit (cs value) if L.null params
then q
else q <> conjunction
where
cols = [ col | col <- params, fst col `notElem` ["order"] ]
conjunction = mconcat $ L.intersperse (" and ",[]) (map wherePred cols)
orderT :: [OrderTerm] -> CompleteQueryT
orderT ts q =
if L.null ts
then q
else q <> (" order by ",[]) <> clause
where
clause = mconcat $ L.intersperse (", ",[]) (map queryTerm ts)
queryTerm :: OrderTerm -> CompleteQuery
queryTerm t =
(" ? ? ",
[EscapeIdentifier (otTerm t), Plain (fromByteString $ otDirection t)]
)
-- order = fromMaybe "" $ join (lookup "order" qs)
-- terms = mapMaybe parseOrderTerm $ splitOn "," $ cs order
-- termPred = mconcat $ L.intersperse ", " (map orderTermSql terms)
wherePred :: Net.QueryItem -> CompleteQuery
wherePred (col, predicate) =
(" ? ? ? ", [EscapeIdentifier col, Plain op, toField value])
where where
opCode:rest = BS.split '.' $ fromMaybe "." predicate opCode:rest = BS.split '.' $ fromMaybe "." predicate
value = BS.intercalate "." rest value = BS.intercalate "." rest
op = case opCode of op = fromByteString $ case opCode of
"eq" -> "=" "eq" -> "="
"gt" -> ">" "gt" -> ">"
"lt" -> "<" "lt" -> "<"
"gte" -> ">=" "gte" -> ">="
"lte" -> "<=" "lte" -> "<="
"neq" -> "<>" "neq" -> "<>"
_ -> "=" _ -> "="
limitClause :: Maybe R.NonnegRange -> Text orderParseTerm :: BS.ByteString -> Maybe OrderTerm
limitClause range = orderParseTerm s =
cs $ " LIMIT " <> limit <> " OFFSET " <> show offset <> " " case split "." s of
[d,c] ->
where if d `elem` ["asc", "desc"]
limit = fromMaybe "ALL" $ show <$> (R.limit =<< range) then Just $ OrderTerm (cs c) $
offset = fromMaybe 0 $ R.offset <$> range if d == "asc" then "asc" else "desc"
else Nothing
globalAndLimitedCounts :: Schema -> Text -> Net.Query -> Text _ -> Nothing
globalAndLimitedCounts schema table qq =
" select "
<> "(select count(1) from " <> pgFmtIdent schema <> "." <> pgFmtIdent table <> " "
<> whereClause qq
<> "), count(t), "
selectStarClause :: Schema -> Text -> Text
selectStarClause schema table =
" select * from " <> pgFmtIdent schema <> "." <> pgFmtIdent table <> " "
jsonArrayRows :: Text -> Text
jsonArrayRows q =
"array_to_json(array_agg(row_to_json(t))) from (" <> q <> ") t"
insert :: Schema -> Text -> SqlRow -> Connection -> IO (M.Map String SqlValue)
insert schema table row conn = do
stmt <- prepare conn $ cs sql
_ <- execute stmt $ sqlRowValues row
Just m <- fetchRowMap stmt
return m
where sql = insertClause schema table row
addUser :: BS.ByteString -> BS.ByteString -> BS.ByteString -> Connection -> IO ()
addUser identity pass role conn = do
Just hashed <- hashPasswordUsingPolicy fastBcryptHashingPolicy $ cs pass
_ <- quickQuery conn
"insert into dbapi.auth (id, pass, rolname) values (?, ?, ?)"
$ map toSql [identity, hashed, role]
return ()
signInRole :: BS.ByteString -> BS.ByteString -> Connection -> IO LoginAttempt
signInRole user pass conn = do
u <- quickQuery conn "select pass, rolname from dbapi.auth where id = ?" [toSql user]
return $ case u of
[[hashed, role]] ->
if checkPass (fromSql hashed) (cs pass)
then LoginSuccess $ fromSql role
else LoginFailed
_ -> LoginFailed
checkPass :: BS.ByteString -> BS.ByteString -> Bool
checkPass = validatePassword
upsert :: Schema -> Text -> SqlRow -> Net.Query -> Connection ->
IO (M.Map String SqlValue)
upsert schema table row qq conn = do
stmt <- prepare conn $ cs $ upsertClause schema table row qq
_ <- execute stmt $ join $ replicate 2 $ sqlRowValues row
m <- fetchRowMap stmt
return $ fromMaybe M.empty m
update :: Schema -> Text -> SqlRow -> Net.Query -> Connection ->
IO (M.Map String SqlValue)
update schema table row qq conn = do
stmt <- prepare conn $ cs $ updateClause schema table row qq
_ <- execute stmt $ sqlRowValues row
m <- fetchRowMap stmt
return $ fromMaybe M.empty m
placeholders :: Text -> SqlRow -> Text
placeholders symbol = intercalate ", " . map (const symbol) . getRow
insertClause :: Schema -> Text -> SqlRow -> Text
insertClause schema table (SqlRow []) =
"insert into " <> pgFmtIdent schema <> "." <> pgFmtIdent table <> " default values returning *"
insertClause schema table row =
"insert into " <> pgFmtIdent schema <> "." <> pgFmtIdent table <> " (" <>
intercalate ", " (map pgFmtIdent (sqlRowColumns row))
<> ") values (" <> placeholders "?" row <> ") returning *"
insertClauseViaSelect :: Schema -> Text -> SqlRow -> Text
insertClauseViaSelect schema table row =
"insert into " <> pgFmtIdent schema <> "." <> pgFmtIdent table <> " (" <>
intercalate ", " (map pgFmtIdent (sqlRowColumns row))
<> ") select " <> placeholders "?" row
updateClause :: Schema -> Text -> SqlRow -> Net.Query -> Text
updateClause schema table row qq =
"update " <> pgFmtIdent schema <> "." <> pgFmtIdent table <> " set (" <>
intercalate ", " (map pgFmtIdent (sqlRowColumns row))
<> ") = (" <> placeholders "?" row <> ")"
<> whereClause qq
upsertClause :: Schema -> Text -> SqlRow -> Net.Query -> Text
upsertClause schema table row qq =
"with upsert as (" <> updateClause schema table row qq
<> " returning *) " <> insertClauseViaSelect schema table row
<> " where not exists (select * from upsert) returning *"
pgFmtIdent :: Text -> Text
pgFmtIdent x =
let escaped = replace "\"" "\"\"" (trimNullChars x) in
if escaped =~ danger
then "\"" <> escaped <> "\""
else escaped
where danger = "^$|^[^a-z_]|[^a-z_0-9]" :: Text
pgFmtLit :: Text -> Text
pgFmtLit 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 conn role = runRaw conn $ "set role " <> cs role
resetRole :: Connection -> IO ()
resetRole conn = runRaw conn "reset role"
+36 -31
View File
@@ -1,8 +1,16 @@
module RangeQuery where module RangeQuery (
rangeParse
, rangeRequested
, rangeLimit
, rangeOffset
, NonnegRange
) where
import Control.Applicative import Control.Applicative
import Network.HTTP.Types.Header import Network.HTTP.Types.Header
import qualified Data.ByteString.Char8 as BS
import Data.Ranged.Boundaries import Data.Ranged.Boundaries
import Data.Ranged.Ranges import Data.Ranged.Ranges
@@ -14,6 +22,33 @@ import Data.Maybe (fromMaybe, listToMaybe)
type NonnegRange = Range Int type NonnegRange = Range Int
rangeParse :: BS.ByteString -> Maybe NonnegRange
rangeParse range = do
let rangeRegex = "^([0-9]+)-([0-9]*)$" :: BS.ByteString
parsedRange <- listToMaybe (range =~ rangeRegex :: [[BS.ByteString]])
let [_, from, to] = readMaybe . cs <$> parsedRange
let lower = fromMaybe emptyRange (rangeGeq <$> from)
let upper = fromMaybe (rangeGeq 0) (rangeLeq <$> to)
return $ rangeIntersection lower upper
rangeRequested :: RequestHeaders -> Maybe NonnegRange
rangeRequested = (rangeParse =<<) . lookup hRange
rangeLimit :: NonnegRange -> Maybe Int
rangeLimit range =
case [rangeLower range, rangeUpper range]
of [BoundaryBelow from, BoundaryAbove to] -> Just (1 + to - from)
_ -> Nothing
rangeOffset :: NonnegRange -> Int
rangeOffset range =
case rangeLower range
of BoundaryBelow from -> from
_ -> error "range without lower bound" -- should never happen
rangeGeq :: Int -> NonnegRange rangeGeq :: Int -> NonnegRange
rangeGeq n = rangeGeq n =
Range (BoundaryBelow n) BoundaryAboveAll Range (BoundaryBelow n) BoundaryAboveAll
@@ -21,33 +56,3 @@ rangeGeq n =
rangeLeq :: Int -> NonnegRange rangeLeq :: Int -> NonnegRange
rangeLeq n = rangeLeq n =
Range BoundaryBelowAll (BoundaryAbove n) Range BoundaryBelowAll (BoundaryAbove n)
parseRange :: String -> Maybe NonnegRange
parseRange range = do
let rangeRegex = "^([0-9]+)-([0-9]*)$" :: String
parsedRange <- listToMaybe (range =~ rangeRegex :: [[String]])
let [_, from, to] = readMaybe <$> parsedRange
let lower = fromMaybe emptyRange (rangeGeq <$> from)
let upper = fromMaybe (rangeGeq 0) (rangeLeq <$> to)
return $ rangeIntersection lower upper
requestedRange :: RequestHeaders -> Maybe NonnegRange
requestedRange hdrs = parseRange =<< cs <$> lookup hRange hdrs
requestedContentRange :: RequestHeaders -> Maybe NonnegRange
requestedContentRange hdrs = parseRange =<< cs <$> lookup "Content-Range" hdrs
limit :: NonnegRange -> Maybe Int
limit range =
case [rangeLower range, rangeUpper range]
of [BoundaryBelow from, BoundaryAbove to] -> Just (1 + to - from)
_ -> Nothing
offset :: NonnegRange -> Int
offset range =
case rangeLower range
of BoundaryBelow from -> from
_ -> error "range without lower bound" -- should never happen