Compare commits

..

1 Commits

Author SHA1 Message Date
Christian Schwarz 17add553d3 WIP runtime-controllable concurrency for replication
Changes done so far:
- signal to route concurrency request up to the stepQueue
    - pretty hacky, no status reporting yet
- stepQueue upsizing (confirmed that it works, no intermediary commit)

Stuck at: stepQueue downsizing
- idea was to have the stepQueue signal to the activated step that it
  should suspend
- ideally, we'd just kill everything associated with the step and track
  it as restartable
    - the driver model doesn't allow for that though within an attempt
    - would need to start a separate run for the step
- less perfect: tell the downsized steps to stop copying, but leave
  all zfs sends + rpc conns open
    - - doesn't give back the resoures that a step aquired before
      being selected as downsize victim (open zfs processe + TCP conn,
      some memory in the pipe)
    - - looks weird to user if the ps aux
    - + re-waking a step is easy: just tell it to proceed with copying
    - (the impl would likely pass a check function down into Step.do
      and have it check that functino periodically. the suspend should
      be acknowledged, and stepQueue should only remove the step from
      the active queue _after_ that step has acknowledged that is
      suspended)
2020-02-17 17:22:52 +01:00
30 changed files with 281 additions and 376 deletions
+10 -3
View File
@@ -9,7 +9,7 @@ import (
)
var SignalCmd = &cli.Subcommand{
Use: "signal [wakeup|reset] JOB",
Use: "signal [wakeup|reset] JOB [DATA]",
Short: "wake up a job from wait state or abort its current invocation",
Run: func(subcommand *cli.Subcommand, args []string) error {
return runSignalCmd(subcommand.Config(), args)
@@ -17,8 +17,13 @@ var SignalCmd = &cli.Subcommand{
}
func runSignalCmd(config *config.Config, args []string) error {
if len(args) != 2 {
return errors.Errorf("Expected 2 arguments: [wakeup|reset] JOB")
if len(args) < 2 || len(args) > 3 {
return errors.Errorf("Expected 2 arguments: [wakeup|reset|set-concurrency] JOB [DATA]")
}
var data string
if len(args) == 3 {
data = args[2]
}
httpc, err := controlHttpClient(config.Global.Control.SockPath)
@@ -30,9 +35,11 @@ func runSignalCmd(config *config.Config, args []string) error {
struct {
Name string
Op string
Data string
}{
Name: args[1],
Op: args[0],
Data: data,
},
struct{}{},
)
+12
View File
@@ -8,6 +8,7 @@ import (
"io"
"net"
"net/http"
"strconv"
"time"
"github.com/pkg/errors"
@@ -47,6 +48,8 @@ func (j *controlJob) OwnedDatasetSubtreeRoot() (p *zfs.DatasetPath, ok bool) { r
func (j *controlJob) SenderConfig() *endpoint.SenderConfig { return nil }
func (j *controlJob) SetConcurrency(concurrency int) error { return errors.Errorf("not supported") }
var promControl struct {
requestBegin *prometheus.CounterVec
requestFinished *prometheus.HistogramVec
@@ -126,6 +129,7 @@ func (j *controlJob) Run(ctx context.Context) {
type reqT struct {
Name string
Op string
Data string
}
var req reqT
if decoder(&req) != nil {
@@ -138,6 +142,14 @@ func (j *controlJob) Run(ctx context.Context) {
err = j.jobs.wakeup(req.Name)
case "reset":
err = j.jobs.reset(req.Name)
case "set-concurrency":
var concurrency int
concurrency, err = strconv.Atoi(req.Data) // shadow
if err != nil {
// fallthrough outer
} else {
err = j.jobs.setConcurrency(req.Name, concurrency)
}
default:
err = fmt.Errorf("operation %q is invalid", req.Op)
}
+12
View File
@@ -176,6 +176,18 @@ func (s *jobs) reset(job string) error {
return wu()
}
func (s *jobs) setConcurrency(jobName string, concurrency int) error {
s.m.RLock()
defer s.m.RUnlock()
job, ok := s.jobs[jobName]
if !ok {
return errors.Errorf("Job %q does not exist", job)
}
return job.SetConcurrency(concurrency)
}
const (
jobNamePrometheus = "_prometheus"
jobNameControl = "_control"
+30 -6
View File
@@ -55,9 +55,12 @@ const (
type activeSideTasks struct {
state ActiveSideState
concurrency int
// valid for state ActiveSideReplicating, ActiveSidePruneSender, ActiveSidePruneReceiver, ActiveSideDone
replicationReport driver.ReportFunc
replicationCancel context.CancelFunc
replicationReport driver.ReportFunc
replicationCancel context.CancelFunc
replicationSetConcurrency driver.SetConcurrencyFunc
// valid for state ActiveSidePruneSender, ActiveSidePruneReceiver, ActiveSideDone
prunerSender, prunerReceiver *pruner.Pruner
@@ -278,6 +281,8 @@ func activeSide(g *config.Global, in *config.ActiveJob, configJob interface{}) (
return nil, errors.Wrap(err, "invalid job name")
}
j.tasks.concurrency = 1 // FIXME
switch v := configJob.(type) {
case *config.PushJob:
j.mode, err = modePushFromConfig(g, v, j.name) // shadow
@@ -375,6 +380,22 @@ func (j *ActiveSide) SenderConfig() *endpoint.SenderConfig {
return push.senderConfig
}
func (j *ActiveSide) SetConcurrency(concurrency int) (err error) {
j.updateTasks(func(tasks *activeSideTasks) {
if tasks.replicationSetConcurrency != nil {
err = tasks.replicationSetConcurrency(concurrency) // no shadow
if err == nil {
tasks.concurrency = concurrency
}
} else {
// FIXME this is not great, should always be able to set it
err = errors.Errorf("cannot set while not replicating")
}
})
return err
}
func (j *ActiveSide) Run(ctx context.Context) {
log := GetLogger(ctx)
ctx = logging.WithSubsystemLoggers(ctx, log)
@@ -436,11 +457,14 @@ func (j *ActiveSide) do(ctx context.Context) {
ctx, repCancel := context.WithCancel(ctx)
var repWait driver.WaitFunc
j.updateTasks(func(tasks *activeSideTasks) {
// reset it
*tasks = activeSideTasks{}
// reset it (almost)
old := *tasks
*tasks = activeSideTasks{
concurrency: old.concurrency,
}
tasks.replicationCancel = repCancel
tasks.replicationReport, repWait = replication.Do(
ctx, logic.NewPlanner(j.promRepStateSecs, j.promBytesReplicated, sender, receiver, j.mode.PlannerPolicy()),
tasks.replicationReport, repWait, tasks.replicationSetConcurrency = replication.Do(
ctx, tasks.concurrency, logic.NewPlanner(j.promRepStateSecs, j.promBytesReplicated, sender, receiver, j.mode.PlannerPolicy()),
)
tasks.state = ActiveSideReplicating
})
+1
View File
@@ -40,6 +40,7 @@ type Job interface {
// must return the root of that subtree as rfs and ok = true
OwnedDatasetSubtreeRoot() (rfs *zfs.DatasetPath, ok bool)
SenderConfig() *endpoint.SenderConfig
SetConcurrency(concurrency int) error
}
type Type string
+4 -31
View File
@@ -4,7 +4,6 @@ import (
"context"
"fmt"
"github.com/golang/protobuf/proto"
"github.com/pkg/errors"
"github.com/prometheus/client_golang/prometheus"
@@ -164,6 +163,8 @@ func (j *PassiveSide) SenderConfig() *endpoint.SenderConfig {
func (*PassiveSide) RegisterMetrics(registerer prometheus.Registerer) {}
func (*PassiveSide) SetConcurrency(concurrency int) error { return errors.Errorf("not supported") }
func (j *PassiveSide) Run(ctx context.Context) {
log := GetLogger(ctx)
@@ -181,39 +182,11 @@ func (j *PassiveSide) Run(ctx context.Context) {
}
ctxInterceptor := func(handlerCtx context.Context) context.Context {
reqLog := log.WithField(logging.ReqIDField, rpc.MustGetRequestID(handlerCtx))
return logging.WithSubsystemLoggers(handlerCtx, reqLog)
}
// ctxInterceptor already injected request loggers (with request id)
pre := func(ctx context.Context, method string, req interface{}) {
l := endpoint.GetLogger(ctx).WithField("method", method).WithField("reqT", fmt.Sprintf("%T", req))
if p, ok := req.(proto.Message); ok {
l = l.WithField("req", proto.CompactTextString(p))
}
ci, ok := ctx.Value(endpoint.ClientIdentityKey).(string)
if !ok {
// like endpoint, we must assume that rpc.Server sets it
panic(method)
}
l = l.WithField("clientIdentity", ci)
l.Debug("incoming")
}
post := func(ctx context.Context, res interface{}, err error) {
l := endpoint.GetLogger(ctx).
WithField("responseT", fmt.Sprintf("%T", res)).
WithField("errT", fmt.Sprintf("%T", err))
if s, ok := res.(fmt.Stringer); ok {
l = l.WithField("response", s.String())
}
if err != nil {
l = l.WithError(err)
}
l.Debug("handler done")
return logging.WithSubsystemLoggers(handlerCtx, log)
}
rpcLoggers := rpc.GetLoggersOrPanic(ctx) // WithSubsystemLoggers above
server := rpc.NewServer(handler, rpcLoggers, ctxInterceptor, pre, post)
server := rpc.NewServer(handler, rpcLoggers, ctxInterceptor)
listener, err := j.listen()
if err != nil {
+2
View File
@@ -117,6 +117,8 @@ outer:
}
}
func (*SnapJob) SetConcurrency(concurrency int) error { return errors.Errorf("not supported") }
// Adaptor that implements pruner.History around a pruner.Target.
// The ReplicationCursor method is Get-op only and always returns
// the filesystem's most recent version's GUID.
-5
View File
@@ -96,11 +96,6 @@ func WithSubsystemLoggers(ctx context.Context, log logger.Logger) context.Contex
Control: log.WithField(SubsysField, SubsysRPCControl),
Data: log.WithField(SubsysField, SubsysRPCData),
},
// rpc.Loggers{
// General: logger.NewNullLogger(),
// Control: logger.NewNullLogger(),
// Data: logger.NewNullLogger(),
// },
)
return ctx
}
+1 -2
View File
@@ -22,7 +22,6 @@ const (
const (
JobField string = "job"
SubsysField string = "subsystem"
ReqIDField string = "reqid"
)
type MetadataFlags int64
@@ -86,7 +85,7 @@ func (f *HumanFormatter) Format(e *logger.Entry) (out []byte, err error) {
fmt.Fprintf(&line, "[%s]", col.Sprint(e.Level.Short()))
}
prefixFields := []string{JobField, SubsysField, ReqIDField}
prefixFields := []string{JobField, SubsysField}
prefixed := make(map[string]bool, len(prefixFields)+2)
for _, field := range prefixFields {
val, ok := e.Fields[field]
+3
View File
@@ -5,6 +5,7 @@ import (
"net"
"net/http"
"github.com/pkg/errors"
"github.com/prometheus/client_golang/prometheus"
"github.com/prometheus/client_golang/prometheus/promhttp"
@@ -51,6 +52,8 @@ func (j *prometheusJob) OwnedDatasetSubtreeRoot() (p *zfs.DatasetPath, ok bool)
func (j *prometheusJob) SenderConfig() *endpoint.SenderConfig { return nil }
func (j *prometheusJob) SetConcurrency(concurrency int) error { return errors.Errorf("not supported") }
func (j *prometheusJob) RegisterMetrics(registerer prometheus.Registerer) {}
func (j *prometheusJob) Run(ctx context.Context) {
-4
View File
@@ -19,10 +19,6 @@ func WithLogger(ctx context.Context, log Logger) context.Context {
return context.WithValue(ctx, contextKeyLogger, log)
}
func GetLogger(ctx context.Context) Logger {
return getLogger(ctx)
}
func getLogger(ctx context.Context) Logger {
if l, ok := ctx.Value(contextKeyLogger).(Logger); ok {
return l
+2 -4
View File
@@ -159,7 +159,6 @@ func sendArgsFromPDUAndValidateExists(ctx context.Context, fs string, fsv *pdu.F
}
func (s *Sender) Send(ctx context.Context, r *pdu.SendReq) (*pdu.SendRes, zfs.StreamCopier, error) {
getLogger(ctx).WithField("req", r.String()).Debug("incoming Send request")
_, err := s.filterCheckFS(r.Filesystem)
if err != nil {
@@ -260,7 +259,6 @@ func (s *Sender) Send(ctx context.Context, r *pdu.SendReq) (*pdu.SendRes, zfs.St
}
func (p *Sender) SendCompleted(ctx context.Context, r *pdu.SendCompletedReq) (*pdu.SendCompletedRes, error) {
getLogger(ctx).WithField("req", r.String()).Debug("incoming SendCompleted request")
orig := r.GetOriginalReq() // may be nil, always use proto getters
fs := orig.GetFilesystem()
@@ -512,7 +510,7 @@ func (s *Receiver) ListFilesystems(ctx context.Context, req *pdu.ListFilesystemR
// present filesystem without the root_fs prefix
fss := make([]*pdu.Filesystem, 0, len(filtered))
for _, a := range filtered {
l := getLogger(ctx).WithField("fs", a.ToString())
l := getLogger(ctx).WithField("fs", a)
ph, err := zfs.ZFSGetFilesystemPlaceholderState(a)
if err != nil {
l.WithError(err).Error("error getting placeholder state")
@@ -599,7 +597,7 @@ func (s *Receiver) Send(ctx context.Context, req *pdu.SendReq) (*pdu.SendRes, zf
var maxConcurrentZFSRecvSemaphore = semaphore.New(envconst.Int64("ZREPL_ENDPOINT_MAX_CONCURRENT_RECV", 10))
func (s *Receiver) Receive(ctx context.Context, req *pdu.ReceiveReq, receive zfs.StreamCopier) (*pdu.ReceiveRes, error) {
getLogger(ctx).WithField("rr", req.String()).Debug("incoming Receive")
getLogger(ctx).Debug("incoming Receive")
defer receive.Close()
root := s.clientRootFromCtx(ctx)
-1
View File
@@ -7,7 +7,6 @@ require (
github.com/gdamore/tcell v1.2.0
github.com/go-logfmt/logfmt v0.4.0
github.com/go-sql-driver/mysql v1.4.1-0.20190907122137-b2c03bcae3d4
github.com/gogo/protobuf v1.1.1
github.com/golang/protobuf v1.3.2
github.com/google/uuid v1.1.1
github.com/jinzhu/copier v0.0.0-20170922082739-db4671f3a9b8
-1
View File
@@ -68,7 +68,6 @@ github.com/go-toolsmith/strparse v1.0.0/go.mod h1:YI2nUKP9YGZnL/L1/DLFBfixrcjslW
github.com/go-toolsmith/typep v0.0.0-20181030061450-d63dc7650676/go.mod h1:JSQCQMUPdRlMZFswiq3TGpNp1GMktqkR2Ns5AIQkATU=
github.com/go-toolsmith/typep v1.0.0/go.mod h1:JSQCQMUPdRlMZFswiq3TGpNp1GMktqkR2Ns5AIQkATU=
github.com/gobwas/glob v0.2.3/go.mod h1:d3Ez4x06l9bZtSvzIay5+Yzi0fmZzPgnTbPcKjJAkT8=
github.com/gogo/protobuf v1.1.1 h1:72R+M5VuhED/KujmZVcIquuo8mBgX4oVda//DQb3PXo=
github.com/gogo/protobuf v1.1.1/go.mod h1:r8qH/GZQm5c6nD/R0oafs1akxWv10x8SbQlK7atdtwQ=
github.com/gogo/protobuf v1.2.1/go.mod h1:hp+jE20tsWTFYpLwKvXlhS1hjn+gTNwPg2I6zVXpSg4=
github.com/golang/glog v0.0.0-20160126235308-23def4e6c14b h1:VKtxabqXZkF25pY9ekfRL6a582T4P37/31XEstQ5p58=
+59 -17
View File
@@ -5,7 +5,6 @@ import (
"errors"
"fmt"
"net"
"os"
"sort"
"strings"
"sync"
@@ -15,7 +14,6 @@ import (
"google.golang.org/grpc/status"
"github.com/zrepl/zrepl/replication/report"
"github.com/zrepl/zrepl/tracing"
"github.com/zrepl/zrepl/util/chainlock"
"github.com/zrepl/zrepl/util/envconst"
)
@@ -91,6 +89,8 @@ type attempt struct {
// if both are nil, it must be assumed that Planner.Plan is active
planErr *timedError
fss []*fs
concurrency int
}
type timedError struct {
@@ -172,11 +172,12 @@ type step struct {
type ReportFunc func() *report.Report
type WaitFunc func(block bool) (done bool)
type SetConcurrencyFunc func(concurrency int) error
var maxAttempts = envconst.Int64("ZREPL_REPLICATION_MAX_ATTEMPTS", 3)
var reconnectHardFailTimeout = envconst.Duration("ZREPL_REPLICATION_RECONNECT_HARD_FAIL_TIMEOUT", 10*time.Minute)
func Do(ctx context.Context, planner Planner) (ReportFunc, WaitFunc) {
func Do(ctx context.Context, initialConcurrency int, planner Planner) (ReportFunc, WaitFunc, SetConcurrencyFunc) {
log := getLog(ctx)
l := chainlock.New()
run := &run{
@@ -184,6 +185,8 @@ func Do(ctx context.Context, planner Planner) (ReportFunc, WaitFunc) {
startedAt: time.Now(),
}
concurrencyChanges := make(chan concurrencyChange)
done := make(chan struct{})
go func() {
defer close(done)
@@ -200,16 +203,21 @@ func Do(ctx context.Context, planner Planner) (ReportFunc, WaitFunc) {
run.waitReconnect.SetZero()
run.waitReconnectError = nil
prevConcurrency := initialConcurrency // FIXME default concurrency
if prev != nil {
prevConcurrency = prev.concurrency
}
// do current attempt
cur := &attempt{
l: l,
startedAt: time.Now(),
planner: planner,
l: l,
startedAt: time.Now(),
planner: planner,
concurrency: prevConcurrency,
}
run.attempts = append(run.attempts, cur)
run.l.DropWhile(func() {
ctx := tracing.Child(ctx, fmt.Sprintf("attempt#%d", ano)) // shadow
cur.do(ctx, prev)
cur.do(ctx, prev, concurrencyChanges)
})
prev = cur
if ctx.Err() != nil {
@@ -280,11 +288,26 @@ func Do(ctx context.Context, planner Planner) (ReportFunc, WaitFunc) {
defer run.l.Lock().Unlock()
return run.report()
}
return report, wait
setConcurrency := func(concurrency int) (reterr error) {
var wg sync.WaitGroup
wg.Add(1)
concurrencyChanges <- concurrencyChange{ concurrency, func(err error) {
defer wg.Done()
reterr = err // shadow
}}
wg.Wait()
return reterr
}
return report, wait, setConcurrency
}
func (a *attempt) do(ctx context.Context, prev *attempt) {
pfss, err := a.planner.Plan(tracing.Child(ctx, "plan"))
type concurrencyChange struct {
value int
resultCallback func(error)
}
func (a *attempt) do(ctx context.Context, prev *attempt, setConcurrency <-chan concurrencyChange) {
pfss, err := a.planner.Plan(ctx)
errTime := time.Now()
defer a.l.Lock().Unlock()
if err != nil {
@@ -357,26 +380,43 @@ func (a *attempt) do(ctx context.Context, prev *attempt) {
}
// invariant: prevs contains an entry for each unambigious correspondence
stepQueue := newStepQueue()
defer stepQueue.Start(1)() // TODO parallel replication
stepQueue := newStepQueue(a.concurrency)
defer stepQueue.Start()()
var fssesDone sync.WaitGroup
for _, f := range a.fss {
fssesDone.Add(1)
go func(f *fs) {
defer fssesDone.Done()
f.do(tracing.Child(ctx, f.fs.ReportInfo().Name), stepQueue, prevs[f])
f.do(ctx, stepQueue, prevs[f])
}(f)
}
changeConcurrencyDone := make(chan struct{})
go func() {
for {
select {
case change := <-setConcurrency:
err := stepQueue.SetConcurrency(change.value)
go change.resultCallback(err)
if err == nil {
a.l.Lock()
a.concurrency = change.value
a.l.Unlock()
}
case <-changeConcurrencyDone:
return
// not waiting for ctx.Done here, the main job are the fsses
}
}
}()
a.l.DropWhile(func() {
fssesDone.Wait()
close(changeConcurrencyDone)
})
a.finishedAt = time.Now()
}
func (fs *fs) do(ctx context.Context, pq *stepQueue, prev *fs) {
fmt.Fprintf(os.Stderr, "CHILD STACK: %v\n", tracing.GetStack(ctx))
defer fs.l.Lock().Unlock()
// get planned steps from replication logic
@@ -466,7 +506,8 @@ func (fs *fs) do(ctx context.Context, pq *stepQueue, prev *fs) {
// lock must not be held while executing step in order for reporting to work
fs.l.DropWhile(func() {
targetDate := s.step.TargetDate()
defer pq.WaitReady(fs, targetDate)()
ctx, done := pq.WaitReady(ctx, fs, targetDate)()
defer done()
err = s.step.Step(ctx) // no shadow
errTime = time.Now() // no shadow
})
@@ -487,6 +528,7 @@ func (r *run) report() *report.Report {
WaitReconnectSince: r.waitReconnect.begin,
WaitReconnectUntil: r.waitReconnect.end,
WaitReconnectError: r.waitReconnectError.IntoReportError(),
// Concurrency: r.concurrency,
}
for i := range report.Attempts {
report.Attempts[i] = r.attempts[i].report()
@@ -151,7 +151,7 @@ func TestReplication(t *testing.T) {
ctx := context.Background()
mp := &mockPlanner{}
getReport, wait := Do(ctx, mp)
getReport, wait, _ := Do(ctx, 1, mp)
begin := time.Now()
fireAt := []time.Duration{
// the following values are relative to the start
+105 -47
View File
@@ -2,6 +2,8 @@ package driver
import (
"container/heap"
"fmt"
"sync"
"time"
"github.com/zrepl/zrepl/util/chainlock"
@@ -11,51 +13,88 @@ type stepQueueRec struct {
ident interface{}
targetDate time.Time
wakeup chan StepCompletedFunc
cancelDueToConcurrencyDownsize interace{}
}
type stepQueue struct {
stop chan struct{}
reqs chan stepQueueRec
// l protects all members except the channels above
l *chainlock.L
pendingCond *sync.Cond
// ident => queueItem
pending *stepQueueHeap
active *stepQueueHeap
queueItems map[interface{}]*stepQueueHeapItem // for tracking used idents in both pending and active
// stopped is used for cancellation of "wake" goroutine
stopped bool
concurrency int
}
type stepQueueHeapItem struct {
idx int
req stepQueueRec
req *stepQueueRec
}
type stepQueueHeap struct {
items []*stepQueueHeapItem
reverse bool // never change after pushing first element
}
type stepQueueHeap []*stepQueueHeapItem
func (h stepQueueHeap) Less(i, j int) bool {
return h[i].req.targetDate.Before(h[j].req.targetDate)
res := h.items[i].req.targetDate.Before(h.items[j].req.targetDate)
if h.reverse {
return !res
}
return res
}
func (h stepQueueHeap) Swap(i, j int) {
h[i], h[j] = h[j], h[i]
h[i].idx = i
h[j].idx = j
h.items[i], h.items[j] = h.items[j], h.items[i]
h.items[i].idx = i
h.items[j].idx = j
}
func (h stepQueueHeap) Len() int {
return len(h)
return len(h.items)
}
func (h *stepQueueHeap) Push(elem interface{}) {
hitem := elem.(*stepQueueHeapItem)
hitem.idx = h.Len()
*h = append(*h, hitem)
h.items = append(h.items, hitem)
}
func (h *stepQueueHeap) Pop() interface{} {
elem := (*h)[h.Len()-1]
elem := h.items[h.Len()-1]
elem.idx = -1
*h = (*h)[:h.Len()-1]
h.items = h.items[:h.Len()-1]
return elem
}
// returned stepQueue must be closed with method Close
func newStepQueue() *stepQueue {
func newStepQueue(concurrency int) *stepQueue {
l := chainlock.New()
q := &stepQueue{
stop: make(chan struct{}),
reqs: make(chan stepQueueRec),
stop: make(chan struct{}),
reqs: make(chan stepQueueRec),
l: l,
pendingCond: l.NewCond(),
// priority queue
pending: &stepQueueHeap{reverse: false},
active: &stepQueueHeap{reverse: true},
// ident => queueItem
queueItems: make(map[interface{}]*stepQueueHeapItem),
// stopped is used for cancellation of "wake" goroutine
stopped: false,
}
err := q.setConcurrencyLocked(concurrency)
if err != nil {
panic(err)
}
return q
}
@@ -65,25 +104,12 @@ func newStepQueue() *stepQueue {
//
// No WaitReady calls must be active at the time done is called
// The behavior of calling WaitReady after done was called is undefined
func (q *stepQueue) Start(concurrency int) (done func()) {
if concurrency < 1 {
panic("concurrency must be >= 1")
}
// l protects pending and queueItems
l := chainlock.New()
pendingCond := l.NewCond()
// priority queue
pending := &stepQueueHeap{}
// ident => queueItem
queueItems := make(map[interface{}]*stepQueueHeapItem)
// stopped is used for cancellation of "wake" goroutine
stopped := false
active := 0
func (q *stepQueue) Start() (done func()) {
go func() { // "stopper" goroutine
<-q.stop
defer l.Lock().Unlock()
stopped = true
pendingCond.Broadcast()
defer q.l.Lock().Unlock()
q.stopped = true
q.pendingCond.Broadcast()
}()
go func() { // "reqs" goroutine
for {
@@ -97,41 +123,52 @@ func (q *stepQueue) Start(concurrency int) (done func()) {
}
case req := <-q.reqs:
func() {
defer l.Lock().Unlock()
if _, ok := queueItems[req.ident]; ok {
defer q.l.Lock().Unlock()
if _, ok := q.queueItems[req.ident]; ok {
panic("WaitReady must not be called twice for the same ident")
}
qitem := &stepQueueHeapItem{
req: req,
}
queueItems[req.ident] = qitem
heap.Push(pending, qitem)
pendingCond.Broadcast()
q.queueItems[req.ident] = qitem
heap.Push(q.pending, qitem)
q.pendingCond.Broadcast()
}()
}
}
}()
go func() { // "wake" goroutine
defer l.Lock().Unlock()
defer q.l.Lock().Unlock()
for {
for !stopped && (active >= concurrency || pending.Len() == 0) {
pendingCond.Wait()
for !q.stopped && (q.active.Len() >= q.concurrency || q.pending.Len() == 0) {
q.pendingCond.Wait()
}
if stopped {
if q.stopped {
return
}
if pending.Len() <= 0 {
if q.pending.Len() <= 0 {
return
}
active++
next := heap.Pop(pending).(*stepQueueHeapItem).req
delete(queueItems, next.ident)
next.wakeup <- func() {
defer l.Lock().Unlock()
active--
pendingCond.Broadcast()
// pop from tracked items
next := heap.Pop(q.pending).(*stepQueueHeapItem)
next.req.cancelDueToConcurrencyDownsize =
heap.Push(q.active, next)
next.req.wakeup <- func() {
defer q.l.Lock().Unlock()
//
qitem := &stepQueueHeapItem{
req: req,
}
// delete(q.queueItems, next.req.ident) // def
q.pendingCond.Broadcast()
}
}
}()
@@ -161,3 +198,24 @@ func (q *stepQueue) WaitReady(ident interface{}, targetDate time.Time) StepCompl
}
return q.sendAndWaitForWakeup(ident, targetDate)
}
// caller must hold lock
func (q *stepQueue) setConcurrencyLocked(newConcurrency int) error {
if !(newConcurrency >= 1) {
return fmt.Errorf("concurrency must be >= 1 but requested %v", newConcurrency)
}
q.concurrency = newConcurrency
q.pendingCond.Broadcast() // wake up waiters who could make progress
for q.active.Len() > q.concurrency {
item := heap.Pop(q.active).(*stepQueueHeapItem)
item.req.cancelDueToConcurrencyDownsize()
heap.Push(q.pending, item)
}
return nil
}
func (q *stepQueue) SetConcurrency(new int) error {
defer q.l.Lock().Unlock()
return q.setConcurrencyLocked(new)
}
@@ -17,7 +17,7 @@ import (
// (relies on scheduler responsivity of < 500ms)
func TestPqNotconcurrent(t *testing.T) {
var ctr uint32
q := newStepQueue()
q := newStepQueue(1)
var wg sync.WaitGroup
wg.Add(4)
go func() {
@@ -29,7 +29,7 @@ func TestPqNotconcurrent(t *testing.T) {
}()
// give goroutine "1" 500ms to enter queue, get the active slot and enter time.Sleep
defer q.Start(1)()
defer q.Start()()
time.Sleep(500 * time.Millisecond)
// while "1" is still running, queue in "2", "3" and "4"
@@ -77,8 +77,9 @@ func (r record) String() string {
// Hence, perform some statistics on the wakeup times and assert that the mean wakeup
// times for each step are close together.
func TestPqConcurrent(t *testing.T) {
q := newStepQueue()
concurrency := 5
q := newStepQueue(concurrency)
var wg sync.WaitGroup
filesystems := 100
stepsPerFS := 20
@@ -104,8 +105,7 @@ func TestPqConcurrent(t *testing.T) {
records <- recs
}(fs)
}
concurrency := 5
defer q.Start(concurrency)()
defer q.Start()()
wg.Wait()
close(records)
t.Logf("loop done")
+5 -6
View File
@@ -606,7 +606,7 @@ func (s *Step) doReplication(ctx context.Context) error {
log := getLogger(ctx)
sr := s.buildSendRequest(false)
log.WithField("sr", sr.String()).Debug("initiate send request")
log.Debug("initiate send request")
sres, sstreamCopier, err := s.sender.Send(ctx, sr)
if err != nil {
log.WithError(err).Error("send request failed")
@@ -633,7 +633,7 @@ func (s *Step) doReplication(ctx context.Context) error {
To: sr.GetTo(),
ClearResumeToken: !sres.UsedResumeToken,
}
log.WithField("rr", rr.String()).Debug("initiate receive request")
log.Debug("initiate receive request")
_, err = s.receiver.Receive(ctx, rr, byteCountingStream)
if err != nil {
log.
@@ -649,11 +649,10 @@ func (s *Step) doReplication(ctx context.Context) error {
}
log.Debug("receive finished")
scr := &pdu.SendCompletedReq{
log.Debug("tell sender replication completed")
_, err = s.sender.SendCompleted(ctx, &pdu.SendCompletedReq{
OriginalReq: sr,
}
log.WithField("scr", scr.String()).Debug("tell sender replication completed")
_, err = s.sender.SendCompleted(ctx, scr)
})
if err != nil {
log.WithError(err).Error("error telling sender that replication completed successfully")
return err
+2 -2
View File
@@ -8,6 +8,6 @@ import (
"github.com/zrepl/zrepl/replication/driver"
)
func Do(ctx context.Context, planner driver.Planner) (driver.ReportFunc, driver.WaitFunc) {
return driver.Do(ctx, planner)
func Do(ctx context.Context, initialConcurrency int, planner driver.Planner) (driver.ReportFunc, driver.WaitFunc, driver.SetConcurrencyFunc) {
return driver.Do(ctx, initialConcurrency, planner)
}
+1
View File
@@ -10,6 +10,7 @@ type Report struct {
WaitReconnectSince, WaitReconnectUntil time.Time
WaitReconnectError *TimedError
Attempts []*AttemptReport
Concurrency int
}
var _, _ = json.Marshal(&Report{})
+1 -1
View File
@@ -75,7 +75,7 @@ func (c *Client) recv(ctx context.Context, conn *stream.Conn, res proto.Message)
if err != nil {
return err
}
header := string(headerBuf) // FIXME
header := string(headerBuf)
if strings.HasPrefix(header, responseHeaderHandlerErrorPrefix) {
// FIXME distinguishable error type
return &RemoteHandlerError{strings.TrimPrefix(header, responseHeaderHandlerErrorPrefix)}
+10 -23
View File
@@ -14,8 +14,8 @@ import (
"github.com/zrepl/zrepl/zfs"
)
// ReqInterceptor has a chance to exchange the context and connection on each request.
type ReqInterceptor func(ctx context.Context, rawConn *transport.AuthConn) (context.Context, *transport.AuthConn)
// WireInterceptor has a chance to exchange the context and connection on each client connection.
type WireInterceptor func(ctx context.Context, rawConn *transport.AuthConn) (context.Context, *transport.AuthConn)
// Handler implements the functionality that is exposed by Server to the Client.
type Handler interface {
@@ -30,27 +30,19 @@ type Handler interface {
PingDataconn(ctx context.Context, r *pdu.PingReq) (*pdu.PingRes, error)
}
type PreHandlerInspector = func(ctx context.Context, endpoint string, req interface{})
type PostHandlerInspector = func(ctx context.Context, response interface{}, err error)
type Logger = logger.Logger
type Server struct {
h Handler
ri ReqInterceptor
log Logger
pre PreHandlerInspector
post PostHandlerInspector
h Handler
wi WireInterceptor
log Logger
}
func NewServer(ri ReqInterceptor, logger Logger, pre PreHandlerInspector, handler Handler, post PostHandlerInspector) *Server {
func NewServer(wi WireInterceptor, logger Logger, handler Handler) *Server {
return &Server{
h: handler,
ri: ri,
wi: wi,
log: logger,
pre: pre,
post: post,
}
}
@@ -91,8 +83,8 @@ func (s *Server) serveConn(nc *transport.AuthConn) {
defer s.log.Debug("serveConn done")
ctx := context.Background()
if s.ri != nil {
ctx, nc = s.ri(ctx, nc)
if s.wi != nil {
ctx, nc = s.wi(ctx, nc)
}
c := stream.Wrap(nc, HeartbeatInterval, HeartbeatPeerTimeout)
@@ -108,7 +100,7 @@ func (s *Server) serveConn(nc *transport.AuthConn) {
s.log.WithError(err).Error("error reading structured part")
return
}
endpoint := string(header) // FIXME
endpoint := string(header)
reqStructured, err := c.ReadStreamedMessage(ctx, RequestStructuredMaxSize, ReqStructured)
if err != nil {
@@ -128,7 +120,6 @@ func (s *Server) serveConn(nc *transport.AuthConn) {
s.log.WithError(err).Error("cannot unmarshal send request")
return
}
s.pre(ctx, endpoint, &req)
res, sendStream, handlerErr = s.h.Send(ctx, &req) // SHADOWING
case EndpointRecv:
var req pdu.ReceiveReq
@@ -136,7 +127,6 @@ func (s *Server) serveConn(nc *transport.AuthConn) {
s.log.WithError(err).Error("cannot unmarshal receive request")
return
}
s.pre(ctx, endpoint, &req)
res, handlerErr = s.h.Receive(ctx, &req, &streamCopier{streamConn: c, closeStreamOnClose: false}) // SHADOWING
case EndpointPing:
var req pdu.PingReq
@@ -144,7 +134,6 @@ func (s *Server) serveConn(nc *transport.AuthConn) {
s.log.WithError(err).Error("cannot unmarshal ping request")
return
}
s.pre(ctx, endpoint, &req)
res, handlerErr = s.h.PingDataconn(ctx, &req) // SHADOWING
default:
s.log.WithField("endpoint", endpoint).Error("unknown endpoint")
@@ -153,8 +142,6 @@ func (s *Server) serveConn(nc *transport.AuthConn) {
s.log.WithField("endpoint", endpoint).WithField("errType", fmt.Sprintf("%T", handlerErr)).Debug("handler returned")
s.post(ctx, res, handlerErr)
// prepare protobuf now to return the protobuf error in the header
// if marshaling fails. We consider failed marshaling a handler error
var protobuf *bytes.Buffer
@@ -101,11 +101,7 @@ func (*transportCredentials) OverrideServerName(string) error {
type ContextInterceptor = func(ctx context.Context) context.Context
type PreHandlerInspector func(ctx context.Context, endpoint string, req interface{})
type PostHandlerInspector func(ctx context.Context, response interface{}, err error)
func NewInterceptors(logger Logger, clientIdentityKey interface{}, ctxInterceptor ContextInterceptor, pre PreHandlerInspector, post PostHandlerInspector) (unary grpc.UnaryServerInterceptor, stream grpc.StreamServerInterceptor) {
func NewInterceptors(logger Logger, clientIdentityKey interface{}, ctxInterceptor ContextInterceptor) (unary grpc.UnaryServerInterceptor, stream grpc.StreamServerInterceptor) {
unary = func(ctx context.Context, req interface{}, info *grpc.UnaryServerInfo, handler grpc.UnaryHandler) (resp interface{}, err error) {
logger.WithField("fullMethod", info.FullMethod).Debug("request")
p, ok := peer.FromContext(ctx)
@@ -122,10 +118,7 @@ func NewInterceptors(logger Logger, clientIdentityKey interface{}, ctxIntercepto
if ctxInterceptor != nil {
ctx = ctxInterceptor(ctx)
}
pre(ctx, info.FullMethod, req)
res, err := handler(ctx, req)
post(ctx, res, err)
return res, err
return handler(ctx, req)
}
stream = func(srv interface{}, ss grpc.ServerStream, info *grpc.StreamServerInfo, handler grpc.StreamHandler) error {
panic("unimplemented")
@@ -4,18 +4,14 @@ package grpchelper
import (
"context"
"fmt"
"strings"
"time"
"google.golang.org/grpc"
"google.golang.org/grpc/keepalive"
"google.golang.org/grpc/metadata"
"github.com/zrepl/zrepl/logger"
"github.com/zrepl/zrepl/rpc/grpcclientidentity"
"github.com/zrepl/zrepl/rpc/netadaptor"
"github.com/zrepl/zrepl/tracing"
"github.com/zrepl/zrepl/transport"
)
@@ -32,40 +28,15 @@ type Logger = logger.Logger
// ClientConn is an easy-to-use wrapper around the Dialer and TransportCredentials interface
// to produce a grpc.ClientConn
func ClientConn(cn transport.Connecter, log Logger, genReqID func() string) *grpc.ClientConn {
func ClientConn(cn transport.Connecter, log Logger) *grpc.ClientConn {
ka := grpc.WithKeepaliveParams(keepalive.ClientParameters{
Time: StartKeepalivesAfterInactivityDuration,
Timeout: KeepalivePeerTimeout,
PermitWithoutStream: true,
})
unaryIntcpt := grpc.WithUnaryInterceptor(func(ctx context.Context, method string, req, reply interface{}, cc *grpc.ClientConn, invoker grpc.UnaryInvoker, opts ...grpc.CallOption) error {
traceStack := tracing.GetStack(ctx)
if len(traceStack) == 0 {
panic("implementation error: expecting trace stack")
}
reqId := genReqID() + " " + strings.Join(traceStack, " <> ")
l := log
l = l.WithField("reqID", reqId)
l = l.WithField("reqT", fmt.Sprintf("%T", req))
l = l.WithField("repT", fmt.Sprintf("%T", reply))
l = l.WithField("ctxT", fmt.Sprintf("%T", ctx))
l.Debug("unary intercepted")
ctx = metadata.AppendToOutgoingContext(ctx, "zrepl-req-id", reqId)
err := invoker(ctx, method, req, reply, cc, opts...)
l = log.WithField("repT", fmt.Sprintf("%T", reply))
l = log.WithField("errT", fmt.Sprintf("%T", err))
l.Debug("reply intercepted")
return err
})
dialerOption := grpc.WithDialer(grpcclientidentity.NewDialer(log, cn))
cred := grpc.WithTransportCredentials(grpcclientidentity.NewTransportCredentials(log))
cc, err := grpc.DialContext(context.Background(), "doesn't matter done by dialer", dialerOption, cred, ka, unaryIntcpt)
cc, err := grpc.DialContext(context.Background(), "doesn't matter done by dialer", dialerOption, cred, ka)
if err != nil {
log.WithError(err).Error("cannot create gRPC client conn (non-blocking)")
// It's ok to panic here: the we call grpc.DialContext without the
@@ -79,7 +50,7 @@ func ClientConn(cn transport.Connecter, log Logger, genReqID func() string) *grp
}
// NewServer is a convenience interface around the TransportCredentials and Interceptors interface.
func NewServer(authListener transport.AuthenticatedListener, clientIdentityKey interface{}, logger grpcclientidentity.Logger, interceptor grpcclientidentity.ContextInterceptor, pre grpcclientidentity.PreHandlerInspector, post grpcclientidentity.PostHandlerInspector) (srv *grpc.Server, serve func() error) {
func NewServer(authListener transport.AuthenticatedListener, clientIdentityKey interface{}, logger grpcclientidentity.Logger, ctxInterceptor grpcclientidentity.ContextInterceptor) (srv *grpc.Server, serve func() error) {
ka := grpc.KeepaliveParams(keepalive.ServerParameters{
Time: StartKeepalivesAfterInactivityDuration,
Timeout: KeepalivePeerTimeout,
@@ -89,17 +60,8 @@ func NewServer(authListener transport.AuthenticatedListener, clientIdentityKey i
PermitWithoutStream: true,
})
tcs := grpcclientidentity.NewTransportCredentials(logger)
unary, stream := grpcclientidentity.NewInterceptors(logger, clientIdentityKey, interceptor, pre, post)
unary2 := func(ctx context.Context, req interface{}, info *grpc.UnaryServerInfo, handler grpc.UnaryHandler) (resp interface{}, err error) {
md, ok := metadata.FromIncomingContext(ctx)
if !ok {
logger.Error("no request id in incoming request")
} else {
logger.Info(fmt.Sprintf("req-id = %v", md.Get("zrepl-req-id")))
}
return unary(ctx, req, info, handler)
}
srv = grpc.NewServer(grpc.Creds(tcs), grpc.UnaryInterceptor(unary2), grpc.StreamInterceptor(stream), ka, ep)
unary, stream := grpcclientidentity.NewInterceptors(logger, clientIdentityKey, ctxInterceptor)
srv = grpc.NewServer(grpc.Creds(tcs), grpc.UnaryInterceptor(unary), grpc.StreamInterceptor(stream), ka, ep)
serve = func() error {
if err := srv.Serve(netadaptor.New(authListener, logger)); err != nil {
+3 -7
View File
@@ -26,7 +26,6 @@ import (
// Client implements the active side of a replication setup.
// It satisfies the Endpoint, Sender and Receiver interface defined by package replication.
type Client struct {
reqIDGen *requestIDGenerator
dataClient *dataconn.Client
controlClient pdu.ReplicationClient // this the grpc client instance, see constructor
controlConn *grpc.ClientConn
@@ -48,13 +47,10 @@ func NewClient(cn transport.Connecter, loggers Loggers) *Client {
muxedConnecter := mux(cn)
c := &Client{
loggers: loggers,
closed: make(chan struct{}),
reqIDGen: newRequestIDGenerator(),
loggers: loggers,
closed: make(chan struct{}),
}
grpcConn := grpchelper.ClientConn(muxedConnecter.control, loggers.Control, func() string {
return c.reqIDGen.newID().String()
})
grpcConn := grpchelper.ClientConn(muxedConnecter.control, loggers.Control)
go func() {
ctx, cancel := context.WithCancel(context.Background())
-58
View File
@@ -1,58 +0,0 @@
package rpc
import (
"context"
"encoding/base64"
"strings"
"github.com/google/uuid"
)
type requestIDGenerator struct{}
func newRequestIDGenerator() *requestIDGenerator {
return &requestIDGenerator{}
}
type RequestID struct{ s string }
func (r RequestID) String() string { return r.s }
type requestIDContextKey int
const (
requestIDContextKeyRequestID requestIDContextKey = 1 + iota
)
func (g *requestIDGenerator) newID() RequestID {
id := uuid.New()
var buf strings.Builder
enc := base64.NewEncoder(base64.RawStdEncoding, &buf)
n, err := enc.Write(id[:])
if err != nil {
panic(err)
} else if n != len(id) {
panic(n)
}
if err := enc.Close(); err != nil {
panic(err)
}
return RequestID{buf.String()}
}
func (g *requestIDGenerator) inject(ctx context.Context) context.Context {
return context.WithValue(ctx, requestIDContextKeyRequestID, g.newID())
}
func GetRequestID(ctx context.Context) (id RequestID, ok bool) {
id, ok = ctx.Value(requestIDContextKeyRequestID).(RequestID)
return id, ok
}
func MustGetRequestID(ctx context.Context) (id RequestID) {
id, ok := GetRequestID(ctx)
if !ok {
panic("calling context expectes request id to bet set")
}
return id
}
+4 -25
View File
@@ -7,7 +7,6 @@ import (
"github.com/zrepl/zrepl/endpoint"
"github.com/zrepl/zrepl/replication/logic/pdu"
"github.com/zrepl/zrepl/rpc/dataconn"
"github.com/zrepl/zrepl/rpc/grpcclientidentity"
"github.com/zrepl/zrepl/rpc/grpcclientidentity/grpchelper"
"github.com/zrepl/zrepl/rpc/versionhandshake"
"github.com/zrepl/zrepl/transport"
@@ -31,35 +30,15 @@ type Server struct {
dataServerServe serveFunc
}
type CtxInterceptor = grpcclientidentity.ContextInterceptor
func chainInterceptors(interceptors []CtxInterceptor) CtxInterceptor {
return func(ctx context.Context) (_ context.Context) {
for _, i := range interceptors {
ctx = i(ctx)
}
return ctx
}
}
type PreHandlerInspector = func(ctx context.Context, endpoint string, req interface{})
type PostHandlerInspector = func(ctx context.Context, response interface{}, err error)
type HandlerContextInterceptor func(ctx context.Context) context.Context
// config must be valid (use its Validate function).
func NewServer(handler Handler, loggers Loggers, ctxInterceptor CtxInterceptor, pre PreHandlerInspector, post PostHandlerInspector) *Server {
reqIdGen := newRequestIDGenerator()
ctxInterceptor = chainInterceptors([]CtxInterceptor{
reqIdGen.inject,
ctxInterceptor,
})
func NewServer(handler Handler, loggers Loggers, ctxInterceptor HandlerContextInterceptor) *Server {
// setup control server
controlServerServe := func(ctx context.Context, controlListener transport.AuthenticatedListener, errOut chan<- error) {
controlServer, serve := grpchelper.NewServer(controlListener, endpoint.ClientIdentityKey, loggers.Control, ctxInterceptor, pre, post)
controlServer, serve := grpchelper.NewServer(controlListener, endpoint.ClientIdentityKey, loggers.Control, ctxInterceptor)
pdu.RegisterReplicationServer(controlServer, handler)
// give time for graceful stop until deadline expires, then hard stop
@@ -84,7 +63,7 @@ func NewServer(handler Handler, loggers Loggers, ctxInterceptor CtxInterceptor,
}
return ctx, wire
}
dataServer := dataconn.NewServer(dataServerClientIdentitySetter, loggers.Data, pre, handler, post)
dataServer := dataconn.NewServer(dataServerClientIdentitySetter, loggers.Data, handler)
dataServerServe := func(ctx context.Context, dataListener transport.AuthenticatedListener, errOut chan<- error) {
dataServer.Serve(ctx, dataListener)
errOut <- nil // TODO bad design of dataServer?
-53
View File
@@ -1,53 +0,0 @@
package tracing
import "context"
type tracingContextKey int
const (
CallerContext tracingContextKey = 1 + iota
)
type jobSubtree struct {
jobid string
}
type ctx struct {
parent *ctx
job *jobSubtree
ident string
}
var root = &ctx{nil, nil, ""}
func getParentOrRoot(c context.Context) *ctx {
parent, ok := c.Value(CallerContext).(*ctx)
if !ok {
parent = root
}
return parent
}
func makeChild(c context.Context, child *ctx) context.Context {
if child.parent == nil {
panic(child)
}
return context.WithValue(c, CallerContext, child)
}
func Child(c context.Context, ident string) context.Context {
parent := getParentOrRoot(c)
return makeChild(c, &ctx{parent: parent, ident: ident})
}
func GetStack(c context.Context) (idents []string) {
ct, ok := c.Value(CallerContext).(*ctx)
if !ok {
return idents
}
for ct.parent != nil {
idents = append(idents, ct.ident)
ct = ct.parent
}
return idents
}
-21
View File
@@ -1,21 +0,0 @@
package tracing
import (
"context"
"strings"
"testing"
"github.com/stretchr/testify/assert"
)
func TestIt(t *testing.T) {
ctx := context.Background()
ctx = Child(ctx, "a")
ctx = Child(ctx, "b")
ctx = Child(ctx, "c")
ctx = Child(ctx, "d")
assert.Equal(t, "dcba", strings.Join(GetStack(ctx), ""))
}