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
3 changes: 2 additions & 1 deletion _examples/basic.go
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
package main

import (
"context"
"fmt"
"time"

Expand All @@ -12,7 +13,7 @@ func main() {
if err != nil {
panic(err)
}
ch, err := c.AddWatch("/", true)
ch, err := c.AddWatch(context.Background(), "/", true)
if err != nil {
panic(err)
}
Expand Down
40 changes: 0 additions & 40 deletions _examples/cached_leaves_walker.go

This file was deleted.

15 changes: 6 additions & 9 deletions _examples/tree_cache.go
Original file line number Diff line number Diff line change
Expand Up @@ -44,15 +44,12 @@ func main() {
}
fmt.Printf("children: %+v, stat: %+v", children, stat)

// Walk leaves of cache.
walker := cache.Walker("/bar") // Effectively: c.Walker("/foo/bar")
err = walker.LeavesOnly().
BreadthFirst().
Walk(func(path string, stat *zk.Stat) error {
fmt.Printf("path: %s, stat: %+v", path, stat)
return nil
})
if err != nil {
// Walk cache breadth-first.
nodes, walkErr := cache.Walker("/bar", zk.BreadthFirstOrder).All(ctx)
for path, stat := range nodes {
fmt.Printf("path: %s, stat: %+v", path, stat)
}
if err := walkErr(); err != nil {
panic(err)
}
}
44 changes: 16 additions & 28 deletions _examples/tree_walker.go
Original file line number Diff line number Diff line change
@@ -1,7 +1,8 @@
package main

import (
"log"
"context"
"log/slog"
"time"

"github.com/Shopify/zk"
Expand All @@ -14,42 +15,29 @@ func main() {
}
go func() {
for e := range events {
log.Printf("SessionEvent: %+v", e)
slog.Info("session event", "event", e)
}
log.Printf("SessionEvent closed")
slog.Info("session event channel closed")
}()

// Walk breath-first.
err = c.TreeWalker("/foo").
BreadthFirst().
Walk(func(p string, stat *zk.Stat) error {
log.Printf("Got %s", p)
return nil
})
if err != nil {
panic(err)
}
ctx := context.Background()

// Walk depth-first and visit leaves only.
err = c.TreeWalker("/foo").
DepthFirst().
LeavesOnly().
Walk(func(p string, stat *zk.Stat) error {
log.Printf("Got %s", p)
// Walk with callback — use when the visitor can fail or needs error propagation.
err = c.Walker("/foo", zk.BreadthFirstOrder).
Walk(ctx, func(_ context.Context, p string, stat *zk.Stat) error {
slog.Info("visited node", "path", p, "version", stat.Version)
return nil
})
if err != nil {
panic(err)
}

// Walk breath-first with parallel traversal and receive events by channel.
ch := c.TreeWalker("/foo").
BreadthFirstParallel().
WalkChan(8) // You can tune the buffer size.
for e := range ch {
if e.Err != nil {
panic(e.Err)
}
log.Printf("Got %s", e.Path)
// Walk with iterator — use for simple collection/iteration.
nodes, walkErr := c.Walker("/foo", zk.DepthFirstOrder).All(ctx)
for p, stat := range nodes {
slog.Info("visited node", "path", p, "version", stat.Version)
}
if err = walkErr(); err != nil {
panic(err)
}
}
67 changes: 25 additions & 42 deletions batch_tree_walker.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,14 +3,12 @@ package zk
import (
"context"
"errors"
"iter"
gopath "path"
)

// BatchVisitorFunc is a function that is called for each batch of nodes visited.
type BatchVisitorFunc func(paths []string) error

// BatchVisitorCtxFunc is like BatchVisitorFunc, but it takes a context.
type BatchVisitorCtxFunc func(ctx context.Context, paths []string) error
type BatchVisitorFunc func(ctx context.Context, paths []string) error

// NewBatchTreeWalker returns a new BatchTreeWalker for the given connection, root path and batch size.
func NewBatchTreeWalker(conn *Conn, path string, batchSize int) *BatchTreeWalker {
Expand All @@ -33,52 +31,37 @@ type BatchTreeWalker struct {
batchSize int
}

// Walk begins traversing the tree and calls the visitor function for each node visited.
func (w *BatchTreeWalker) Walk(visitor BatchVisitorFunc) error {
vc := func(ctx context.Context, paths []string) error {
return visitor(paths)
// All returns an iterator over all node paths in the tree and an error function.
// The caller can stop iteration early by breaking out of the range loop.
// After iteration, call the returned error function to check if the walk
// was interrupted by an error (as opposed to completing or being broken out of).
func (w *BatchTreeWalker) All(ctx context.Context) (iter.Seq[string], func() error) {
var walkErr error
seq := func(yield func(string) bool) {
walkErr = w.Walk(ctx, func(_ context.Context, paths []string) error {
for _, p := range paths {
if !yield(p) {
return errBreak
}
}
return nil
})
if errors.Is(walkErr, errBreak) {
walkErr = nil // Break is not an error.
}
}
return w.WalkCtx(context.Background(), vc)
return seq, func() error { return walkErr }
}

func (w *BatchTreeWalker) WalkCtx(ctx context.Context, visitor BatchVisitorCtxFunc) error {
// Walk traverses the tree and calls the visitor function for each batch of nodes visited.
func (w *BatchTreeWalker) Walk(ctx context.Context, visitor BatchVisitorFunc) error {
return w.walkBatch(ctx, []string{w.path}, visitor)
}

// WalkChan begins traversing the tree and sends the results to the returned channel.
// The channel will be buffered with the given size.
// The channel is closed when the traversal is complete.
// If an error occurs, an error event will be sent to the channel before it is closed.
func (w *BatchTreeWalker) WalkChan(bufferSize int) <-chan VisitEvent {
return w.WalkChanCtx(context.Background(), bufferSize)
}

// WalkChanCtx is like WalkChan, but it takes a context that can be used to cancel the walk.
func (w *BatchTreeWalker) WalkChanCtx(ctx context.Context, bufferSize int) <-chan VisitEvent {
ch := make(chan VisitEvent, bufferSize)
visitor := func(ctx context.Context, paths []string) error {
for _, p := range paths {
select {
case <-ctx.Done():
return ctx.Err()
case ch <- VisitEvent{Path: p}:
}
}
return nil
}
go func() {
defer close(ch)
if err := w.WalkCtx(ctx, visitor); err != nil {
ch <- VisitEvent{Err: err}
}
}()
return ch
}

// walkBatch recursively walks the tree in batches.
// It calls the visitor function for each batch of nodes visited.
// It fetches children in batches to reduce the number of round trips.
func (w *BatchTreeWalker) walkBatch(ctx context.Context, paths []string, visitor BatchVisitorCtxFunc) error {
func (w *BatchTreeWalker) walkBatch(ctx context.Context, paths []string, visitor BatchVisitorFunc) error {
// Execute the visitor function on all paths.
if err := visitor(ctx, paths); err != nil {
return err
Expand Down Expand Up @@ -123,7 +106,7 @@ func (w *BatchTreeWalker) fetchChildrenBatch(ctx context.Context, paths []string
requests[i] = &GetChildrenRequest{Path: p}
}

responses, err := w.conn.MultiReadCtx(ctx, requests...)
responses, err := w.conn.MultiRead(ctx, requests...)
if err != nil && !errors.Is(err, ErrNoNode) { // Treat ErrNoNode as empty children.
return nil, err
}
Expand Down
26 changes: 14 additions & 12 deletions cluster_test.go
Original file line number Diff line number Diff line change
@@ -1,7 +1,9 @@
package zk

import (
"context"
"errors"
"log/slog"
"sync"
"testing"
"time"
Expand Down Expand Up @@ -32,15 +34,15 @@ func TestBasicCluster(t *testing.T) {

time.Sleep(time.Second * 5)

if _, err := c1.Create("/gozk-test", []byte("foo-cluster"), 0, WorldACL(PermAll)); err != nil {
if _, err := c1.Create(context.Background(), "/gozk-test", []byte("foo-cluster"), 0, WorldACL(PermAll)); err != nil {
t.Fatalf("Create failed on node 1: %+v", err)
}

if _, err := c2.Sync("/gozk-test"); err != nil {
if _, err := c2.Sync(context.Background(), "/gozk-test"); err != nil {
t.Fatalf("Sync failed on node 2: %+v", err)
}

if by, _, err := c2.Get("/gozk-test"); err != nil {
if by, _, err := c2.Get(context.Background(), "/gozk-test"); err != nil {
t.Fatalf("Get failed on node 2: %+v", err)
} else if string(by) != "foo-cluster" {
t.Fatal("Wrong data for node 2")
Expand All @@ -64,7 +66,7 @@ func TestClientClusterFailover(t *testing.T) {
t.Fatalf("Failed to connect and get session")
}

if _, err := c.Create("/gozk-test", []byte("foo-cluster"), 0, WorldACL(PermAll)); err != nil {
if _, err := c.Create(context.Background(), "/gozk-test", []byte("foo-cluster"), 0, WorldACL(PermAll)); err != nil {
t.Fatalf("Create failed on node 1: %+v", err)
}

Expand All @@ -78,7 +80,7 @@ func TestClientClusterFailover(t *testing.T) {
t.Fatalf("Failover failed")
}

if by, _, err := c.Get("/gozk-test"); err != nil {
if by, _, err := c.Get(context.Background(), "/gozk-test"); err != nil {
t.Fatalf("Get failed on node 2: %+v", err)
} else if string(by) != "foo-cluster" {
t.Fatal("Wrong data for node 2")
Expand All @@ -103,10 +105,10 @@ func TestNoQuorum(t *testing.T) {
t.Fatalf("Failed to connect and get session")
}
initialSessionID := c.sessionID.Load()
DefaultLogger.Printf(" Session established: id=%d, timeout=%d", c.sessionID.Load(), c.sessionTimeoutMs)
slog.Info("session established", "id", c.sessionID.Load(), "timeout", c.sessionTimeoutMs)

// Kill the ZooKeeper leader and wait for the session to reconnect.
DefaultLogger.Printf(" Kill the leader")
slog.Info(" Kill the leader")
disconnectWatcher1 := sl.NewWatcher(sessionStateMatcher(StateDisconnected))
hasSessionWatcher2 := sl.NewWatcher(sessionStateMatcher(StateHasSession))
tc.StopServer(hasSessionEvent1.Server)
Expand All @@ -126,7 +128,7 @@ func TestNoQuorum(t *testing.T) {
}

// Kill the ZooKeeper leader leaving the cluster without quorum.
DefaultLogger.Printf(" Kill the leader")
slog.Info(" Kill the leader")
disconnectWatcher2 := sl.NewWatcher(sessionStateMatcher(StateDisconnected))
tc.StopServer(hasSessionEvent2.Server)

Expand All @@ -142,7 +144,7 @@ func TestNoQuorum(t *testing.T) {
// Make sure that we keep retrying connecting to the only remaining
// ZooKeeper server, but the attempts are being dropped because there is
// no quorum.
DefaultLogger.Printf(" Retrying no luck...")
slog.Info(" Retrying no luck...")
var firstDisconnect *Event
begin := time.Now()
for time.Since(begin) < 6*time.Second {
Expand Down Expand Up @@ -222,14 +224,14 @@ func TestBadSession(t *testing.T) {
}
defer c.Close()

if err := c.Delete("/gozk-test", -1); err != nil && !errors.Is(err, ErrNoNode) {
if err := c.Delete(context.Background(), "/gozk-test", -1); err != nil && !errors.Is(err, ErrNoNode) {
t.Fatalf("Delete returned error: %+v", err)
}

c.conn.Close()
time.Sleep(time.Millisecond * 100)

if err := c.Delete("/gozk-test", -1); err != nil && !errors.Is(err, ErrNoNode) {
if err := c.Delete(context.Background(), "/gozk-test", -1); err != nil && !errors.Is(err, ErrNoNode) {
t.Fatalf("Delete returned error: %+v", err)
}
})
Expand All @@ -253,7 +255,7 @@ func NewStateLogger(eventCh <-chan Event) *EventLogger {
sw.matchCh <- event
}
}
DefaultLogger.Printf(" event received: %v\n", event)
slog.Info("event received", "event", event)
el.events = append(el.events, event)
el.lock.Unlock()
}
Expand Down
Loading
Loading