Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 10 additions & 0 deletions Database/MySQL/Base.hs
Original file line number Diff line number Diff line change
Expand Up @@ -82,6 +82,10 @@ module Database.MySQL.Base
, initLibrary
, initThread
, endThread
-- * Prepared Statements
, prepare
, Statement.bindParams
, Statement.execute
) where

import Control.Applicative ((<$>), (<*>))
Expand All @@ -96,6 +100,8 @@ import Data.Int (Int64)
import Data.List (foldl')
import Data.Typeable (Typeable)
import Data.Word (Word, Word16, Word64)
import Data.Text (Text)
import qualified Database.MySQL.PreparedStatement as Statement
import Database.MySQL.Base.C
import Database.MySQL.Base.Types
import Foreign.C.String (CString, peekCString, withCString)
Expand Down Expand Up @@ -648,3 +654,7 @@ connectionError_ func ptr =do
errno <- mysql_errno ptr
msg <- peekCString =<< mysql_error ptr
throw $ ConnectionError func (fromIntegral errno) msg

prepare :: Connection -> Text -> IO Statement.Statement
prepare conn q = do
Statement.prepare (connFP conn) q
64 changes: 64 additions & 0 deletions Database/MySQL/Base/C.hsc
Original file line number Diff line number Diff line change
Expand Up @@ -48,6 +48,7 @@ module Database.MySQL.Base.C
-- * Working with results
, mysql_free_result
, mysql_free_result_nonblock
, mysql_num_fields
, mysql_fetch_fields
, mysql_fetch_fields_nonblock
, mysql_data_seek
Expand All @@ -68,6 +69,19 @@ module Database.MySQL.Base.C
, mysql_library_init
, mysql_thread_init
, mysql_thread_end


, mysql_stmt_init
, mysql_stmt_close
, mysql_stmt_prepare
, mysql_stmt_bind_param
, mysql_stmt_bind_result
, mysql_stmt_execute
, mysql_stmt_fetch
, mysql_stmt_store_result
, mysql_stmt_result_metadata
, mysql_stmt_free_result
, mysql_stmt_error
) where

#include "mysql_signals.h"
Expand Down Expand Up @@ -242,6 +256,9 @@ foreign import ccall safe "mysql_signals.h _hs_mysql_free_result" mysql_free_res
foreign import ccall safe "mysql.h mysql_free_result" mysql_free_result_nonblock
:: Ptr MYSQL_RES -> IO ()

foreign import ccall safe mysql_num_fields
:: Ptr MYSQL_RES -> IO CUInt

foreign import ccall safe mysql_fetch_fields
:: Ptr MYSQL_RES -> IO (Ptr Field)

Expand Down Expand Up @@ -299,3 +316,50 @@ foreign import ccall safe "mysql.h" mysql_thread_init

foreign import ccall safe "mysql.h" mysql_thread_end
:: IO ()


-- Prepared Statements
foreign import ccall safe mysql_stmt_init
:: Ptr MYSQL -> IO (Ptr MYSQL_STMT)

foreign import ccall safe mysql_stmt_close
:: Ptr MYSQL_STMT -> IO MyBool

foreign import ccall safe mysql_stmt_bind_param
:: Ptr MYSQL_STMT -> Ptr MYSQL_BIND -> IO MyBool

foreign import ccall safe mysql_stmt_bind_result
:: Ptr MYSQL_STMT -> Ptr MYSQL_BIND -> IO MyBool

foreign import ccall safe mysql_stmt_errno
:: Ptr MYSQL_STMT -> IO CInt

foreign import ccall safe mysql_stmt_error
:: Ptr MYSQL_STMT -> IO CString

foreign import ccall safe mysql_stmt_prepare
:: Ptr MYSQL_STMT -> CString -> CULong -> IO CInt

foreign import ccall safe mysql_stmt_execute
:: Ptr MYSQL_STMT -> IO CInt

foreign import ccall safe mysql_stmt_fetch
:: Ptr MYSQL_STMT -> IO CInt

