diff --git a/client/config.go b/client/config.go index 9f18a6861..cde8eec94 100644 --- a/client/config.go +++ b/client/config.go @@ -18,6 +18,7 @@ package client import ( "net/url" + "strconv" "strings" ) @@ -28,11 +29,17 @@ type Config struct { // additional configs should be filled // such as read/write/exec timeout // currently no timeout is supported. + + // UseLeader use leader nodes to do queries + UseLeader bool + + // UseFollower use follower nodes to do queries + UseFollower bool } // NewConfig creates a new config with default value. func NewConfig() *Config { - return &Config{} + return &Config{UseLeader: true} } // FormatDSN formats the given Config into a DSN string which can be passed to the driver. @@ -45,6 +52,8 @@ func (cfg *Config) FormatDSN() string { } newQuery := u.Query() + newQuery.Add("use_leader", strconv.FormatBool(cfg.UseLeader)) + newQuery.Add("use_follower", strconv.FormatBool(cfg.UseFollower)) u.RawQuery = newQuery.Encode() return u.String() @@ -58,11 +67,19 @@ func ParseDSN(dsn string) (cfg *Config, err error) { var u *url.URL if u, err = url.Parse(dsn); err != nil { - return + return nil, err } cfg = NewConfig() cfg.DatabaseID = u.Host - return + q := u.Query() + // option: use_leader, use_follower + cfg.UseLeader, _ = strconv.ParseBool(q.Get("use_leader")) + cfg.UseFollower, _ = strconv.ParseBool(q.Get("use_follower")) + if !cfg.UseLeader && !cfg.UseFollower { + cfg.UseLeader = true + } + + return cfg, nil } diff --git a/client/config_test.go b/client/config_test.go index f7a3d501b..cbfb7a60d 100644 --- a/client/config_test.go +++ b/client/config_test.go @@ -19,28 +19,56 @@ package client import ( "testing" - "github.com/CovenantSQL/CovenantSQL/proto" . "github.com/smartystreets/goconvey/convey" ) func TestConfig(t *testing.T) { - Convey("test config", t, func() { - var cfg *Config - var err error + Convey("test config without additional options", t, func() { + cfg, err := ParseDSN("covenantsql://db") + So(err, ShouldBeNil) + So(cfg, ShouldResemble, &Config{ + DatabaseID: "db", + UseLeader: true, + UseFollower: false, + }) - cfg, err = ParseDSN("covenantsql://db") + recoveredCfg, err := ParseDSN(cfg.FormatDSN()) So(err, ShouldBeNil) - So(cfg.DatabaseID, ShouldEqual, proto.DatabaseID("db")) - So(cfg.FormatDSN(), ShouldEqual, "covenantsql://db") + So(cfg, ShouldResemble, recoveredCfg) }) + Convey("test invalid config", t, func() { - _, err := ParseDSN("invalid dsn") + cfg, err := ParseDSN("invalid dsn") So(err, ShouldNotBeNil) + So(cfg, ShouldBeNil) }) + Convey("test dsn with only database id", t, func() { dbIDStr := "00000bef611d346c0cbe1beaa76e7f0ed705a194fdf9ac3a248ec70e9c198bf9" cfg, err := ParseDSN(dbIDStr) So(err, ShouldBeNil) - So(cfg.DatabaseID, ShouldEqual, dbIDStr) + So(cfg, ShouldResemble, &Config{ + DatabaseID: dbIDStr, + UseLeader: true, + UseFollower: false, + }) + + recoveredCfg, err := ParseDSN(cfg.FormatDSN()) + So(err, ShouldBeNil) + So(cfg, ShouldResemble, recoveredCfg) + }) + + Convey("test dsn with additional options", t, func() { + cfg, err := ParseDSN("covenantsql://db?use_leader=0&use_follower=true") + So(err, ShouldBeNil) + So(cfg, ShouldResemble, &Config{ + DatabaseID: "db", + UseLeader: false, + UseFollower: true, + }) + + recoveredCfg, err := ParseDSN(cfg.FormatDSN()) + So(err, ShouldBeNil) + So(cfg, ShouldResemble, recoveredCfg) }) } diff --git a/client/conn.go b/client/conn.go index eeec4d1db..b0ef613af 100644 --- a/client/conn.go +++ b/client/conn.go @@ -31,6 +31,7 @@ import ( "github.com/CovenantSQL/CovenantSQL/rpc" "github.com/CovenantSQL/CovenantSQL/types" "github.com/CovenantSQL/CovenantSQL/utils/log" + "github.com/pkg/errors" ) // conn implements an interface sql.Conn. @@ -41,10 +42,18 @@ type conn struct { localNodeID proto.NodeID privKey *asymmetric.PrivateKey - ackCh chan *types.Ack inTransaction bool closed int32 - pCaller *rpc.PersistentCaller + + leader *pconn + follower *pconn +} + +// pconn represents a connection to a peer +type pconn struct { + parent *conn + ackCh chan *types.Ack + pCaller *rpc.PersistentCaller } func newConn(cfg *Config) (c *conn, err error) { @@ -67,27 +76,53 @@ func newConn(cfg *Config) (c *conn, err error) { queries: make([]types.Query, 0), } - var peers *proto.Peers // get peers from BP + var peers *proto.Peers if peers, err = cacheGetPeers(c.dbID, c.privKey); err != nil { - log.WithError(err).Error("cacheGetPeers failed") - c = nil - return + return nil, errors.WithMessage(err, "cacheGetPeers failed") } - c.pCaller = rpc.NewPersistentCaller(peers.Leader) - err = c.startAckWorkers(2) - if err != nil { - log.WithError(err).Error("startAckWorkers failed") - c = nil - return + if cfg.UseLeader { + c.leader = &pconn{ + parent: c, + pCaller: rpc.NewPersistentCaller(peers.Leader), + } } - log.WithField("db", c.dbID).Debug("new connection to database") + // choose a random follower node + if cfg.UseFollower && len(peers.Servers) > 1 { + for { + node := peers.Servers[randSource.Intn(len(peers.Servers))] + if node != peers.Leader { + c.follower = &pconn{ + parent: c, + pCaller: rpc.NewPersistentCaller(node), + } + break + } + } + } + + if c.leader == nil && c.follower == nil { + return nil, errors.New("no follower peers found") + } + + if c.leader != nil { + if err := c.leader.startAckWorkers(2); err != nil { + return nil, errors.WithMessage(err, "leader startAckWorkers failed") + } + } + if c.follower != nil { + if err := c.follower.startAckWorkers(2); err != nil { + return nil, errors.WithMessage(err, "follower startAckWorkers failed") + } + } + + log.WithField("db", c.dbID).Debug("new connection to database") return } -func (c *conn) startAckWorkers(workerCount int) (err error) { +func (c *pconn) startAckWorkers(workerCount int) (err error) { c.ackCh = make(chan *types.Ack, workerCount*4) for i := 0; i < workerCount; i++ { go c.ackWorker() @@ -95,50 +130,52 @@ func (c *conn) startAckWorkers(workerCount int) (err error) { return } -func (c *conn) stopAckWorkers() { +func (c *pconn) stopAckWorkers() { close(c.ackCh) } -func (c *conn) ackWorker() { - if rawPeers, ok := peerList.Load(c.dbID); ok { - if peers, ok := rawPeers.(*proto.Peers); ok { - var ( - oneTime sync.Once - pc *rpc.PersistentCaller - err error - ) - - ackWorkerLoop: - for { - ack, got := <-c.ackCh - if !got { //closed and empty - break ackWorkerLoop - } - oneTime.Do(func() { - pc = rpc.NewPersistentCaller(peers.Leader) - }) - if err = ack.Sign(c.privKey, false); err != nil { - log.WithField("target", pc.TargetID).WithError(err).Error("failed to sign ack") - continue - } +func (c *pconn) ackWorker() { + var ( + oneTime sync.Once + pc *rpc.PersistentCaller + err error + ) + +ackWorkerLoop: + for { + ack, got := <-c.ackCh + if !got { // closed and empty + break ackWorkerLoop + } + oneTime.Do(func() { + pc = rpc.NewPersistentCaller(c.pCaller.TargetID) + }) + if err = ack.Sign(c.parent.privKey, false); err != nil { + log.WithField("target", pc.TargetID).WithError(err).Error("failed to sign ack") + continue + } - var ackRes types.AckResponse - // send ack back - if err = pc.Call(route.DBSAck.String(), ack, &ackRes); err != nil { - log.WithError(err).Warning("send ack failed") - continue - } - } - if pc != nil { - pc.CloseStream() - } - log.Debug("ack worker quiting") - return + var ackRes types.AckResponse + // send ack back + if err = pc.Call(route.DBSAck.String(), ack, &ackRes); err != nil { + log.WithError(err).Warning("send ack failed") + continue } } - log.Fatal("must GetPeers first") - return + if pc != nil { + pc.CloseStream() + } + + log.Debug("ack worker quiting") +} + +func (c *pconn) close() error { + c.stopAckWorkers() + if c.pCaller != nil { + c.pCaller.CloseStream() + } + return nil } // Prepare implements the driver.Conn.Prepare method. @@ -152,8 +189,12 @@ func (c *conn) Close() error { if atomic.CompareAndSwapInt32(&c.closed, 0, 1) { log.WithField("db", c.dbID).Debug("closed connection") } - c.stopAckWorkers() - c.pCaller.CloseStream() + if c.leader != nil { + c.leader.close() + } + if c.follower != nil { + c.follower.close() + } return nil } @@ -307,9 +348,15 @@ func (c *conn) addQuery(queryType types.QueryType, query *types.Query) (affected } func (c *conn) sendQuery(queryType types.QueryType, queries []types.Query) (affectedRows int64, lastInsertID int64, rows driver.Rows, err error) { - var peers *proto.Peers - if peers, err = cacheGetPeers(c.dbID, c.privKey); err != nil { - return + var uc *pconn // peer connection used to execute the queries + + uc = c.leader + // use follower pconn only when the query is readonly + if queryType == types.ReadQuery && c.follower != nil { + uc = c.follower + } + if uc == nil { + uc = c.follower } // allocate sequence @@ -322,7 +369,7 @@ func (c *conn) sendQuery(queryType types.QueryType, queries []types.Query) (affe "type": queryType.String(), "connID": connID, "seqNo": seqNo, - "target": peers.Leader, + "target": uc.pCaller.TargetID, "source": c.localNodeID, }).WithError(err).Debug("send query") }() @@ -349,7 +396,7 @@ func (c *conn) sendQuery(queryType types.QueryType, queries []types.Query) (affe } var response types.Response - if err = c.pCaller.Call(route.DBSQuery.String(), req, &response); err != nil { + if err = uc.pCaller.Call(route.DBSQuery.String(), req, &response); err != nil { return } @@ -365,7 +412,7 @@ func (c *conn) sendQuery(queryType types.QueryType, queries []types.Query) (affe } // build ack - c.ackCh <- &types.Ack{ + uc.ackCh <- &types.Ack{ Header: types.SignedAckHeader{ AckHeader: types.AckHeader{ Response: response.Header, diff --git a/client/driver_test.go b/client/driver_test.go index 018cd9665..4fade57e5 100644 --- a/client/driver_test.go +++ b/client/driver_test.go @@ -61,7 +61,13 @@ func TestCreate(t *testing.T) { var dsn string dsn, err = Create(ResourceMeta{}) So(err, ShouldBeNil) - So(dsn, ShouldEqual, "covenantsql://db") + + recoveredCfg, err := ParseDSN(dsn) + So(err, ShouldBeNil) + So(recoveredCfg, ShouldResemble, &Config{ + DatabaseID: "db", + UseLeader: true, + }) }) }