cql-io 0.16.0 → 1.0.0
raw patch · 15 files changed
+968/−236 lines, 15 filesdep +Decimaldep +cql-iodep +raw-strings-qqdep ~basedep ~containersdep ~cqlsetup-changednew-uploaderPVP ok
version bump matches the API change (PVP)
Dependencies added: Decimal, cql-io, raw-strings-qq, tasty, tasty-hunit, unordered-containers
Dependency ranges changed: base, containers, cql, iproute, mtl, text, time, tinylog, uuid
API changes (from Hackage documentation)
+ Database.CQL.IO: All :: Consistency
+ Database.CQL.IO: Any :: Consistency
+ Database.CQL.IO: AuthMechanism :: Text -> AuthMechanism
+ Database.CQL.IO: AuthPass :: Text -> AuthPass
+ Database.CQL.IO: AuthUser :: Text -> AuthUser
+ Database.CQL.IO: AuthenticationRequired :: !AuthMechanism -> AuthenticationError
+ Database.CQL.IO: Authenticator :: !AuthMechanism -> (AuthContext -> IO (AuthResponse, s)) -> Maybe (s -> AuthChallenge -> IO (AuthResponse, s)) -> (s -> AuthSuccess -> IO ()) -> Authenticator
+ Database.CQL.IO: EachQuorum :: Consistency
+ Database.CQL.IO: LocalOne :: Consistency
+ Database.CQL.IO: LocalQuorum :: Consistency
+ Database.CQL.IO: LocalSerial :: Consistency
+ Database.CQL.IO: LocalSerialConsistency :: SerialConsistency
+ Database.CQL.IO: One :: Consistency
+ Database.CQL.IO: QueryParams :: !Consistency -> !Bool -> a -> Maybe Int32 -> Maybe PagingState -> Maybe SerialConsistency -> Maybe Bool -> QueryParams a
+ Database.CQL.IO: QueryString :: Text -> QueryString k a b
+ Database.CQL.IO: Quorum :: Consistency
+ Database.CQL.IO: Serial :: Consistency
+ Database.CQL.IO: SerialConsistency :: SerialConsistency
+ Database.CQL.IO: Three :: Consistency
+ Database.CQL.IO: Two :: Consistency
+ Database.CQL.IO: UnexpectedAuthenticationChallenge :: !AuthMechanism -> !AuthChallenge -> AuthenticationError
+ Database.CQL.IO: [authMechanism] :: Authenticator -> !AuthMechanism
+ Database.CQL.IO: [authOnChallenge] :: Authenticator -> Maybe (s -> AuthChallenge -> IO (AuthResponse, s))
+ Database.CQL.IO: [authOnRequest] :: Authenticator -> AuthContext -> IO (AuthResponse, s)
+ Database.CQL.IO: [authOnSuccess] :: Authenticator -> s -> AuthSuccess -> IO ()
+ Database.CQL.IO: [consistency] :: QueryParams a -> !Consistency
+ Database.CQL.IO: [enableTracing] :: QueryParams a -> Maybe Bool
+ Database.CQL.IO: [pageSize] :: QueryParams a -> Maybe Int32
+ Database.CQL.IO: [queryPagingState] :: QueryParams a -> Maybe PagingState
+ Database.CQL.IO: [serialConsistency] :: QueryParams a -> Maybe SerialConsistency
+ Database.CQL.IO: [skipMetaData] :: QueryParams a -> !Bool
+ Database.CQL.IO: [unQueryString] :: QueryString k a b -> Text
+ Database.CQL.IO: [values] :: QueryParams a -> a
+ Database.CQL.IO: authConnId :: Lens' AuthContext ConnId
+ Database.CQL.IO: authHost :: Lens' AuthContext InetAddr
+ Database.CQL.IO: data AuthContext
+ Database.CQL.IO: data AuthenticationError
+ Database.CQL.IO: data Authenticator
+ Database.CQL.IO: data ConnId
+ Database.CQL.IO: data Consistency :: *
+ Database.CQL.IO: data QueryParams a :: * -> *
+ Database.CQL.IO: data R :: *
+ Database.CQL.IO: data Row :: *
+ Database.CQL.IO: data S :: *
+ Database.CQL.IO: data SerialConsistency :: *
+ Database.CQL.IO: data W :: *
+ Database.CQL.IO: defQueryParams :: Consistency -> a -> QueryParams a
+ Database.CQL.IO: fromRow :: Cql a => Int -> Row -> Either String a
+ Database.CQL.IO: newtype AuthMechanism
+ Database.CQL.IO: newtype AuthPass
+ Database.CQL.IO: newtype AuthUser
+ Database.CQL.IO: newtype QueryString k a b :: * -> * -> * -> *
+ Database.CQL.IO: once :: MonadClient m => m a -> m a
+ Database.CQL.IO: passwordAuthenticator :: AuthUser -> AuthPass -> Authenticator
+ Database.CQL.IO: setAuthentication :: [Authenticator] -> Settings -> Settings
- Database.CQL.IO: [UnexpectedResponse] :: UnexpectedResponse
+ Database.CQL.IO: [UnexpectedResponse] :: !(Response k a b) -> UnexpectedResponse
Files
- AUTHORS +4/−1
- CHANGELOG +9/−0
- Setup.hs +0/−2
- cql-io.cabal +29/−5
- src/Database/CQL/IO.hs +123/−69
- src/Database/CQL/IO/Batch.hs +6/−8
- src/Database/CQL/IO/Client.hs +124/−83
- src/Database/CQL/IO/Cluster/Policies.hs +1/−1
- src/Database/CQL/IO/Connection.hs +66/−49
- src/Database/CQL/IO/Connection/Settings.hs +140/−0
- src/Database/CQL/IO/Connection/Socket.hs +1/−1
- src/Database/CQL/IO/Protocol.hs +7/−1
- src/Database/CQL/IO/Settings.hs +19/−8
- src/Database/CQL/IO/Types.hs +75/−8
- test/Main.hs +364/−0
AUTHORS view
@@ -2,6 +2,7 @@ ----------- - Toralf Wittner <tw@dtex.org>+- Roman S. Borschel <roman@pkaboo.org> Original Authors ----------------@@ -13,4 +14,6 @@ - Ricardo Catalinas Jiménez <jimenezrick@gmail.com> - Robert J. Macomber <robert.macomber@socrata.com>-- Roman S. Borschel <roman.borschel@googlemail.com>+- Roman S. Borschel <roman@pkaboo.org>+- Steve Severance <sseverance@alphaheavy.com>+- Ewout Van Troostenberghe <e@ewout.name>
CHANGELOG view
@@ -1,3 +1,12 @@+1.0.0+------+- Add support for CQL V4 binary protocol.+- Remove support for CQL V2 binary protocol.+- Add support for SASL-based authentication handlers.+- Bugfix: Retries for error responses were not handled correctly.+- Update and extend test suite.+- Require `cql >= 4.0`+ 0.16.0 ------ - Update Cabal settings to allow `cql` >= 3.1
− Setup.hs
@@ -1,2 +0,0 @@-import Distribution.Simple-main = defaultMain
cql-io.cabal view
@@ -1,11 +1,11 @@ name: cql-io-version: 0.16.0+version: 1.0.0 synopsis: Cassandra CQL client.-stability: experimental license: MPL-2.0 license-file: LICENSE author: Toralf Wittner-maintainer: Toralf Wittner <tw@dtex.org>+maintainer: Toralf Wittner <tw@dtex.org>,+ Roman S. Borschel <roman@pkaboo.org> copyright: (C) 2014-2016 Toralf Wittner homepage: https://gitlab.com/twittner/cql-io/ bug-reports: https://gitlab.com/twittner/cql-io/issues@@ -17,7 +17,7 @@ AUTHORS description:- CQL Cassandra driver supporting native protocol versions 2 and 3.+ CQL Cassandra driver supporting native protocol versions 3 and 4. . This library uses the <http://hackage.haskell.org/package/cql cql> library which implements Cassandra's CQL protocol and complements it with the@@ -63,6 +63,7 @@ Database.CQL.IO.Cluster.Policies Database.CQL.IO.Connection Database.CQL.IO.Connection.Socket+ Database.CQL.IO.Connection.Settings Database.CQL.IO.Hexdump Database.CQL.IO.Jobs Database.CQL.IO.Pool@@ -81,7 +82,7 @@ , base >= 4.7 && < 5.0 , bytestring >= 0.10 , containers >= 0.5- , cql >= 3.0+ , cql >= 4.0 , cryptohash >= 0.11 , data-default-class , exceptions >= 0.4@@ -101,5 +102,28 @@ , time >= 1.4 , transformers >= 0.3 , transformers-base >= 0.4+ , unordered-containers >= 0.2 , uuid >= 1.2.6 , vector >= 0.10++test-suite cql-io-tests+ type: exitcode-stdio-1.0+ default-language: Haskell2010+ main-is: Main.hs+ hs-source-dirs: test+ ghc-options: -threaded -Wall -O2 -fwarn-tabs+ build-depends:+ base >= 4.7+ , containers+ , cql+ , cql-io+ , Decimal+ , iproute >= 1.7+ , mtl+ , tasty >= 0.11+ , tasty-hunit >= 0.9+ , text+ , raw-strings-qq >= 1.1+ , time+ , tinylog+ , uuid >= 1.3
src/Database/CQL/IO.hs view
@@ -13,14 +13,14 @@ -- > import Data.Text (Text) -- > import Data.Functor.Identity -- > import Database.CQL.IO as Client--- > import Database.CQL.Protocol -- > import qualified System.Logger as Logger -- > -- > g <- Logger.new Logger.defSettings -- > c <- Client.init g defSettings--- > let p = QueryParams One False () Nothing Nothing Nothing--- > runClient c $ query ("SELECT cql_version from system.local" :: QueryString R () (Identity Text)) p--- [Identity "3.2.0"]+-- > let q = "SELECT cql_version from system.local" :: QueryString R () (Identity Text)+-- > let p = defQueryParams One ()+-- > runClient c (query q p)+-- [Identity "3.4.4"] -- > shutdown c -- @ --@@ -52,12 +52,12 @@ -- action using 'withPrepareStrategy'. {-# LANGUAGE DeriveFunctor #-}+{-# LANGUAGE LambdaCase #-} module Database.CQL.IO- ( -- * Client settings+ ( -- * Client Settings Settings- , PrepareStrategy (..)- , defSettings+ , S.defSettings , addContact , setCompression , setConnectTimeout@@ -70,6 +70,7 @@ , setPolicy , setPoolStripes , setPortNumber+ , PrepareStrategy (..) , setPrepareStrategy , setProtocolVersion , setResponseTimeout@@ -78,6 +79,18 @@ , setMaxRecvBuffer , setSSLContext + -- ** Authentication+ , setAuthentication+ , Authenticator (..)+ , AuthContext+ , ConnId+ , authConnId+ , authHost+ , AuthMechanism (..)+ , AuthUser (..)+ , AuthPass (..)+ , passwordAuthenticator+ -- ** Retry Settings , RetrySettings , noRetry@@ -90,36 +103,59 @@ , adjustSendTimeout , adjustResponseTimeout - -- * Query Runner- , RunQ (..)+ -- ** Load-balancing+ , Policy (..)+ , random+ , roundRobin - -- * Client monad+ -- *** Hosts+ , Host+ , HostEvent (..)+ , InetAddr (..)+ , hostAddr+ , dataCentre+ , rack++ -- * Client Monad , Client , MonadClient (..) , ClientState , DebugInfo (..) , init , runClient- , retry , shutdown , debugInfo + -- * Queries+ , R, W, S+ , QueryParams (..)+ , defQueryParams+ , Consistency (..)+ , SerialConsistency (..)+ , QueryString (..)++ -- ** Basic Queries , query , query1 , write , schema- , trans - , Page (..)- , emptyPage- , paginate-- -- * Prepared Queries+ -- ** Prepared Queries , PrepQuery , prepared , queryString - -- * Batch+ -- ** Paging+ , Page (..)+ , emptyPage+ , paginate++ -- ** Lightweight Transactions+ , Row+ , fromRow+ , trans++ -- ** Batch Queries , BatchM , addQuery , addPrepQuery@@ -128,34 +164,30 @@ , setSerialConsistency , batch - -- ** low-level- , request-- -- * Policies- , Policy (..)- , random- , roundRobin+ -- ** Retries+ , retry+ , once - -- ** Hosts- , Host- , HostEvent (..)- , InetAddr (..)- , hostAddr- , dataCentre- , rack+ -- ** Low-Level Queries+ --+ -- | Note: Use of these low-level functions may require additional imports from+ -- @Database.CQL.Protocol@ or its submodules in order to construct+ -- 'Request's and evaluate 'Response's.+ , RunQ (..)+ , request - -- * Exceptions- , InvalidSettings (..)- , InternalError (..)- , HostError (..)- , ConnectionError (..)- , UnexpectedResponse (..)- , Timeout (..)- , HashCollision (..)+ -- * Exceptions+ , InvalidSettings (..)+ , InternalError (..)+ , HostError (..)+ , ConnectionError (..)+ , UnexpectedResponse (..)+ , Timeout (..)+ , HashCollision (..)+ , AuthenticationError (..) ) where import Control.Applicative-import Control.Monad (void) import Control.Monad.Catch import Data.Maybe (isJust, listToMaybe) import Database.CQL.Protocol@@ -163,59 +195,80 @@ import Database.CQL.IO.Client import Database.CQL.IO.Cluster.Host import Database.CQL.IO.Cluster.Policies+import Database.CQL.IO.Connection.Settings as C import Database.CQL.IO.PrepQuery-import Database.CQL.IO.Settings+import Database.CQL.IO.Settings as S import Database.CQL.IO.Types import Prelude hiding (init) import qualified Database.CQL.IO.Batch as B --- | A type which can run a query against Cassandra.+-- | A type which can be run as a query. class RunQ q where runQ :: (MonadClient m, Tuple a, Tuple b) => q k a b -> QueryParams a -> m (Response k a b) instance RunQ QueryString where- runQ q p = do- r <- request (RqQuery (Query q p))- case r of- RsError _ e -> throwM e- _ -> return r+ runQ q p = request (RqQuery (Query q p)) instance RunQ PrepQuery where runQ q = liftClient . execute q --- | Run a CQL read-only query against a Cassandra node.+-- | Construct default 'QueryParams' for the given consistency+-- and bound values. In particular, no page size, paging state+-- or serial consistency will be set.+defQueryParams :: Consistency -> a -> QueryParams a+defQueryParams c a = QueryParams+ { consistency = c+ , values = a+ , skipMetaData = False+ , pageSize = Nothing+ , queryPagingState = Nothing+ , serialConsistency = Nothing+ , enableTracing = Nothing+ }++-- | Run a CQL read-only query returning a list of results. query :: (MonadClient m, Tuple a, Tuple b, RunQ q) => q R a b -> QueryParams a -> m [b] query q p = do r <- runQ q p- case r of- RsResult _ (RowsResult _ b) -> return b- _ -> throwM UnexpectedResponse+ getResult r >>= \case+ RowsResult _ b -> return b+ _ -> throwM $ UnexpectedResponse r --- | Run a CQL read-only query against a Cassandra node.+-- | Run a CQL read-only query returning a single result. query1 :: (MonadClient m, Tuple a, Tuple b, RunQ q) => q R a b -> QueryParams a -> m (Maybe b) query1 q p = listToMaybe <$> query q p --- | Run a CQL insert/update query against a Cassandra node.+-- | Run a CQL write-only query (e.g. insert\/update\/delete),+-- returning no result.+--+-- /Note: If the write operation is conditional, i.e. is in fact a "lightweight+-- transaction" returning a result, 'trans' must be used instead./ write :: (MonadClient m, Tuple a, RunQ q) => q W a () -> QueryParams a -> m ()-write q p = void $ runQ q p+write q p = do+ r <- runQ q p+ getResult r >>= \case+ VoidResult -> return ()+ _ -> throwM $ UnexpectedResponse r --- | Run a CQL insert/update query as a \"lightweight transaction\" against a Cassandra node.+-- | Run a CQL conditional write query (e.g. insert\/update\/delete) as a+-- "lightweight transaction", returning the result 'Row's describing the+-- outcome. trans :: (MonadClient m, Tuple a, RunQ q) => q W a Row -> QueryParams a -> m [Row] trans q p = do r <- runQ q p- case r of- RsResult _ (RowsResult _ b) -> return b- _ -> throwM UnexpectedResponse+ getResult r >>= \case+ RowsResult _ b -> return b+ _ -> throwM $ UnexpectedResponse' r --- | Run a CQL schema query against a Cassandra node.+-- | Run a CQL schema query, returning 'SchemaChange' information, if any. schema :: (MonadClient m, Tuple a, RunQ q) => q S a () -> QueryParams a -> m (Maybe SchemaChange)-schema x y = do- r <- runQ x y- case r of- RsResult _ (SchemaChangeResult s) -> return $ Just s- RsResult _ VoidResult -> return Nothing- _ -> throwM UnexpectedResponse+schema q p = do+ r <- runQ q p+ getResult r >>= \case+ SchemaChangeResult s -> return $ Just s+ VoidResult -> return Nothing+ _ -> throwM $ UnexpectedResponse r -- | Run a batch query against a Cassandra node. batch :: MonadClient m => BatchM () -> m ()@@ -247,10 +300,11 @@ paginate q p = do let p' = p { pageSize = pageSize p <|> Just 10000 } r <- runQ q p'- case r of- RsResult _ (RowsResult m b) ->+ getResult r >>= \case+ RowsResult m b -> if isJust (pagingState m) then return $ Page True b (paginate q p' { queryPagingState = pagingState m }) else return $ Page False b (return emptyPage)- _ -> throwM UnexpectedResponse+ _ -> throwM $ UnexpectedResponse r+
src/Database/CQL/IO/Batch.hs view
@@ -4,6 +4,7 @@ {-# LANGUAGE CPP #-} {-# LANGUAGE GeneralizedNewtypeDeriving #-}+{-# LANGUAGE LambdaCase #-} module Database.CQL.IO.Batch ( BatchM@@ -35,14 +36,11 @@ -- | Execute the complete 'Batch' statement. batch :: BatchM a -> Client () batch m = do- b <- execStateT (unBatchM m) s- checkRs =<< executeWithPrepare Nothing (RqBatch b :: Raw Request)- where- checkRs (RsResult _ VoidResult) = return ()- checkRs (RsError _ e) = throwM e- checkRs _ = throwM UnexpectedResponse-- s = Batch BatchLogged [] Quorum Nothing+ b <- execStateT (unBatchM m) (Batch BatchLogged [] Quorum Nothing)+ r <- executeWithPrepare Nothing (RqBatch b :: Raw Request)+ getResult r >>= \case+ VoidResult -> return ()+ _ -> throwM $ UnexpectedResponse' r -- | Add a query to this batch. addQuery :: (Show a, Tuple a, Tuple b) => QueryString W a b -> a -> BatchM ()
src/Database/CQL/IO/Client.hs view
@@ -11,8 +11,10 @@ {-# LANGUAGE OverloadedStrings #-} {-# LANGUAGE ScopedTypeVariables #-} {-# LANGUAGE TemplateHaskell #-}+{-# LANGUAGE TupleSections #-} {-# LANGUAGE TypeFamilies #-} {-# LANGUAGE UndecidableInstances #-}+{-# LANGUAGE ViewPatterns #-} module Database.CQL.IO.Client ( Client@@ -20,19 +22,20 @@ , ClientState , DebugInfo (..) , runClient- , Database.CQL.IO.Client.init+ , init , shutdown , request , requestN , request1- , mkRequest , execute , executeWithPrepare , prepare , retry+ , once , debugInfo , preparedQueries , withPrepareStrategy+ , getResult ) where import Control.Applicative@@ -51,7 +54,8 @@ #if MIN_VERSION_transformers(0,4,0) import Control.Monad.Trans.Except #endif-import Control.Retry (recovering, capDelay, exponentialBackoff, rsIterNumber)+import Control.Retry (capDelay, exponentialBackoff, rsIterNumber, applyPolicy)+import Control.Retry (recovering) import Data.Foldable (for_, foldrM) import Data.List (find) import Data.List.NonEmpty (NonEmpty (..))@@ -63,6 +67,7 @@ import Database.CQL.IO.Cluster.Host import Database.CQL.IO.Cluster.Policies import Database.CQL.IO.Connection hiding (request)+import Database.CQL.IO.Connection.Settings import Database.CQL.IO.Jobs (Jobs) import Database.CQL.IO.Pool import Database.CQL.IO.PrepQuery (PrepQuery, PreparedQueries)@@ -75,7 +80,7 @@ import Network.Socket (SockAddr (..), PortNumber) import OpenSSL.Session (SomeSSLException) import System.Logger.Class hiding (Settings, new, settings, create)-import Prelude+import Prelude hiding (init) import qualified Control.Monad.Reader as Reader import qualified Control.Monad.State.Strict as S@@ -86,6 +91,7 @@ import qualified Database.CQL.IO.Jobs as Jobs import qualified Database.CQL.IO.PrepQuery as PQ import qualified Database.CQL.IO.Timeouts as TM+import qualified Database.CQL.Protocol as Cql import qualified System.Logger as Logger data ControlState@@ -127,8 +133,7 @@ -- 'Database.CQL.IO.Client.init' and after finishing operation it should be -- terminated with 'shutdown'. ----- Actual CQL queries are handled by invoking 'request'.--- Additionally 'debugInfo' returns an internal cluster view.+-- To lift 'Client' actions into another monad, see 'MonadClient'. newtype Client a = Client { client :: ReaderT ClientState IO a } deriving ( Functor@@ -164,7 +169,7 @@ -- | Monads in which 'Client' actions may be embedded. class (Functor m, Applicative m, Monad m, MonadIO m, MonadCatch m) => MonadClient m where- -- | Lift a computation to the 'Client' monad.+ -- | Lift a computation from the 'Client' monad. liftClient :: Client a -> m a -- | Execute an action with a modified 'ClientState'. localState :: (ClientState -> ClientState) -> m a -> m a@@ -202,11 +207,21 @@ retry :: MonadClient m => RetrySettings -> m a -> m a retry r = localState (set (context.settings.retrySettings) r) +-- | Execute a client action once, without retries, i.e.+--+-- @once action = retry noRetry action@.+--+-- Primarily for use in applications where global 'RetrySettings'+-- are configured and need to be selectively disabled for individual+-- queries.+once :: MonadClient m => m a -> m a+once = retry noRetry+ -- | Change the default 'PrepareStrategy' for the given client action. withPrepareStrategy :: MonadClient m => PrepareStrategy -> m a -> m a withPrepareStrategy s = localState (set (context.settings.prepStrategy) s) --- | Send a CQL 'Request' to the server and return a 'Response'.+-- | Send a 'Request' to the server and return a 'Response'. -- -- This function will first ask the clients load-balancing 'Policy' for -- some host and use its connection pool to acquire a connection for@@ -215,53 +230,12 @@ -- If all available hosts are busy (i.e. their connection pools are fully -- utilised), the function will block until a connection becomes available -- or the maximum wait-queue length has been reached.+--+-- The request is retried according to the configured 'RetrySettings'. request :: (MonadClient m, Tuple a, Tuple b) => Request k a b -> m (Response k a b) request a = liftClient $ do n <- liftIO . hostCount =<< view policy- snd <$> mkRequest (requestN n) a--mkRequest :: (Tuple a, Tuple b)- => (Request k a b -> ClientState -> Client (Maybe (Host, Response k a b)))- -> Request k a b- -> Client (Host, Response k a b)-mkRequest fn a = do- s <- ask- recovering (s^.context.settings.retrySettings.retryPolicy) recoverFrom $ \i -> do- r <- if rsIterNumber i == 0- then fn a s- else fn (newRequest s) (adjust s)- maybe (throwM HostsBusy) return r- where- adjust s =- let x = s^.context.settings.retrySettings.sendTimeoutChange- y = s^.context.settings.retrySettings.recvTimeoutChange- in over (context.settings.connSettings.sendTimeout) (+ x)- . over (context.settings.connSettings.responseTimeout) (+ y)- $ s-- newRequest s =- case s^.context.settings.retrySettings.reducedConsistency of- Nothing -> a- Just c ->- case a of- RqQuery (Query q p) -> RqQuery (Query q p { consistency = c })- RqExecute (Execute q p) -> RqExecute (Execute q p { consistency = c })- RqBatch b -> RqBatch b { batchConsistency = c }- _ -> a-- recoverFrom =- [ const $ Handler $ \e -> case e of- ReadTimeout {} -> return True- WriteTimeout {} -> return True- Overloaded {} -> return True- Unavailable {} -> return True- ServerError {} -> return True- _ -> return False- , const $ Handler $ \(_ :: ConnectionError) -> return True- , const $ Handler $ \(_ :: IOException) -> return True- , const $ Handler $ \(_ :: HostError) -> return True- , const $ Handler $ \(_ :: SomeSSLException) -> return True- ]+ snd <$> withRetries (requestN n) a -- | Invoke 'request1' up to @n@ times with different hosts if no -- connection is available. May return 'Nothing' if no connection@@ -282,7 +256,12 @@ request1 h a s = do p <- Map.lookup h <$> readTVarIO' (s^.hostmap) case p of- Just x -> with x transaction `catches` handlers+ Just x -> do+ result <- with x transaction `catches` handlers+ for_ result $ \(_, r) ->+ for_ (Cql.warnings r) $ \w ->+ warn $ msg (val "server warning: " +++ w)+ return result Nothing -> do err $ msg (val "no pool for host " +++ h) p' <- mkPool (s^.context) (h^.hostAddr)@@ -307,9 +286,9 @@ executeWithPrepare :: (Tuple b, Tuple a) => Maybe Host -> Request k a b -> Client (Response k a b) executeWithPrepare h q = do f <- selectAction h- r <- mkRequest f q+ r <- withRetries f q case snd r of- RsError _ (Unprepared _ i) -> do+ RsError _ _ (Unprepared _ i) -> do pq <- preparedQueries qs <- atomically' (PQ.lookupQueryString (QueryId i) pq) case qs of@@ -317,8 +296,8 @@ Just s -> do (g, _) <- prepare (Just LazyPrepare) (s :: Raw QueryString) executeWithPrepare (Just g) q- RsError _ e -> throwM e- x -> return x+ RsError _ _ e -> throwM e+ x -> return x where selectAction Nothing = view policy >>= liftIO . hostCount >>= return . requestN selectAction (Just x) = return (request1 x)@@ -330,24 +309,21 @@ prepare (Just LazyPrepare) qs = do s <- ask n <- liftIO $ hostCount (s^.policy)- (h, r) <- mkRequest (requestN n) (RqPrepare (Prepare qs))- case r of- RsResult _ (PreparedResult i _ _) -> return (h, i)- RsError _ e -> throwM e- _ -> throwM UnexpectedResponse+ (h, r) <- withRetries (requestN n) (RqPrepare (Prepare qs))+ (h,) <$> getPreparedQueryId r+ prepare (Just EagerPrepare) qs = view policy >>= liftIO . current >>= mapM (action (RqPrepare (Prepare qs))) >>= first where action rq h = do- r <- mkRequest (request1 h) rq- case snd r of- RsResult _ (PreparedResult i _ _) -> return (h, i)- RsError _ e -> throwM e- _ -> throwM UnexpectedResponse+ r <- snd <$> withRetries (request1 h) rq+ (h,) <$> getPreparedQueryId r+ first (x:_) = return x first [] = throwM NoHostAvailable+ prepare Nothing qs = do ps <- view (context.settings.prepStrategy) prepare (Just ps) qs@@ -407,7 +383,7 @@ <*> newTVarIO Map.empty <*> Jobs.new e^.sigMonit |-> onEvent p- runClient x (initialise c)+ runClient x (initialise c) `onException` liftIO (C.close c) return x where mkConnection t h = do@@ -524,6 +500,56 @@ ----------------------------------------------------------------------------- -- Exception handling +withRetries+ :: (Tuple a, Tuple b)+ => (Request k a b -> ClientState -> Client (Maybe (Host, Response k a b)))+ -> Request k a b+ -> Client (Host, Response k a b)+withRetries fn a = do+ s <- ask+ let p = s^.context.settings.retrySettings.retryPolicy+ recovering p recoverFrom $ \i -> do+ r <- if rsIterNumber i == 0+ then fn a s+ else fn (newRequest s) (adjust s)+ case r of+ Nothing -> throwM HostsBusy+ Just hr -> case snd hr of+ RsError _ _ e -> applyPolicy p i+ >>= maybe (return hr) (const (throwM e))+ _ -> return hr+ where+ adjust s =+ let x = s^.context.settings.retrySettings.sendTimeoutChange+ y = s^.context.settings.retrySettings.recvTimeoutChange+ in over (context.settings.connSettings.sendTimeout) (+ x)+ . over (context.settings.connSettings.responseTimeout) (+ y)+ $ s++ newRequest s =+ case s^.context.settings.retrySettings.reducedConsistency of+ Nothing -> a+ Just c ->+ case a of+ RqQuery (Query q p) -> RqQuery (Query q p { consistency = c })+ RqExecute (Execute q p) -> RqExecute (Execute q p { consistency = c })+ RqBatch b -> RqBatch b { batchConsistency = c }+ _ -> a++ recoverFrom =+ [ const $ Handler $ \e -> return $ case e of+ ReadTimeout {} -> True+ WriteTimeout {} -> True+ Overloaded {} -> True+ Unavailable {} -> True+ ServerError {} -> True+ _ -> False+ , const $ Handler $ \(_ :: ConnectionError) -> return True+ , const $ Handler $ \(_ :: IOException) -> return True+ , const $ Handler $ \(_ :: HostError) -> return True+ , const $ Handler $ \(_ :: SomeSSLException) -> return True+ ]+ onConnectionError :: Exception e => Host -> e -> Client () onConnectionError h exc = do warn $ "exception" .= show exc@@ -572,33 +598,34 @@ ctl <- view control let s = ctx^.settings c <- C.connect (s^.connSettings) (ctx^.timeouts) (s^.protoVersion) (ctx^.logger) a- initialise c+ initialise c `onException` liftIO (C.close c) atomically' $ writeTVar ctl (Control Connected c) info $ msg (val "new control connection: " +++ c) +-----------------------------------------------------------------------------+-- Event handling+ onCqlEvent :: Event -> Client () onCqlEvent x = do info $ "client.event" .= show x pol <- view policy prt <- view (context.settings.portnumber) case x of- StatusEvent Down sa -> do- liftIO $ onEvent pol $ HostDown (InetAddr $ mapPort prt sa)- TopologyEvent RemovedNode sa -> do- let a = InetAddr $ mapPort prt sa+ StatusEvent Down (mapAddr prt -> a) ->+ liftIO $ onEvent pol (HostDown a)+ TopologyEvent RemovedNode (mapAddr prt -> a) -> do hmap <- view hostmap atomically' $ modifyTVar' hmap (Map.filterWithKey (\h _ -> h^.hostAddr /= a)) liftIO $ onEvent pol $ HostGone a- StatusEvent Up sa -> do+ StatusEvent Up (mapAddr prt -> a) -> do s <- ask- startMonitor s $ (InetAddr $ mapPort prt sa)- TopologyEvent NewNode sa -> do+ startMonitor s a+ TopologyEvent NewNode (mapAddr prt -> a) -> do s <- ask let ctx = s^.context let hmap = s^.hostmap ctrl <- readTVarIO' (s^.control)- let a = InetAddr $ mapPort prt sa let c = ctrl^.connection h <- fromMaybe (Host a "" "") . find ((a == ) . view hostAddr) <$> discoverPeers' ctx c okay <- liftIO $ acceptable pol h@@ -609,11 +636,11 @@ Jobs.add (s^.jobs) a False $ runClient s (prepareAllQueries h) SchemaEvent _ -> return () where- mapPort i (SockAddrInet _ a) = SockAddrInet i a- mapPort i (SockAddrInet6 _ f a b) = SockAddrInet6 i f a b- mapPort _ unix = unix+ mapAddr i (SockAddrInet _ a) = InetAddr (SockAddrInet i a)+ mapAddr i (SockAddrInet6 _ f a b) = InetAddr (SockAddrInet6 i f a b)+ mapAddr _ unix = InetAddr unix - discoverPeers' ctx c = discoverPeers ctx c `catchAll` (const $ return [])+ discoverPeers' ctx c = discoverPeers ctx c `catchAll` const (return []) startMonitor s a = do hmp <- readTVarIO' (s^.hostmap)@@ -629,11 +656,25 @@ qs <- atomically' $ PQ.queryStrings pq for_ qs $ \q -> let qry = QueryString q :: Raw QueryString in- mkRequest (request1 h) (RqPrepare (Prepare qry))+ withRetries (request1 h) (RqPrepare (Prepare qry)) ----------------------------------------------------------------------------- -- Utilities +getResult :: MonadThrow m => Response k a b -> m (Result k a b)+getResult (RsResult _ _ r) = return r+getResult (RsError _ _ e) = throwM e+getResult r = throwM $ UnexpectedResponse r+{-# INLINE getResult #-}++getPreparedQueryId :: MonadThrow m => Response k a b -> m (QueryId k a b)+getPreparedQueryId r = do+ rs <- getResult r+ case rs of+ PreparedResult i _ _ -> return i+ _ -> throwM $ UnexpectedResponse r+{-# INLINE getPreparedQueryId #-}+ peer2Host :: PortNumber -> Peer -> Host peer2Host i p = Host (ip2inet i (peerRPC p)) (peerDC p) (peerRack p) @@ -646,7 +687,7 @@ tryAll :: NonEmpty a -> (a -> IO b) -> IO b tryAll (a :| []) f = f a-tryAll (a :| aa) f = f a `catchAll` (const $ tryAll (NE.fromList aa) f)+tryAll (a :| aa) f = f a `catchAll` const (tryAll (NE.fromList aa) f) atomically' :: STM a -> Client a atomically' = liftIO . atomically
src/Database/CQL/IO/Cluster/Policies.hs view
@@ -37,7 +37,7 @@ -- through this function. , select :: IO (Maybe Host) -- ^ Host selection. The driver will ask for a host to use in a query- -- through this function. A policy which has no available nodes my+ -- through this function. A policy which has no available nodes may -- return Nothing. , current :: IO [Host] -- ^ Return all currently alive hosts.
src/Database/CQL/IO/Connection.hs view
@@ -2,12 +2,15 @@ -- License, v. 2.0. If a copy of the MPL was not distributed with this -- file, You can obtain one at http://mozilla.org/MPL/2.0/. +{-# LANGUAGE LambdaCase #-} {-# LANGUAGE OverloadedStrings #-} {-# LANGUAGE ScopedTypeVariables #-} {-# LANGUAGE TemplateHaskell #-}+{-# LANGUAGE ViewPatterns #-} module Database.CQL.IO.Connection ( Connection+ , ConnId , resolve , ping , connect@@ -20,17 +23,6 @@ , address , protocol , eventSig-- , ConnectionSettings- , defSettings- , connectTimeout- , sendTimeout- , responseTimeout- , maxStreams- , compression- , defKeyspace- , maxRecvBuffer- , tlsContext ) where import Control.Applicative@@ -52,6 +44,7 @@ import Data.Vector (Vector, (!)) import Database.CQL.Protocol import Database.CQL.IO.Connection.Socket (Socket)+import Database.CQL.IO.Connection.Settings import Database.CQL.IO.Hexdump import Database.CQL.IO.Protocol import Database.CQL.IO.Signal hiding (connect)@@ -60,7 +53,6 @@ import Database.CQL.IO.Tickets (Pool, toInt, markAvailable) import Database.CQL.IO.Timeouts (TimeoutManager, withTimeout) import Network.Socket hiding (Socket, close, connect, send)-import OpenSSL.Session (SSLContext) import System.IO (nativeNewline, Newline (..)) import System.Logger hiding (Settings, close, defSettings, settings) import System.Timeout@@ -68,23 +60,13 @@ import qualified Data.ByteString.Lazy as L import qualified Data.ByteString.Lazy.Char8 as Char8+import qualified Data.HashMap.Strict as HashMap import qualified Data.Vector as Vector import qualified Database.CQL.IO.Connection.Socket as Socket import qualified Database.CQL.IO.Sync as Sync import qualified Database.CQL.IO.Tickets as Tickets import qualified Network.Socket as S -data ConnectionSettings = ConnectionSettings- { _connectTimeout :: !Milliseconds- , _sendTimeout :: !Milliseconds- , _responseTimeout :: !Milliseconds- , _maxStreams :: !Int- , _compression :: !Compression- , _defKeyspace :: !(Maybe Keyspace)- , _maxRecvBuffer :: !Int- , _tlsContext :: !(Maybe SSLContext)- }- type Streams = Vector (Sync (Header, ByteString)) data Connection = Connection@@ -100,10 +82,9 @@ , _tickets :: !Pool , _logger :: !Logger , _eventSig :: !(Signal Event)- , _ident :: !Unique+ , _ident :: !ConnId } -makeLenses ''ConnectionSettings makeLenses ''Connection instance Eq Connection where@@ -115,17 +96,6 @@ instance ToBytes Connection where bytes c = bytes (c^.address) +++ val "#" +++ c^.sock -defSettings :: ConnectionSettings-defSettings =- ConnectionSettings 5000 -- connect timeout- 3000 -- send timeout- 10000 -- response timeout- 128 -- max streams per connection- noCompression -- compression- Nothing -- keyspace- 16384 -- receive buffer size- Nothing -- no tls by default- resolve :: String -> PortNumber -> IO [InetAddr] resolve host port = map (InetAddr . addrAddress) <$> getAddrInfo (Just hints) (Just host) (Just (show port))@@ -141,7 +111,7 @@ sta <- newTVarIO True sig <- signal rdr <- async (readLoop v g t tck a s syn sig sta lck)- Connection t a m v s sta syn lck rdr tck g sig <$> newUnique+ Connection t a m v s sta syn lck rdr tck g sig . ConnId <$> newUnique validateSettings c `onException` close c return c @@ -169,9 +139,9 @@ case fromStreamId $ streamId (fst x) of -1 -> case parse (set^.compression) x :: Raw Response of- RsError _ e -> throwM e- RsEvent _ e -> emit s e- r -> throwM (UnexpectedResponse' r)+ RsError _ _ e -> throwM e+ RsEvent _ _ e -> emit s e+ r -> throwM (UnexpectedResponse' r) sid -> do ok <- Sync.put x (syn ! sid) unless ok $@@ -216,13 +186,13 @@ let e = TimeoutRead (show c ++ ":" ++ show i) tid <- myThreadId withTimeout (c^.tmanager) (c^.settings.responseTimeout) (throwTo tid e) $ do- x <- Sync.get (view streams c ! i) `onException` (Sync.kill e) (view streams c ! i)+ x <- Sync.get (view streams c ! i) `onException` Sync.kill e (view streams c ! i) markAvailable (c^.tickets) i return x readSocket :: Version -> Logger -> InetAddr -> Socket -> Int -> IO (Header, ByteString) readSocket v g i s n = do- b <- Socket.recv n i s (if v == V3 then 9 else 8)+ b <- Socket.recv n i s 9 h <- case header v b of Left e -> throwM $ InternalError ("response header reading: " ++ e) Right h -> return h@@ -246,15 +216,62 @@ let req = RqStartup (Startup Cqlv300 (algorithm cmp)) let enc = serialise (c^.protocol) cmp (req :: Raw Request) res <- request c enc- (parse cmp res :: Raw Response) `seq` return ()+ case parse cmp res :: Raw Response of+ RsReady _ _ Ready -> checkAuth c+ RsAuthenticate _ _ auth -> authenticate c auth+ RsError _ _ e -> throwM e+ other -> throwM $ UnexpectedResponse' other +checkAuth :: Connection -> IO ()+checkAuth c = unless (null (c^.settings.authenticators)) $+ warn (_logger c) $ msg $ val+ "Authentication configured but none required by server."++authenticate :: (MonadIO m, MonadThrow m) => Connection -> Authenticate -> m ()+authenticate c (Authenticate (AuthMechanism -> m)) =+ case HashMap.lookup m (c^.settings.authenticators) of+ Nothing -> throwM $ AuthenticationRequired m+ Just Authenticator {+ authOnRequest = onR+ , authOnChallenge = onC+ , authOnSuccess = onS+ } -> liftIO $ do+ (rs, s) <- onR context+ case onC of+ Just f -> loop f onS (rs, s)+ Nothing -> authResponse c rs >>= either+ (throwM . UnexpectedAuthenticationChallenge m)+ (onS s)+ where+ context = AuthContext (c^.ident) (c^.address)++ loop onC onS (rs, s) =+ authResponse c rs >>= either+ (onC s >=> loop onC onS)+ (onS s)++authResponse :: MonadIO m+ => Connection+ -> AuthResponse+ -> m (Either AuthChallenge AuthSuccess)+authResponse c resp = liftIO $ do+ let cmp = c^.settings.compression+ let req = RqAuthResp resp+ let enc = serialise (c^.protocol) cmp (req :: Raw Request)+ res <- request c enc+ case parse cmp res :: Raw Response of+ RsAuthSuccess _ _ success -> return $ Right success+ RsAuthChallenge _ _ challenge -> return $ Left challenge+ RsError _ _ e -> throwM e+ other -> throwM $ UnexpectedResponse' other+ register :: MonadIO m => Connection -> [EventType] -> EventHandler -> m () register c e f = liftIO $ do let req = RqRegister (Register e) :: Raw Request let enc = serialise (c^.protocol) (c^.settings.compression) req res <- request c enc case parse (c^.settings.compression) res :: Raw Response of- RsReady _ Ready -> c^.eventSig |-> f+ RsReady _ _ Ready -> c^.eventSig |-> f other -> throwM (UnexpectedResponse' other) validateSettings :: MonadIO m => Connection -> m ()@@ -269,18 +286,18 @@ let options = RqOptions Options :: Raw Request res <- request c (serialise (c^.protocol) noCompression options) case parse noCompression res :: Raw Response of- RsSupported _ x -> return x+ RsSupported _ _ x -> return x other -> throwM (UnexpectedResponse' other) useKeyspace :: MonadIO m => Connection -> Keyspace -> m () useKeyspace c ks = liftIO $ do let cmp = c^.settings.compression- params = QueryParams One False () Nothing Nothing Nothing+ params = QueryParams One False () Nothing Nothing Nothing Nothing kspace = quoted (fromStrict $ unKeyspace ks) req = RqQuery (Query (QueryString $ "use " <> kspace) params) res <- request c (serialise (c^.protocol) cmp req) case parse cmp res :: Raw Response of- RsResult _ (SetKeyspaceResult _) -> return ()+ RsResult _ _ (SetKeyspaceResult _) -> return () other -> throwM (UnexpectedResponse' other) query :: forall k a b m. (Tuple a, Tuple b, Show b, MonadIO m)@@ -294,10 +311,10 @@ let enc = serialise (c^.protocol) (c^.settings.compression) req res <- request c enc case parse (c^.settings.compression) res :: Response k a b of- RsResult _ (RowsResult _ b) -> return b+ RsResult _ _ (RowsResult _ b) -> return b other -> throwM (UnexpectedResponse' other) where- params = QueryParams cons False p Nothing Nothing Nothing+ params = QueryParams cons False p Nothing Nothing Nothing Nothing -- logging helpers:
+ src/Database/CQL/IO/Connection/Settings.hs view
@@ -0,0 +1,140 @@+-- This Source Code Form is subject to the terms of the Mozilla Public+-- License, v. 2.0. If a copy of the MPL was not distributed with this+-- file, You can obtain one at http://mozilla.org/MPL/2.0/.++{-# LANGUAGE ExistentialQuantification #-}+{-# LANGUAGE GeneralizedNewtypeDeriving #-}+{-# LANGUAGE OverloadedStrings #-}+{-# LANGUAGE ScopedTypeVariables #-}+{-# LANGUAGE TemplateHaskell #-}++module Database.CQL.IO.Connection.Settings+ ( ConnectionSettings+ , defSettings+ , connectTimeout+ , sendTimeout+ , responseTimeout+ , maxStreams+ , compression+ , defKeyspace+ , maxRecvBuffer+ , tlsContext+ , authenticators++ -- * Authentication+ , AuthMechanism (..)+ , Authenticator (..)+ , AuthContext (..)+ , authConnId+ , authHost+ , passwordAuthenticator+ , AuthUser (..)+ , AuthPass (..)+ ) where++import Control.Lens (makeLenses)+import Control.Monad+import Data.HashMap.Strict (HashMap)+import Data.Int+import Database.CQL.Protocol+import Database.CQL.IO.Types+import OpenSSL.Session (SSLContext)+import Prelude++import qualified Data.ByteString.Lazy.Char8 as Char8+import qualified Data.HashMap.Strict as HashMap+import qualified Data.Text.Lazy as Lazy+import qualified Data.Text.Lazy.Encoding as Lazy++data ConnectionSettings = ConnectionSettings+ { _connectTimeout :: !Milliseconds+ , _sendTimeout :: !Milliseconds+ , _responseTimeout :: !Milliseconds+ , _maxStreams :: !Int+ , _compression :: !Compression+ , _defKeyspace :: !(Maybe Keyspace)+ , _maxRecvBuffer :: !Int+ , _tlsContext :: !(Maybe SSLContext)+ , _authenticators :: !(HashMap AuthMechanism Authenticator)+ }++-- | Context information given to 'Authenticator's when+-- the server requests authentication on a connection.+-- See 'authOnRequest'.+data AuthContext = AuthContext+ { _authConnId :: !ConnId+ , _authHost :: !InetAddr+ }++-- | A client authentication handler.+--+-- The fields of an 'Authenticator' must implement the client-side+-- of an (SASL) authentication mechanism as follows:+--+-- * When a Cassandra server requests authentication on a new connection,+-- 'authOnRequest' is called with the 'AuthContext' of the+-- connection.+--+-- * If additional challenges are posed by the server,+-- 'authOnChallenge' is called, if available, otherwise an+-- 'AuthenticationError' is thrown, i.e. every challenge must be+-- answered.+--+-- * Upon successful authentication 'authOnSuccess' is called.+--+-- The existential type @s@ is chosen by an implementation and can+-- be used to thread arbitrary state through the sequence of callback+-- invocations during an authentication exchange.+--+-- See also:+-- <https://tools.ietf.org/html/rfc4422 RFC4422>+-- <https://docs.datastax.com/en/cassandra/latest/cassandra/configuration/secureInternalAuthenticationTOC.html Authentication>+data Authenticator = forall s. Authenticator+ { authMechanism :: !AuthMechanism+ -- ^ The (unique) name of the (SASL) mechanism that the callbacks+ -- implement.+ , authOnRequest :: AuthContext -> IO (AuthResponse, s)+ -- ^ Callback for initiating an authentication exchange.+ , authOnChallenge :: Maybe (s -> AuthChallenge -> IO (AuthResponse, s))+ -- ^ Optional callback for additional challenges posed by the server.+ -- If the authentication mechanism does not require additional+ -- challenges, it should be set to 'Nothing'. Otherwise every+ -- challenge must be answered with a response.+ , authOnSuccess :: s -> AuthSuccess -> IO ()+ -- ^ Callback for successful completion of an authentication exchange.+ }++makeLenses ''AuthContext+makeLenses ''ConnectionSettings++newtype AuthUser = AuthUser Lazy.Text+newtype AuthPass = AuthPass Lazy.Text++-- | A password authentication handler for use with Cassandra's+-- @PasswordAuthenticator@.+--+-- See: <https://docs.datastax.com/en/cassandra/latest/cassandra/configuration/secureConfigNativeAuth.html Configuring Authentication>+passwordAuthenticator :: AuthUser -> AuthPass -> Authenticator+passwordAuthenticator (AuthUser u) (AuthPass p) = Authenticator+ { authMechanism = "org.apache.cassandra.auth.PasswordAuthenticator"+ , authOnChallenge = Nothing+ , authOnSuccess = \() _ -> return ()+ , authOnRequest = \_ctx ->+ let user = Lazy.encodeUtf8 u+ pass = Lazy.encodeUtf8 p+ resp = AuthResponse (Char8.concat ["\0", user, "\0", pass])+ in return (resp, ())+ }++defSettings :: ConnectionSettings+defSettings =+ ConnectionSettings 5000 -- connect timeout+ 3000 -- send timeout+ 10000 -- response timeout+ 128 -- max streams per connection+ noCompression -- compression+ Nothing -- keyspace+ 16384 -- receive buffer size+ Nothing -- no tls by default+ HashMap.empty -- no authentication+
src/Database/CQL/IO/Connection/Socket.hs view
@@ -66,7 +66,7 @@ familyOf (SockAddrInet _ _) = AF_INET familyOf (SockAddrInet6 _ _ _ _) = AF_INET6 familyOf (SockAddrUnix _) = AF_UNIX-#if MIN_VERSION_network(2,6,1)+#if MIN_VERSION_network(2,6,1) && !MIN_VERSION_network(3,0,0) familyOf (SockAddrCan _ ) = AF_CAN #endif
src/Database/CQL/IO/Protocol.hs view
@@ -8,6 +8,7 @@ import Control.Exception (throw) import Data.ByteString.Lazy (ByteString)+import Data.Maybe (fromMaybe) import Data.Monoid ((<>)) import Database.CQL.Protocol import Database.CQL.IO.Types@@ -27,7 +28,12 @@ OcOptions -> noCompression _ -> f s = mkStreamId i- in either (throw $ InternalError "request creation") id (pack v c False s r)+ in either (throw $ InternalError "request creation") id (pack v c (isTracing r) s r)+ where+ isTracing :: Request k a b -> Bool+ isTracing (RqQuery (Query _ p)) = fromMaybe False $ enableTracing p+ isTracing (RqExecute (Execute _ p)) = fromMaybe False $ enableTracing p+ isTracing _ = False quoted :: LT.Text -> LT.Text quoted s = "\"" <> LT.replace "\"" "\"\"" s <> "\""
src/Database/CQL/IO/Settings.hs view
@@ -15,15 +15,16 @@ import Data.Time import Data.Word import Database.CQL.Protocol-import Database.CQL.IO.Connection import Database.CQL.IO.Cluster.Policies (Policy, random)-import Database.CQL.IO.Connection as C+import Database.CQL.IO.Connection.Settings as C import Database.CQL.IO.Pool as P import Database.CQL.IO.Types (Milliseconds (..)) import Network.Socket (PortNumber (..)) import OpenSSL.Session (SSLContext) import Prelude +import qualified Data.HashMap.Strict as HashMap+ data PrepareStrategy = EagerPrepare -- ^ cluster-wide preparation | LazyPrepare -- ^ on-demand per node preparation@@ -56,7 +57,7 @@ -- -- * load-balancing policy is 'random' ----- * binary protocol version is 3 (supported by Cassandra >= 2.1.0)+-- * binary protocol version is 3 -- -- * connection idle timeout is 60s --@@ -82,7 +83,7 @@ C.defSettings noRetry V3- (fromInteger 9042)+ 9042 ("localhost" :| []) random LazyPrepare@@ -153,10 +154,9 @@ -- binary protocol at most 128 streams can be used. Version 3 supports up -- to 32768 streams. setMaxStreams :: Int -> Settings -> Settings-setMaxStreams v s = case s^.protoVersion of- V2 | v < 1 || v > 128 -> error "cql-io settings: max. streams must be within [1, 128]"- V3 | v < 1 || v > 32768 -> error "cql-io settings: max. streams must be within [1, 32768]"- _ -> set (connSettings.maxStreams) v s+setMaxStreams v s+ | v < 1 || v > 32768 = error "cql-io settings: max. streams must be within [1, 32768]"+ | otherwise = set (connSettings.maxStreams) v s -- | Set the connect timeout of a connection. setConnectTimeout :: NominalDiffTime -> Settings -> Settings@@ -194,6 +194,17 @@ -- This will make client server queries use TLS. setSSLContext :: SSLContext -> Settings -> Settings setSSLContext v = set (connSettings.tlsContext) (Just v)++-- | Set the supported authentication mechanisms.+--+-- When a Cassandra server requests authentication on a connection,+-- it specifies the requested 'AuthMechanism'. The client 'Authenticator'+-- is chosen based that name. If no authenticator with a matching+-- name is configured, an 'AuthenticationError' is thrown.+setAuthentication :: [C.Authenticator] -> Settings -> Settings+setAuthentication = set (connSettings.authenticators)+ . HashMap.fromList+ . map (\a -> (authMechanism a, a)) ----------------------------------------------------------------------------- -- Retry Settings
src/Database/CQL/IO/Types.hs view
@@ -12,13 +12,18 @@ module Database.CQL.IO.Types where import Control.Monad.Catch+import Data.Hashable import Data.IP-import Data.Text.Lazy (Text)+import Data.String+import Data.Text (Text) import Data.Typeable-import Database.CQL.Protocol (Event, Response, CompressionAlgorithm)+import Data.Unique+import Database.CQL.Protocol import Network.Socket (SockAddr (..), PortNumber) import System.Logger.Message +import qualified Data.Text.Lazy as Lazy+ type EventHandler = Event -> IO () newtype Milliseconds = Ms { ms :: Int } deriving (Eq, Show, Num)@@ -26,6 +31,14 @@ type Raw a = a () () () -----------------------------------------------------------------------------+-- ConnId++newtype ConnId = ConnId Unique deriving (Eq, Ord)++instance Hashable ConnId where+ hashWithSalt _ (ConnId u) = hashUnique u++----------------------------------------------------------------------------- -- InetAddr newtype InetAddr = InetAddr { sockAddr :: SockAddr } deriving (Eq, Ord)@@ -38,7 +51,7 @@ let i = fromIntegral p :: Int in shows (fromHostAddress6 a) . showString ":" . shows i $ "" show (InetAddr (SockAddrUnix unix)) = unix-#if MIN_VERSION_network(2,6,1)+#if MIN_VERSION_network(2,6,1) && !MIN_VERSION_network(3,0,0) show (InetAddr (SockAddrCan int32)) = show int32 #endif @@ -50,7 +63,7 @@ let i = fromIntegral p :: Int in show (fromHostAddress6 a) +++ val ":" +++ i bytes (InetAddr (SockAddrUnix unix)) = bytes unix-#if MIN_VERSION_network(2,6,1)+#if MIN_VERSION_network(2,6,1) && !MIN_VERSION_network(3,0,0) bytes (InetAddr (SockAddrCan int32)) = bytes int32 #endif @@ -130,21 +143,44 @@ ----------------------------------------------------------------------------- -- UnexpectedResponse +-- | Placeholder for parts of a 'Response' that are not 'Show'able.+data NoShow = NoShow deriving Show+ data UnexpectedResponse where- UnexpectedResponse :: UnexpectedResponse+ UnexpectedResponse :: !(Response k a b) -> UnexpectedResponse UnexpectedResponse' :: Show b => !(Response k a b) -> UnexpectedResponse deriving instance Typeable UnexpectedResponse instance Exception UnexpectedResponse instance Show UnexpectedResponse where- show UnexpectedResponse = "cql-io: unexpected response"- show (UnexpectedResponse' r) = "cql-io: unexpected response: " ++ show r+ show x = showString "cql-io: unexpected response: "+ . case x of+ UnexpectedResponse r -> shows (f r)+ UnexpectedResponse' r -> shows r+ $ ""+ where+ f :: Response k a b -> Response k a NoShow+ f (RsError a b c) = RsError a b c+ f (RsReady a b c) = RsReady a b c+ f (RsAuthenticate a b c) = RsAuthenticate a b c+ f (RsAuthChallenge a b c) = RsAuthChallenge a b c+ f (RsAuthSuccess a b c) = RsAuthSuccess a b c+ f (RsSupported a b c) = RsSupported a b c+ f (RsResult a b c) = RsResult a b (g c)+ f (RsEvent a b c) = RsEvent a b c + g :: Result k a b -> Result k a NoShow+ g VoidResult = VoidResult+ g (RowsResult a b ) = RowsResult a (map (const NoShow) b)+ g (SetKeyspaceResult a ) = SetKeyspaceResult a+ g (SchemaChangeResult a ) = SchemaChangeResult a+ g (PreparedResult (QueryId a) b c) = PreparedResult (QueryId a) b c+ ----------------------------------------------------------------------------- -- HashCollision -data HashCollision = HashCollision !Text !Text+data HashCollision = HashCollision !Lazy.Text !Lazy.Text deriving Typeable instance Exception HashCollision@@ -155,6 +191,37 @@ . showString " " . shows b $ ""++-----------------------------------------------------------------------------+-- Authentication++-- | The (unique) name of a SASL authentication mechanism.+--+-- In the case of Cassandra, this is currently always the fully-qualified+-- Java class name of the configured server-side @IAuthenticator@+-- implementation.+newtype AuthMechanism = AuthMechanism Text+ deriving (Eq, Ord, Show, IsString, Hashable)++data AuthenticationError+ = AuthenticationRequired !AuthMechanism+ | UnexpectedAuthenticationChallenge !AuthMechanism !AuthChallenge++instance Exception AuthenticationError++instance Show AuthenticationError where+ show (AuthenticationRequired a)+ = showString "cql-io: authentication required: "+ . shows a+ $ ""++ show (UnexpectedAuthenticationChallenge n c)+ = showString "cql-io: unexpected authentication challenge: '"+ . shows c+ . showString "' using mechanism '"+ . shows n+ . showString "'"+ $ "" ignore :: IO () -> IO () ignore a = catchAll a (const $ return ())
+ test/Main.hs view
@@ -0,0 +1,364 @@+{-# LANGUAGE OverloadedStrings #-}+{-# LANGUAGE QuasiQuotes #-}+{-# LANGUAGE TypeFamilies #-}++module Main (main) where++import Control.Monad+import Control.Monad.Identity+import Control.Monad.IO.Class+import Data.Decimal+import Data.Int+import Data.IP+import Data.List (sort)+import Data.Maybe+import Data.Text (Text)+import Data.Time+import Data.UUID+import Database.CQL.Protocol+import Database.CQL.IO as Client+import System.Environment+import Test.Tasty+import Test.Tasty.HUnit+import Text.RawString.QQ++import qualified Data.Set as Set+import qualified System.Logger as Log++-----------------------------------------------------------------------------+-- Test Setup++type TestHost = String++main :: IO ()+main = do+ h <- fromMaybe "localhost" <$> lookupEnv "CASSANDRA_HOST"+ g <- Log.new (Log.setLogLevel Log.Warn Log.defSettings)+ initSchema g h+ defaultMain . testGroup "cql-io" =<<+ forM versions (\v -> do+ c <- Client.init g (settings h v)+ return $ testGroup (show v) (tests c))++versions :: [Version]+versions = [V3, V4]++settings :: TestHost -> Version -> Settings+settings h v = setContacts h []+ . setProtocolVersion v+ $ defSettings++initSchema :: Log.Logger -> TestHost -> IO ()+initSchema g h = do+ c <- Client.init g (settings h V4)+ runClient c $ do+ dropKeyspace+ createKeyspace+ createTables+ shutdown c++test :: ClientState -> String -> Client () -> TestTree+test c name runTest =+ testCase name $+ runClient c $ do+ truncateTables+ runTest++-----------------------------------------------------------------------------+-- Test Schema++-- Columns of cqltest.table1+type Ty1 =+ ( Int64+ , Ascii+ , Blob+ , Bool+ , Decimal+ , Double+ , Float+ , Int32+ , UTCTime+ , UUID+ , Text+ , Integer+ , TimeUuid+ , IP+ )++-- Columns of cqltest.table2+type Ty2 =+ ( Int64+ , [Int32]+ , Set Ascii+ , Map Ascii Int32+ , Maybe Int32+ , (Bool, Ascii, Int32)+ , Map Int32 (Map Int32 (Set Ascii))+ )++createKeyspace :: Client ()+createKeyspace = void $ schema cql (params ())+ where+ cql :: QueryString S () ()+ cql = [r| create keyspace if not exists cqltest+ with replication = {+ 'class': 'SimpleStrategy',+ 'replication_factor': '1'+ } |]++dropKeyspace :: Client ()+dropKeyspace = void $ schema cql (params ())+ where+ cql :: QueryString S () ()+ cql = "drop keyspace if exists cqltest"++createTables :: Client ()+createTables = forM_ [cql1, cql2, cql3] $ \q ->+ void $ schema q (params ())+ where+ cql1, cql2, cql3 :: QueryString S () ()+ cql1 = [r|+ create table if not exists cqltest.test1+ ( a bigint+ , b ascii+ , c blob+ , d boolean+ , e decimal+ , f double+ , g float+ , h int+ , i timestamp+ , j uuid+ , k varchar+ , l varint+ , m timeuuid+ , n inet+ , primary key (a)+ ) |]++ cql2 = [r|+ create table if not exists cqltest.test2+ ( a bigint+ , b list<int>+ , c set<ascii>+ , d map<ascii,int>+ , e int+ , f tuple<boolean,ascii,int>+ , g map<int,frozen<map<int,set<ascii>>>>+ , primary key (a)+ ) |]++ cql3 = [r|+ create table if not exists cqltest.counters+ ( a bigint+ , n counter+ , primary key (a)+ ) |]++truncateTables :: Client ()+truncateTables = forM_ [cql1, cql2, cql3] $ \q ->+ void $ schema q (params ())+ where+ cql1, cql2, cql3 :: QueryString S () ()+ cql1 = "truncate table cqltest.test1"+ cql2 = "truncate table cqltest.test2"+ cql3 = "truncate table cqltest.counters"++-----------------------------------------------------------------------------+-- Tests++tests :: ClientState -> [TestTree]+tests c =+ [ test c "write-read" testWriteRead+ , test c "write-read-ttl" testWriteReadTtl+ , test c "trans" testTrans+ , test c "paging" testPaging+ , test c "batch" testBatch+ , test c "batch-counter" testBatchCounter+ ]++testWriteRead :: Client ()+testWriteRead = do+ t <- liftIO $ fmap (\x -> x { utctDayTime = secondsToDiffTime 3600 }) getCurrentTime+ let a = ( 4835637638+ , "hello world"+ , Blob "blooooooooooooooooooooooob"+ , False+ , 1.2342342342423423423423423442+ , 433243.13+ , 1.23+ , 2342342+ , t+ , fromJust (fromString "af93aafe-dea5-4427-bea4-8d7872507efb")+ , "sdfsdžȢぴせそぼξλж҈Ҵאבג"+ , 8763847563478568734687345683765873458734+ , TimeUuid . fromJust $ fromString "559ab19e-52d8-11e3-a847-270bf6910c08"+ , read "127.0.0.1"+ )+ let b = ( 4835637638+ , [1,2,3]+ , Set ["peter", "paul", "mary"]+ , Map [("peter", 1), ("paul", 2), ("mary", 3)]+ , Just 42+ , (True, "ascii", 42)+ , Map [(1, Map [(1, Set ["ascii"])])+ ,(2, Map [(2, Set ["ascii", "text"])])+ ]+ )+ write ins1 (params a)+ write ins2 (params b)+ x <- fromJust <$> query1 get1 (params (Identity 4835637638))+ y <- fromJust <$> query1 get2 (params (Identity 4835637638))+ liftIO $ do+ a @=? x+ b @=? y+ where+ ins1 :: PrepQuery W Ty1 ()+ ins1 = [r|+ insert into cqltest.test1+ (a,b,c,d,e,f,g,h,i,j,k,l,m,n)+ values+ (?,?,?,?,?,?,?,?,?,?,?,?,?,?) |]++ ins2 :: PrepQuery W Ty2 ()+ ins2 = [r|+ insert into cqltest.test2+ (a,b,c,d,e,f,g)+ values+ (?,?,?,?,?,?,?) |]++ get1 :: PrepQuery R (Identity Int64) Ty1+ get1 = "select a,b,c,d,e,f,g,h,i,j,k,l,m,n from cqltest.test1 where a = ?"++ get2 :: PrepQuery R (Identity Int64) Ty2+ get2 = "select a,b,c,d,e,f,g from cqltest.test2 where a = ?"++testWriteReadTtl :: Client ()+testWriteReadTtl = do+ write ins (params (1000, True))+ (True, Just ttl) <- fromJust <$> query1 get (params (Identity 1000))+ liftIO $ assertBool "TTL > 0" (ttl > 0)+ where+ ins :: PrepQuery W (Int64, Bool) ()+ ins = "insert into cqltest.test1 (a,d) values (?,?) using ttl 3600"++ get :: PrepQuery R (Identity Int64) (Bool, Maybe Int32)+ get = "select d, ttl(d) from cqltest.test1 where a = ?"++testTrans :: Client ()+testTrans = do+ -- 1st insert (success)+ [_row] <- trans ins (params (1, "ascii-1"))+ assertApplied _row++ -- 2nd insert (conflict)+ [_row] <- trans ins (params (1, "ascii-1"))+ liftIO $ do+ rowLength _row @?= 15 -- [applied] + full existing row+ fromRow 0 _row @?= Right (Just False) -- [applied]+ fromRow 1 _row @?= Right (Just (1 :: Int64)) -- a+ fromRow 2 _row @?= Right (Just (Ascii "ascii-1")) -- b+ -- remaining columns with null values (since none were inserted)+ let vnull = Nothing :: Maybe Blob -- type irrelevant+ map (($ _row) . fromRow) [3..14] @?= replicate 12 (Right vnull)++ -- 1st update (success)+ [_row] <- trans upd (params ("ascii-2", 1, "ascii-1"))+ assertApplied _row++ -- 2nd update (conflict)+ [_row] <- trans upd (params ("ascii-2", 1, "ascii-1"))+ liftIO $ do+ rowLength _row @?= 2 -- [applied] + conflicting value+ fromRow 0 _row @?= Right (Just False) -- [applied]+ fromRow 1 _row @?= Right (Just (Ascii "ascii-2")) -- b+ where+ ins :: PrepQuery W (Int64, Text) Row+ ins = "insert into cqltest.test1 (a,b) values (?,?) if not exists"++ upd :: PrepQuery W (Text, Int64, Text) Row+ upd = "update cqltest.test1 set b = ? where a = ? if b = ?"++ assertApplied row = liftIO $ do+ rowLength row @?= 1 -- [applied]+ fromRow 0 row @?= Right (Just True)++testPaging :: Client ()+testPaging = do+ let dat = zip [1..101] (repeat "b")+ mapM_ (write ins . params) dat+ p <- paginate qry $ (params ()) { pageSize = Just 10 }+ assertPages 11 (Set.fromList dat) p+ where+ ins :: PrepQuery W (Int64, Ascii) ()+ ins = "insert into cqltest.test1 (a,b) values (?,?)"++ qry :: PrepQuery R () (Int64, Ascii)+ qry = "select a,b from cqltest.test1"++testBatch :: Client ()+testBatch = do+ exec $ setType BatchLogged+ exec $ setType BatchUnLogged+ exec $ setSerialConsistency SerialConsistency+ exec $ setSerialConsistency LocalSerialConsistency+ where+ exec configure = do+ batch $ configure >> forM_ dat (addQuery ins)+ rs <- query qry (params ())+ liftIO $ sort rs @?= dat+ truncateTables++ dat :: [(Int64, Ascii)]+ dat = [(1, "1"), (2, "2"), (3, "3")]++ ins :: QueryString W (Int64, Ascii) ()+ ins = "insert into cqltest.test1 (a,b) values (?,?)"++ qry :: PrepQuery R () (Int64, Ascii)+ qry = "select a,b from cqltest.test1"++testBatchCounter :: Client ()+testBatchCounter = exec >> total 3 >> exec >> total 6+ where+ exec = batch $ do+ setType BatchCounter+ addQuery upd (Identity 1)+ addQuery upd (Identity 2)+ addQuery upd (Identity 3)++ total n = do+ rs <- query qry (params ())+ let n' = sum (map (fromCounter . runIdentity) rs)+ liftIO $ n @=? n'++ upd :: QueryString W (Identity Int64) ()+ upd = "update cqltest.counters set n = n + 1 where a = ?"++ qry :: PrepQuery R () (Identity Counter)+ qry = "select n from cqltest.counters"++-----------------------------------------------------------------------------+-- Utilities++assertPages :: (Ord a, Show a) => Int -> Set.Set a -> Page a -> Client ()+assertPages numPages expected p = do+ let got = Set.fromList (result p)+ let remaining = Set.difference expected got+ liftIO $ hasMore p @?= numPages > 1+ liftIO $ got `Set.isSubsetOf` expected @?= True+ if numPages > 1+ then nextPage p >>= assertPages (numPages - 1) remaining+ else liftIO $ remaining @?= Set.empty++params :: Tuple a => a -> QueryParams a+params p = QueryParams+ { consistency = One+ , skipMetaData = False+ , values = p+ , pageSize = Nothing+ , queryPagingState = Nothing+ , serialConsistency = Nothing+ , enableTracing = Nothing+ }+