foreign import ccall safe mysql_stmt_store_result
:: Ptr MYSQL_STMT -> IO CInt

foreign import ccall safe mysql_stmt_insert_id
:: Ptr MYSQL_STMT -> IO CULLong

foreign import ccall safe mysql_stmt_field_count
:: Ptr MYSQL_STMT -> IO CUInt

foreign import ccall safe mysql_stmt_affected_rows
:: Ptr MYSQL_STMT -> IO CULLong

foreign import ccall safe mysql_stmt_result_metadata
:: Ptr MYSQL_STMT -> IO (Ptr MYSQL_RES)

foreign import ccall safe mysql_stmt_free_result
:: Ptr MYSQL_STMT -> IO MyBool
105 changes: 104 additions & 1 deletion Database/MySQL/Base/Types.hsc
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,9 @@ module Database.MySQL.Base.Types
, MYSQL_ROW
, MYSQL_ROWS
, MYSQL_ROW_OFFSET
, MYSQL_STMT
, MYSQL_BIND(..)
, MYSQL_TIME(..)
, MyBool
-- * Field flags
, hasAllFlags
Expand Down Expand Up @@ -60,10 +63,12 @@ import Data.Word (Word, Word8)
import Foreign.C.Types (CChar, CInt, CUInt, CULong)
import Foreign.Ptr (Ptr)
import Foreign.Storable (Storable(..), peekByteOff)
import qualified Foreign as Foreign
import qualified Data.IntMap as IntMap

data MYSQL
data MYSQL_RES
data MYSQL_STMT
data MYSQL_ROWS
type MYSQL_ROW = Ptr (Ptr CChar)
type MYSQL_ROW_OFFSET = Ptr MYSQL_ROWS
Expand Down Expand Up @@ -109,6 +114,39 @@ data Type = Decimal
| Json
deriving (Enum, Eq, Show, Typeable)

fromType :: Type -> CInt
fromType t =
case t of
Decimal -> #const MYSQL_TYPE_DECIMAL
Tiny -> #const MYSQL_TYPE_TINY
Short -> #const MYSQL_TYPE_SHORT
Int24 -> #const MYSQL_TYPE_INT24
Long -> #const MYSQL_TYPE_LONG
Float -> #const MYSQL_TYPE_FLOAT
Double -> #const MYSQL_TYPE_DOUBLE
Null -> #const MYSQL_TYPE_NULL
Timestamp -> #const MYSQL_TYPE_TIMESTAMP
LongLong -> #const MYSQL_TYPE_LONGLONG
Date -> #const MYSQL_TYPE_DATE
Time -> #const MYSQL_TYPE_TIME
DateTime -> #const MYSQL_TYPE_DATETIME
Year -> #const MYSQL_TYPE_YEAR
NewDate -> #const MYSQL_TYPE_NEWDATE
VarChar -> #const MYSQL_TYPE_VARCHAR
Bit -> #const MYSQL_TYPE_BIT
NewDecimal -> #const MYSQL_TYPE_NEWDECIMAL
Enum -> #const MYSQL_TYPE_ENUM
Set -> #const MYSQL_TYPE_SET
TinyBlob -> #const MYSQL_TYPE_TINY_BLOB
MediumBlob -> #const MYSQL_TYPE_MEDIUM_BLOB
LongBlob -> #const MYSQL_TYPE_LONG_BLOB
Blob -> #const MYSQL_TYPE_BLOB
VarString -> #const MYSQL_TYPE_VAR_STRING
String -> #const MYSQL_TYPE_STRING
Geometry -> #const MYSQL_TYPE_GEOMETRY
Json -> mysql_type_json


toType :: CInt -> Type
toType v = IntMap.findWithDefault oops (fromIntegral v) typeMap
where
Expand Down Expand Up @@ -144,6 +182,71 @@ toType v = IntMap.findWithDefault oops (fromIntegral v) typeMap
, (mysql_type_json, Json)
]

