Pass escaping function as argument to ERaw (instead of Connection).

This commit is contained in:
Felipe Lessa 2012-09-03 15:53:38 -03:00
parent a4bd0268aa
commit fe7a32e7e4

View File

@ -86,7 +86,9 @@ idents _ =
-- | An expression on the SQL backend. -- | An expression on the SQL backend.
data SqlExpr a where data SqlExpr a where
EEntity :: Ident -> SqlExpr (Entity val) EEntity :: Ident -> SqlExpr (Entity val)
ERaw :: (Connection -> (TLB.Builder, [PersistValue])) -> SqlExpr (Single a) ERaw :: (Escape -> (TLB.Builder, [PersistValue])) -> SqlExpr (Single a)
type Escape = DBName -> TLB.Builder
instance Esqueleto SqlQuery SqlExpr SqlPersist where instance Esqueleto SqlQuery SqlExpr SqlPersist where
fromSingle = Q $ do fromSingle = Q $ do
@ -100,13 +102,13 @@ instance Esqueleto SqlQuery SqlExpr SqlPersist where
where_ expr = Q $ W.tell mempty { sdWhereClause = Where expr } where_ expr = Q $ W.tell mempty { sdWhereClause = Where expr }
EEntity (I ident) ^. field = ERaw $ \conn -> (ident <> ("." <> name conn field), []) EEntity (I ident) ^. field = ERaw $ \esc -> (ident <> ("." <> name esc field), [])
where name conn = fromDBName conn . fieldDB . persistFieldDef where name esc = esc . fieldDB . persistFieldDef
val = ERaw . const . (,) "?" . return . toPersistValue val = ERaw . const . (,) "?" . return . toPersistValue
not_ (ERaw f) = ERaw $ \conn -> let (b, vals) = f conn not_ (ERaw f) = ERaw $ \esc -> let (b, vals) = f esc
in ("NOT " <> parens b, vals) in ("NOT " <> parens b, vals)
(==.) = binop " = " (==.) = binop " = "
(>=.) = binop " >= " (>=.) = binop " >= "
@ -128,10 +130,10 @@ fromDBName conn = TLB.fromText . escapeName conn
binop :: TLB.Builder -> SqlExpr (Single a) -> SqlExpr (Single b) -> SqlExpr (Single c) binop :: TLB.Builder -> SqlExpr (Single a) -> SqlExpr (Single b) -> SqlExpr (Single c)
binop op (ERaw f1) (ERaw f2) = ERaw f binop op (ERaw f1) (ERaw f2) = ERaw f
where where
f conn = let (b1, vals1) = f1 conn f esc = let (b1, vals1) = f1 esc
(b2, vals2) = f2 conn (b2, vals2) = f2 esc
in ( parens b1 <> op <> parens b2 in ( parens b1 <> op <> parens b2
, vals1 <> vals2 ) , vals1 <> vals2 )
-- | TODO -- | TODO
@ -142,7 +144,7 @@ select :: ( SqlSelect a
=> SqlQuery a -> SqlPersist m [SqlSelectRet r] => SqlQuery a -> SqlPersist m [SqlSelectRet r]
select query = do select query = do
conn <- getConnection conn <- getConnection
uncurry rawSql $ toRawSelectSql conn query uncurry rawSql $ toRawSelectSql (fromDBName conn) query
-- | Get current database 'Connection'. -- | Get current database 'Connection'.
@ -151,22 +153,22 @@ getConnection = SqlPersist R.ask
-- | Pretty prints a 'SqlQuery' into a SQL query. -- | Pretty prints a 'SqlQuery' into a SQL query.
toRawSelectSql :: SqlSelect a => Connection -> SqlQuery a -> (T.Text, [PersistValue]) toRawSelectSql :: SqlSelect a => Escape -> SqlQuery a -> (T.Text, [PersistValue])
toRawSelectSql conn query = toRawSelectSql esc query =
let (ret, SideData fromClauses whereClauses) = let (ret, SideData fromClauses whereClauses) =
flip S.evalSupply (idents ()) $ flip S.evalSupply (idents ()) $
W.runWriterT $ W.runWriterT $
unQ query unQ query
(selectText, selectVars) = makeSelect conn ret (selectText, selectVars) = makeSelect esc ret
(whereText, whereVars) = makeWhere conn whereClauses (whereText, whereVars) = makeWhere esc whereClauses
text = TL.toStrict $ text = TL.toStrict $
TLB.toLazyText $ TLB.toLazyText $
mconcat mconcat
[ "SELECT " [ "SELECT "
, selectText , selectText
, makeFrom conn fromClauses , makeFrom esc fromClauses
, whereText , whereText
] ]
@ -175,27 +177,27 @@ toRawSelectSql conn query =
class RawSql (SqlSelectRet a) => SqlSelect a where class RawSql (SqlSelectRet a) => SqlSelect a where
type SqlSelectRet a :: * type SqlSelectRet a :: *
makeSelect :: Connection -> a -> (TLB.Builder, [PersistValue]) makeSelect :: Escape -> a -> (TLB.Builder, [PersistValue])
instance RawSql a => SqlSelect (SqlExpr a) where instance RawSql a => SqlSelect (SqlExpr a) where
type SqlSelectRet (SqlExpr a) = a type SqlSelectRet (SqlExpr a) = a
makeSelect _ (EEntity _) = ("??", mempty) makeSelect _ (EEntity _) = ("??", mempty)
makeSelect conn (ERaw f) = first parens (f conn) makeSelect esc (ERaw f) = first parens (f esc)
instance (SqlSelect a, SqlSelect b) => SqlSelect (a, b) where instance (SqlSelect a, SqlSelect b) => SqlSelect (a, b) where
type SqlSelectRet (a, b) = (SqlSelectRet a, SqlSelectRet b) type SqlSelectRet (a, b) = (SqlSelectRet a, SqlSelectRet b)
makeSelect conn (a, b) = uncommas' [makeSelect conn a, makeSelect conn b] makeSelect esc (a, b) = uncommas' [makeSelect esc a, makeSelect esc b]
instance (SqlSelect a, SqlSelect b, SqlSelect c) => SqlSelect (a, b, c) where instance (SqlSelect a, SqlSelect b, SqlSelect c) => SqlSelect (a, b, c) where
type SqlSelectRet (a, b, c) = type SqlSelectRet (a, b, c) =
( SqlSelectRet a ( SqlSelectRet a
, SqlSelectRet b , SqlSelectRet b
, SqlSelectRet c , SqlSelectRet c
) )
makeSelect conn (a, b, c) = makeSelect esc (a, b, c) =
uncommas' uncommas'
[ makeSelect conn a [ makeSelect esc a
, makeSelect conn b , makeSelect esc b
, makeSelect conn c , makeSelect esc c
] ]
instance ( SqlSelect a instance ( SqlSelect a
, SqlSelect b , SqlSelect b
@ -208,12 +210,12 @@ instance ( SqlSelect a
, SqlSelectRet c , SqlSelectRet c
, SqlSelectRet d , SqlSelectRet d
) )
makeSelect conn (a, b, c, d) = makeSelect esc (a, b, c, d) =
uncommas' uncommas'
[ makeSelect conn a [ makeSelect esc a
, makeSelect conn b , makeSelect esc b
, makeSelect conn c , makeSelect esc c
, makeSelect conn d , makeSelect esc d
] ]
instance ( SqlSelect a instance ( SqlSelect a
, SqlSelect b , SqlSelect b
@ -228,13 +230,13 @@ instance ( SqlSelect a
, SqlSelectRet d , SqlSelectRet d
, SqlSelectRet e , SqlSelectRet e
) )
makeSelect conn (a, b, c, d, e) = makeSelect esc (a, b, c, d, e) =
uncommas' uncommas'
[ makeSelect conn a [ makeSelect esc a
, makeSelect conn b , makeSelect esc b
, makeSelect conn c , makeSelect esc c
, makeSelect conn d , makeSelect esc d
, makeSelect conn e , makeSelect esc e
] ]
instance ( SqlSelect a instance ( SqlSelect a
, SqlSelect b , SqlSelect b
@ -251,14 +253,14 @@ instance ( SqlSelect a
, SqlSelectRet e , SqlSelectRet e
, SqlSelectRet f , SqlSelectRet f
) )
makeSelect conn (a, b, c, d, e, f) = makeSelect esc (a, b, c, d, e, f) =
uncommas' uncommas'
[ makeSelect conn a [ makeSelect esc a
, makeSelect conn b , makeSelect esc b
, makeSelect conn c , makeSelect esc c
, makeSelect conn d , makeSelect esc d
, makeSelect conn e , makeSelect esc e
, makeSelect conn f , makeSelect esc f
] ]
instance ( SqlSelect a instance ( SqlSelect a
, SqlSelect b , SqlSelect b
@ -277,15 +279,15 @@ instance ( SqlSelect a
, SqlSelectRet f , SqlSelectRet f
, SqlSelectRet g , SqlSelectRet g
) )
makeSelect conn (a, b, c, d, e, f, g) = makeSelect esc (a, b, c, d, e, f, g) =
uncommas' uncommas'
[ makeSelect conn a [ makeSelect esc a
, makeSelect conn b , makeSelect esc b
, makeSelect conn c , makeSelect esc c
, makeSelect conn d , makeSelect esc d
, makeSelect conn e , makeSelect esc e
, makeSelect conn f , makeSelect esc f
, makeSelect conn g , makeSelect esc g
] ]
instance ( SqlSelect a instance ( SqlSelect a
, SqlSelect b , SqlSelect b
@ -306,16 +308,16 @@ instance ( SqlSelect a
, SqlSelectRet g , SqlSelectRet g
, SqlSelectRet h , SqlSelectRet h
) )
makeSelect conn (a, b, c, d, e, f, g, h) = makeSelect esc (a, b, c, d, e, f, g, h) =
uncommas' uncommas'
[ makeSelect conn a [ makeSelect esc a
, makeSelect conn b , makeSelect esc b
, makeSelect conn c , makeSelect esc c
, makeSelect conn d , makeSelect esc d
, makeSelect conn e , makeSelect esc e
, makeSelect conn f , makeSelect esc f
, makeSelect conn g , makeSelect esc g
, makeSelect conn h , makeSelect esc h
] ]
@ -326,15 +328,15 @@ uncommas' :: Monoid a => [(TLB.Builder, a)] -> (TLB.Builder, a)
uncommas' = uncommas . map fst &&& mconcat . map snd uncommas' = uncommas . map fst &&& mconcat . map snd
makeFrom :: Connection -> [FromClause] -> TLB.Builder makeFrom :: Escape -> [FromClause] -> TLB.Builder
makeFrom conn = uncommas . map mk makeFrom esc = uncommas . map mk
where where
mk (From (I i) def) = fromDBName conn (entityDB def) <> (" AS " <> i) mk (From (I i) def) = esc (entityDB def) <> (" AS " <> i)
makeWhere :: Connection -> WhereClause -> (TLB.Builder, [PersistValue]) makeWhere :: Escape -> WhereClause -> (TLB.Builder, [PersistValue])
makeWhere _ NoWhere = mempty makeWhere _ NoWhere = mempty
makeWhere conn (Where (ERaw f)) = first ("\nWHERE " <>) (f conn) makeWhere esc (Where (ERaw f)) = first ("\nWHERE " <>) (f esc)
parens :: TLB.Builder -> TLB.Builder parens :: TLB.Builder -> TLB.Builder