Skip to content
Merged
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
23 changes: 20 additions & 3 deletions client/config.go
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@ package client

import (
"net/url"
"strconv"
"strings"
)

Expand All @@ -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.
Expand All @@ -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()
Expand All @@ -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
}
46 changes: 37 additions & 9 deletions client/config_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
})
}
165 changes: 106 additions & 59 deletions client/conn.go
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand All @@ -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) {
Expand All @@ -67,78 +76,106 @@ 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()
}
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.
Expand All @@ -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
}

Expand Down Expand Up @@ -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
Expand All @@ -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")
}()
Expand All @@ -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
}

Expand All @@ -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,
Expand Down
8 changes: 7 additions & 1 deletion client/driver_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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,
})
})
}

Expand Down