Thread IdentState through subqueries (fixes #28).

There used to be name clashes if a subquery referenced
an entity that was already being used on the outer query.
Now we thread the outer query's IdentState to its subqueries,
which use it instead of initialIdentState.

Note that clashes still may occur between subqueries of
a query, but I think that's harmless.
This commit is contained in:
Felipe Lessa 2013-09-15 04:16:35 -03:00
parent c5c76959bd
commit 33b1fafc2d

View File

@ -35,6 +35,9 @@ module Database.Esqueleto.Internal.Sql
, rawEsqueleto , rawEsqueleto
, toRawSql , toRawSql
, Mode(..) , Mode(..)
, IdentState
, initialIdentState
, IdentInfo
, SqlSelect , SqlSelect
, veryUnsafeCoerceSqlExprValue , veryUnsafeCoerceSqlExprValue
) where ) where
@ -219,9 +222,13 @@ newIdentFor = Q . lift . try . unDBName
return (I t) return (I t)
-- | Information needed to escape and use identifiers.
type IdentInfo = (Connection, IdentState)
-- | Use an identifier. -- | Use an identifier.
useIdent :: Connection -> Ident -> TLB.Builder useIdent :: IdentInfo -> Ident -> TLB.Builder
useIdent conn (I ident) = fromDBName conn $ DBName ident useIdent info (I ident) = fromDBName info $ DBName ident
---------------------------------------------------------------------- ----------------------------------------------------------------------
@ -240,7 +247,7 @@ data SqlExpr a where
-- connection (mainly for escaping names) and returns both an -- connection (mainly for escaping names) and returns both an
-- string ('TLB.Builder') and a list of values to be -- string ('TLB.Builder') and a list of values to be
-- interpolated by the SQL backend. -- interpolated by the SQL backend.
ERaw :: NeedParens -> (Connection -> (TLB.Builder, [PersistValue])) -> SqlExpr (Value a) ERaw :: NeedParens -> (IdentInfo -> (TLB.Builder, [PersistValue])) -> SqlExpr (Value a)
-- | 'EList' and 'EEmptyList' are used by list operators. -- | 'EList' and 'EEmptyList' are used by list operators.
EList :: SqlExpr (Value a) -> SqlExpr (ValueList a) EList :: SqlExpr (Value a) -> SqlExpr (ValueList a)
@ -256,7 +263,7 @@ data SqlExpr a where
EPreprocessedFrom :: a -> FromClause -> SqlExpr (PreprocessedFrom a) EPreprocessedFrom :: a -> FromClause -> SqlExpr (PreprocessedFrom a)
-- | Used by 'insertSelect'. -- | Used by 'insertSelect'.
EInsert :: Proxy a -> (Connection -> (TLB.Builder, [PersistValue])) -> SqlExpr (Insertion a) EInsert :: Proxy a -> (IdentInfo -> (TLB.Builder, [PersistValue])) -> SqlExpr (Insertion a)
data NeedParens = Parens | Never data NeedParens = Parens | Never
@ -317,7 +324,7 @@ instance Esqueleto SqlQuery SqlExpr SqlBackend where
sub_selectDistinct = sub SELECT_DISTINCT sub_selectDistinct = sub SELECT_DISTINCT
EEntity ident ^. field = EEntity ident ^. field =
ERaw Never $ \conn -> (useIdent conn ident <> ("." <> fieldName conn field), []) ERaw Never $ \info -> (useIdent info ident <> ("." <> fieldName info field), [])
EMaybe r ?. field = maybelize (r ^. field) EMaybe r ?. field = maybelize (r ^. field)
where where
@ -331,10 +338,10 @@ instance Esqueleto SqlQuery SqlExpr SqlBackend where
nothing = unsafeSqlValue "NULL" nothing = unsafeSqlValue "NULL"
joinV (ERaw p f) = ERaw p f joinV (ERaw p f) = ERaw p f
countRows = unsafeSqlValue "COUNT(*)" countRows = unsafeSqlValue "COUNT(*)"
count (ERaw _ f) = ERaw Never $ \conn -> let (b, vals) = f conn count (ERaw _ f) = ERaw Never $ \info -> let (b, vals) = f info
in ("COUNT" <> parens b, vals) in ("COUNT" <> parens b, vals)
not_ (ERaw p f) = ERaw Never $ \conn -> let (b, vals) = f conn not_ (ERaw p f) = ERaw Never $ \info -> let (b, vals) = f info
in ("NOT " <> parensM p b, vals) in ("NOT " <> parensM p b, vals)
(==.) = unsafeSqlBinOp " = " (==.) = unsafeSqlBinOp " = "
@ -399,21 +406,21 @@ instance ToSomeValues SqlExpr (SqlExpr (Value a)) where
toSomeValues a = [SomeValue a] toSomeValues a = [SomeValue a]
fieldName :: (PersistEntity val, PersistField typ) fieldName :: (PersistEntity val, PersistField typ)
=> Connection -> EntityField val typ -> TLB.Builder => IdentInfo -> EntityField val typ -> TLB.Builder
fieldName conn = fromDBName conn . fieldDB . persistFieldDef fieldName info = fromDBName info . fieldDB . persistFieldDef
setAux :: (PersistEntity val, PersistField typ) setAux :: (PersistEntity val, PersistField typ)
=> EntityField val typ => EntityField val typ
-> (SqlExpr (Entity val) -> SqlExpr (Value typ)) -> (SqlExpr (Entity val) -> SqlExpr (Value typ))
-> SqlExpr (Update val) -> SqlExpr (Update val)
setAux field mkVal = ESet $ \ent -> unsafeSqlBinOp " = " name (mkVal ent) setAux field mkVal = ESet $ \ent -> unsafeSqlBinOp " = " name (mkVal ent)
where name = ERaw Never $ \conn -> (fieldName conn field, mempty) where name = ERaw Never $ \info -> (fieldName info field, mempty)
sub :: PersistField a => Mode -> SqlQuery (SqlExpr (Value a)) -> SqlExpr (Value a) sub :: PersistField a => Mode -> SqlQuery (SqlExpr (Value a)) -> SqlExpr (Value a)
sub mode query = ERaw Parens $ \conn -> toRawSql mode pureQuery conn query sub mode query = ERaw Parens $ \info -> toRawSql mode pureQuery info query
fromDBName :: Connection -> DBName -> TLB.Builder fromDBName :: IdentInfo -> DBName -> TLB.Builder
fromDBName conn = TLB.fromText . connEscapeName conn fromDBName (conn, _) = TLB.fromText . connEscapeName conn
existsHelper :: SqlQuery () -> SqlExpr (Value Bool) existsHelper :: SqlQuery () -> SqlExpr (Value Bool)
existsHelper = sub SELECT . (>> return true) existsHelper = sub SELECT . (>> return true)
@ -444,8 +451,8 @@ ifNotEmptyList (EList _) _ x = x
unsafeSqlBinOp :: TLB.Builder -> SqlExpr (Value a) -> SqlExpr (Value b) -> SqlExpr (Value c) unsafeSqlBinOp :: TLB.Builder -> SqlExpr (Value a) -> SqlExpr (Value b) -> SqlExpr (Value c)
unsafeSqlBinOp op (ERaw p1 f1) (ERaw p2 f2) = ERaw Parens f unsafeSqlBinOp op (ERaw p1 f1) (ERaw p2 f2) = ERaw Parens f
where where
f conn = let (b1, vals1) = f1 conn f info = let (b1, vals1) = f1 info
(b2, vals2) = f2 conn (b2, vals2) = f2 info
in ( parensM p1 b1 <> op <> parensM p2 b2 in ( parensM p1 b1 <> op <> parensM p2 b2
, vals1 <> vals2 ) , vals1 <> vals2 )
{-# INLINE unsafeSqlBinOp #-} {-# INLINE unsafeSqlBinOp #-}
@ -463,9 +470,9 @@ unsafeSqlValue v = ERaw Never $ \_ -> (v, mempty)
unsafeSqlFunction :: UnsafeSqlFunctionArgument a => unsafeSqlFunction :: UnsafeSqlFunctionArgument a =>
TLB.Builder -> a -> SqlExpr (Value b) TLB.Builder -> a -> SqlExpr (Value b)
unsafeSqlFunction name arg = unsafeSqlFunction name arg =
ERaw Never $ \conn -> ERaw Never $ \info ->
let (argsTLB, argsVals) = let (argsTLB, argsVals) =
uncommas' $ map (\(ERaw _ f) -> f conn) $ toArgList arg uncommas' $ map (\(ERaw _ f) -> f info) $ toArgList arg
in (name <> parens argsTLB, argsVals) in (name <> parens argsTLB, argsVals)
class UnsafeSqlFunctionArgument a where class UnsafeSqlFunctionArgument a where
@ -527,7 +534,7 @@ rawSelectSource mode query = src
run conn = run conn =
uncurry rawQuery $ uncurry rawQuery $
first builderToText $ first builderToText $
toRawSql mode pureQuery conn query toRawSql mode pureQuery (conn, initialIdentState) query
massage = do massage = do
mrow <- C.await mrow <- C.await
@ -639,7 +646,7 @@ rawEsqueleto mode query = do
conn <- SqlPersistT R.ask conn <- SqlPersistT R.ask
uncurry rawExecuteCount $ uncurry rawExecuteCount $
first builderToText $ first builderToText $
toRawSql mode pureQuery conn query toRawSql mode pureQuery (conn, initialIdentState) query
-- | Execute an @esqueleto@ @DELETE@ query inside @persistent@'s -- | Execute an @esqueleto@ @DELETE@ query inside @persistent@'s
@ -723,24 +730,37 @@ builderToText = TL.toStrict . TLB.toLazyTextWith defaultChunkSize
-- @esqueleto@, instead of manually using this function (which is -- @esqueleto@, instead of manually using this function (which is
-- possible but tedious), you may just turn on query logging of -- possible but tedious), you may just turn on query logging of
-- @persistent@. -- @persistent@.
toRawSql :: SqlSelect a r => Mode -> QueryType a -> Connection -> SqlQuery a -> (TLB.Builder, [PersistValue]) toRawSql :: SqlSelect a r => Mode -> QueryType a -> IdentInfo -> SqlQuery a -> (TLB.Builder, [PersistValue])
toRawSql mode qt conn query = toRawSql mode qt (conn, firstIdentState) query =
let (ret, SideData fromClauses setClauses whereClauses groupByClause havingClause orderByClauses limitClause) = let ((ret, sd), finalIdentState) =
flip S.evalState initialIdentState $ flip S.runState firstIdentState $
W.runWriterT $ W.runWriterT $
unQ query unQ query
SideData fromClauses
setClauses
whereClauses
groupByClause
havingClause
orderByClauses
limitClause = sd
-- Pass the finalIdentState (containing all identifiers
-- that were used) to the subsequent calls. This ensures
-- that no name clashes will occur on subqueries that may
-- appear on the expressions below.
info = (conn, finalIdentState)
in mconcat in mconcat
[ makeInsert qt ret [ makeInsert qt ret
, makeSelect conn mode ret , makeSelect info mode ret
, makeFrom conn mode fromClauses , makeFrom info mode fromClauses
, makeSet conn setClauses , makeSet info setClauses
, makeWhere conn whereClauses , makeWhere info whereClauses
, makeGroupBy conn groupByClause , makeGroupBy info groupByClause
, makeHaving conn havingClause , makeHaving info havingClause
, makeOrderBy conn orderByClauses , makeOrderBy info orderByClauses
, makeLimit conn limitClause , makeLimit info limitClause
] ]
-- | (Internal) Mode of query being converted by 'toRawSql'. -- | (Internal) Mode of query being converted by 'toRawSql'.
data Mode = SELECT | SELECT_DISTINCT | DELETE | UPDATE data Mode = SELECT | SELECT_DISTINCT | DELETE | UPDATE
@ -767,21 +787,21 @@ uncommas' :: Monoid a => [(TLB.Builder, a)] -> (TLB.Builder, a)
uncommas' = (uncommas *** mconcat) . unzip uncommas' = (uncommas *** mconcat) . unzip
makeSelect :: SqlSelect a r => Connection -> Mode -> a -> (TLB.Builder, [PersistValue]) makeSelect :: SqlSelect a r => IdentInfo -> Mode -> a -> (TLB.Builder, [PersistValue])
makeSelect conn mode ret = makeSelect info mode ret =
case mode of case mode of
SELECT -> withCols "SELECT " SELECT -> withCols "SELECT "
SELECT_DISTINCT -> withCols "SELECT DISTINCT " SELECT_DISTINCT -> withCols "SELECT DISTINCT "
DELETE -> plain "DELETE " DELETE -> plain "DELETE "
UPDATE -> plain "UPDATE " UPDATE -> plain "UPDATE "
where where
withCols v = first (v <>) (sqlSelectCols conn ret) withCols v = first (v <>) (sqlSelectCols info ret)
plain v = (v, []) plain v = (v, [])
makeFrom :: Connection -> Mode -> [FromClause] -> (TLB.Builder, [PersistValue]) makeFrom :: IdentInfo -> Mode -> [FromClause] -> (TLB.Builder, [PersistValue])
makeFrom _ _ [] = mempty makeFrom _ _ [] = mempty
makeFrom conn mode fs = ret makeFrom info mode fs = ret
where where
ret = case collectOnClauses fs of ret = case collectOnClauses fs of
Left expr -> throw $ mkExc expr Left expr -> throw $ mkExc expr
@ -802,8 +822,8 @@ makeFrom conn mode fs = ret
base ident@(I identText) def = base ident@(I identText) def =
let db@(DBName dbText) = entityDB def let db@(DBName dbText) = entityDB def
in ( if dbText == identText in ( if dbText == identText
then fromDBName conn db then fromDBName info db
else fromDBName conn db <> (" AS " <> useIdent conn ident) else fromDBName info db <> (" AS " <> useIdent info ident)
, mempty ) , mempty )
fromKind InnerJoinKind = " INNER JOIN " fromKind InnerJoinKind = " INNER JOIN "
@ -812,56 +832,56 @@ makeFrom conn mode fs = ret
fromKind RightOuterJoinKind = " RIGHT OUTER JOIN " fromKind RightOuterJoinKind = " RIGHT OUTER JOIN "
fromKind FullOuterJoinKind = " FULL OUTER JOIN " fromKind FullOuterJoinKind = " FULL OUTER JOIN "
makeOnClause (ERaw _ f) = first (" ON " <>) (f conn) makeOnClause (ERaw _ f) = first (" ON " <>) (f info)
mkExc :: SqlExpr (Value Bool) -> OnClauseWithoutMatchingJoinException mkExc :: SqlExpr (Value Bool) -> OnClauseWithoutMatchingJoinException
mkExc (ERaw _ f) = mkExc (ERaw _ f) =
OnClauseWithoutMatchingJoinException $ OnClauseWithoutMatchingJoinException $
TL.unpack $ TLB.toLazyText $ fst (f conn) TL.unpack $ TLB.toLazyText $ fst (f info)
makeSet :: Connection -> [SetClause] -> (TLB.Builder, [PersistValue]) makeSet :: IdentInfo -> [SetClause] -> (TLB.Builder, [PersistValue])
makeSet _ [] = mempty makeSet _ [] = mempty
makeSet conn os = first ("\nSET " <>) $ uncommas' (map mk os) makeSet info os = first ("\nSET " <>) $ uncommas' (map mk os)
where where
mk (SetClause (ERaw _ f)) = f conn mk (SetClause (ERaw _ f)) = f info
makeWhere :: Connection -> WhereClause -> (TLB.Builder, [PersistValue]) makeWhere :: IdentInfo -> WhereClause -> (TLB.Builder, [PersistValue])
makeWhere _ NoWhere = mempty makeWhere _ NoWhere = mempty
makeWhere conn (Where (ERaw _ f)) = first ("\nWHERE " <>) (f conn) makeWhere info (Where (ERaw _ f)) = first ("\nWHERE " <>) (f info)
makeGroupBy :: Connection -> GroupByClause -> (TLB.Builder, [PersistValue]) makeGroupBy :: IdentInfo -> GroupByClause -> (TLB.Builder, [PersistValue])
makeGroupBy _ (GroupBy []) = (mempty, []) makeGroupBy _ (GroupBy []) = (mempty, [])
makeGroupBy conn (GroupBy fields) = first ("\nGROUP BY " <>) build makeGroupBy info (GroupBy fields) = first ("\nGROUP BY " <>) build
where where
build = uncommas' $ map (\(SomeValue (ERaw _ f)) -> f conn) fields build = uncommas' $ map (\(SomeValue (ERaw _ f)) -> f info) fields
makeHaving :: Connection -> WhereClause -> (TLB.Builder, [PersistValue]) makeHaving :: IdentInfo -> WhereClause -> (TLB.Builder, [PersistValue])
makeHaving _ NoWhere = mempty makeHaving _ NoWhere = mempty
makeHaving conn (Where (ERaw _ f)) = first ("\nHAVING " <>) (f conn) makeHaving info (Where (ERaw _ f)) = first ("\nHAVING " <>) (f info)
makeOrderBy :: Connection -> [OrderByClause] -> (TLB.Builder, [PersistValue]) makeOrderBy :: IdentInfo -> [OrderByClause] -> (TLB.Builder, [PersistValue])
makeOrderBy _ [] = mempty makeOrderBy _ [] = mempty
makeOrderBy conn os = first ("\nORDER BY " <>) $ uncommas' (map mk os) makeOrderBy info os = first ("\nORDER BY " <>) $ uncommas' (map mk os)
where where
mk (EOrderBy t (ERaw p f)) = first ((<> orderByType t) . parensM p) (f conn) mk (EOrderBy t (ERaw p f)) = first ((<> orderByType t) . parensM p) (f info)
orderByType ASC = " ASC" orderByType ASC = " ASC"
orderByType DESC = " DESC" orderByType DESC = " DESC"
makeLimit :: Connection -> LimitClause -> (TLB.Builder, [PersistValue]) makeLimit :: IdentInfo -> LimitClause -> (TLB.Builder, [PersistValue])
makeLimit _ (Limit Nothing Nothing) = mempty makeLimit _ (Limit Nothing Nothing) = mempty
makeLimit _ (Limit Nothing (Just 0)) = mempty makeLimit _ (Limit Nothing (Just 0)) = mempty
makeLimit conn (Limit ml mo) = (ret, mempty) makeLimit info (Limit ml mo) = (ret, mempty)
where where
ret = TLB.singleton '\n' <> (limitTLB <> offsetTLB) ret = TLB.singleton '\n' <> (limitTLB <> offsetTLB)
limitTLB = limitTLB =
case ml of case ml of
Just l -> "LIMIT " <> TLBI.decimal l Just l -> "LIMIT " <> TLBI.decimal l
Nothing -> TLB.fromText (connNoLimit conn) Nothing -> TLB.fromText (connNoLimit $ fst info)
offsetTLB = offsetTLB =
case mo of case mo of
@ -886,7 +906,7 @@ class SqlSelect a r | a -> r, r -> a where
-- | Creates the variable part of the @SELECT@ query and -- | Creates the variable part of the @SELECT@ query and
-- returns the list of 'PersistValue's that will be given to -- returns the list of 'PersistValue's that will be given to
-- 'rawQuery'. -- 'rawQuery'.
sqlSelectCols :: Connection -> a -> (TLB.Builder, [PersistValue]) sqlSelectCols :: IdentInfo -> a -> (TLB.Builder, [PersistValue])
-- | Number of columns that will be consumed. -- | Number of columns that will be consumed.
sqlSelectColCount :: Proxy a -> Int sqlSelectColCount :: Proxy a -> Int
@ -897,7 +917,7 @@ class SqlSelect a r | a -> r, r -> a where
-- | You may return an insertion of some PersistEntity -- | You may return an insertion of some PersistEntity
instance PersistEntity a => SqlSelect (SqlExpr (Insertion a)) (Insertion a) where instance PersistEntity a => SqlSelect (SqlExpr (Insertion a)) (Insertion a) where
sqlSelectCols conn (EInsert _ f) = f conn sqlSelectCols info (EInsert _ f) = f info
sqlSelectColCount = const 0 sqlSelectColCount = const 0
sqlSelectProcessRow = const (Right (error msg)) sqlSelectProcessRow = const (Right (error msg))
where where
@ -913,10 +933,10 @@ instance SqlSelect () () where
-- | You may return an 'Entity' from a 'select' query. -- | You may return an 'Entity' from a 'select' query.
instance PersistEntity a => SqlSelect (SqlExpr (Entity a)) (Entity a) where instance PersistEntity a => SqlSelect (SqlExpr (Entity a)) (Entity a) where
sqlSelectCols conn expr@(EEntity ident) = ret sqlSelectCols info expr@(EEntity ident) = ret
where where
process ed = uncommas $ process ed = uncommas $
map ((name <>) . fromDBName conn) $ map ((name <>) . fromDBName info) $
(entityID ed:) $ (entityID ed:) $
map fieldDB $ map fieldDB $
entityFields ed entityFields ed
@ -926,7 +946,7 @@ instance PersistEntity a => SqlSelect (SqlExpr (Entity a)) (Entity a) where
-- clause), while 'rawSql' assumes that it's just the -- clause), while 'rawSql' assumes that it's just the
-- name of the table (which doesn't allow self-joins, for -- name of the table (which doesn't allow self-joins, for
-- example). -- example).
name = useIdent conn ident <> "." name = useIdent info ident <> "."
ret = let ed = entityDef $ getEntityVal $ return expr ret = let ed = entityDef $ getEntityVal $ return expr
in (process ed, mempty) in (process ed, mempty)
sqlSelectColCount = (+1) . length . entityFields . entityDef . getEntityVal sqlSelectColCount = (+1) . length . entityFields . entityDef . getEntityVal
@ -941,7 +961,7 @@ getEntityVal = const Proxy
-- | You may return a possibly-@NULL@ 'Entity' from a 'select' query. -- | You may return a possibly-@NULL@ 'Entity' from a 'select' query.
instance PersistEntity a => SqlSelect (SqlExpr (Maybe (Entity a))) (Maybe (Entity a)) where instance PersistEntity a => SqlSelect (SqlExpr (Maybe (Entity a))) (Maybe (Entity a)) where
sqlSelectCols conn (EMaybe ent) = sqlSelectCols conn ent sqlSelectCols info (EMaybe ent) = sqlSelectCols info ent
sqlSelectColCount = sqlSelectColCount . fromEMaybe sqlSelectColCount = sqlSelectColCount . fromEMaybe
where where
fromEMaybe :: Proxy (SqlExpr (Maybe e)) -> Proxy (SqlExpr e) fromEMaybe :: Proxy (SqlExpr (Maybe e)) -> Proxy (SqlExpr e)
@ -954,8 +974,8 @@ instance PersistEntity a => SqlSelect (SqlExpr (Maybe (Entity a))) (Maybe (Entit
-- | You may return any single value (i.e. a single column) from -- | You may return any single value (i.e. a single column) from
-- a 'select' query. -- a 'select' query.
instance PersistField a => SqlSelect (SqlExpr (Value a)) (Value a) where instance PersistField a => SqlSelect (SqlExpr (Value a)) (Value a) where
sqlSelectCols esc (ERaw p f) = let (b, vals) = f esc sqlSelectCols info (ERaw p f) = let (b, vals) = f info
in (parensM p b, vals) in (parensM p b, vals)
sqlSelectColCount = const 1 sqlSelectColCount = const 1
sqlSelectProcessRow [pv] = Value <$> fromPersistValue pv sqlSelectProcessRow [pv] = Value <$> fromPersistValue pv
sqlSelectProcessRow _ = Left "SqlSelect (Value a): wrong number of columns." sqlSelectProcessRow _ = Left "SqlSelect (Value a): wrong number of columns."
@ -1468,4 +1488,4 @@ insertGeneralSelect :: (MonadLogger m, MonadResourceBase m, SqlSelect (SqlExpr (
Mode -> SqlQuery (SqlExpr (Insertion a)) -> SqlPersistT m () Mode -> SqlQuery (SqlExpr (Insertion a)) -> SqlPersistT m ()
insertGeneralSelect mode query = do insertGeneralSelect mode query = do
conn <- SqlPersistT R.ask conn <- SqlPersistT R.ask
uncurry rawExecute $ first builderToText $ toRawSql mode insertQuery conn query uncurry rawExecute $ first builderToText $ toRawSql mode insertQuery (conn, initialIdentState) query