This change makes the API surface between MainTx and App smaller. Currently, App reconstructs a database transaction by unpacking the isolation level, transaction mode, DbHandler, and transaction runner returned by MainTx. That exposes MainTx internals at the call site even though MainTx already owns query setup, execution, decoding, and rollback behavior. The goal is to keep transaction assembly in MainTx while App remains responsible for pool execution, database error mapping, and response orchestration. DbTx now carries the assembled SQL session, and App passes that session directly to the connection pool.
274 lines
12 KiB
Haskell
274 lines
12 KiB
Haskell
{-# LANGUAGE NamedFieldPuns #-}
|
|
{-# LANGUAGE RecordWildCards #-}
|
|
{-|
|
|
Module : PostgREST.MainTx
|
|
Description : PostgREST transaction executor
|
|
|
|
This module parametrizes, prepares, executes SQL queries and decodes their results.
|
|
-}
|
|
module PostgREST.MainTx
|
|
( MainTx (..)
|
|
, DbResult (..)
|
|
, ResultSet (..)
|
|
, mainTx
|
|
) where
|
|
|
|
import Control.Lens ((^?))
|
|
import Control.Monad.Extra (whenJust)
|
|
import qualified Data.Aeson.Lens as L
|
|
import qualified Data.ByteString as BS hiding
|
|
(break)
|
|
import qualified Data.ByteString.Char8 as BS
|
|
import qualified Data.HashMap.Strict as HM
|
|
import qualified Data.Set as S
|
|
import qualified Hasql.Decoders as HD
|
|
import qualified Hasql.DynamicStatements.Statement as SQL
|
|
import qualified Hasql.Session as SQL (Session)
|
|
import qualified Hasql.Transaction as SQL
|
|
import qualified Hasql.Transaction.Sessions as SQL
|
|
|
|
import qualified PostgREST.Error as Error
|
|
import qualified PostgREST.SchemaCache as SchemaCache
|
|
|
|
|
|
import PostgREST.ApiRequest (ApiRequest (..))
|
|
import PostgREST.ApiRequest.Preferences (PreferCount (..),
|
|
PreferHandling (..),
|
|
PreferMaxAffected (..),
|
|
PreferTransaction (..),
|
|
Preferences (..))
|
|
import PostgREST.ApiRequest.Types (Mutation (..))
|
|
import PostgREST.Auth.Types (AuthResult (..))
|
|
import PostgREST.Config (AppConfig (..),
|
|
OpenAPIMode (..))
|
|
import PostgREST.Error (Error)
|
|
import PostgREST.MediaType (MediaType (..))
|
|
import PostgREST.Plan (ActionPlan (..),
|
|
CrudPlan (..),
|
|
DbActionPlan (..),
|
|
InfoPlan (..),
|
|
InspectPlan (..))
|
|
import PostgREST.Query (MainQuery (..))
|
|
import PostgREST.SchemaCache (SchemaCache (..))
|
|
import PostgREST.SchemaCache.Identifiers (QualifiedIdentifier (..))
|
|
import PostgREST.SchemaCache.Routine (Routine (..), RoutineMap)
|
|
import PostgREST.SchemaCache.Table (TablesMap)
|
|
|
|
import Protolude hiding (Handler)
|
|
|
|
type DbHandler = ExceptT Error SQL.Transaction
|
|
|
|
data MainTx
|
|
= DbTx (SQL.Session (Either Error DbResult))
|
|
| NoDbTx DbResult
|
|
|
|
data DbResult
|
|
= DbCrudResult CrudPlan ResultSet
|
|
| DbPlanResult MediaType BS.ByteString
|
|
| MaybeDbResult InspectPlan (Maybe (TablesMap, RoutineMap, Maybe Text))
|
|
| NoDbResult InfoPlan
|
|
|
|
-- | Standard result set format used for the mqMain query
|
|
data ResultSet
|
|
= RSStandard
|
|
{ rsTableTotal :: Maybe Int64
|
|
-- ^ count of all the table rows
|
|
, rsQueryTotal :: Int64
|
|
-- ^ count of the query rows
|
|
, rsLocation :: [(BS.ByteString, BS.ByteString)]
|
|
-- ^ The Location header(only used for inserts) is represented as a list of strings containing
|
|
-- variable bindings like @"k1=eq.42"@, or the empty list if there is no location header.
|
|
, rsBody :: BS.ByteString
|
|
-- ^ the aggregated body of the query
|
|
, rsGucHeaders :: Maybe BS.ByteString
|
|
-- ^ the HTTP headers to be added to the response
|
|
, rsGucStatus :: Maybe Text
|
|
-- ^ the HTTP status to be added to the response
|
|
, rsInserted :: Maybe Int64
|
|
-- ^ the number of rows inserted (Only used for upserts)
|
|
}
|
|
|
|
mainTx :: MainQuery -> AppConfig -> AuthResult -> ApiRequest -> ActionPlan -> SchemaCache -> MainTx
|
|
mainTx _ _ _ _ (NoDb x) _ = NoDbTx $ NoDbResult x
|
|
mainTx genQ@MainQuery{..} conf@AppConfig{..} AuthResult{..} apiReq (Db plan) sCache =
|
|
DbTx $ SQL.transactionNoRetry isoLvl txMode $ runExceptT dbHandler
|
|
where
|
|
isoLvl = planIsoLvl conf authRole plan
|
|
txMode = planTxMode plan
|
|
dbHandler = do
|
|
lift $ SQL.statement mempty $ SQL.dynamicallyParameterized mqTxVars
|
|
HD.noResult configDbPreparedStatements
|
|
lift $ whenJust mqPreReq $ \q ->
|
|
SQL.statement mempty $ SQL.dynamicallyParameterized q
|
|
HD.noResult configDbPreparedStatements
|
|
actionResult genQ plan conf apiReq sCache
|
|
|
|
planTxMode :: DbActionPlan -> SQL.Mode
|
|
planTxMode (DbCrud _ x) = pTxMode x
|
|
planTxMode (MayUseDb x) = ipTxmode x
|
|
|
|
planIsoLvl :: AppConfig -> ByteString -> DbActionPlan -> SQL.IsolationLevel
|
|
planIsoLvl AppConfig{configRoleIsoLvl} role actPlan = case actPlan of
|
|
DbCrud _ CallReadPlan{crProc} -> fromMaybe roleIsoLvl $ pdIsoLvl crProc
|
|
_ -> roleIsoLvl
|
|
where
|
|
roleIsoLvl = HM.findWithDefault SQL.ReadCommitted role configRoleIsoLvl
|
|
|
|
actionResult :: MainQuery -> DbActionPlan -> AppConfig -> ApiRequest -> SchemaCache -> ExceptT Error SQL.Transaction DbResult
|
|
actionResult MainQuery{..} (DbCrud True plan) conf@AppConfig{..} apiReq _ = do
|
|
explRes <- lift $ SQL.statement mempty $ SQL.dynamicallyParameterized mqMain planRow configDbPreparedStatements
|
|
optionalRollback conf apiReq
|
|
pure $ DbPlanResult (pMedia plan) explRes
|
|
|
|
actionResult MainQuery{..} (DbCrud _ plan@WrappedReadPlan{..}) conf@AppConfig{..} apiReq@ApiRequest{iPreferences=Preferences{..}} _ = do
|
|
resultSet@RSStandard{rsTableTotal=tableTotal} <- lift $ SQL.statement mempty $ dynStmt (HD.singleRow $ standardRow True)
|
|
failNotSingular pMedia resultSet
|
|
optionalRollback conf apiReq
|
|
explainTotal <- lift . fmap join $ traverse (\snip ->
|
|
SQL.statement mempty $ SQL.dynamicallyParameterized snip decodeExplain configDbPreparedStatements)
|
|
mqExplain
|
|
|
|
pure $ DbCrudResult plan
|
|
resultSet{rsTableTotal=case preferCount of
|
|
Just PlannedCount -> explainTotal
|
|
Just EstimatedCount -> if tableTotal > (fromIntegral <$> configDbMaxRows)
|
|
then max <$> tableTotal <*> explainTotal
|
|
else tableTotal
|
|
_ -> tableTotal}
|
|
where
|
|
dynStmt decod = SQL.dynamicallyParameterized mqMain decod configDbPreparedStatements
|
|
|
|
decodeExplain :: HD.Result (Maybe Int64)
|
|
decodeExplain =
|
|
let row = HD.singleRow $ column HD.bytea in
|
|
(^? L.nth 0 . L.key "Plan" . L.key "Plan Rows" . L._Integral) <$> row
|
|
|
|
actionResult MainQuery{..} (DbCrud _ plan@MutateReadPlan{..}) conf@AppConfig{..} apiReq@ApiRequest{iPreferences=Preferences{..}} _ = do
|
|
resultSet <- lift $ SQL.statement mempty $ dynStmt decodeRow
|
|
failMutation resultSet
|
|
optionalRollback conf apiReq
|
|
pure $ DbCrudResult plan resultSet
|
|
where
|
|
dynStmt decod = SQL.dynamicallyParameterized mqMain decod configDbPreparedStatements
|
|
failMutation resultSet = case mrMutation of
|
|
MutationCreate -> do
|
|
failNotSingular pMedia resultSet
|
|
MutationUpdate -> do
|
|
failNotSingular pMedia resultSet
|
|
failExceedsMaxAffectedPref (preferMaxAffected,preferHandling) resultSet
|
|
MutationSingleUpsert -> do
|
|
failPut resultSet
|
|
MutationDelete -> do
|
|
failNotSingular pMedia resultSet
|
|
failExceedsMaxAffectedPref (preferMaxAffected,preferHandling) resultSet
|
|
decodeRow = fromMaybe (RSStandard Nothing 0 mempty mempty Nothing Nothing Nothing) <$> HD.rowMaybe (standardRow False)
|
|
|
|
actionResult MainQuery{..} (DbCrud _ plan@CallReadPlan{..}) conf@AppConfig{..} apiReq@ApiRequest{iPreferences=Preferences{..}} _ = do
|
|
resultSet <- lift $ SQL.statement mempty $ dynStmt decodeRow
|
|
optionalRollback conf apiReq
|
|
failNotSingular pMedia resultSet
|
|
failExceedsMaxAffectedPref (preferMaxAffected,preferHandling) resultSet
|
|
pure $ DbCrudResult plan resultSet
|
|
where
|
|
dynStmt decod = SQL.dynamicallyParameterized mqMain decod configDbPreparedStatements
|
|
decodeRow = fromMaybe (RSStandard (Just 0) 0 mempty mempty Nothing Nothing Nothing) <$> HD.rowMaybe (standardRow True)
|
|
|
|
actionResult MainQuery{mqOpenAPI=(tblsQ, funcsQ, schQ)} (MayUseDb plan@InspectPlan{ipSchema=tSchema}) AppConfig{..} _ sCache =
|
|
mainActionQuery
|
|
where
|
|
mainActionQuery = lift $
|
|
case configOpenApiMode of
|
|
OAFollowPriv -> do
|
|
tableAccess <- SQL.statement mempty $ SQL.dynamicallyParameterized tblsQ decodeAccessibleIdentifiers configDbPreparedStatements
|
|
accFuncs <- SQL.statement mempty $ SQL.dynamicallyParameterized funcsQ SchemaCache.decodeFuncs configDbPreparedStatements
|
|
schDesc <- SQL.statement mempty $ SQL.dynamicallyParameterized schQ decodeSchemaDesc configDbPreparedStatements
|
|
let tbls = HM.filterWithKey (\qi _ -> S.member qi tableAccess) $ SchemaCache.dbTables sCache
|
|
|
|
pure $ MaybeDbResult plan (Just (tbls, accFuncs, schDesc))
|
|
OAIgnorePriv -> do
|
|
schDesc <- SQL.statement mempty (SQL.dynamicallyParameterized schQ decodeSchemaDesc configDbPreparedStatements)
|
|
|
|
let tbls = HM.filterWithKey (\(QualifiedIdentifier sch _) _ -> sch == tSchema) (SchemaCache.dbTables sCache)
|
|
routs = HM.filterWithKey (\(QualifiedIdentifier sch _) _ -> sch == tSchema) (SchemaCache.dbRoutines sCache)
|
|
|
|
pure $ MaybeDbResult plan (Just (tbls, routs, schDesc))
|
|
OADisabled ->
|
|
pure $ MaybeDbResult plan Nothing
|
|
|
|
decodeSchemaDesc :: HD.Result (Maybe Text)
|
|
decodeSchemaDesc = join <$> HD.rowMaybe (nullableColumn HD.text)
|
|
|
|
decodeAccessibleIdentifiers :: HD.Result (S.Set QualifiedIdentifier)
|
|
decodeAccessibleIdentifiers =
|
|
let
|
|
row = QualifiedIdentifier
|
|
<$> column HD.text
|
|
<*> column HD.text
|
|
in
|
|
S.fromList <$> HD.rowList row
|
|
|
|
-- Makes sure the querystring pk matches the payload pk
|
|
-- e.g. PUT /items?id=eq.1 { "id" : 1, .. } is accepted,
|
|
-- PUT /items?id=eq.14 { "id" : 2, .. } is rejected.
|
|
-- If this condition is not satisfied then nothing is inserted,
|
|
-- check the WHERE for INSERT in QueryBuilder.hs to see how it's done
|
|
failPut :: ResultSet -> DbHandler ()
|
|
failPut RSStandard{rsQueryTotal=queryTotal} =
|
|
when (queryTotal /= 1) $ do
|
|
lift SQL.condemn
|
|
throwError $ Error.ApiRequestErr Error.PutMatchingPkError
|
|
|
|
-- |
|
|
-- Fail a response if a single JSON object was requested and not exactly one
|
|
-- was found.
|
|
failNotSingular :: MediaType -> ResultSet -> DbHandler ()
|
|
failNotSingular mediaType RSStandard{rsQueryTotal=queryTotal} =
|
|
when (elem mediaType [MTVndSingularJSON True, MTVndSingularJSON False] && queryTotal /= 1) $ do
|
|
lift SQL.condemn
|
|
throwError $ Error.ApiRequestErr . Error.SingularityError $ toInteger queryTotal
|
|
|
|
failExceedsMaxAffectedPref :: (Maybe PreferMaxAffected, Maybe PreferHandling) -> ResultSet -> DbHandler ()
|
|
failExceedsMaxAffectedPref (Nothing,_) _ = pure ()
|
|
failExceedsMaxAffectedPref (Just (PreferMaxAffected n), handling) RSStandard{rsQueryTotal=queryTotal} = when ((queryTotal > n) && (handling == Just Strict)) $ do
|
|
lift SQL.condemn
|
|
throwError $ Error.ApiRequestErr . Error.MaxAffectedViolationError $ toInteger queryTotal
|
|
|
|
-- | Set a transaction to roll back if requested
|
|
optionalRollback :: AppConfig -> ApiRequest -> DbHandler ()
|
|
optionalRollback AppConfig{..} ApiRequest{iPreferences=Preferences{..}} = do
|
|
lift $ when (shouldRollback || (configDbTxRollbackAll && not shouldCommit)) $ do
|
|
SQL.sql "SET CONSTRAINTS ALL IMMEDIATE"
|
|
SQL.condemn
|
|
where
|
|
shouldCommit =
|
|
preferTransaction == Just Commit
|
|
shouldRollback =
|
|
preferTransaction == Just Rollback
|
|
|
|
-- | We use rowList because when doing EXPLAIN (FORMAT TEXT), the result comes as many rows. FORMAT JSON comes as one.
|
|
planRow :: HD.Result BS.ByteString
|
|
planRow = BS.unlines <$> HD.rowList (column HD.bytea)
|
|
|
|
column :: HD.Value a -> HD.Row a
|
|
column = HD.column . HD.nonNullable
|
|
|
|
nullableColumn :: HD.Value a -> HD.Row (Maybe a)
|
|
nullableColumn = HD.column . HD.nullable
|
|
|
|
arrayColumn :: HD.Value a -> HD.Row [a]
|
|
arrayColumn = column . HD.listArray . HD.nonNullable
|
|
|
|
standardRow :: Bool -> HD.Row ResultSet
|
|
standardRow noLocation =
|
|
RSStandard <$> nullableColumn HD.int8 <*> column HD.int8
|
|
<*> (if noLocation then pure mempty else fmap splitKeyValue <$> arrayColumn HD.bytea)
|
|
<*> (fromMaybe mempty <$> nullableColumn HD.bytea)
|
|
<*> nullableColumn HD.bytea
|
|
<*> nullableColumn HD.text
|
|
<*> nullableColumn HD.int8
|
|
where
|
|
splitKeyValue :: ByteString -> (ByteString, ByteString)
|
|
splitKeyValue kv =
|
|
let (k, v) = BS.break (== '=') kv in
|
|
(k, BS.tail v)
|