WIP: cleaner PgQuery functions
This commit is contained in:
+7
-2
@@ -14,8 +14,9 @@ executable dbapi
|
||||
ghc-options: -Wall -W -Werror -O2
|
||||
default-language: Haskell2010
|
||||
default-extensions: OverloadedStrings
|
||||
other-extensions: QuasiQuotes
|
||||
build-depends: base >=4.6 && <5
|
||||
, HDBC, HDBC-postgresql
|
||||
, postgresql-simple >= 0.4.7.0
|
||||
, warp >= 3.0.2, wai >= 3.0.1
|
||||
, wai-extra, wai-cors
|
||||
, wai-middleware-static >= 0.6.0
|
||||
@@ -24,6 +25,7 @@ executable dbapi
|
||||
, scientific, time
|
||||
, aeson, network >= 2.6
|
||||
, bytestring, text, split, string-conversions
|
||||
, stringsearch
|
||||
, containers, unordered-containers
|
||||
, optparse-applicative >= 0.9.1 && < 0.10
|
||||
, regex-base, regex-tdfa
|
||||
@@ -33,6 +35,7 @@ executable dbapi
|
||||
, bcrypt, base64-string
|
||||
, network-uri >= 2.6
|
||||
, resource-pool, process
|
||||
, blaze-builder
|
||||
Other-Modules: Dbapi
|
||||
, PgStructure
|
||||
, PgQuery
|
||||
@@ -51,7 +54,7 @@ Test-Suite spec
|
||||
Other-Modules: Dbapi, Spec, SpecHelper
|
||||
Build-Depends: base, hspec2, QuickCheck
|
||||
, 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
|
||||
, HTTP, convertible
|
||||
, case-insensitive
|
||||
@@ -60,6 +63,7 @@ Test-Suite spec
|
||||
, http-types, scientific, time
|
||||
, bytestring, aeson, network >= 2.6
|
||||
, text, optparse-applicative
|
||||
, stringsearch
|
||||
, unordered-containers
|
||||
, regex-base
|
||||
, string-conversions
|
||||
@@ -72,3 +76,4 @@ Test-Suite spec
|
||||
, split
|
||||
, network-uri >= 2.6
|
||||
, resource-pool
|
||||
, blaze-builder
|
||||
|
||||
+73
-242
@@ -1,262 +1,93 @@
|
||||
-- {{{ Imports
|
||||
module PgQuery (
|
||||
getRows
|
||||
, insert
|
||||
, update
|
||||
, upsert
|
||||
, addUser
|
||||
, signInRole
|
||||
, setRole
|
||||
, resetRole
|
||||
, checkPass
|
||||
, pgFmtIdent
|
||||
, pgFmtLit
|
||||
, RangedResult(..)
|
||||
, LoginAttempt(..)
|
||||
, DbRole
|
||||
) where
|
||||
module PgQuery where
|
||||
|
||||
import Data.Text (Text, splitOn, intercalate, replace, takeWhile)
|
||||
import Data.String.Conversions (cs)
|
||||
import Data.Functor ( (<$>) )
|
||||
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 RangeQuery
|
||||
import Database.PostgreSQL.Simple
|
||||
import Database.PostgreSQL.Simple.ToField
|
||||
import qualified Data.ByteString.Char8 as BS
|
||||
import qualified Data.ByteString.Lazy as BL
|
||||
import qualified Data.List as L
|
||||
|
||||
import Database.HDBC hiding (colType, colNullable)
|
||||
import Database.HDBC.PostgreSQL
|
||||
|
||||
import Data.ByteString.Search (split)
|
||||
import qualified Network.HTTP.Types.URI as Net
|
||||
|
||||
import Types (SqlRow(..), getRow, sqlRowColumns, sqlRowValues)
|
||||
import Crypto.BCrypt (hashPasswordUsingPolicy, fastBcryptHashingPolicy, validatePassword)
|
||||
|
||||
-- }}}
|
||||
import Blaze.ByteString.Builder.ByteString (fromByteString)
|
||||
import Data.Text hiding (map, intersperse, split)
|
||||
import Data.Monoid
|
||||
import Data.Maybe (fromMaybe)
|
||||
import Data.Functor ( (<$>) )
|
||||
import Data.String.Conversions (cs)
|
||||
import qualified Data.List as L
|
||||
|
||||
data RangedResult = RangedResult {
|
||||
rrFrom :: Int
|
||||
, rrTo :: Int
|
||||
, rrTotal :: Int
|
||||
, rrBody :: BL.ByteString
|
||||
, rrBody :: BS.ByteString
|
||||
} deriving (Show)
|
||||
|
||||
type Schema = Text
|
||||
type DbRole = BS.ByteString
|
||||
|
||||
data LoginAttempt =
|
||||
NoCredentials
|
||||
| MalformedAuth
|
||||
| LoginFailed
|
||||
| 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
|
||||
|
||||
type CompleteQuery = (Query, [Action])
|
||||
type CompleteQueryT = CompleteQuery -> CompleteQuery
|
||||
type JsonQuery = CompleteQuery
|
||||
data QualifiedTable = QualifiedTable {
|
||||
qtSchema :: Text
|
||||
, qtName :: Text
|
||||
} deriving (Show)
|
||||
|
||||
data OrderTerm = OrderTerm {
|
||||
otDirection :: Text
|
||||
, otColumn :: Text
|
||||
otTerm :: BS.ByteString
|
||||
, 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
|
||||
wherePred (column, predicate) =
|
||||
pgFmtIdent (cs column) <> " " <> op <> " " <> pgFmtLit (cs value)
|
||||
whereT :: Net.Query -> CompleteQueryT
|
||||
whereT params q =
|
||||
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
|
||||
opCode:rest = BS.split '.' $ fromMaybe "." predicate
|
||||
value = BS.intercalate "." rest
|
||||
op = case opCode of
|
||||
"eq" -> "="
|
||||
"gt" -> ">"
|
||||
"lt" -> "<"
|
||||
"gte" -> ">="
|
||||
"lte" -> "<="
|
||||
"neq" -> "<>"
|
||||
_ -> "="
|
||||
op = fromByteString $ case opCode of
|
||||
"eq" -> "="
|
||||
"gt" -> ">"
|
||||
"lt" -> "<"
|
||||
"gte" -> ">="
|
||||
"lte" -> "<="
|
||||
"neq" -> "<>"
|
||||
_ -> "="
|
||||
|
||||
limitClause :: Maybe R.NonnegRange -> Text
|
||||
limitClause range =
|
||||
cs $ " LIMIT " <> limit <> " OFFSET " <> show offset <> " "
|
||||
|
||||
where
|
||||
limit = fromMaybe "ALL" $ show <$> (R.limit =<< range)
|
||||
offset = fromMaybe 0 $ R.offset <$> range
|
||||
|
||||
globalAndLimitedCounts :: Schema -> Text -> Net.Query -> Text
|
||||
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"
|
||||
orderParseTerm :: BS.ByteString -> Maybe OrderTerm
|
||||
orderParseTerm s =
|
||||
case split "." s of
|
||||
[d,c] ->
|
||||
if d `elem` ["asc", "desc"]
|
||||
then Just $ OrderTerm (cs c) $
|
||||
if d == "asc" then "asc" else "desc"
|
||||
else Nothing
|
||||
_ -> Nothing
|
||||
|
||||
+36
-31
@@ -1,8 +1,16 @@
|
||||
module RangeQuery where
|
||||
module RangeQuery (
|
||||
rangeParse
|
||||
, rangeRequested
|
||||
, rangeLimit
|
||||
, rangeOffset
|
||||
, NonnegRange
|
||||
) where
|
||||
|
||||
import Control.Applicative
|
||||
import Network.HTTP.Types.Header
|
||||
|
||||
import qualified Data.ByteString.Char8 as BS
|
||||
|
||||
import Data.Ranged.Boundaries
|
||||
import Data.Ranged.Ranges
|
||||
|
||||
@@ -14,6 +22,33 @@ import Data.Maybe (fromMaybe, listToMaybe)
|
||||
|
||||
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 n =
|
||||
Range (BoundaryBelow n) BoundaryAboveAll
|
||||
@@ -21,33 +56,3 @@ rangeGeq n =
|
||||
rangeLeq :: Int -> NonnegRange
|
||||
rangeLeq 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
|
||||
|
||||
Reference in New Issue
Block a user