data MYSQL_BIND = MYSQL_BIND
{ mysqlBindBufferType :: Type
, mysqlBindBuffer :: Ptr ()
, mysqlBindBufferLength :: CULong
, mysqlBindLength :: Ptr CULong
, mysqlBindIsNull :: Ptr MyBool
, mysqlBindIsUnsigned :: MyBool
, mysqlBindError :: Ptr MyBool
}

instance Storable MYSQL_BIND where
sizeOf _ = #{size MYSQL_BIND}
alignment _ = #{alignment MYSQL_BIND}
peek ptr =
MYSQL_BIND
<$> (toType <$> (#peek MYSQL_BIND, buffer_type) ptr)
<*> (#peek MYSQL_BIND, buffer) ptr
<*> (#peek MYSQL_BIND, buffer_length) ptr
<*> (#peek MYSQL_BIND, length) ptr
<*> (#peek MYSQL_BIND, is_null) ptr
<*> (#peek MYSQL_BIND, is_unsigned) ptr
<*> (#peek MYSQL_BIND, error) ptr
poke ptr val = do
(#poke MYSQL_BIND, buffer_type) ptr $ fromType $ mysqlBindBufferType val
(#poke MYSQL_BIND, buffer) ptr $ mysqlBindBuffer val
(#poke MYSQL_BIND, buffer_length) ptr $ mysqlBindBufferLength val
(#poke MYSQL_BIND, length) ptr $ mysqlBindLength val
(#poke MYSQL_BIND, is_null) ptr $ mysqlBindIsNull val
(#poke MYSQL_BIND, is_unsigned) ptr $ mysqlBindIsUnsigned val
(#poke MYSQL_BIND, error) ptr $ mysqlBindError val

data MYSQL_TIME = MYSQL_TIME
{ mysqlTimeYear :: CUInt
, mysqlTimeMonth :: CUInt
, mysqlTimeDay :: CUInt
, mysqlTimeHour :: CUInt
, mysqlTimeMinute :: CUInt
, mysqlTimeSecond :: CUInt
, mysqlTimeNeg :: Bool
, mysqlTimeSecondPart :: CULong
}

instance Storable MYSQL_TIME where
sizeOf _ = #{size MYSQL_TIME}
alignment _ = #{alignment MYSQL_TIME}
peek ptr =
MYSQL_TIME
<$> (#peek MYSQL_TIME, year) ptr
<*> (#peek MYSQL_TIME, month) ptr
<*> (#peek MYSQL_TIME, day) ptr
<*> (#peek MYSQL_TIME, hour) ptr
<*> (#peek MYSQL_TIME, minute) ptr
<*> (#peek MYSQL_TIME, second) ptr
<*> (Foreign.toBool <$> ((#peek MYSQL_TIME, neg) ptr :: IO MyBool))
<*> (#peek MYSQL_TIME, second_part) ptr
poke ptr val = do
(#poke MYSQL_TIME, year) ptr $ mysqlTimeYear val
(#poke MYSQL_TIME, month) ptr $ mysqlTimeMonth val
(#poke MYSQL_TIME, day) ptr $ mysqlTimeDay val
(#poke MYSQL_TIME, hour) ptr $ mysqlTimeHour val
(#poke MYSQL_TIME, minute) ptr $ mysqlTimeMinute val
(#poke MYSQL_TIME, second) ptr $ mysqlTimeSecond val
(#poke MYSQL_TIME, neg) ptr $ ((Foreign.fromBool $ mysqlTimeNeg val) :: MyBool)
(#poke MYSQL_TIME, second_part) ptr $ mysqlTimeSecondPart val

-- | A description of a field (column) of a table.
data Field = Field {
fieldName :: ByteString -- ^ Name of column.
Expand Down Expand Up @@ -239,7 +342,7 @@ peekField ptr = do

instance Storable Field where
sizeOf _ = #{size MYSQL_FIELD}
alignment _ = alignment (undefined :: Ptr CChar)
alignment _ = #{alignment MYSQL_FIELD}
peek = peekField
poke _ _ = return () -- Unused, but define it to avoid a warning

Expand Down
Loading