From 15e0680d095b012672b03f585fed8f4c65ad6afe Mon Sep 17 00:00:00 2001 From: mkbeh Date: Wed, 26 Aug 2026 23:47:35 +0300 Subject: [PATCH 1/6] refactor: cluster and shard architecture --- cluster/cluster.go | 33 +- examples/shard/main.go | 14 +- examples/shard/setup.go | 39 ++- shard/errors.go | 23 +- shard/errors_test.go | 38 --- shard/foreach.go | 50 ++- shard/foreach_test.go | 356 ---------------------- shard/group.go | 20 +- shard/group_test.go | 248 --------------- shard/helpers_test.go | 73 ----- shard/resolver.go | 5 +- shard/resolver/custom.go | 23 +- shard/resolver/custom_test.go | 177 ----------- shard/resolver/encoder_test.go | 148 --------- shard/resolver/hash_test.go | 255 ---------------- shard/resolver/helpers_test.go | 59 ---- shard/resolver/range.go | 44 ++- shard/resolver/range_test.go | 275 ----------------- shard/resolver/{hash.go => rendezvous.go} | 77 ++--- shard/resolver/time_range.go | 44 +-- shard/resolver/time_range_test.go | 237 -------------- shard/resolver/validation.go | 3 +- shard/shard.go | 28 +- shard/shard_test.go | 97 ------ shard/topology.go | 102 +++---- shard/topology_test.go | 162 ---------- 26 files changed, 237 insertions(+), 2393 deletions(-) delete mode 100644 shard/errors_test.go delete mode 100644 shard/foreach_test.go delete mode 100644 shard/group_test.go delete mode 100644 shard/helpers_test.go delete mode 100644 shard/resolver/custom_test.go delete mode 100644 shard/resolver/encoder_test.go delete mode 100644 shard/resolver/hash_test.go delete mode 100644 shard/resolver/helpers_test.go delete mode 100644 shard/resolver/range_test.go rename shard/resolver/{hash.go => rendezvous.go} (59%) delete mode 100644 shard/resolver/time_range_test.go delete mode 100644 shard/shard_test.go delete mode 100644 shard/topology_test.go diff --git a/cluster/cluster.go b/cluster/cluster.go index dbe925d..071622e 100644 --- a/cluster/cluster.go +++ b/cluster/cluster.go @@ -15,9 +15,9 @@ type ID string // Config configures a Cluster from independently created pools. // -// ID and Labels are optional metadata. New takes ownership of Primary and -// Replicas only after it returns successfully. Cluster.Close closes all owned -// pools. +// ID is required. Labels are optional metadata. New takes ownership of Primary +// and Replicas only after it returns successfully. Cluster.Close closes the +// owned pools. type Config struct { ID ID Labels map[string]string @@ -27,8 +27,8 @@ type Config struct { Selector ReplicaSelector } -// Cluster represents a logical PostgreSQL cluster composed of an optional -// primary pool and zero or more replica pools. +// Cluster routes operations across an optional primary pool and zero or more +// replica pools. // // A deployment with one PostgreSQL endpoint is represented by a Cluster with // one Primary and no Replicas. A read-only deployment may omit Primary and @@ -52,9 +52,13 @@ type Cluster struct { // New creates a Cluster from independently configured pools. // -// At least one pool is required. When Selector is nil, replicas are selected -// using round-robin. +// ID and at least one pool are required. When Selector is nil, replicas are +// selected using round-robin. func New(config Config) (*Cluster, error) { + if config.ID == "" { + return nil, errors.New("xpg/cluster: cluster ID must not be empty") + } + if config.Primary != nil && config.Primary.Raw() == nil { return nil, errors.New("xpg/cluster: primary pool is invalid") } @@ -96,7 +100,7 @@ func New(config Config) (*Cluster, error) { }, nil } -// ID returns the logical cluster ID. +// ID returns the stable logical cluster ID. func (c *Cluster) ID() ID { if c == nil { return "" @@ -149,21 +153,24 @@ func (c *Cluster) ReplicaCount() int { // ReplicaAt returns the replica at index in registration order. // // The returned pool is borrowed and must not be closed separately. ReplicaAt -// panics when c is nil or index is out of range. +// panics when c is nil or index is outside the replica set, matching ordinary +// slice indexing semantics. func (c *Cluster) ReplicaAt(index int) *xpg.Pool { return c.replicas[index] } -// Close closes replicas in reverse registration order, then closes the primary -// when one is configured. Close is safe to call multiple times. +// Close closes all replica pools in reverse registration order and then closes +// the primary pool when one is configured. +// +// Close is safe to call multiple times. func (c *Cluster) Close() { if c == nil { return } c.closeOnce.Do(func() { - for _, replica := range slices.Backward(c.replicas) { - replica.Close() + for _, v := range slices.Backward(c.replicas) { + v.Close() } if c.primary != nil { diff --git a/examples/shard/main.go b/examples/shard/main.go index 9f45891..b2c000a 100644 --- a/examples/shard/main.go +++ b/examples/shard/main.go @@ -33,6 +33,9 @@ func run(ctx context.Context) error { } defer topology.Close() + // A resolver describes how one dataset is distributed across the topology. + // NewRange resolves shard IDs once and stores the resulting Shard handles in + // an immutable routing table used by subsequent lookups. userResolver, err := resolver.NewRange( topology, []resolver.Range[uint64]{ @@ -60,14 +63,17 @@ func run(ctx context.Context) error { fmt.Println("range routing:") for _, current := range users { - resolved, err := userResolver.Resolve(current.ID) + targetShard, err := userResolver.Resolve(current.ID) if err != nil { return fmt.Errorf("resolve user %d: %w", current.ID, err) } - primary := resolved.Primary() + primary := targetShard.Primary() if primary == nil { - return fmt.Errorf("shard %q has no primary", resolved.ID()) + return fmt.Errorf( + "shard %q has no primary", + targetShard.ID(), + ) } if _, err := primary.Exec( @@ -85,7 +91,7 @@ func run(ctx context.Context) error { fmt.Printf( "- user_id=%d shard=%s pool=%s\n", current.ID, - resolved.ID(), + targetShard.ID(), primary.Name(), ) } diff --git a/examples/shard/setup.go b/examples/shard/setup.go index 2a43d9d..a50b049 100644 --- a/examples/shard/setup.go +++ b/examples/shard/setup.go @@ -14,12 +14,12 @@ const ( defaultShardADatabaseURL = "postgres://postgres:postgres@localhost:56431/postgres?sslmode=disable" defaultShardBDatabaseURL = "postgres://postgres:postgres@localhost:56432/postgres?sslmode=disable" - shardAID shard.ID = "shard-a" - shardBID shard.ID = "shard-b" + shardAID cluster.ID = "shard-a" + shardBID cluster.ID = "shard-b" ) func openTopology(ctx context.Context) (*shard.Topology, error) { - shardA, err := openShard( + clusterA, err := openCluster( ctx, shardAID, "shard.shard-a.primary", @@ -29,10 +29,10 @@ func openTopology(ctx context.Context) (*shard.Topology, error) { ), ) if err != nil { - return nil, fmt.Errorf("open shard-a: %w", err) + return nil, fmt.Errorf("open shard-a cluster: %w", err) } - shardB, err := openShard( + clusterB, err := openCluster( ctx, shardBID, "shard.shard-b.primary", @@ -42,18 +42,20 @@ func openTopology(ctx context.Context) (*shard.Topology, error) { ), ) if err != nil { - shardA.Close() + clusterA.Close() - return nil, fmt.Errorf("open shard-b: %w", err) + return nil, fmt.Errorf("open shard-b cluster: %w", err) } - topology, err := shard.NewTopology([]shard.Config{ - {Cluster: shardA}, - {Cluster: shardB}, - }) + // After NewTopology succeeds, the topology owns both clusters and closes + // them through Topology.Close. On constructor failure ownership remains here. + topology, err := shard.NewTopology( + clusterA, + clusterB, + ) if err != nil { - shardB.Close() - shardA.Close() + clusterB.Close() + clusterA.Close() return nil, fmt.Errorf("create topology: %w", err) } @@ -61,7 +63,12 @@ func openTopology(ctx context.Context) (*shard.Topology, error) { return topology, nil } -func openShard(ctx context.Context, id shard.ID, name, databaseURL string) (*cluster.Cluster, error) { +func openCluster( + ctx context.Context, + id cluster.ID, + name string, + databaseURL string, +) (*cluster.Cluster, error) { pool, err := xpg.Open( ctx, databaseURL, @@ -79,7 +86,7 @@ func openShard(ctx context.Context, id shard.ID, name, databaseURL string) (*clu return nil, fmt.Errorf("ping pool: %w", err) } - shardCluster, err := cluster.New(cluster.Config{ + dbCluster, err := cluster.New(cluster.Config{ ID: id, Primary: pool, }) @@ -89,7 +96,7 @@ func openShard(ctx context.Context, id shard.ID, name, databaseURL string) (*clu return nil, fmt.Errorf("create cluster: %w", err) } - return shardCluster, nil + return dbCluster, nil } func environment(name, fallback string) string { diff --git a/shard/errors.go b/shard/errors.go index 908af96..a6ae31d 100644 --- a/shard/errors.go +++ b/shard/errors.go @@ -3,25 +3,26 @@ package shard import ( "errors" "fmt" + + "github.com/mkbeh/xpg/cluster" ) var ( - // ErrNoShard is returned when a resolver cannot map a key to any shard. + // ErrNoShard indicates that a resolver could not map a key to any shard. ErrNoShard = errors.New("xpg/shard: no shard resolved") - // ErrUnknownShard is returned when routing references a shard that does not - // exist in the topology. + // ErrUnknownShard indicates that routing configuration or custom routing + // logic referenced a shard that does not exist in the topology. ErrUnknownShard = errors.New("xpg/shard: unknown shard") - // ErrShardMismatch is returned when keys expected to be colocated resolve to + // ErrShardMismatch indicates that keys expected to be colocated resolved to // different shards. ErrShardMismatch = errors.New("xpg/shard: keys resolve to different shards") ) -// UnknownShardError identifies a shard referenced by routing that does not -// exist in the topology. +// UnknownShardError identifies a shard that does not exist in a topology. type UnknownShardError struct { - ShardID ID + ShardID cluster.ID } func (e *UnknownShardError) Error() string { @@ -32,11 +33,11 @@ func (e *UnknownShardError) Unwrap() error { return ErrUnknownShard } -// MismatchError describes the first key whose resolved shard differs from the -// shard of the first key. +// MismatchError describes the first key that resolved to a different shard +// than the first key. type MismatchError struct { - Expected ID - Actual ID + Expected cluster.ID + Actual cluster.ID Index int } diff --git a/shard/errors_test.go b/shard/errors_test.go deleted file mode 100644 index 5523408..0000000 --- a/shard/errors_test.go +++ /dev/null @@ -1,38 +0,0 @@ -package shard - -import ( - "errors" - "testing" -) - -func TestUnknownShardError(t *testing.T) { - t.Parallel() - - err := &UnknownShardError{ShardID: "missing"} - - if got, want := err.Error(), `xpg/shard: unknown shard "missing"`; got != want { - t.Fatalf("Error() = %q, want %q", got, want) - } - - if !errors.Is(err, ErrUnknownShard) { - t.Fatal("errors.Is() = false, want ErrUnknownShard") - } -} - -func TestMismatchError(t *testing.T) { - t.Parallel() - - err := &MismatchError{ - Expected: "shard-a", - Actual: "shard-b", - Index: 2, - } - - if got, want := err.Error(), `xpg/shard: key 2 resolved to shard "shard-b" instead of "shard-a"`; got != want { - t.Fatalf("Error() = %q, want %q", got, want) - } - - if !errors.Is(err, ErrShardMismatch) { - t.Fatal("errors.Is() = false, want ErrShardMismatch") - } -} diff --git a/shard/foreach.go b/shard/foreach.go index 9effadc..9eef711 100644 --- a/shard/foreach.go +++ b/shard/foreach.go @@ -5,11 +5,13 @@ import ( "errors" "fmt" "sync" + + "github.com/mkbeh/xpg/cluster" ) -// ForEachShardResult contains the result associated with one shard. +// ForEachShardResult contains the result of one shard callback invocation. type ForEachShardResult struct { - ShardID ID + ShardID cluster.ID Err error } @@ -28,7 +30,7 @@ func (results ForEachShardResults) Err() error { errs = append( errs, fmt.Errorf( - "xpg/shard: shard %q: %w", + "xpg/shard: shard %q callback: %w", result.ShardID, result.Err, ), @@ -38,29 +40,20 @@ func (results ForEachShardResults) Err() error { return errors.Join(errs...) } -// ForEachShard invokes fn across the topology with at most concurrency -// callbacks running at once. Results are returned in topology registration -// order; callback execution order is not guaranteed. -// -// Callback failures and context cancellation are stored in the corresponding -// results and can be joined with ForEachShardResults.Err. The returned error is -// reserved for invalid invocation arguments. +// ForEachShard invokes fn for each shard with at most concurrency callbacks +// running at once. Results are returned in topology registration order. The +// returned error joins all per-shard failures and is equivalent to results.Err(). // // Once context cancellation is observed, callbacks that have not started are // skipped and their results contain ctx.Err(). Callbacks already running are -// responsible for observing ctx. ForEachShard waits for all started callbacks -// to finish before returning. +// responsible for observing ctx. func (t *Topology) ForEachShard( ctx context.Context, concurrency int, fn func(context.Context, Shard) error, ) (ForEachShardResults, error) { - if t == nil { - return nil, errors.New("xpg/shard: topology is nil") - } - - if len(t.shards) == 0 { - return nil, errors.New("xpg/shard: topology is empty") + if t == nil || len(t.shards) == 0 { + return nil, errors.New("xpg/shard: topology is nil or empty") } if concurrency <= 0 { @@ -72,8 +65,16 @@ func (t *Topology) ForEachShard( } results := make(ForEachShardResults, len(t.shards)) - for index, shard := range t.shards { - results[index].ShardID = shard.ID() + for index, current := range t.shards { + results[index].ShardID = current.ID() + } + + if err := ctx.Err(); err != nil { + for index := range results { + results[index].Err = err + } + + return results, results.Err() } workerCount := min(concurrency, len(t.shards)) @@ -87,8 +88,6 @@ func (t *Topology) ForEachShard( defer workers.Done() for index := range jobs { - // An index may have been scheduled immediately before context - // cancellation. Skip callbacks that have not started yet. if err := ctx.Err(); err != nil { results[index].Err = err continue @@ -102,9 +101,6 @@ func (t *Topology) ForEachShard( nextIndex := 0 for nextIndex < len(t.shards) && ctx.Err() == nil { - // The explicit context check above prevents scheduling new work after - // cancellation has already been observed. The select still handles - // cancellation that happens while waiting for a worker. select { case jobs <- nextIndex: nextIndex++ @@ -115,13 +111,11 @@ func (t *Topology) ForEachShard( close(jobs) workers.Wait() - // Workers own results for scheduled indexes [0, nextIndex). After all - // workers finish, remaining indexes can be marked canceled without races. if err := ctx.Err(); err != nil { for index := nextIndex; index < len(results); index++ { results[index].Err = err } } - return results, nil + return results, results.Err() } diff --git a/shard/foreach_test.go b/shard/foreach_test.go deleted file mode 100644 index ceb8060..0000000 --- a/shard/foreach_test.go +++ /dev/null @@ -1,356 +0,0 @@ -package shard - -import ( - "context" - "errors" - "sync/atomic" - "testing" - "time" -) - -func TestForEachShardValidatesArguments(t *testing.T) { - t.Parallel() - - topology := newTestTopology(t, "shard-a") - - tests := []struct { - name string - topology *Topology - concurrency int - fn func(context.Context, Shard) error - wantError string - }{ - { - name: "nil topology", - topology: nil, - concurrency: 1, - fn: func(context.Context, Shard) error { return nil }, - wantError: "xpg/shard: topology is nil", - }, - { - name: "empty topology", - topology: &Topology{}, - concurrency: 1, - fn: func(context.Context, Shard) error { return nil }, - wantError: "xpg/shard: topology is empty", - }, - { - name: "zero concurrency", - topology: topology, - concurrency: 0, - fn: func(context.Context, Shard) error { return nil }, - wantError: "xpg/shard: concurrency must be positive", - }, - { - name: "nil callback", - topology: topology, - concurrency: 1, - fn: nil, - wantError: "xpg/shard: callback is nil", - }, - } - - for _, test := range tests { - t.Run(test.name, func(t *testing.T) { - _, err := test.topology.ForEachShard( - t.Context(), - test.concurrency, - test.fn, - ) - if err == nil { - t.Fatal("expected error") - } - - if got := err.Error(); got != test.wantError { - t.Fatalf("error = %q, want %q", got, test.wantError) - } - }) - } -} - -func TestForEachShardPreservesRegistrationOrder(t *testing.T) { - t.Parallel() - - topology := newTestTopology(t, "shard-c", "shard-a", "shard-b") - - results, err := topology.ForEachShard( - t.Context(), - 2, - func(context.Context, Shard) error { return nil }, - ) - if err != nil { - t.Fatalf("ForEachShard() error = %v", err) - } - - want := []ID{"shard-c", "shard-a", "shard-b"} - for index, wantID := range want { - if got := results[index].ShardID; got != wantID { - t.Fatalf("results[%d].ShardID = %q, want %q", index, got, wantID) - } - - if results[index].Err != nil { - t.Fatalf("results[%d].Err = %v", index, results[index].Err) - } - } -} - -func TestForEachShardHonorsConcurrencyLimit(t *testing.T) { - t.Parallel() - - topology := newTestTopology( - t, - "shard-a", - "shard-b", - "shard-c", - "shard-d", - "shard-e", - "shard-f", - ) - - ctx, cancel := context.WithTimeout(t.Context(), 5*time.Second) - defer cancel() - - started := make(chan struct{}, topology.Len()) - release := make(chan struct{}) - - var active atomic.Int32 - var maximum atomic.Int32 - var calls atomic.Int32 - - done := make(chan struct { - results ForEachShardResults - err error - }, 1) - - go func() { - results, err := topology.ForEachShard( - ctx, - 2, - func(ctx context.Context, _ Shard) error { - current := active.Add(1) - defer active.Add(-1) - - calls.Add(1) - - for { - observed := maximum.Load() - if current <= observed || maximum.CompareAndSwap(observed, current) { - break - } - } - - started <- struct{}{} - - select { - case <-release: - return nil - case <-ctx.Done(): - return ctx.Err() - } - }, - ) - - done <- struct { - results ForEachShardResults - err error - }{ - results: results, - err: err, - } - }() - - for range 2 { - select { - case <-started: - case <-ctx.Done(): - close(release) - t.Fatal("two callbacks did not start concurrently") - } - } - - close(release) - - var outcome struct { - results ForEachShardResults - err error - } - - select { - case outcome = <-done: - case <-ctx.Done(): - t.Fatal("ForEachShard() did not finish") - } - - if outcome.err != nil { - t.Fatalf("ForEachShard() error = %v", outcome.err) - } - - if got, want := calls.Load(), int32(topology.Len()); got != want { - t.Fatalf("callback calls = %d, want %d", got, want) - } - - if got, want := maximum.Load(), int32(2); got != want { - t.Fatalf("maximum concurrent callbacks = %d, want %d", got, want) - } - - if err := outcome.results.Err(); err != nil { - t.Fatalf("results.Err() = %v", err) - } -} - -func TestForEachShardCallbackErrorsDoNotStopOtherShards(t *testing.T) { - t.Parallel() - - topology := newTestTopology(t, "shard-a", "shard-b", "shard-c") - sentinel := errors.New("callback failed") - var calls atomic.Int32 - - results, err := topology.ForEachShard( - t.Context(), - 2, - func(_ context.Context, current Shard) error { - calls.Add(1) - if current.ID() == "shard-b" { - return sentinel - } - - return nil - }, - ) - if err != nil { - t.Fatalf("ForEachShard() error = %v", err) - } - - if got, want := calls.Load(), int32(3); got != want { - t.Fatalf("callback calls = %d, want %d", got, want) - } - - if results[0].Err != nil || !errors.Is(results[1].Err, sentinel) || results[2].Err != nil { - t.Fatalf("results = %+v", results) - } - - joined := results.Err() - if !errors.Is(joined, sentinel) { - t.Fatalf("results.Err() = %v, want wrapped sentinel", joined) - } - - if got, want := joined.Error(), - `xpg/shard: shard "shard-b": callback failed`; got != want { - t.Fatalf("results.Err() = %q, want %q", got, want) - } -} - -func TestForEachShardCanceledBeforeScheduling(t *testing.T) { - t.Parallel() - - topology := newTestTopology(t, "shard-a", "shard-b", "shard-c") - ctx, cancel := context.WithCancel(t.Context()) - cancel() - - var calls atomic.Int32 - - results, err := topology.ForEachShard( - ctx, - 2, - func(context.Context, Shard) error { - calls.Add(1) - return nil - }, - ) - if err != nil { - t.Fatalf("ForEachShard() error = %v", err) - } - - if got := calls.Load(); got != 0 { - t.Fatalf("callback calls = %d, want 0", got) - } - - for index, result := range results { - if !errors.Is(result.Err, context.Canceled) { - t.Fatalf("results[%d].Err = %v, want context.Canceled", index, result.Err) - } - } -} - -func TestForEachShardCancellationSkipsCallbacksNotStarted(t *testing.T) { - t.Parallel() - - topology := newTestTopology(t, "shard-a", "shard-b", "shard-c", "shard-d") - ctx, cancel := context.WithCancel(t.Context()) - defer cancel() - - started := make(chan struct{}) - var calls atomic.Int32 - - done := make(chan struct { - results ForEachShardResults - err error - }, 1) - - go func() { - results, err := topology.ForEachShard( - ctx, - 1, - func(ctx context.Context, _ Shard) error { - if calls.Add(1) == 1 { - close(started) - } - - <-ctx.Done() - return ctx.Err() - }, - ) - - done <- struct { - results ForEachShardResults - err error - }{results: results, err: err} - }() - - <-started - cancel() - - outcome := <-done - if outcome.err != nil { - t.Fatalf("ForEachShard() error = %v", outcome.err) - } - - if got := calls.Load(); got != 1 { - t.Fatalf("callback calls = %d, want 1", got) - } - - for index, result := range outcome.results { - if !errors.Is(result.Err, context.Canceled) { - t.Fatalf("results[%d].Err = %v, want context.Canceled", index, result.Err) - } - } -} - -func TestForEachShardResultsErr(t *testing.T) { - t.Parallel() - - first := errors.New("first") - second := errors.New("second") - - results := ForEachShardResults{ - {ShardID: "shard-a", Err: first}, - {ShardID: "shard-b"}, - {ShardID: "shard-c", Err: second}, - } - - err := results.Err() - if !errors.Is(err, first) || !errors.Is(err, second) { - t.Fatalf("Err() = %v, want both failures", err) - } - - want := "xpg/shard: shard \"shard-a\": first\n" + - "xpg/shard: shard \"shard-c\": second" - - if got := err.Error(); got != want { - t.Fatalf("Err() = %q, want %q", got, want) - } - - if err := (ForEachShardResults{{ShardID: "shard-a"}}).Err(); err != nil { - t.Fatalf("Err() = %v, want nil", err) - } -} diff --git a/shard/group.go b/shard/group.go index 5d04c0d..14796c4 100644 --- a/shard/group.go +++ b/shard/group.go @@ -3,12 +3,12 @@ package shard import ( "errors" "fmt" + + "github.com/mkbeh/xpg/cluster" ) -// SameShard resolves keys and verifies that they all belong to the same shard. -// -// It returns ErrNoShard when no keys are provided and MismatchError when a key -// resolves to a different shard. +// SameShard resolves the keys and verifies that they all belong to the same +// shard. It returns that shard when all keys are colocated. func SameShard[K any](resolver Resolver[K], keys ...K) (Shard, error) { if resolver == nil { return Shard{}, errors.New("xpg/shard: resolver is nil") @@ -46,24 +46,22 @@ func SameShard[K any](resolver Resolver[K], keys ...K) (Shard, error) { return expected, nil } -// Group contains keys that resolve to the same shard. -// Keys preserve their original relative order. +// Group contains input keys that resolve to one shard. Keys preserve their +// original relative order. type Group[K any] struct { Shard Shard Keys []K } -// GroupByShard resolves each key once and groups keys by shard. -// -// Groups are returned in order of each shard's first appearance in keys. -// Keys within each group preserve their original relative order. +// GroupByShard resolves every key once and returns groups in order of each +// shard's first appearance in the input. func GroupByShard[K any](resolver Resolver[K], keys []K) ([]Group[K], error) { if resolver == nil { return nil, errors.New("xpg/shard: resolver is nil") } groups := make([]Group[K], 0) - indexByID := make(map[ID]int) + indexByID := make(map[cluster.ID]int) for keyIndex, key := range keys { resolved, err := resolver.Resolve(key) diff --git a/shard/group_test.go b/shard/group_test.go deleted file mode 100644 index 5d548f5..0000000 --- a/shard/group_test.go +++ /dev/null @@ -1,248 +0,0 @@ -package shard - -import ( - "errors" - "slices" - "testing" -) - -func TestSameShard(t *testing.T) { - t.Parallel() - - topology := newTestTopology(t, "shard-a", "shard-b") - shardA := topology.At(0) - shardB := topology.At(1) - - resolver := testResolverFunc[int](func(key int) (Shard, error) { - if key < 100 { - return shardA, nil - } - - return shardB, nil - }) - - resolved, err := SameShard(resolver, 1, 2, 3) - if err != nil { - t.Fatalf("SameShard() error = %v", err) - } - - if got, want := resolved.ID(), ID("shard-a"); got != want { - t.Fatalf("SameShard().ID() = %q, want %q", got, want) - } -} - -func TestSameShardRejectsNilResolver(t *testing.T) { - t.Parallel() - - _, err := SameShard[int](nil, 1) - if err == nil { - t.Fatal("expected error") - } - - if got, want := err.Error(), "xpg/shard: resolver is nil"; got != want { - t.Fatalf("error = %q, want %q", got, want) - } -} - -func TestSameShardRequiresKey(t *testing.T) { - t.Parallel() - - resolver := testResolverFunc[int](func(int) (Shard, error) { - t.Fatal("resolver should not be called") - return Shard{}, nil - }) - - _, err := SameShard(resolver) - if !errors.Is(err, ErrNoShard) { - t.Fatalf("error = %v, want ErrNoShard", err) - } -} - -func TestSameShardWrapsFirstResolveError(t *testing.T) { - t.Parallel() - - sentinel := errors.New("resolve failed") - resolver := testResolverFunc[int](func(int) (Shard, error) { - return Shard{}, sentinel - }) - - _, err := SameShard(resolver, 1) - if !errors.Is(err, sentinel) { - t.Fatalf("error = %v, want wrapped sentinel", err) - } - - if got, want := err.Error(), "xpg/shard: resolve key 0: resolve failed"; got != want { - t.Fatalf("error = %q, want %q", got, want) - } -} - -func TestSameShardWrapsResolveErrorWithIndex(t *testing.T) { - t.Parallel() - - sentinel := errors.New("resolve failed") - resolver := testResolverFunc[int](func(key int) (Shard, error) { - if key == 2 { - return Shard{}, sentinel - } - - return Shard{}, nil - }) - - _, err := SameShard(resolver, 1, 2) - if !errors.Is(err, sentinel) { - t.Fatalf("error = %v, want wrapped sentinel", err) - } - - if got, want := err.Error(), "xpg/shard: resolve key 1: resolve failed"; got != want { - t.Fatalf("error = %q, want %q", got, want) - } -} - -func TestSameShardReturnsMismatchDetails(t *testing.T) { - t.Parallel() - - topology := newTestTopology(t, "shard-a", "shard-b") - shardA := topology.At(0) - shardB := topology.At(1) - - resolver := testResolverFunc[int](func(key int) (Shard, error) { - if key == 3 { - return shardB, nil - } - - return shardA, nil - }) - - _, err := SameShard(resolver, 1, 2, 3) - if !errors.Is(err, ErrShardMismatch) { - t.Fatalf("error = %v, want ErrShardMismatch", err) - } - - var mismatch *MismatchError - if !errors.As(err, &mismatch) { - t.Fatalf("error = %T, want *MismatchError", err) - } - - if mismatch.Expected != "shard-a" || mismatch.Actual != "shard-b" || mismatch.Index != 2 { - t.Fatalf("mismatch = %+v", mismatch) - } -} - -func TestGroupByShardPreservesGroupAndKeyOrder(t *testing.T) { - t.Parallel() - - topology := newTestTopology(t, "shard-a", "shard-b") - shardA := topology.At(0) - shardB := topology.At(1) - - resolver := testResolverFunc[int](func(key int) (Shard, error) { - if key < 100 { - return shardA, nil - } - - return shardB, nil - }) - - groups, err := GroupByShard(resolver, []int{142, 42, 143, 43}) - if err != nil { - t.Fatalf("GroupByShard() error = %v", err) - } - - if got, want := len(groups), 2; got != want { - t.Fatalf("len(groups) = %d, want %d", got, want) - } - - if got, want := groups[0].Shard.ID(), ID("shard-b"); got != want { - t.Fatalf("groups[0].Shard.ID() = %q, want %q", got, want) - } - if got, want := groups[0].Keys, []int{142, 143}; !slices.Equal(got, want) { - t.Fatalf("groups[0].Keys = %v, want %v", got, want) - } - - if got, want := groups[1].Shard.ID(), ID("shard-a"); got != want { - t.Fatalf("groups[1].Shard.ID() = %q, want %q", got, want) - } - if got, want := groups[1].Keys, []int{42, 43}; !slices.Equal(got, want) { - t.Fatalf("groups[1].Keys = %v, want %v", got, want) - } -} - -func TestGroupByShardResolvesEachKeyOnce(t *testing.T) { - t.Parallel() - - topology := newTestTopology(t, "shard-a") - resolved := topology.At(0) - calls := 0 - - resolver := testResolverFunc[int](func(int) (Shard, error) { - calls++ - return resolved, nil - }) - - keys := []int{1, 2, 3, 4} - groups, err := GroupByShard(resolver, keys) - if err != nil { - t.Fatalf("GroupByShard() error = %v", err) - } - - if got, want := calls, len(keys); got != want { - t.Fatalf("resolve calls = %d, want %d", got, want) - } - - if len(groups) != 1 { - t.Fatalf("len(groups) = %d, want 1", len(groups)) - } -} - -func TestGroupByShardRejectsNilResolver(t *testing.T) { - t.Parallel() - - _, err := GroupByShard[int](nil, []int{1}) - if err == nil { - t.Fatal("expected error") - } - - if got, want := err.Error(), "xpg/shard: resolver is nil"; got != want { - t.Fatalf("error = %q, want %q", got, want) - } -} - -func TestGroupByShardEmptyKeys(t *testing.T) { - t.Parallel() - - resolver := testResolverFunc[int](func(int) (Shard, error) { - t.Fatal("resolver should not be called") - return Shard{}, nil - }) - - groups, err := GroupByShard(resolver, nil) - if err != nil { - t.Fatalf("GroupByShard() error = %v", err) - } - - if len(groups) != 0 { - t.Fatalf("len(groups) = %d, want 0", len(groups)) - } -} - -func TestGroupByShardWrapsResolveErrorWithIndex(t *testing.T) { - t.Parallel() - - sentinel := errors.New("resolve failed") - resolver := testResolverFunc[int](func(key int) (Shard, error) { - if key == 3 { - return Shard{}, sentinel - } - - return Shard{}, nil - }) - - _, err := GroupByShard(resolver, []int{1, 2, 3}) - if !errors.Is(err, sentinel) { - t.Fatalf("error = %v, want wrapped sentinel", err) - } - - if got, want := err.Error(), "xpg/shard: resolve key 2: resolve failed"; got != want { - t.Fatalf("error = %q, want %q", got, want) - } -} diff --git a/shard/helpers_test.go b/shard/helpers_test.go deleted file mode 100644 index 17aab99..0000000 --- a/shard/helpers_test.go +++ /dev/null @@ -1,73 +0,0 @@ -package shard - -import ( - "testing" - - "github.com/jackc/pgx/v5/pgxpool" - "github.com/mkbeh/xpg" - "github.com/mkbeh/xpg/cluster" -) - -const testDatabaseURL = "postgres://postgres@127.0.0.1:1/postgres?sslmode=disable" - -func newTestCluster(t *testing.T, id ID, labels map[string]string) *cluster.Cluster { - t.Helper() - - config, err := pgxpool.ParseConfig(testDatabaseURL) - if err != nil { - t.Fatalf("pgxpool.ParseConfig() error = %v", err) - } - - config.MinConns = 0 - config.MaxConns = 1 - - pool, err := xpg.New( - t.Context(), - config, - xpg.WithName("shard."+string(id)+".primary"), - ) - if err != nil { - t.Fatalf("xpg.New() error = %v", err) - } - - shardCluster, err := cluster.New(cluster.Config{ - ID: id, - Labels: labels, - Primary: pool, - }) - if err != nil { - pool.Close() - t.Fatalf("cluster.New() error = %v", err) - } - - t.Cleanup(shardCluster.Close) - - return shardCluster -} - -func newTestTopology(t *testing.T, ids ...ID) *Topology { - t.Helper() - - configs := make([]Config, len(ids)) - - for index, id := range ids { - configs[index] = Config{ - Cluster: newTestCluster(t, id, nil), - } - } - - topology, err := NewTopology(configs) - if err != nil { - t.Fatalf("NewTopology() error = %v", err) - } - - t.Cleanup(topology.Close) - - return topology -} - -type testResolverFunc[K any] func(K) (Shard, error) - -func (resolve testResolverFunc[K]) Resolve(key K) (Shard, error) { - return resolve(key) -} diff --git a/shard/resolver.go b/shard/resolver.go index 500fcb8..800cd39 100644 --- a/shard/resolver.go +++ b/shard/resolver.go @@ -1,9 +1,10 @@ package shard -// Resolver maps a typed application key to one shard. +// Resolver maps a typed application key to the shard owning that key. // // Resolve should return ErrNoShard when the key cannot be mapped to a shard. -// Implementations must be safe for concurrent use. +// When Resolve returns nil error, it must return a valid Shard. Implementations +// shared by concurrent callers must be concurrency-safe. type Resolver[K any] interface { Resolve(key K) (Shard, error) } diff --git a/shard/resolver/custom.go b/shard/resolver/custom.go index cae2096..0e47de9 100644 --- a/shard/resolver/custom.go +++ b/shard/resolver/custom.go @@ -3,25 +3,24 @@ package resolver import ( "errors" + "github.com/mkbeh/xpg/cluster" "github.com/mkbeh/xpg/shard" ) -// ResolveFunc maps an application key to a shard ID within topology. +// ResolveFunc maps an application key to a shard ID. // -// Resolve functions must be deterministic and safe for concurrent use. They -// should return shard.ErrNoShard when a key cannot be mapped and should not -// perform hidden I/O or modify topology. -type ResolveFunc[K any] func(key K, topology *shard.Topology) (shard.ID, error) +// Resolve functions should return shard.ErrNoShard when a key cannot be mapped +// to a shard. Implementations shared by concurrent callers must be deterministic +// and concurrency-safe. They should not perform hidden I/O. +type ResolveFunc[K any] func(key K) (cluster.ID, error) // CustomResolver adapts ResolveFunc to shard.Resolver. -// -// CustomResolver borrows its topology and must not outlive it. type CustomResolver[K any] struct { topology *shard.Topology resolve ResolveFunc[K] } -// NewCustom binds custom routing logic to an immutable topology. +// NewCustom binds custom routing logic to one immutable topology. func NewCustom[K any](topology *shard.Topology, resolve ResolveFunc[K]) (*CustomResolver[K], error) { if err := requireTopology(topology); err != nil { return nil, err @@ -37,22 +36,20 @@ func NewCustom[K any](topology *shard.Topology, resolve ResolveFunc[K]) (*Custom }, nil } -// Resolve maps key to a shard in the bound topology. +// Resolve maps key to a shard and rejects IDs absent from the bound topology. func (resolver *CustomResolver[K]) Resolve(key K) (shard.Shard, error) { if resolver == nil || resolver.topology == nil || resolver.resolve == nil { return shard.Shard{}, errors.New("xpg/shard/resolver: custom resolver is not initialized") } - id, err := resolver.resolve(key, resolver.topology) + id, err := resolver.resolve(key) if err != nil { return shard.Shard{}, err } resolved, ok := resolver.topology.Shard(id) if !ok { - return shard.Shard{}, &shard.UnknownShardError{ - ShardID: id, - } + return shard.Shard{}, &shard.UnknownShardError{ShardID: id} } return resolved, nil diff --git a/shard/resolver/custom_test.go b/shard/resolver/custom_test.go deleted file mode 100644 index 764c4f4..0000000 --- a/shard/resolver/custom_test.go +++ /dev/null @@ -1,177 +0,0 @@ -package resolver - -import ( - "errors" - "testing" - - "github.com/mkbeh/xpg/shard" -) - -func TestNewCustomValidatesArguments(t *testing.T) { - t.Parallel() - - topology := newTestTopology(t, "shard-a") - - validResolve := ResolveFunc[int]( - func(int, *shard.Topology) (shard.ID, error) { - return "shard-a", nil - }, - ) - - tests := []struct { - name string - topology *shard.Topology - resolve ResolveFunc[int] - wantError string - }{ - { - name: "nil topology", - resolve: validResolve, - wantError: "xpg/shard/resolver: topology is nil or empty", - }, - { - name: "nil resolve function", - topology: topology, - wantError: "xpg/shard/resolver: custom resolve function is nil", - }, - } - - for _, test := range tests { - t.Run(test.name, func(t *testing.T) { - t.Parallel() - - _, err := NewCustom( - test.topology, - test.resolve, - ) - if err == nil { - t.Fatal("expected error") - } - - if got := err.Error(); got != test.wantError { - t.Fatalf("error = %q, want %q", got, test.wantError) - } - }) - } -} - -func TestCustomResolverResolve(t *testing.T) { - t.Parallel() - - topology := newTestTopology(t, "shard-a", "shard-b") - - resolver, err := NewCustom( - topology, - func(key int, gotTopology *shard.Topology) (shard.ID, error) { - if gotTopology != topology { - t.Fatal("resolve function received a different topology") - } - - if key < 100 { - return "shard-a", nil - } - - return "shard-b", nil - }, - ) - if err != nil { - t.Fatalf("NewCustom() error = %v", err) - } - - resolved, err := resolver.Resolve(142) - if err != nil { - t.Fatalf("Resolve() error = %v", err) - } - - if got, want := resolved.ID(), shard.ID("shard-b"); got != want { - t.Fatalf("Resolve().ID() = %q, want %q", got, want) - } -} - -func TestCustomResolverPropagatesResolveError(t *testing.T) { - t.Parallel() - - topology := newTestTopology(t, "shard-a") - sentinel := errors.New("resolve failed") - - resolver, err := NewCustom( - topology, - func(int, *shard.Topology) (shard.ID, error) { - return "", sentinel - }, - ) - if err != nil { - t.Fatalf("NewCustom() error = %v", err) - } - - _, err = resolver.Resolve(1) - if !errors.Is(err, sentinel) { - t.Fatalf("Resolve() error = %v, want sentinel", err) - } -} - -func TestCustomResolverPropagatesErrNoShard(t *testing.T) { - t.Parallel() - - topology := newTestTopology(t, "shard-a") - - resolver, err := NewCustom( - topology, - func(int, *shard.Topology) (shard.ID, error) { - return "", shard.ErrNoShard - }, - ) - if err != nil { - t.Fatalf("NewCustom() error = %v", err) - } - - _, err = resolver.Resolve(1) - if !errors.Is(err, shard.ErrNoShard) { - t.Fatalf("Resolve() error = %v, want ErrNoShard", err) - } -} - -func TestCustomResolverRejectsUnknownShard(t *testing.T) { - t.Parallel() - - topology := newTestTopology(t, "shard-a") - - resolver, err := NewCustom( - topology, - func(int, *shard.Topology) (shard.ID, error) { - return "missing", nil - }, - ) - if err != nil { - t.Fatalf("NewCustom() error = %v", err) - } - - _, err = resolver.Resolve(1) - if !errors.Is(err, shard.ErrUnknownShard) { - t.Fatalf("Resolve() error = %v, want ErrUnknownShard", err) - } - - var unknown *shard.UnknownShardError - if !errors.As(err, &unknown) { - t.Fatalf("Resolve() error = %T, want *shard.UnknownShardError", err) - } - - if got, want := unknown.ShardID, shard.ID("missing"); got != want { - t.Fatalf("ShardID = %q, want %q", got, want) - } -} - -func TestCustomResolverUninitialized(t *testing.T) { - t.Parallel() - - var resolver *CustomResolver[int] - - _, err := resolver.Resolve(1) - if err == nil { - t.Fatal("expected error") - } - - if got, want := err.Error(), "xpg/shard/resolver: custom resolver is not initialized"; got != want { - t.Fatalf("error = %q, want %q", got, want) - } -} diff --git a/shard/resolver/encoder_test.go b/shard/resolver/encoder_test.go deleted file mode 100644 index 067187b..0000000 --- a/shard/resolver/encoder_test.go +++ /dev/null @@ -1,148 +0,0 @@ -package resolver - -import ( - "bytes" - "encoding/hex" - "errors" - "testing" -) - -func TestKeyEncoderFuncNil(t *testing.T) { - t.Parallel() - - var encoder KeyEncoderFunc[int] - - _, err := encoder.Encode(1) - if err == nil { - t.Fatal("expected error") - } - - if got, want := err.Error(), "xpg/shard/resolver: key encoder function is nil"; got != want { - t.Fatalf("error = %q, want %q", got, want) - } -} - -func TestKeyEncoderFunc(t *testing.T) { - t.Parallel() - - sentinel := errors.New("encode failed") - encoder := KeyEncoderFunc[int](func(key int) ([]byte, error) { - if key < 0 { - return nil, sentinel - } - - return []byte{byte(key)}, nil - }) - - encoded, err := encoder.Encode(7) - if err != nil { - t.Fatalf("Encode() error = %v", err) - } - if got, want := encoded, []byte{7}; !bytes.Equal(got, want) { - t.Fatalf("Encode() = %v, want %v", got, want) - } - - if _, err := encoder.Encode(-1); !errors.Is(err, sentinel) { - t.Fatalf("Encode() error = %v, want sentinel", err) - } -} - -func TestStringKeyEncoder(t *testing.T) { - t.Parallel() - - encoded, err := StringKeyEncoder().Encode("a\x00b") - if err != nil { - t.Fatalf("Encode() error = %v", err) - } - - if got, want := string(encoded), "a\x00b"; got != want { - t.Fatalf("Encode() = %q, want %q", got, want) - } -} - -func TestBytesKeyEncoderReturnsDefensiveCopy(t *testing.T) { - t.Parallel() - - key := []byte{1, 2, 3} - encoded, err := BytesKeyEncoder().Encode(key) - if err != nil { - t.Fatalf("Encode() error = %v", err) - } - - encoded[0] = 9 - - if got, want := key[0], byte(1); got != want { - t.Fatalf("input key changed to %d, want %d", got, want) - } -} - -func TestIntegerKeyEncoders(t *testing.T) { - t.Parallel() - - tests := []struct { - name string - got func() ([]byte, error) - want string - }{ - { - name: "int64 positive", - got: func() ([]byte, error) { return Int64KeyEncoder().Encode(1) }, - want: "0000000000000001", - }, - { - name: "int64 negative", - got: func() ([]byte, error) { return Int64KeyEncoder().Encode(-1) }, - want: "ffffffffffffffff", - }, - { - name: "uint64", - got: func() ([]byte, error) { return Uint64KeyEncoder().Encode(0x0102030405060708) }, - want: "0102030405060708", - }, - } - - for _, test := range tests { - t.Run(test.name, func(t *testing.T) { - t.Parallel() - - encoded, err := test.got() - if err != nil { - t.Fatalf("Encode() error = %v", err) - } - - if got := hex.EncodeToString(encoded); got != test.want { - t.Fatalf("Encode() = %s, want %s", got, test.want) - } - }) - } -} - -func TestFixedSizeKeyEncoders(t *testing.T) { - t.Parallel() - - var key16 [16]byte - for index := range key16 { - key16[index] = byte(index) - } - - encoded16, err := Bytes16KeyEncoder().Encode(key16) - if err != nil { - t.Fatalf("Bytes16KeyEncoder.Encode() error = %v", err) - } - if got, want := encoded16, key16[:]; !bytes.Equal(got, want) { - t.Fatalf("Bytes16KeyEncoder.Encode() = %v, want %v", got, want) - } - - var key32 [32]byte - for index := range key32 { - key32[index] = byte(31 - index) - } - - encoded32, err := Bytes32KeyEncoder().Encode(key32) - if err != nil { - t.Fatalf("Bytes32KeyEncoder.Encode() error = %v", err) - } - if got, want := encoded32, key32[:]; !bytes.Equal(got, want) { - t.Fatalf("Bytes32KeyEncoder.Encode() = %v, want %v", got, want) - } -} diff --git a/shard/resolver/hash_test.go b/shard/resolver/hash_test.go deleted file mode 100644 index f29c5c4..0000000 --- a/shard/resolver/hash_test.go +++ /dev/null @@ -1,255 +0,0 @@ -package resolver - -import ( - "errors" - "fmt" - "testing" - - "github.com/mkbeh/xpg/shard" -) - -func TestNewHashValidatesArguments(t *testing.T) { - t.Parallel() - - topology := newTestTopology(t, "shard-a") - - tests := []struct { - name string - topology *shard.Topology - namespace string - encoder KeyEncoder[string] - wantError string - }{ - { - name: "nil topology", - namespace: "users", - encoder: StringKeyEncoder(), - wantError: "xpg/shard/resolver: topology is nil or empty", - }, - { - name: "nil encoder", - topology: topology, - namespace: "users", - wantError: "xpg/shard/resolver: key encoder is nil", - }, - { - name: "empty namespace", - topology: topology, - encoder: StringKeyEncoder(), - wantError: "xpg/shard/resolver: hash namespace must not be empty", - }, - } - - for _, test := range tests { - t.Run(test.name, func(t *testing.T) { - t.Parallel() - - _, err := NewHash( - test.topology, - test.namespace, - test.encoder, - ) - if err == nil { - t.Fatal("expected error") - } - - if got := err.Error(); got != test.wantError { - t.Fatalf("error = %q, want %q", got, test.wantError) - } - }) - } -} - -func TestNewHashTreatsNamespaceAsOpaqueNonEmptyString(t *testing.T) { - t.Parallel() - - topology := newTestTopology(t, "shard-a") - - for _, namespace := range []string{"users", " users ", " "} { - resolver, err := NewHash(topology, namespace, StringKeyEncoder()) - if err != nil { - t.Fatalf("NewHash(%q) error = %v", namespace, err) - } - - resolved, err := resolver.Resolve("alice") - if err != nil { - t.Fatalf("Resolve() error = %v", err) - } - - if got, want := resolved.ID(), shard.ID("shard-a"); got != want { - t.Fatalf("Resolve().ID() = %q, want %q", got, want) - } - } -} - -func TestHashResolverStablePlacementVectors(t *testing.T) { - t.Parallel() - - topology := newTestTopology(t, "shard-a", "shard-b", "shard-c") - resolver, err := NewHash(topology, "users", StringKeyEncoder()) - if err != nil { - t.Fatalf("NewHash() error = %v", err) - } - - tests := []struct { - key string - want shard.ID - }{ - {key: "alice", want: "shard-a"}, - {key: "bob", want: "shard-b"}, - {key: "carol", want: "shard-b"}, - {key: "dave", want: "shard-b"}, - {key: "eve", want: "shard-c"}, - {key: "0", want: "shard-a"}, - {key: "1", want: "shard-b"}, - {key: "2", want: "shard-c"}, - } - - for _, test := range tests { - t.Run(test.key, func(t *testing.T) { - t.Parallel() - - resolved, err := resolver.Resolve(test.key) - if err != nil { - t.Fatalf("Resolve() error = %v", err) - } - - if got := resolved.ID(); got != test.want { - t.Fatalf( - "Resolve(%q).ID() = %q, want %q", - test.key, - got, - test.want, - ) - } - }) - } -} - -func TestHashResolverPlacementDoesNotDependOnTopologyOrder(t *testing.T) { - t.Parallel() - - first := newTestTopology(t, "shard-a", "shard-b", "shard-c") - second := newTestTopology(t, "shard-c", "shard-a", "shard-b") - - firstResolver, err := NewHash(first, "users", StringKeyEncoder()) - if err != nil { - t.Fatalf("NewHash(first) error = %v", err) - } - - secondResolver, err := NewHash(second, "users", StringKeyEncoder()) - if err != nil { - t.Fatalf("NewHash(second) error = %v", err) - } - - for _, key := range []string{"alice", "bob", "carol", "dave", "eve", "user-123"} { - firstShard, err := firstResolver.Resolve(key) - if err != nil { - t.Fatalf("first Resolve(%q) error = %v", key, err) - } - - secondShard, err := secondResolver.Resolve(key) - if err != nil { - t.Fatalf("second Resolve(%q) error = %v", key, err) - } - - if firstShard.ID() != secondShard.ID() { - t.Fatalf( - "Resolve(%q) = %q and %q for different topology orders", - key, - firstShard.ID(), - secondShard.ID(), - ) - } - } -} - -func TestHashResolverAddingShardOnlyMovesKeysToNewShard(t *testing.T) { - t.Parallel() - - before := newTestTopology(t, "shard-a", "shard-b") - after := newTestTopology(t, "shard-a", "shard-b", "shard-c") - - beforeResolver, err := NewHash(before, "users", StringKeyEncoder()) - if err != nil { - t.Fatalf("NewHash(before) error = %v", err) - } - - afterResolver, err := NewHash(after, "users", StringKeyEncoder()) - if err != nil { - t.Fatalf("NewHash(after) error = %v", err) - } - - moved := 0 - - for index := range 256 { - key := fmt.Sprintf("user-%d", index) - - previous, err := beforeResolver.Resolve(key) - if err != nil { - t.Fatalf("before Resolve(%q) error = %v", key, err) - } - - current, err := afterResolver.Resolve(key) - if err != nil { - t.Fatalf("after Resolve(%q) error = %v", key, err) - } - - if previous.ID() == current.ID() { - continue - } - - moved++ - - if current.ID() != "shard-c" { - t.Fatalf( - "Resolve(%q) moved from %q to existing shard %q", - key, - previous.ID(), - current.ID(), - ) - } - } - - if moved == 0 { - t.Fatal("expected at least one key to move to the new shard") - } -} - -func TestHashResolverWrapsEncoderError(t *testing.T) { - t.Parallel() - - topology := newTestTopology(t, "shard-a") - sentinel := errors.New("encode failed") - - resolver, err := NewHash( - topology, - "users", - KeyEncoderFunc[string](func(string) ([]byte, error) { - return nil, sentinel - }), - ) - if err != nil { - t.Fatalf("NewHash() error = %v", err) - } - - _, err = resolver.Resolve("alice") - if !errors.Is(err, sentinel) { - t.Fatalf("Resolve() error = %v, want wrapped sentinel", err) - } -} - -func TestHashResolverUninitialized(t *testing.T) { - t.Parallel() - - var resolver *HashResolver[string] - - _, err := resolver.Resolve("alice") - if err == nil { - t.Fatal("expected error") - } - - if got, want := err.Error(), "xpg/shard/resolver: hash resolver is not initialized"; got != want { - t.Fatalf("error = %q, want %q", got, want) - } -} diff --git a/shard/resolver/helpers_test.go b/shard/resolver/helpers_test.go deleted file mode 100644 index a12991b..0000000 --- a/shard/resolver/helpers_test.go +++ /dev/null @@ -1,59 +0,0 @@ -package resolver - -import ( - "testing" - - "github.com/jackc/pgx/v5/pgxpool" - "github.com/mkbeh/xpg" - "github.com/mkbeh/xpg/cluster" - "github.com/mkbeh/xpg/shard" -) - -const testDatabaseURL = "postgres://postgres@127.0.0.1:1/postgres?sslmode=disable" - -func newTestTopology(t *testing.T, ids ...shard.ID) *shard.Topology { - t.Helper() - - configs := make([]shard.Config, len(ids)) - - for index, id := range ids { - poolConfig, err := pgxpool.ParseConfig(testDatabaseURL) - if err != nil { - t.Fatalf("pgxpool.ParseConfig() error = %v", err) - } - - poolConfig.MinConns = 0 - poolConfig.MaxConns = 1 - - pool, err := xpg.New( - t.Context(), - poolConfig, - xpg.WithName("shard."+string(id)+".primary"), - ) - if err != nil { - t.Fatalf("xpg.New() error = %v", err) - } - - shardCluster, err := cluster.New(cluster.Config{ - ID: id, - Primary: pool, - }) - if err != nil { - pool.Close() - t.Fatalf("cluster.New() error = %v", err) - } - - t.Cleanup(shardCluster.Close) - - configs[index] = shard.Config{Cluster: shardCluster} - } - - topology, err := shard.NewTopology(configs) - if err != nil { - t.Fatalf("shard.NewTopology() error = %v", err) - } - - t.Cleanup(topology.Close) - - return topology -} diff --git a/shard/resolver/range.go b/shard/resolver/range.go index 9714777..3a339cf 100644 --- a/shard/resolver/range.go +++ b/shard/resolver/range.go @@ -7,6 +7,7 @@ import ( "slices" "sort" + "github.com/mkbeh/xpg/cluster" "github.com/mkbeh/xpg/shard" ) @@ -17,18 +18,27 @@ import ( type Range[K cmp.Ordered] struct { Start K End K - ShardID shard.ID + ShardID cluster.ID } -// RangeResolver routes ordered keys through non-overlapping ranges. +// RangeResolver resolves ordered keys through bounded, non-overlapping ranges. type RangeResolver[K cmp.Ordered] struct { ranges []rangeEntry[K] } -// NewRange creates a resolver from non-overlapping half-open ranges. +type rangeEntry[K cmp.Ordered] struct { + start K + end K + shard shard.Shard + + sourceIndex int +} + +// NewRange creates a resolver from bounded, non-overlapping half-open ranges. // -// The supplied ranges may be unordered. NewRange copies and sorts them by Start, -// validates their boundaries and overlap, and leaves the caller's slice unchanged. +// NewRange resolves every ShardID to its immutable Shard handle once, copies +// the routing data into an internal representation, sorts it by Start, and +// validates that ranges do not overlap. The caller's slice is not modified. func NewRange[K cmp.Ordered](topology *shard.Topology, ranges []Range[K]) (*RangeResolver[K], error) { if err := requireTopology(topology); err != nil { return nil, err @@ -45,8 +55,6 @@ func NewRange[K cmp.Ordered](topology *shard.Topology, ranges []Range[K]) (*Rang return nil, fmt.Errorf("xpg/shard/resolver: range %d: %w", index, err) } - // Using < intentionally rejects empty and reversed ranges as well as - // ranges with NaN boundaries for floating-point key types. if !(valueRange.Start < valueRange.End) { //nolint:staticcheck // Negated comparison intentionally rejects NaN boundaries. return nil, fmt.Errorf("xpg/shard/resolver: range %d must satisfy start < end", index) } @@ -56,9 +64,7 @@ func NewRange[K cmp.Ordered](topology *shard.Topology, ranges []Range[K]) (*Rang return nil, fmt.Errorf( "xpg/shard/resolver: range %d: %w", index, - &shard.UnknownShardError{ - ShardID: valueRange.ShardID, - }, + &shard.UnknownShardError{ShardID: valueRange.ShardID}, ) } @@ -77,15 +83,10 @@ func NewRange[K cmp.Ordered](topology *shard.Topology, ranges []Range[K]) (*Rang }, ) - // Once ranges are sorted by Start, checking adjacent entries is sufficient - // to detect every overlap. for index := 1; index < len(entries); index++ { previous := entries[index-1] current := entries[index] - // Adjacent half-open ranges are valid: - // - // [0, 100) and [100, 200) if previous.end <= current.start { continue } @@ -102,7 +103,10 @@ func NewRange[K cmp.Ordered](topology *shard.Topology, ranges []Range[K]) (*Rang }, nil } -// Resolve returns the shard whose configured range contains key. +// Resolve returns the shard whose range contains key. +// +// Resolve performs only an in-memory lookup. It does not consult topology, +// acquire a connection, or execute a PostgreSQL query. func (resolver *RangeResolver[K]) Resolve(key K) (shard.Shard, error) { if resolver == nil || len(resolver.ranges) == 0 { return shard.Shard{}, errors.New("xpg/shard/resolver: range resolver is not initialized") @@ -131,11 +135,3 @@ func (resolver *RangeResolver[K]) Resolve(key K) (shard.Shard, error) { return entry.shard, nil } - -type rangeEntry[K cmp.Ordered] struct { - start K - end K - shard shard.Shard - - sourceIndex int -} diff --git a/shard/resolver/range_test.go b/shard/resolver/range_test.go deleted file mode 100644 index d001386..0000000 --- a/shard/resolver/range_test.go +++ /dev/null @@ -1,275 +0,0 @@ -package resolver - -import ( - "errors" - "math" - "slices" - "testing" - - "github.com/mkbeh/xpg/shard" -) - -func TestNewRangeValidatesArguments(t *testing.T) { - t.Parallel() - - topology := newTestTopology(t, "shard-a") - - tests := []struct { - name string - topology *shard.Topology - ranges []Range[int] - wantError string - }{ - { - name: "nil topology", - ranges: []Range[int]{ - {Start: 0, End: 10, ShardID: "shard-a"}, - }, - wantError: "xpg/shard/resolver: topology is nil or empty", - }, - { - name: "empty ranges", - topology: topology, - wantError: "xpg/shard/resolver: range resolver requires at least one range", - }, - { - name: "empty shard ID", - topology: topology, - ranges: []Range[int]{ - {Start: 0, End: 10}, - }, - wantError: "xpg/shard/resolver: range 0: shard ID must not be empty", - }, - { - name: "empty interval", - topology: topology, - ranges: []Range[int]{ - {Start: 10, End: 10, ShardID: "shard-a"}, - }, - wantError: "xpg/shard/resolver: range 0 must satisfy start < end", - }, - { - name: "reversed interval", - topology: topology, - ranges: []Range[int]{ - {Start: 20, End: 10, ShardID: "shard-a"}, - }, - wantError: "xpg/shard/resolver: range 0 must satisfy start < end", - }, - { - name: "unknown shard", - topology: topology, - ranges: []Range[int]{ - {Start: 0, End: 10, ShardID: "missing"}, - }, - wantError: `xpg/shard/resolver: range 0: xpg/shard: unknown shard "missing"`, - }, - } - - for _, test := range tests { - t.Run(test.name, func(t *testing.T) { - t.Parallel() - - _, err := NewRange( - test.topology, - test.ranges, - ) - if err == nil { - t.Fatal("expected error") - } - - if got := err.Error(); got != test.wantError { - t.Fatalf("error = %q, want %q", got, test.wantError) - } - }) - } -} - -func TestNewRangeRejectsNaNBoundaries(t *testing.T) { - t.Parallel() - - topology := newTestTopology(t, "shard-a") - - tests := []struct { - name string - valueRange Range[float64] - }{ - { - name: "NaN start", - valueRange: Range[float64]{ - Start: math.NaN(), - End: 10, - ShardID: "shard-a", - }, - }, - { - name: "NaN end", - valueRange: Range[float64]{ - Start: 0, - End: math.NaN(), - ShardID: "shard-a", - }, - }, - } - - for _, test := range tests { - t.Run(test.name, func(t *testing.T) { - t.Parallel() - - _, err := NewRange( - topology, - []Range[float64]{test.valueRange}, - ) - if err == nil { - t.Fatal("expected error") - } - - if got, want := err.Error(), - "xpg/shard/resolver: range 0 must satisfy start < end"; got != want { - t.Fatalf("error = %q, want %q", got, want) - } - }) - } -} - -func TestNewRangeRejectsOverlapUsingSourceIndexes(t *testing.T) { - t.Parallel() - - topology := newTestTopology(t, "shard-a", "shard-b") - - _, err := NewRange(topology, []Range[int]{ - {Start: 100, End: 200, ShardID: "shard-b"}, - {Start: 50, End: 150, ShardID: "shard-a"}, - }) - if err == nil { - t.Fatal("expected overlap error") - } - - if got, want := err.Error(), "xpg/shard/resolver: ranges 1 and 0 overlap"; got != want { - t.Fatalf("error = %q, want %q", got, want) - } -} - -func TestNewRangeDoesNotModifyInputOrder(t *testing.T) { - t.Parallel() - - topology := newTestTopology(t, "shard-a", "shard-b") - ranges := []Range[int]{ - {Start: 100, End: 200, ShardID: "shard-b"}, - {Start: 0, End: 100, ShardID: "shard-a"}, - } - want := slices.Clone(ranges) - - if _, err := NewRange(topology, ranges); err != nil { - t.Fatalf("NewRange() error = %v", err) - } - - if !slices.Equal(ranges, want) { - t.Fatalf("ranges = %+v, want unchanged %+v", ranges, want) - } -} - -func TestRangeResolverHalfOpenBoundariesAndGaps(t *testing.T) { - t.Parallel() - - topology := newTestTopology(t, "shard-a", "shard-b") - resolver, err := NewRange(topology, []Range[int]{ - {Start: 20, End: 30, ShardID: "shard-b"}, - {Start: 0, End: 10, ShardID: "shard-a"}, - }) - if err != nil { - t.Fatalf("NewRange() error = %v", err) - } - - tests := []struct { - key int - wantID shard.ID - wantErr error - }{ - {key: 0, wantID: "shard-a"}, - {key: 9, wantID: "shard-a"}, - {key: 10, wantErr: shard.ErrNoShard}, - {key: 19, wantErr: shard.ErrNoShard}, - {key: 20, wantID: "shard-b"}, - {key: 29, wantID: "shard-b"}, - {key: 30, wantErr: shard.ErrNoShard}, - } - - for _, test := range tests { - resolved, err := resolver.Resolve(test.key) - if test.wantErr != nil { - if !errors.Is(err, test.wantErr) { - t.Fatalf("Resolve(%d) error = %v, want %v", test.key, err, test.wantErr) - } - - continue - } - - if err != nil { - t.Fatalf("Resolve(%d) error = %v", test.key, err) - } - - if got := resolved.ID(); got != test.wantID { - t.Fatalf("Resolve(%d).ID() = %q, want %q", test.key, got, test.wantID) - } - } -} - -func TestRangeResolverAllowsAdjacentRanges(t *testing.T) { - t.Parallel() - - topology := newTestTopology(t, "shard-a", "shard-b") - resolver, err := NewRange(topology, []Range[int]{ - {Start: 0, End: 100, ShardID: "shard-a"}, - {Start: 100, End: 200, ShardID: "shard-b"}, - }) - if err != nil { - t.Fatalf("NewRange() error = %v", err) - } - - resolved, err := resolver.Resolve(100) - if err != nil { - t.Fatalf("Resolve() error = %v", err) - } - - if got, want := resolved.ID(), shard.ID("shard-b"); got != want { - t.Fatalf("Resolve().ID() = %q, want %q", got, want) - } -} - -func TestRangeResolverSupportsStrings(t *testing.T) { - t.Parallel() - - topology := newTestTopology(t, "shard-a", "shard-b") - resolver, err := NewRange(topology, []Range[string]{ - {Start: "a", End: "m", ShardID: "shard-a"}, - {Start: "m", End: "z", ShardID: "shard-b"}, - }) - if err != nil { - t.Fatalf("NewRange() error = %v", err) - } - - resolved, err := resolver.Resolve("m") - if err != nil { - t.Fatalf("Resolve() error = %v", err) - } - - if got, want := resolved.ID(), shard.ID("shard-b"); got != want { - t.Fatalf("Resolve().ID() = %q, want %q", got, want) - } -} - -func TestRangeResolverUninitialized(t *testing.T) { - t.Parallel() - - var resolver *RangeResolver[int] - - _, err := resolver.Resolve(1) - if err == nil { - t.Fatal("expected error") - } - - if got, want := err.Error(), "xpg/shard/resolver: range resolver is not initialized"; got != want { - t.Fatalf("error = %q, want %q", got, want) - } -} diff --git a/shard/resolver/hash.go b/shard/resolver/rendezvous.go similarity index 59% rename from shard/resolver/hash.go rename to shard/resolver/rendezvous.go index a631afe..741b579 100644 --- a/shard/resolver/hash.go +++ b/shard/resolver/rendezvous.go @@ -8,6 +8,7 @@ import ( "fmt" "math" + "github.com/mkbeh/xpg/cluster" "github.com/mkbeh/xpg/shard" ) @@ -20,27 +21,25 @@ const ( rendezvousLengthSize = 4 ) -// HashResolver routes keys using rendezvous/HRW hashing with SHA-256. -// -// HashResolver captures the shard set when it is created. The shards remain -// borrowed from the topology, so the resolver must not outlive it. -type HashResolver[K any] struct { - shards []shard.Shard - prefix []byte - encoder KeyEncoder[K] - maxShardIDLength int +// RendezvousResolver implements rendezvous/HRW routing with SHA-256 and stable +// named shard IDs. +type RendezvousResolver[K any] struct { + shards []shard.Shard + prefix []byte + encoder KeyEncoder[K] + maxIDLength int } -// NewHash creates a rendezvous hash resolver bound to topology. +// NewRendezvous creates the version-1 rendezvous resolver bound to topology. // -// Namespace is an opaque non-empty string and part of the persistent placement -// contract. Changing the namespace, key encoder, shard IDs, or placement format -// changes shard placement and may require data migration. -func NewHash[K any]( +// Namespace is part of the persistent placement contract. Changing it changes +// shard placement and may require data migration. The topology's shard slice is +// copied once; Resolve does not consult topology. +func NewRendezvous[K any]( topology *shard.Topology, namespace string, encoder KeyEncoder[K], -) (*HashResolver[K], error) { +) (*RendezvousResolver[K], error) { if err := requireTopology(topology); err != nil { return nil, err } @@ -50,24 +49,24 @@ func NewHash[K any]( } if namespace == "" { - return nil, errors.New("xpg/shard/resolver: hash namespace must not be empty") + return nil, errors.New("xpg/shard/resolver: rendezvous namespace must not be empty") } - if len(namespace) > math.MaxUint32 { - return nil, errors.New("xpg/shard/resolver: hash namespace is too large") + if uint64(len(namespace)) > uint64(math.MaxUint32) { + return nil, errors.New("xpg/shard/resolver: rendezvous namespace is too large") } shards := topology.Shards() - maxShardIDLength := 0 + maxIDLength := 0 for _, candidate := range shards { id := candidate.ID() - if len(id) > math.MaxUint32 { + if uint64(len(id)) > uint64(math.MaxUint32) { return nil, errors.New("xpg/shard/resolver: shard ID is too large") } - maxShardIDLength = max(maxShardIDLength, len(id)) + maxIDLength = max(maxIDLength, len(id)) } prefixSize := len(rendezvousDomain) + rendezvousLengthSize + len(namespace) @@ -85,26 +84,26 @@ func NewHash[K any]( copy(prefix[namespaceOffset:], namespace) - return &HashResolver[K]{ - shards: shards, - prefix: prefix, - encoder: encoder, - maxShardIDLength: maxShardIDLength, + return &RendezvousResolver[K]{ + shards: shards, + prefix: prefix, + encoder: encoder, + maxIDLength: maxIDLength, }, nil } // Resolve maps key to a shard using rendezvous hashing. -func (resolver *HashResolver[K]) Resolve(key K) (shard.Shard, error) { +func (resolver *RendezvousResolver[K]) Resolve(key K) (shard.Shard, error) { if resolver == nil || len(resolver.shards) == 0 || resolver.encoder == nil { - return shard.Shard{}, errors.New("xpg/shard/resolver: hash resolver is not initialized") + return shard.Shard{}, errors.New("xpg/shard/resolver: rendezvous resolver is not initialized") } encoded, err := resolver.encoder.Encode(key) if err != nil { - return shard.Shard{}, fmt.Errorf("xpg/shard/resolver: encode hash key: %w", err) + return shard.Shard{}, fmt.Errorf("xpg/shard/resolver: encode rendezvous key: %w", err) } - if len(encoded) > math.MaxUint32 { + if uint64(len(encoded)) > uint64(math.MaxUint32) { return shard.Shard{}, errors.New("xpg/shard/resolver: encoded key is too large") } @@ -113,17 +112,7 @@ func (resolver *HashResolver[K]) Resolve(key K) (shard.Shard, error) { idLengthOffset := keyOffset + len(encoded) idOffset := idLengthOffset + rendezvousLengthSize - // Persistent placement format: - // - // domain || namespace_length || namespace || - // key_length || key || shard_id_length || shard_id - // - // The candidate-independent prefix and key are written once. Only the shard - // ID suffix is overwritten while evaluating candidates. - scoreInput := make( - []byte, - idOffset+resolver.maxShardIDLength, - ) + scoreInput := make([]byte, idOffset+resolver.maxIDLength) copy(scoreInput, resolver.prefix) @@ -137,7 +126,7 @@ func (resolver *HashResolver[K]) Resolve(key K) (shard.Shard, error) { var ( selected shard.Shard best [sha256.Size]byte - bestID shard.ID + bestID cluster.ID hasBest bool ) @@ -157,9 +146,7 @@ func (resolver *HashResolver[K]) Resolve(key K) (shard.Shard, error) { // Shard ID is the deterministic tie-breaker, so placement does not // depend on topology registration order when scores are equal. - if !hasBest || - comparison > 0 || - (comparison == 0 && candidateID < bestID) { + if !hasBest || comparison > 0 || (comparison == 0 && candidateID < bestID) { selected = candidate best = score bestID = candidateID diff --git a/shard/resolver/time_range.go b/shard/resolver/time_range.go index 84d329b..52b5495 100644 --- a/shard/resolver/time_range.go +++ b/shard/resolver/time_range.go @@ -7,6 +7,7 @@ import ( "sort" "time" + "github.com/mkbeh/xpg/cluster" "github.com/mkbeh/xpg/shard" ) @@ -17,19 +18,29 @@ import ( type TimeRange struct { Start time.Time End time.Time - ShardID shard.ID + ShardID cluster.ID } -// TimeRangeResolver routes time instants through non-overlapping ranges. +// TimeRangeResolver resolves time instants through bounded, non-overlapping +// ranges. type TimeRangeResolver struct { ranges []timeRangeEntry } -// NewTimeRange creates a resolver from non-overlapping half-open time ranges. +type timeRangeEntry struct { + start time.Time + end time.Time + shard shard.Shard + + sourceIndex int +} + +// NewTimeRange creates a resolver from bounded, non-overlapping half-open time +// ranges. // -// The supplied ranges may be unordered. NewTimeRange normalizes boundaries to -// UTC, sorts ranges by Start, validates their overlap, and leaves the caller's -// slice unchanged. +// Range boundaries are normalized to UTC. Every ShardID is resolved to its +// immutable Shard handle once. The caller's slice and time values are not +// modified. func NewTimeRange(topology *shard.Topology, ranges []TimeRange) (*TimeRangeResolver, error) { if err := requireTopology(topology); err != nil { return nil, err @@ -58,9 +69,7 @@ func NewTimeRange(topology *shard.Topology, ranges []TimeRange) (*TimeRangeResol return nil, fmt.Errorf( "xpg/shard/resolver: time range %d: %w", index, - &shard.UnknownShardError{ - ShardID: valueRange.ShardID, - }, + &shard.UnknownShardError{ShardID: valueRange.ShardID}, ) } @@ -102,7 +111,10 @@ func NewTimeRange(topology *shard.Topology, ranges []TimeRange) (*TimeRangeResol }, nil } -// Resolve returns the shard whose configured time range contains key. +// Resolve returns the shard whose time range contains key. +// +// Resolve performs only an in-memory lookup. It does not consult topology, +// acquire a connection, or execute a PostgreSQL query. func (resolver *TimeRangeResolver) Resolve(key time.Time) (shard.Shard, error) { if resolver == nil || len(resolver.ranges) == 0 { return shard.Shard{}, errors.New("xpg/shard/resolver: time range resolver is not initialized") @@ -134,14 +146,6 @@ func (resolver *TimeRangeResolver) Resolve(key time.Time) (shard.Shard, error) { return entry.shard, nil } -type timeRangeEntry struct { - start time.Time - end time.Time - shard shard.Shard - - sourceIndex int -} - -func timeToUTC(value time.Time) time.Time { - return value.UTC() +func timeToUTC(t time.Time) time.Time { + return t.UTC() } diff --git a/shard/resolver/time_range_test.go b/shard/resolver/time_range_test.go deleted file mode 100644 index 246c34d..0000000 --- a/shard/resolver/time_range_test.go +++ /dev/null @@ -1,237 +0,0 @@ -package resolver - -import ( - "errors" - "slices" - "testing" - "time" - - "github.com/mkbeh/xpg/shard" -) - -func TestNewTimeRangeValidatesArguments(t *testing.T) { - t.Parallel() - - topology := newTestTopology(t, "shard-a") - start := time.Date(2026, 1, 1, 0, 0, 0, 0, time.UTC) - end := start.Add(time.Hour) - - tests := []struct { - name string - topology *shard.Topology - ranges []TimeRange - wantError string - }{ - { - name: "nil topology", - ranges: []TimeRange{ - {Start: start, End: end, ShardID: "shard-a"}, - }, - wantError: "xpg/shard/resolver: topology is nil or empty", - }, - { - name: "empty ranges", - topology: topology, - wantError: "xpg/shard/resolver: time range resolver requires at least one range", - }, - { - name: "empty shard ID", - topology: topology, - ranges: []TimeRange{ - {Start: start, End: end}, - }, - wantError: "xpg/shard/resolver: time range 0: shard ID must not be empty", - }, - { - name: "empty interval", - topology: topology, - ranges: []TimeRange{ - {Start: start, End: start, ShardID: "shard-a"}, - }, - wantError: "xpg/shard/resolver: time range 0 must satisfy start < end", - }, - { - name: "reversed interval", - topology: topology, - ranges: []TimeRange{ - {Start: end, End: start, ShardID: "shard-a"}, - }, - wantError: "xpg/shard/resolver: time range 0 must satisfy start < end", - }, - { - name: "unknown shard", - topology: topology, - ranges: []TimeRange{ - {Start: start, End: end, ShardID: "missing"}, - }, - wantError: `xpg/shard/resolver: time range 0: xpg/shard: unknown shard "missing"`, - }, - } - - for _, test := range tests { - t.Run(test.name, func(t *testing.T) { - t.Parallel() - - _, err := NewTimeRange( - test.topology, - test.ranges, - ) - if err == nil { - t.Fatal("expected error") - } - - if got := err.Error(); got != test.wantError { - t.Fatalf("error = %q, want %q", got, test.wantError) - } - }) - } -} - -func TestNewTimeRangeRejectsOverlapUsingSourceIndexes(t *testing.T) { - t.Parallel() - - topology := newTestTopology(t, "shard-a", "shard-b") - base := time.Date(2026, 1, 1, 0, 0, 0, 0, time.UTC) - - _, err := NewTimeRange(topology, []TimeRange{ - {Start: base.Add(2 * time.Hour), End: base.Add(4 * time.Hour), ShardID: "shard-b"}, - {Start: base.Add(time.Hour), End: base.Add(3 * time.Hour), ShardID: "shard-a"}, - }) - if err == nil { - t.Fatal("expected overlap error") - } - - if got, want := err.Error(), "xpg/shard/resolver: time ranges 1 and 0 overlap"; got != want { - t.Fatalf("error = %q, want %q", got, want) - } -} - -func TestNewTimeRangeDoesNotModifyInput(t *testing.T) { - t.Parallel() - - topology := newTestTopology(t, "shard-a", "shard-b") - location := time.FixedZone("UTC+3", 3*60*60) - base := time.Date(2026, 1, 1, 0, 0, 0, 0, location) - ranges := []TimeRange{ - {Start: base.Add(time.Hour), End: base.Add(2 * time.Hour), ShardID: "shard-b"}, - {Start: base, End: base.Add(time.Hour), ShardID: "shard-a"}, - } - want := slices.Clone(ranges) - - if _, err := NewTimeRange(topology, ranges); err != nil { - t.Fatalf("NewTimeRange() error = %v", err) - } - - if !slices.Equal(ranges, want) { - t.Fatalf("ranges = %+v, want unchanged %+v", ranges, want) - } -} - -func TestTimeRangeResolverHalfOpenBoundariesAndGaps(t *testing.T) { - t.Parallel() - - topology := newTestTopology(t, "shard-a", "shard-b") - base := time.Date(2026, 1, 1, 0, 0, 0, 0, time.UTC) - resolver, err := NewTimeRange(topology, []TimeRange{ - {Start: base.Add(2 * time.Hour), End: base.Add(3 * time.Hour), ShardID: "shard-b"}, - {Start: base, End: base.Add(time.Hour), ShardID: "shard-a"}, - }) - if err != nil { - t.Fatalf("NewTimeRange() error = %v", err) - } - - tests := []struct { - key time.Time - wantID shard.ID - wantErr error - }{ - {key: base, wantID: "shard-a"}, - {key: base.Add(time.Hour - time.Nanosecond), wantID: "shard-a"}, - {key: base.Add(time.Hour), wantErr: shard.ErrNoShard}, - {key: base.Add(2 * time.Hour), wantID: "shard-b"}, - {key: base.Add(3 * time.Hour), wantErr: shard.ErrNoShard}, - } - - for _, test := range tests { - resolved, err := resolver.Resolve(test.key) - if test.wantErr != nil { - if !errors.Is(err, test.wantErr) { - t.Fatalf("Resolve(%v) error = %v, want %v", test.key, err, test.wantErr) - } - - continue - } - - if err != nil { - t.Fatalf("Resolve(%v) error = %v", test.key, err) - } - - if got := resolved.ID(); got != test.wantID { - t.Fatalf("Resolve(%v).ID() = %q, want %q", test.key, got, test.wantID) - } - } -} - -func TestTimeRangeResolverNormalizesToUTC(t *testing.T) { - t.Parallel() - - topology := newTestTopology(t, "shard-a") - start := time.Date(2026, 1, 1, 0, 0, 0, 0, time.UTC) - end := start.Add(time.Hour) - resolver, err := NewTimeRange(topology, []TimeRange{ - {Start: start, End: end, ShardID: "shard-a"}, - }) - if err != nil { - t.Fatalf("NewTimeRange() error = %v", err) - } - - location := time.FixedZone("UTC+3", 3*60*60) - key := start.Add(30 * time.Minute).In(location) - - resolved, err := resolver.Resolve(key) - if err != nil { - t.Fatalf("Resolve() error = %v", err) - } - - if got, want := resolved.ID(), shard.ID("shard-a"); got != want { - t.Fatalf("Resolve().ID() = %q, want %q", got, want) - } -} - -func TestTimeRangeResolverAllowsAdjacentRanges(t *testing.T) { - t.Parallel() - - topology := newTestTopology(t, "shard-a", "shard-b") - base := time.Date(2026, 1, 1, 0, 0, 0, 0, time.UTC) - resolver, err := NewTimeRange(topology, []TimeRange{ - {Start: base, End: base.Add(time.Hour), ShardID: "shard-a"}, - {Start: base.Add(time.Hour), End: base.Add(2 * time.Hour), ShardID: "shard-b"}, - }) - if err != nil { - t.Fatalf("NewTimeRange() error = %v", err) - } - - resolved, err := resolver.Resolve(base.Add(time.Hour)) - if err != nil { - t.Fatalf("Resolve() error = %v", err) - } - - if got, want := resolved.ID(), shard.ID("shard-b"); got != want { - t.Fatalf("Resolve().ID() = %q, want %q", got, want) - } -} - -func TestTimeRangeResolverUninitialized(t *testing.T) { - t.Parallel() - - var resolver *TimeRangeResolver - - _, err := resolver.Resolve(time.Now()) - if err == nil { - t.Fatal("expected error") - } - - if got, want := err.Error(), "xpg/shard/resolver: time range resolver is not initialized"; got != want { - t.Fatalf("error = %q, want %q", got, want) - } -} diff --git a/shard/resolver/validation.go b/shard/resolver/validation.go index c70966f..5a9b484 100644 --- a/shard/resolver/validation.go +++ b/shard/resolver/validation.go @@ -3,6 +3,7 @@ package resolver import ( "errors" + "github.com/mkbeh/xpg/cluster" "github.com/mkbeh/xpg/shard" ) @@ -14,7 +15,7 @@ func requireTopology(topology *shard.Topology) error { return nil } -func requireShardID(id shard.ID) error { +func requireShardID(id cluster.ID) error { if id == "" { return errors.New("shard ID must not be empty") } diff --git a/shard/shard.go b/shard/shard.go index c3afb65..32f2bb7 100644 --- a/shard/shard.go +++ b/shard/shard.go @@ -8,19 +8,19 @@ import ( "github.com/mkbeh/xpg/cluster" ) -// ID identifies one logical shard. -type ID = cluster.ID - -// Shard is a borrowed handle to one cluster registered in a Topology. +// Shard is an immutable, restricted view of one Cluster registered in a +// Topology. // -// Shard exposes shard-local operations without exposing cluster lifecycle or -// replica-set management. +// Shard exposes shard-local data access without exposing cluster lifecycle or +// replica-set management. A Shard does not own the underlying Cluster and is +// valid only for the lifetime of its owning Topology; it must not be used after +// Topology.Close. type Shard struct { cluster *cluster.Cluster } -// ID returns the logical shard ID. -func (s Shard) ID() ID { +// ID returns the stable logical shard ID inherited from the underlying Cluster. +func (s Shard) ID() cluster.ID { if s.cluster == nil { return "" } @@ -46,10 +46,9 @@ func (s Shard) Labels() map[string]string { return s.cluster.Labels() } -// Primary returns the shard primary pool. -// -// Primary returns nil when no primary is configured. The returned pool is -// borrowed and must not be closed separately. +// Primary returns the shard primary pool, or nil when the underlying cluster +// has no primary configured. The returned pool is borrowed and must not be +// closed separately. func (s Shard) Primary() *xpg.Pool { if s.cluster == nil { return nil @@ -58,7 +57,7 @@ func (s Shard) Primary() *xpg.Pool { return s.cluster.Primary() } -// ReadPool returns a borrowed pool according to policy. +// ReadPool returns a borrowed pool for a read operation according to policy. func (s Shard) ReadPool( ctx context.Context, policy cluster.ReadPolicy, @@ -83,8 +82,7 @@ func (s Shard) InPrimaryTx( return s.cluster.InPrimaryTx(ctx, options, fn) } -// InReadTx executes fn in a read-only transaction on a pool selected according -// to policy within the shard. +// InReadTx executes fn in a read-only transaction resolved within this shard. func (s Shard) InReadTx( ctx context.Context, policy cluster.ReadPolicy, diff --git a/shard/shard_test.go b/shard/shard_test.go deleted file mode 100644 index 400822b..0000000 --- a/shard/shard_test.go +++ /dev/null @@ -1,97 +0,0 @@ -package shard - -import ( - "errors" - "testing" - - "github.com/jackc/pgx/v5" - "github.com/mkbeh/xpg/cluster" -) - -func TestShardZeroValue(t *testing.T) { - t.Parallel() - - var shard Shard - - if got := shard.ID(); got != "" { - t.Fatalf("ID() = %q, want empty", got) - } - - if value, ok := shard.Label("region"); ok || value != "" { - t.Fatalf("Label() = %q, %v; want empty, false", value, ok) - } - - if labels := shard.Labels(); labels != nil { - t.Fatalf("Labels() = %#v, want nil", labels) - } - - if primary := shard.Primary(); primary != nil { - t.Fatalf("Primary() = %p, want nil", primary) - } - - if _, err := shard.ReadPool(t.Context(), cluster.ReadPrimary); !errors.Is(err, ErrNoShard) { - t.Fatalf("ReadPool() error = %v, want ErrNoShard", err) - } - - if err := shard.InPrimaryTx(t.Context(), pgx.TxOptions{}, nil); !errors.Is(err, ErrNoShard) { - t.Fatalf("InPrimaryTx() error = %v, want ErrNoShard", err) - } - - if err := shard.InReadTx( - t.Context(), - cluster.ReadPrimary, - cluster.ReadTxOptions{}, - nil, - ); !errors.Is(err, ErrNoShard) { - t.Fatalf("InReadTx() error = %v, want ErrNoShard", err) - } -} - -func TestShardDelegatesClusterMetadataAndRouting(t *testing.T) { - t.Parallel() - - shardCluster := newTestCluster(t, "shard-a", map[string]string{ - "region": "eu-west", - "role": "", - }) - - topology, err := NewTopology([]Config{{Cluster: shardCluster}}) - if err != nil { - t.Fatalf("NewTopology() error = %v", err) - } - t.Cleanup(topology.Close) - - resolved := topology.At(0) - - if got, want := resolved.ID(), ID("shard-a"); got != want { - t.Fatalf("ID() = %q, want %q", got, want) - } - - if got, ok := resolved.Label("region"); !ok || got != "eu-west" { - t.Fatalf("Label(region) = %q, %v", got, ok) - } - - if got, ok := resolved.Label("role"); !ok || got != "" { - t.Fatalf("Label(role) = %q, %v", got, ok) - } - - labels := resolved.Labels() - labels["region"] = "changed" - - if got, _ := resolved.Label("region"); got != "eu-west" { - t.Fatalf("Label(region) after mutation = %q, want eu-west", got) - } - - if resolved.Primary() != shardCluster.Primary() { - t.Fatal("Primary() did not return cluster primary") - } - - pool, err := resolved.ReadPool(t.Context(), cluster.ReadPrimary) - if err != nil { - t.Fatalf("ReadPool() error = %v", err) - } - - if pool != shardCluster.Primary() { - t.Fatal("ReadPool() did not return cluster primary") - } -} diff --git a/shard/topology.go b/shard/topology.go index db7468a..3ee59ad 100644 --- a/shard/topology.go +++ b/shard/topology.go @@ -9,73 +9,49 @@ import ( "github.com/mkbeh/xpg/cluster" ) -// Config registers one Cluster as a logical shard. +// Topology is an immutable ordered set of logical shards. // -// The shard ID and labels are provided by Cluster. -type Config struct { - Cluster *cluster.Cluster -} - -// Topology is an immutable ordered set of logical shards and their PostgreSQL -// clusters. -// -// NewTopology takes ownership of all configured clusters only after it returns +// NewTopology takes ownership of all clusters only after it returns // successfully. Close closes every owned cluster exactly once. type Topology struct { - shards []Shard - indexByID map[ID]int + shards []Shard + shardsByID map[cluster.ID]Shard closeOnce sync.Once } -// NewTopology validates and creates an immutable topology. -// -// Shards retain their registration order. Every cluster must have a non-empty -// and unique ID. -func NewTopology(configs []Config) (*Topology, error) { - if len(configs) == 0 { - return nil, errors.New( - "xpg/shard: topology must contain at least one shard", - ) +// NewTopology validates and creates an immutable topology. Shards retain the +// cluster registration order. Every cluster must have a unique, non-empty ID. +func NewTopology(clusters ...*cluster.Cluster) (*Topology, error) { + if len(clusters) == 0 { + return nil, errors.New("xpg/shard: topology must contain at least one shard") } - shards := make([]Shard, len(configs)) - indexByID := make(map[ID]int, len(configs)) + shards := make([]Shard, len(clusters)) + shardsByID := make(map[cluster.ID]Shard, len(clusters)) - for index, config := range configs { - if config.Cluster == nil { - return nil, fmt.Errorf( - "xpg/shard: shard %d: cluster is nil", - index, - ) + for index, candidate := range clusters { + if candidate == nil { + return nil, fmt.Errorf("xpg/shard: shard %d: cluster is nil", index) } - id := config.Cluster.ID() + id := candidate.ID() if id == "" { - return nil, fmt.Errorf( - "xpg/shard: shard %d: cluster ID must not be empty", - index, - ) + return nil, fmt.Errorf("xpg/shard: shard %d: cluster ID must not be empty", index) } - if previousIndex, exists := indexByID[id]; exists { - return nil, fmt.Errorf( - "xpg/shard: duplicate shard ID %q at indexes %d and %d", - id, - previousIndex, - index, - ) + if _, exists := shardsByID[id]; exists { + return nil, fmt.Errorf("xpg/shard: duplicate shard ID %q", id) } - shards[index] = Shard{ - cluster: config.Cluster, - } - indexByID[id] = index + current := Shard{cluster: candidate} + shards[index] = current + shardsByID[id] = current } return &Topology{ - shards: shards, - indexByID: indexByID, + shards: shards, + shardsByID: shardsByID, }, nil } @@ -88,34 +64,30 @@ func (t *Topology) Len() int { return len(t.shards) } -// At returns the shard at index in registration order. -// -// At panics when t is nil or index is out of range. +// At returns the shard at index in registration order. At panics when index is +// outside the topology, matching ordinary slice indexing semantics. func (t *Topology) At(index int) Shard { return t.shards[index] } -// Shard returns the shard with id. -func (t *Topology) Shard(id ID) (Shard, bool) { +// Shards returns a defensive copy of shards in registration order. +func (t *Topology) Shards() []Shard { if t == nil { - return Shard{}, false - } - - index, ok := t.indexByID[id] - if !ok { - return Shard{}, false + return nil } - return t.shards[index], true + return slices.Clone(t.shards) } -// Shards returns a defensive copy of shards in registration order. -func (t *Topology) Shards() []Shard { +// Shard returns one shard by stable cluster ID. +func (t *Topology) Shard(id cluster.ID) (Shard, bool) { if t == nil { - return nil + return Shard{}, false } - return slices.Clone(t.shards) + resolved, ok := t.shardsByID[id] + + return resolved, ok } // Close closes owned clusters in reverse registration order. Close is safe to @@ -126,8 +98,8 @@ func (t *Topology) Close() { } t.closeOnce.Do(func() { - for _, shard := range slices.Backward(t.shards) { - shard.cluster.Close() + for _, current := range slices.Backward(t.shards) { + current.cluster.Close() } }) } diff --git a/shard/topology_test.go b/shard/topology_test.go deleted file mode 100644 index 81712e5..0000000 --- a/shard/topology_test.go +++ /dev/null @@ -1,162 +0,0 @@ -package shard - -import ( - "testing" -) - -func TestNewTopologyRequiresShard(t *testing.T) { - t.Parallel() - - topology, err := NewTopology(nil) - if err == nil { - topology.Close() - t.Fatal("expected error") - } - - if got, want := err.Error(), "xpg/shard: topology must contain at least one shard"; got != want { - t.Fatalf("error = %q, want %q", got, want) - } -} - -func TestNewTopologyRejectsNilCluster(t *testing.T) { - t.Parallel() - - topology, err := NewTopology([]Config{{}}) - if err == nil { - topology.Close() - t.Fatal("expected error") - } - - if got, want := err.Error(), "xpg/shard: shard 0: cluster is nil"; got != want { - t.Fatalf("error = %q, want %q", got, want) - } -} - -func TestNewTopologyRejectsEmptyClusterID(t *testing.T) { - t.Parallel() - - shardCluster := newTestCluster(t, "", nil) - topology, err := NewTopology([]Config{{Cluster: shardCluster}}) - if err == nil { - topology.Close() - t.Fatal("expected error") - } - - if got, want := err.Error(), "xpg/shard: shard 0: cluster ID must not be empty"; got != want { - t.Fatalf("error = %q, want %q", got, want) - } -} - -func TestNewTopologyRejectsDuplicateIDs(t *testing.T) { - t.Parallel() - - first := newTestCluster(t, "shard-a", nil) - second := newTestCluster(t, "shard-a", nil) - topology, err := NewTopology([]Config{ - {Cluster: first}, - {Cluster: second}, - }) - if err == nil { - topology.Close() - t.Fatal("expected error") - } - - if got, want := err.Error(), - `xpg/shard: duplicate shard ID "shard-a" at indexes 0 and 1`; got != want { - t.Fatalf("error = %q, want %q", got, want) - } -} - -func TestTopologyPreservesRegistrationOrder(t *testing.T) { - t.Parallel() - - topology := newTestTopology(t, "shard-b", "shard-a", "shard-c") - - if got, want := topology.Len(), 3; got != want { - t.Fatalf("Len() = %d, want %d", got, want) - } - - want := []ID{"shard-b", "shard-a", "shard-c"} - - for index, wantID := range want { - if got := topology.At(index).ID(); got != wantID { - t.Fatalf("At(%d).ID() = %q, want %q", index, got, wantID) - } - - resolved, ok := topology.Shard(wantID) - if !ok { - t.Fatalf("Shard(%q) not found", wantID) - } - - if got := resolved.ID(); got != wantID { - t.Fatalf("Shard(%q).ID() = %q", wantID, got) - } - } -} - -func TestTopologyShardsReturnsDefensiveCopy(t *testing.T) { - t.Parallel() - - topology := newTestTopology(t, "shard-a", "shard-b") - - shards := topology.Shards() - shards[0] = Shard{} - - if got, want := topology.At(0).ID(), ID("shard-a"); got != want { - t.Fatalf("At(0).ID() = %q, want %q", got, want) - } -} - -func TestTopologyShardUnknownID(t *testing.T) { - t.Parallel() - - topology := newTestTopology(t, "shard-a") - - resolved, ok := topology.Shard("missing") - if ok { - t.Fatalf("Shard() = %+v, true; want false", resolved) - } -} - -func TestTopologyAtPanicsOutOfRange(t *testing.T) { - t.Parallel() - - topology := newTestTopology(t, "shard-a") - - defer func() { - if recover() == nil { - t.Fatal("expected panic") - } - }() - - _ = topology.At(1) -} - -func TestTopologyNilReceiver(t *testing.T) { - t.Parallel() - - var topology *Topology - - if got := topology.Len(); got != 0 { - t.Fatalf("Len() = %d, want 0", got) - } - - if shards := topology.Shards(); shards != nil { - t.Fatalf("Shards() = %#v, want nil", shards) - } - - if resolved, ok := topology.Shard("shard-a"); ok || resolved.ID() != "" { - t.Fatalf("Shard() = %+v, %v; want zero, false", resolved, ok) - } - - topology.Close() -} - -func TestTopologyCloseIsIdempotent(t *testing.T) { - t.Parallel() - - topology := newTestTopology(t, "shard-a", "shard-b") - - topology.Close() - topology.Close() -} From 6d2c58f607c5aa2ac288b9f75f6e36e18ce49e36 Mon Sep 17 00:00:00 2001 From: mkbeh Date: Thu, 27 Aug 2026 13:52:05 +0300 Subject: [PATCH 2/6] refactor: topology package layout --- examples/cluster/main.go | 2 +- examples/cluster/setup.go | 2 +- examples/shard/main.go | 4 ++-- examples/shard/setup.go | 10 +++++----- examples/shard_geo/main.go | 2 +- examples/shard_geo/resolver.go | 6 +++--- examples/shard_geo/setup.go | 12 ++++++------ {cluster => topology/cluster}/cluster.go | 0 {cluster => topology/cluster}/cluster_test.go | 0 {cluster => topology/cluster}/doc.go | 0 {cluster => topology/cluster}/errors.go | 4 ++-- {cluster => topology/cluster}/helpers_test.go | 0 {cluster => topology/cluster}/resolver.go | 0 {cluster => topology/cluster}/routing_test.go | 0 {cluster => topology/cluster}/selector.go | 0 {cluster => topology/cluster}/selector_test.go | 0 {cluster => topology/cluster}/tx.go | 0 {cluster => topology/cluster}/tx_test.go | 0 {shard => topology/shard}/doc.go | 0 {shard => topology/shard}/errors.go | 14 ++++++-------- {shard => topology/shard}/foreach.go | 4 +--- {shard => topology/shard}/group.go | 4 +--- {shard => topology/shard}/resolver.go | 0 {shard => topology/shard}/resolver/custom.go | 5 ++--- {shard => topology/shard}/resolver/doc.go | 0 {shard => topology/shard}/resolver/encoder.go | 0 {shard => topology/shard}/resolver/range.go | 5 ++--- {shard => topology/shard}/resolver/rendezvous.go | 5 ++--- {shard => topology/shard}/resolver/time_range.go | 5 ++--- {shard => topology/shard}/resolver/validation.go | 5 ++--- {shard => topology/shard}/shard.go | 10 ++++++++-- {shard => topology/shard}/topology.go | 8 ++++---- 32 files changed, 51 insertions(+), 56 deletions(-) rename {cluster => topology/cluster}/cluster.go (100%) rename {cluster => topology/cluster}/cluster_test.go (100%) rename {cluster => topology/cluster}/doc.go (100%) rename {cluster => topology/cluster}/errors.go (61%) rename {cluster => topology/cluster}/helpers_test.go (100%) rename {cluster => topology/cluster}/resolver.go (100%) rename {cluster => topology/cluster}/routing_test.go (100%) rename {cluster => topology/cluster}/selector.go (100%) rename {cluster => topology/cluster}/selector_test.go (100%) rename {cluster => topology/cluster}/tx.go (100%) rename {cluster => topology/cluster}/tx_test.go (100%) rename {shard => topology/shard}/doc.go (100%) rename {shard => topology/shard}/errors.go (78%) rename {shard => topology/shard}/foreach.go (97%) rename {shard => topology/shard}/group.go (96%) rename {shard => topology/shard}/resolver.go (100%) rename {shard => topology/shard}/resolver/custom.go (92%) rename {shard => topology/shard}/resolver/doc.go (100%) rename {shard => topology/shard}/resolver/encoder.go (100%) rename {shard => topology/shard}/resolver/range.go (97%) rename {shard => topology/shard}/resolver/rendezvous.go (98%) rename {shard => topology/shard}/resolver/time_range.go (97%) rename {shard => topology/shard}/resolver/validation.go (74%) rename {shard => topology/shard}/shard.go (90%) rename {shard => topology/shard}/topology.go (92%) diff --git a/examples/cluster/main.go b/examples/cluster/main.go index 0d614cc..7782275 100644 --- a/examples/cluster/main.go +++ b/examples/cluster/main.go @@ -6,7 +6,7 @@ import ( "log" "github.com/jackc/pgx/v5" - "github.com/mkbeh/xpg/cluster" + "github.com/mkbeh/xpg/topology/cluster" ) type nodeInfo struct { diff --git a/examples/cluster/setup.go b/examples/cluster/setup.go index fe62855..f6db7cf 100644 --- a/examples/cluster/setup.go +++ b/examples/cluster/setup.go @@ -6,7 +6,7 @@ import ( "os" "github.com/mkbeh/xpg" - "github.com/mkbeh/xpg/cluster" + "github.com/mkbeh/xpg/topology/cluster" ) const ( diff --git a/examples/shard/main.go b/examples/shard/main.go index b2c000a..455c385 100644 --- a/examples/shard/main.go +++ b/examples/shard/main.go @@ -5,8 +5,8 @@ import ( "fmt" "log" - "github.com/mkbeh/xpg/shard" - "github.com/mkbeh/xpg/shard/resolver" + "github.com/mkbeh/xpg/topology/shard" + "github.com/mkbeh/xpg/topology/shard/resolver" ) const ( diff --git a/examples/shard/setup.go b/examples/shard/setup.go index a50b049..a2b6ba0 100644 --- a/examples/shard/setup.go +++ b/examples/shard/setup.go @@ -6,16 +6,16 @@ import ( "os" "github.com/mkbeh/xpg" - "github.com/mkbeh/xpg/cluster" - "github.com/mkbeh/xpg/shard" + "github.com/mkbeh/xpg/topology/cluster" + "github.com/mkbeh/xpg/topology/shard" ) const ( defaultShardADatabaseURL = "postgres://postgres:postgres@localhost:56431/postgres?sslmode=disable" defaultShardBDatabaseURL = "postgres://postgres:postgres@localhost:56432/postgres?sslmode=disable" - shardAID cluster.ID = "shard-a" - shardBID cluster.ID = "shard-b" + shardAID shard.ID = "shard-a" + shardBID shard.ID = "shard-b" ) func openTopology(ctx context.Context) (*shard.Topology, error) { @@ -65,7 +65,7 @@ func openTopology(ctx context.Context) (*shard.Topology, error) { func openCluster( ctx context.Context, - id cluster.ID, + id shard.ID, name string, databaseURL string, ) (*cluster.Cluster, error) { diff --git a/examples/shard_geo/main.go b/examples/shard_geo/main.go index f2d75fd..2df5352 100644 --- a/examples/shard_geo/main.go +++ b/examples/shard_geo/main.go @@ -7,7 +7,7 @@ import ( "log" "github.com/jackc/pgx/v5" - "github.com/mkbeh/xpg/shard" + "github.com/mkbeh/xpg/topology/shard" ) type tenantKey struct { diff --git a/examples/shard_geo/resolver.go b/examples/shard_geo/resolver.go index dca5e14..b314a14 100644 --- a/examples/shard_geo/resolver.go +++ b/examples/shard_geo/resolver.go @@ -3,8 +3,8 @@ package main import ( "fmt" - "github.com/mkbeh/xpg/shard" - "github.com/mkbeh/xpg/shard/resolver" + "github.com/mkbeh/xpg/topology/shard" + "github.com/mkbeh/xpg/topology/shard/resolver" ) func newTenantResolver(topology *shard.Topology) (shard.Resolver[tenantKey], error) { @@ -31,7 +31,7 @@ func newTenantResolver(topology *shard.Topology) (shard.Resolver[tenantKey], err } resolve := resolver.ResolveFunc[tenantKey]( - func(key tenantKey, _ *shard.Topology) (shard.ID, error) { + func(key tenantKey) (shard.ID, error) { id, ok := shardByRegion[key.Region] if !ok { return "", fmt.Errorf( diff --git a/examples/shard_geo/setup.go b/examples/shard_geo/setup.go index e189457..13d8f36 100644 --- a/examples/shard_geo/setup.go +++ b/examples/shard_geo/setup.go @@ -6,8 +6,8 @@ import ( "os" "github.com/mkbeh/xpg" - "github.com/mkbeh/xpg/cluster" - "github.com/mkbeh/xpg/shard" + "github.com/mkbeh/xpg/topology/cluster" + "github.com/mkbeh/xpg/topology/shard" ) const ( @@ -49,10 +49,10 @@ func openTopology(ctx context.Context) (*shard.Topology, error) { return nil, fmt.Errorf("open shard-us: %w", err) } - topology, err := shard.NewTopology([]shard.Config{ - {Cluster: shardEU}, - {Cluster: shardUS}, - }) + topology, err := shard.NewTopology( + shardEU, + shardUS, + ) if err != nil { shardUS.Close() shardEU.Close() diff --git a/cluster/cluster.go b/topology/cluster/cluster.go similarity index 100% rename from cluster/cluster.go rename to topology/cluster/cluster.go diff --git a/cluster/cluster_test.go b/topology/cluster/cluster_test.go similarity index 100% rename from cluster/cluster_test.go rename to topology/cluster/cluster_test.go diff --git a/cluster/doc.go b/topology/cluster/doc.go similarity index 100% rename from cluster/doc.go rename to topology/cluster/doc.go diff --git a/cluster/errors.go b/topology/cluster/errors.go similarity index 61% rename from cluster/errors.go rename to topology/cluster/errors.go index 03adc17..d99bcbf 100644 --- a/cluster/errors.go +++ b/topology/cluster/errors.go @@ -5,9 +5,9 @@ import "errors" var ( // ErrNoPrimary is returned when an operation requires a primary but none // is configured. - ErrNoPrimary = errors.New("xpg/cluster: no primary available") + ErrNoPrimary = errors.New("xpg/topology/cluster: no primary available") // ErrNoReplica is returned when an operation requires a replica but none // can be selected. - ErrNoReplica = errors.New("xpg/cluster: no replica available") + ErrNoReplica = errors.New("xpg/topology/cluster: no replica available") ) diff --git a/cluster/helpers_test.go b/topology/cluster/helpers_test.go similarity index 100% rename from cluster/helpers_test.go rename to topology/cluster/helpers_test.go diff --git a/cluster/resolver.go b/topology/cluster/resolver.go similarity index 100% rename from cluster/resolver.go rename to topology/cluster/resolver.go diff --git a/cluster/routing_test.go b/topology/cluster/routing_test.go similarity index 100% rename from cluster/routing_test.go rename to topology/cluster/routing_test.go diff --git a/cluster/selector.go b/topology/cluster/selector.go similarity index 100% rename from cluster/selector.go rename to topology/cluster/selector.go diff --git a/cluster/selector_test.go b/topology/cluster/selector_test.go similarity index 100% rename from cluster/selector_test.go rename to topology/cluster/selector_test.go diff --git a/cluster/tx.go b/topology/cluster/tx.go similarity index 100% rename from cluster/tx.go rename to topology/cluster/tx.go diff --git a/cluster/tx_test.go b/topology/cluster/tx_test.go similarity index 100% rename from cluster/tx_test.go rename to topology/cluster/tx_test.go diff --git a/shard/doc.go b/topology/shard/doc.go similarity index 100% rename from shard/doc.go rename to topology/shard/doc.go diff --git a/shard/errors.go b/topology/shard/errors.go similarity index 78% rename from shard/errors.go rename to topology/shard/errors.go index a6ae31d..0ab6c86 100644 --- a/shard/errors.go +++ b/topology/shard/errors.go @@ -3,26 +3,24 @@ package shard import ( "errors" "fmt" - - "github.com/mkbeh/xpg/cluster" ) var ( // ErrNoShard indicates that a resolver could not map a key to any shard. - ErrNoShard = errors.New("xpg/shard: no shard resolved") + ErrNoShard = errors.New("xpg/topology/shard: no shard resolved") // ErrUnknownShard indicates that routing configuration or custom routing // logic referenced a shard that does not exist in the topology. - ErrUnknownShard = errors.New("xpg/shard: unknown shard") + ErrUnknownShard = errors.New("xpg/topology/shard: unknown shard") // ErrShardMismatch indicates that keys expected to be colocated resolved to // different shards. - ErrShardMismatch = errors.New("xpg/shard: keys resolve to different shards") + ErrShardMismatch = errors.New("xpg/topology/shard: keys resolve to different shards") ) // UnknownShardError identifies a shard that does not exist in a topology. type UnknownShardError struct { - ShardID cluster.ID + ShardID ID } func (e *UnknownShardError) Error() string { @@ -36,8 +34,8 @@ func (e *UnknownShardError) Unwrap() error { // MismatchError describes the first key that resolved to a different shard // than the first key. type MismatchError struct { - Expected cluster.ID - Actual cluster.ID + Expected ID + Actual ID Index int } diff --git a/shard/foreach.go b/topology/shard/foreach.go similarity index 97% rename from shard/foreach.go rename to topology/shard/foreach.go index 9eef711..bb006e6 100644 --- a/shard/foreach.go +++ b/topology/shard/foreach.go @@ -5,13 +5,11 @@ import ( "errors" "fmt" "sync" - - "github.com/mkbeh/xpg/cluster" ) // ForEachShardResult contains the result of one shard callback invocation. type ForEachShardResult struct { - ShardID cluster.ID + ShardID ID Err error } diff --git a/shard/group.go b/topology/shard/group.go similarity index 96% rename from shard/group.go rename to topology/shard/group.go index 14796c4..66c4513 100644 --- a/shard/group.go +++ b/topology/shard/group.go @@ -3,8 +3,6 @@ package shard import ( "errors" "fmt" - - "github.com/mkbeh/xpg/cluster" ) // SameShard resolves the keys and verifies that they all belong to the same @@ -61,7 +59,7 @@ func GroupByShard[K any](resolver Resolver[K], keys []K) ([]Group[K], error) { } groups := make([]Group[K], 0) - indexByID := make(map[cluster.ID]int) + indexByID := make(map[ID]int) for keyIndex, key := range keys { resolved, err := resolver.Resolve(key) diff --git a/shard/resolver.go b/topology/shard/resolver.go similarity index 100% rename from shard/resolver.go rename to topology/shard/resolver.go diff --git a/shard/resolver/custom.go b/topology/shard/resolver/custom.go similarity index 92% rename from shard/resolver/custom.go rename to topology/shard/resolver/custom.go index 0e47de9..834c4cc 100644 --- a/shard/resolver/custom.go +++ b/topology/shard/resolver/custom.go @@ -3,8 +3,7 @@ package resolver import ( "errors" - "github.com/mkbeh/xpg/cluster" - "github.com/mkbeh/xpg/shard" + "github.com/mkbeh/xpg/topology/shard" ) // ResolveFunc maps an application key to a shard ID. @@ -12,7 +11,7 @@ import ( // Resolve functions should return shard.ErrNoShard when a key cannot be mapped // to a shard. Implementations shared by concurrent callers must be deterministic // and concurrency-safe. They should not perform hidden I/O. -type ResolveFunc[K any] func(key K) (cluster.ID, error) +type ResolveFunc[K any] func(key K) (shard.ID, error) // CustomResolver adapts ResolveFunc to shard.Resolver. type CustomResolver[K any] struct { diff --git a/shard/resolver/doc.go b/topology/shard/resolver/doc.go similarity index 100% rename from shard/resolver/doc.go rename to topology/shard/resolver/doc.go diff --git a/shard/resolver/encoder.go b/topology/shard/resolver/encoder.go similarity index 100% rename from shard/resolver/encoder.go rename to topology/shard/resolver/encoder.go diff --git a/shard/resolver/range.go b/topology/shard/resolver/range.go similarity index 97% rename from shard/resolver/range.go rename to topology/shard/resolver/range.go index 3a339cf..510e2d1 100644 --- a/shard/resolver/range.go +++ b/topology/shard/resolver/range.go @@ -7,8 +7,7 @@ import ( "slices" "sort" - "github.com/mkbeh/xpg/cluster" - "github.com/mkbeh/xpg/shard" + "github.com/mkbeh/xpg/topology/shard" ) // Range maps the bounded half-open interval [Start, End) to one shard. @@ -18,7 +17,7 @@ import ( type Range[K cmp.Ordered] struct { Start K End K - ShardID cluster.ID + ShardID shard.ID } // RangeResolver resolves ordered keys through bounded, non-overlapping ranges. diff --git a/shard/resolver/rendezvous.go b/topology/shard/resolver/rendezvous.go similarity index 98% rename from shard/resolver/rendezvous.go rename to topology/shard/resolver/rendezvous.go index 741b579..a564f0d 100644 --- a/shard/resolver/rendezvous.go +++ b/topology/shard/resolver/rendezvous.go @@ -8,8 +8,7 @@ import ( "fmt" "math" - "github.com/mkbeh/xpg/cluster" - "github.com/mkbeh/xpg/shard" + "github.com/mkbeh/xpg/topology/shard" ) const ( @@ -126,7 +125,7 @@ func (resolver *RendezvousResolver[K]) Resolve(key K) (shard.Shard, error) { var ( selected shard.Shard best [sha256.Size]byte - bestID cluster.ID + bestID shard.ID hasBest bool ) diff --git a/shard/resolver/time_range.go b/topology/shard/resolver/time_range.go similarity index 97% rename from shard/resolver/time_range.go rename to topology/shard/resolver/time_range.go index 52b5495..b522401 100644 --- a/shard/resolver/time_range.go +++ b/topology/shard/resolver/time_range.go @@ -7,8 +7,7 @@ import ( "sort" "time" - "github.com/mkbeh/xpg/cluster" - "github.com/mkbeh/xpg/shard" + "github.com/mkbeh/xpg/topology/shard" ) // TimeRange maps the bounded half-open interval [Start, End) to one shard. @@ -18,7 +17,7 @@ import ( type TimeRange struct { Start time.Time End time.Time - ShardID cluster.ID + ShardID shard.ID } // TimeRangeResolver resolves time instants through bounded, non-overlapping diff --git a/shard/resolver/validation.go b/topology/shard/resolver/validation.go similarity index 74% rename from shard/resolver/validation.go rename to topology/shard/resolver/validation.go index 5a9b484..b0913a0 100644 --- a/shard/resolver/validation.go +++ b/topology/shard/resolver/validation.go @@ -3,8 +3,7 @@ package resolver import ( "errors" - "github.com/mkbeh/xpg/cluster" - "github.com/mkbeh/xpg/shard" + "github.com/mkbeh/xpg/topology/shard" ) func requireTopology(topology *shard.Topology) error { @@ -15,7 +14,7 @@ func requireTopology(topology *shard.Topology) error { return nil } -func requireShardID(id cluster.ID) error { +func requireShardID(id shard.ID) error { if id == "" { return errors.New("shard ID must not be empty") } diff --git a/shard/shard.go b/topology/shard/shard.go similarity index 90% rename from shard/shard.go rename to topology/shard/shard.go index 32f2bb7..a0ce6bb 100644 --- a/shard/shard.go +++ b/topology/shard/shard.go @@ -5,9 +5,15 @@ import ( "github.com/jackc/pgx/v5" "github.com/mkbeh/xpg" - "github.com/mkbeh/xpg/cluster" + "github.com/mkbeh/xpg/topology/cluster" ) +// ID identifies one logical shard. +// +// ID is an alias of cluster.ID because a shard inherits the stable identity of +// its underlying cluster. +type ID = cluster.ID + // Shard is an immutable, restricted view of one Cluster registered in a // Topology. // @@ -20,7 +26,7 @@ type Shard struct { } // ID returns the stable logical shard ID inherited from the underlying Cluster. -func (s Shard) ID() cluster.ID { +func (s Shard) ID() ID { if s.cluster == nil { return "" } diff --git a/shard/topology.go b/topology/shard/topology.go similarity index 92% rename from shard/topology.go rename to topology/shard/topology.go index 3ee59ad..9fd1971 100644 --- a/shard/topology.go +++ b/topology/shard/topology.go @@ -6,7 +6,7 @@ import ( "slices" "sync" - "github.com/mkbeh/xpg/cluster" + "github.com/mkbeh/xpg/topology/cluster" ) // Topology is an immutable ordered set of logical shards. @@ -15,7 +15,7 @@ import ( // successfully. Close closes every owned cluster exactly once. type Topology struct { shards []Shard - shardsByID map[cluster.ID]Shard + shardsByID map[ID]Shard closeOnce sync.Once } @@ -28,7 +28,7 @@ func NewTopology(clusters ...*cluster.Cluster) (*Topology, error) { } shards := make([]Shard, len(clusters)) - shardsByID := make(map[cluster.ID]Shard, len(clusters)) + shardsByID := make(map[ID]Shard, len(clusters)) for index, candidate := range clusters { if candidate == nil { @@ -80,7 +80,7 @@ func (t *Topology) Shards() []Shard { } // Shard returns one shard by stable cluster ID. -func (t *Topology) Shard(id cluster.ID) (Shard, bool) { +func (t *Topology) Shard(id ID) (Shard, bool) { if t == nil { return Shard{}, false } From 57da7b823ab37c0a283e69a3c523d9a1cb1f1d15 Mon Sep 17 00:00:00 2001 From: mkbeh Date: Thu, 27 Aug 2026 15:07:13 +0300 Subject: [PATCH 3/6] refactor: simplify topology examples --- examples/basic/main.go | 11 +---- examples/cluster/setup.go | 87 +++++++++++++++++++++++-------------- examples/shard/main.go | 5 +-- examples/shard/setup.go | 78 ++++++++++++++++++++------------- examples/shard_geo/setup.go | 82 +++++++++++++++++++++------------- 5 files changed, 156 insertions(+), 107 deletions(-) diff --git a/examples/basic/main.go b/examples/basic/main.go index b3c1244..a965199 100644 --- a/examples/basic/main.go +++ b/examples/basic/main.go @@ -104,11 +104,7 @@ func upsertUsers(ctx context.Context, pool *xpg.Pool) (int64, error) { return tag.RowsAffected(), nil } -func loadUser( - ctx context.Context, - pool *xpg.Pool, - userID int64, -) (user, error) { +func loadUser(ctx context.Context, pool *xpg.Pool, userID int64) (user, error) { var selected user err := pool.QueryRow( @@ -134,10 +130,7 @@ func loadUser( return selected, nil } -func listActiveUsers( - ctx context.Context, - pool *xpg.Pool, -) ([]user, error) { +func listActiveUsers(ctx context.Context, pool *xpg.Pool) ([]user, error) { rows, err := pool.Query( ctx, `SELECT diff --git a/examples/cluster/setup.go b/examples/cluster/setup.go index f6db7cf..ac711c0 100644 --- a/examples/cluster/setup.go +++ b/examples/cluster/setup.go @@ -4,6 +4,7 @@ import ( "context" "fmt" "os" + "slices" "github.com/mkbeh/xpg" "github.com/mkbeh/xpg/topology/cluster" @@ -16,50 +17,64 @@ const ( ) func openCluster(ctx context.Context) (*cluster.Cluster, error) { - primary, err := openPool( - ctx, - environment("XPG_PRIMARY_DATABASE_URL", defaultPrimaryDatabaseURL), - "cluster.primary", - "primary", - ) - if err != nil { - return nil, fmt.Errorf("open primary pool: %w", err) + type poolConfig struct { + databaseURL string + name string + role string } - replicaOne, err := openPool( - ctx, - environment("XPG_REPLICA_ONE_DATABASE_URL", defaultReplicaOneURL), - "cluster.replica-one", - "replica", - ) - if err != nil { - primary.Close() - - return nil, fmt.Errorf("open replica-one pool: %w", err) + configs := []poolConfig{ + { + databaseURL: environment( + "XPG_PRIMARY_DATABASE_URL", + defaultPrimaryDatabaseURL, + ), + name: "cluster.primary", + role: "primary", + }, + { + databaseURL: environment( + "XPG_REPLICA_ONE_DATABASE_URL", + defaultReplicaOneURL, + ), + name: "cluster.replica-one", + role: "replica", + }, + { + databaseURL: environment( + "XPG_REPLICA_TWO_DATABASE_URL", + defaultReplicaTwoURL, + ), + name: "cluster.replica-two", + role: "replica", + }, } - replicaTwo, err := openPool( - ctx, - environment("XPG_REPLICA_TWO_DATABASE_URL", defaultReplicaTwoURL), - "cluster.replica-two", - "replica", - ) - if err != nil { - replicaOne.Close() - primary.Close() + pools := make([]*xpg.Pool, 0, len(configs)) + + for _, config := range configs { + pool, err := openPool( + ctx, + config.databaseURL, + config.name, + config.role, + ) + if err != nil { + closePools(pools) + + return nil, fmt.Errorf("open %s pool: %w", config.name, err) + } - return nil, fmt.Errorf("open replica-two pool: %w", err) + pools = append(pools, pool) } dbCluster, err := cluster.New(cluster.Config{ ID: "cluster-example", - Primary: primary, - Replicas: []*xpg.Pool{replicaOne, replicaTwo}, + Primary: pools[0], + Replicas: pools[1:], }) if err != nil { - replicaTwo.Close() - replicaOne.Close() - primary.Close() + closePools(pools) return nil, fmt.Errorf("create cluster: %w", err) } @@ -87,6 +102,12 @@ func openPool(ctx context.Context, databaseURL, name, role string) (*xpg.Pool, e return pool, nil } +func closePools(pools []*xpg.Pool) { + for _, pool := range slices.Backward(pools) { + pool.Close() + } +} + func environment(key, fallback string) string { if value := os.Getenv(key); value != "" { return value diff --git a/examples/shard/main.go b/examples/shard/main.go index 455c385..8706674 100644 --- a/examples/shard/main.go +++ b/examples/shard/main.go @@ -70,10 +70,7 @@ func run(ctx context.Context) error { primary := targetShard.Primary() if primary == nil { - return fmt.Errorf( - "shard %q has no primary", - targetShard.ID(), - ) + return fmt.Errorf("shard %q has no primary", targetShard.ID()) } if _, err := primary.Exec( diff --git a/examples/shard/setup.go b/examples/shard/setup.go index a2b6ba0..ebf6907 100644 --- a/examples/shard/setup.go +++ b/examples/shard/setup.go @@ -4,6 +4,7 @@ import ( "context" "fmt" "os" + "slices" "github.com/mkbeh/xpg" "github.com/mkbeh/xpg/topology/cluster" @@ -19,43 +20,54 @@ const ( ) func openTopology(ctx context.Context) (*shard.Topology, error) { - clusterA, err := openCluster( - ctx, - shardAID, - "shard.shard-a.primary", - environment( - "XPG_SHARD_A_DATABASE_URL", - defaultShardADatabaseURL, - ), - ) - if err != nil { - return nil, fmt.Errorf("open shard-a cluster: %w", err) + type clusterConfig struct { + id cluster.ID + name string + databaseURL string } - clusterB, err := openCluster( - ctx, - shardBID, - "shard.shard-b.primary", - environment( - "XPG_SHARD_B_DATABASE_URL", - defaultShardBDatabaseURL, - ), - ) - if err != nil { - clusterA.Close() + configs := []clusterConfig{ + { + id: shardAID, + name: "shard.shard-a.primary", + databaseURL: environment( + "XPG_SHARD_A_DATABASE_URL", + defaultShardADatabaseURL, + ), + }, + { + id: shardBID, + name: "shard.shard-b.primary", + databaseURL: environment( + "XPG_SHARD_B_DATABASE_URL", + defaultShardBDatabaseURL, + ), + }, + } - return nil, fmt.Errorf("open shard-b cluster: %w", err) + clusters := make([]*cluster.Cluster, 0, len(configs)) + + for _, config := range configs { + dbCluster, err := openCluster( + ctx, + config.id, + config.name, + config.databaseURL, + ) + if err != nil { + closeClusters(clusters) + + return nil, fmt.Errorf("open %s cluster: %w", config.id, err) + } + + clusters = append(clusters, dbCluster) } - // After NewTopology succeeds, the topology owns both clusters and closes + // After NewTopology succeeds, the topology owns all clusters and closes // them through Topology.Close. On constructor failure ownership remains here. - topology, err := shard.NewTopology( - clusterA, - clusterB, - ) + topology, err := shard.NewTopology(clusters...) if err != nil { - clusterB.Close() - clusterA.Close() + closeClusters(clusters) return nil, fmt.Errorf("create topology: %w", err) } @@ -99,6 +111,12 @@ func openCluster( return dbCluster, nil } +func closeClusters(clusters []*cluster.Cluster) { + for _, cluster := range slices.Backward(clusters) { + cluster.Close() + } +} + func environment(name, fallback string) string { if value := os.Getenv(name); value != "" { return value diff --git a/examples/shard_geo/setup.go b/examples/shard_geo/setup.go index 13d8f36..f850b38 100644 --- a/examples/shard_geo/setup.go +++ b/examples/shard_geo/setup.go @@ -4,6 +4,7 @@ import ( "context" "fmt" "os" + "slices" "github.com/mkbeh/xpg" "github.com/mkbeh/xpg/topology/cluster" @@ -19,43 +20,56 @@ const ( ) func openTopology(ctx context.Context) (*shard.Topology, error) { - shardEU, err := openShard( - ctx, - shardEUID, - "eu", - "geo.shard-eu.primary", - environment( - "XPG_SHARD_EU_DATABASE_URL", - defaultShardEUDatabaseURL, - ), - ) - if err != nil { - return nil, fmt.Errorf("open shard-eu: %w", err) + type shardConfig struct { + id cluster.ID + region string + name string + databaseURL string } - shardUS, err := openShard( - ctx, - shardUSID, - "us", - "geo.shard-us.primary", - environment( - "XPG_SHARD_US_DATABASE_URL", - defaultShardUSDatabaseURL, - ), - ) - if err != nil { - shardEU.Close() + configs := []shardConfig{ + { + id: shardEUID, + region: "eu", + name: "geo.shard-eu.primary", + databaseURL: environment( + "XPG_SHARD_EU_DATABASE_URL", + defaultShardEUDatabaseURL, + ), + }, + { + id: shardUSID, + region: "us", + name: "geo.shard-us.primary", + databaseURL: environment( + "XPG_SHARD_US_DATABASE_URL", + defaultShardUSDatabaseURL, + ), + }, + } + + clusters := make([]*cluster.Cluster, 0, len(configs)) + + for _, config := range configs { + dbCluster, err := openShard( + ctx, + config.id, + config.region, + config.name, + config.databaseURL, + ) + if err != nil { + closeClusters(clusters) + + return nil, fmt.Errorf("open %s shard: %w", config.id, err) + } - return nil, fmt.Errorf("open shard-us: %w", err) + clusters = append(clusters, dbCluster) } - topology, err := shard.NewTopology( - shardEU, - shardUS, - ) + topology, err := shard.NewTopology(clusters...) if err != nil { - shardUS.Close() - shardEU.Close() + closeClusters(clusters) return nil, fmt.Errorf("create topology: %w", err) } @@ -101,6 +115,12 @@ func openShard( return shardCluster, nil } +func closeClusters(clusters []*cluster.Cluster) { + for _, cluster := range slices.Backward(clusters) { + cluster.Close() + } +} + func environment(name, fallback string) string { if value := os.Getenv(name); value != "" { return value From 1efb1564816c69e3ce54e76dfb40b29fc5de581c Mon Sep 17 00:00:00 2001 From: mkbeh Date: Thu, 27 Aug 2026 16:12:08 +0300 Subject: [PATCH 4/6] refactor: cluster and shard topology --- topology/cluster/cluster.go | 10 +- topology/cluster/cluster_test.go | 35 +- topology/cluster/helpers_test.go | 9 +- topology/cluster/resolver.go | 10 +- topology/cluster/routing_test.go | 10 +- topology/cluster/selector.go | 2 +- topology/cluster/selector_test.go | 2 +- topology/cluster/tx.go | 2 +- topology/cluster/tx_test.go | 2 +- topology/shard/errors.go | 4 +- topology/shard/errors_test.go | 38 +++ topology/shard/foreach.go | 8 +- topology/shard/foreach_test.go | 361 +++++++++++++++++++++ topology/shard/group.go | 10 +- topology/shard/group_test.go | 254 +++++++++++++++ topology/shard/helpers_test.go | 70 ++++ topology/shard/resolver/custom.go | 4 +- topology/shard/resolver/custom_test.go | 173 ++++++++++ topology/shard/resolver/encoder.go | 2 +- topology/shard/resolver/encoder_test.go | 148 +++++++++ topology/shard/resolver/helpers_test.go | 67 ++++ topology/shard/resolver/range.go | 12 +- topology/shard/resolver/range_test.go | 275 ++++++++++++++++ topology/shard/resolver/rendezvous.go | 14 +- topology/shard/resolver/rendezvous_test.go | 255 +++++++++++++++ topology/shard/resolver/time_range.go | 12 +- topology/shard/resolver/time_range_test.go | 237 ++++++++++++++ topology/shard/resolver/validation.go | 2 +- topology/shard/shard_test.go | 97 ++++++ topology/shard/topology.go | 8 +- topology/shard/topology_test.go | 145 +++++++++ 31 files changed, 2215 insertions(+), 63 deletions(-) create mode 100644 topology/shard/errors_test.go create mode 100644 topology/shard/foreach_test.go create mode 100644 topology/shard/group_test.go create mode 100644 topology/shard/helpers_test.go create mode 100644 topology/shard/resolver/custom_test.go create mode 100644 topology/shard/resolver/encoder_test.go create mode 100644 topology/shard/resolver/helpers_test.go create mode 100644 topology/shard/resolver/range_test.go create mode 100644 topology/shard/resolver/rendezvous_test.go create mode 100644 topology/shard/resolver/time_range_test.go create mode 100644 topology/shard/shard_test.go create mode 100644 topology/shard/topology_test.go diff --git a/topology/cluster/cluster.go b/topology/cluster/cluster.go index 071622e..254127f 100644 --- a/topology/cluster/cluster.go +++ b/topology/cluster/cluster.go @@ -56,19 +56,19 @@ type Cluster struct { // selected using round-robin. func New(config Config) (*Cluster, error) { if config.ID == "" { - return nil, errors.New("xpg/cluster: cluster ID must not be empty") + return nil, errors.New("xpg/topology/cluster: cluster ID must not be empty") } if config.Primary != nil && config.Primary.Raw() == nil { - return nil, errors.New("xpg/cluster: primary pool is invalid") + return nil, errors.New("xpg/topology/cluster: primary pool is invalid") } if config.Primary == nil && len(config.Replicas) == 0 { - return nil, errors.New("xpg/cluster: at least one pool is required") + return nil, errors.New("xpg/topology/cluster: at least one pool is required") } if err := validateLabels(config.Labels); err != nil { - return nil, fmt.Errorf("xpg/cluster: %w", err) + return nil, fmt.Errorf("xpg/topology/cluster: %w", err) } replicas := slices.Clone(config.Replicas) @@ -76,7 +76,7 @@ func New(config Config) (*Cluster, error) { for index, replica := range replicas { if replica == nil || replica.Raw() == nil { - return nil, fmt.Errorf("xpg/cluster: replica %d is invalid", index) + return nil, fmt.Errorf("xpg/topology/cluster: replica %d is invalid", index) } metadata[index] = ReplicaInfo{ diff --git a/topology/cluster/cluster_test.go b/topology/cluster/cluster_test.go index 403d64b..17978ed 100644 --- a/topology/cluster/cluster_test.go +++ b/topology/cluster/cluster_test.go @@ -7,16 +7,34 @@ import ( "github.com/mkbeh/xpg" ) +func TestNewRequiresID(t *testing.T) { + t.Parallel() + + primary := newTestPool(t, "primary", nil) + + cluster, err := New(Config{ + Primary: primary, + }) + if err == nil { + cluster.Close() + t.Fatal("expected error") + } + + if got, want := err.Error(), "xpg/topology/cluster: cluster ID must not be empty"; got != want { + t.Fatalf("error = %q, want %q", got, want) + } +} + func TestNewRequiresPool(t *testing.T) { t.Parallel() - cluster, err := New(Config{}) + cluster, err := New(Config{ID: testClusterID}) if err == nil { cluster.Close() t.Fatal("expected error") } - if got, want := err.Error(), "xpg/cluster: at least one pool is required"; got != want { + if got, want := err.Error(), "xpg/topology/cluster: at least one pool is required"; got != want { t.Fatalf("error = %q, want %q", got, want) } } @@ -25,6 +43,7 @@ func TestNewRejectsInvalidPrimary(t *testing.T) { t.Parallel() cluster, err := New(Config{ + ID: testClusterID, Primary: &xpg.Pool{}, }) if err == nil { @@ -32,7 +51,7 @@ func TestNewRejectsInvalidPrimary(t *testing.T) { t.Fatal("expected error") } - if got, want := err.Error(), "xpg/cluster: primary pool is invalid"; got != want { + if got, want := err.Error(), "xpg/topology/cluster: primary pool is invalid"; got != want { t.Fatalf("error = %q, want %q", got, want) } } @@ -59,6 +78,7 @@ func TestNewRejectsInvalidReplica(t *testing.T) { t.Parallel() cluster, err := New(Config{ + ID: testClusterID, Replicas: []*xpg.Pool{test.replica}, }) if err == nil { @@ -66,7 +86,7 @@ func TestNewRejectsInvalidReplica(t *testing.T) { t.Fatal("expected error") } - if got, want := err.Error(), "xpg/cluster: replica 0 is invalid"; got != want { + if got, want := err.Error(), "xpg/topology/cluster: replica 0 is invalid"; got != want { t.Fatalf("error = %q, want %q", got, want) } }) @@ -130,6 +150,7 @@ func TestNewRejectsEmptyLabelKey(t *testing.T) { primary := newTestPool(t, "primary", nil) cluster, err := New(Config{ + ID: testClusterID, Labels: map[string]string{ "": "value", }, @@ -140,7 +161,7 @@ func TestNewRejectsEmptyLabelKey(t *testing.T) { t.Fatal("expected error") } - if got, want := err.Error(), "xpg/cluster: label key must not be empty"; got != want { + if got, want := err.Error(), "xpg/topology/cluster: label key must not be empty"; got != want { t.Fatalf("error = %q, want %q", got, want) } } @@ -153,6 +174,7 @@ func TestNewClonesReplicaSlice(t *testing.T) { replicas := []*xpg.Pool{replicaA} cluster, err := New(Config{ + ID: testClusterID, Replicas: replicas, }) if err != nil { @@ -173,6 +195,7 @@ func TestNewAllowsDuplicatePools(t *testing.T) { pool := newTestPool(t, "shared", nil) cluster, err := New(Config{ + ID: testClusterID, Primary: pool, Replicas: []*xpg.Pool{ pool, @@ -234,6 +257,7 @@ func TestNewCapturesReplicaMetadata(t *testing.T) { }) cluster, err := New(Config{ + ID: testClusterID, Replicas: []*xpg.Pool{replica}, Selector: selector, }) @@ -308,6 +332,7 @@ func TestCloseIsIdempotent(t *testing.T) { replica := newTestPool(t, "replica", nil) cluster, err := New(Config{ + ID: testClusterID, Primary: primary, Replicas: []*xpg.Pool{replica}, }) diff --git a/topology/cluster/helpers_test.go b/topology/cluster/helpers_test.go index 092a235..469163e 100644 --- a/topology/cluster/helpers_test.go +++ b/topology/cluster/helpers_test.go @@ -7,7 +7,10 @@ import ( "github.com/mkbeh/xpg" ) -const testDatabaseURL = "postgres://postgres:postgres@127.0.0.1:1/postgres?sslmode=disable" //nolint:gosec // Test-only DSN with non-production credentials. +const ( + testDatabaseURL = "postgres://postgres:postgres@127.0.0.1:1/postgres?sslmode=disable" //nolint:gosec // Test-only DSN with non-production credentials. + testClusterID = ID("test-cluster") +) func newTestPool(t *testing.T, name string, labels map[string]string) *xpg.Pool { t.Helper() @@ -40,6 +43,10 @@ func newTestPool(t *testing.T, name string, labels map[string]string) *xpg.Pool func newTestCluster(t *testing.T, config Config) *Cluster { t.Helper() + if config.ID == "" { + config.ID = testClusterID + } + cluster, err := New(config) if err != nil { t.Fatalf("New() error = %v", err) diff --git a/topology/cluster/resolver.go b/topology/cluster/resolver.go index 0d0a902..19ac9a8 100644 --- a/topology/cluster/resolver.go +++ b/topology/cluster/resolver.go @@ -41,7 +41,7 @@ func ParseReadPolicy(value string) (ReadPolicy, error) { return ReadReplicaRequired, nil default: return 0, fmt.Errorf( - "xpg/cluster: unknown read policy %q", + "xpg/topology/cluster: unknown read policy %q", value, ) } @@ -67,7 +67,7 @@ func (policy ReadPolicy) String() string { // selected. Other selector errors are returned to the caller. func (c *Cluster) ReadPool(ctx context.Context, policy ReadPolicy) (*xpg.Pool, error) { if c == nil { - return nil, errors.New("xpg/cluster: cluster is nil") + return nil, errors.New("xpg/topology/cluster: cluster is nil") } switch policy { @@ -90,7 +90,7 @@ func (c *Cluster) ReadPool(ctx context.Context, policy ReadPolicy) (*xpg.Pool, e return c.resolveReplica(ctx) default: - return nil, fmt.Errorf("xpg/cluster: unsupported read policy %d", policy) + return nil, fmt.Errorf("xpg/topology/cluster: unsupported read policy %d", policy) } } @@ -109,12 +109,12 @@ func (c *Cluster) resolveReplica(ctx context.Context) (*xpg.Pool, error) { index, err := c.selector.Select(ctx, c.metadata) if err != nil { - return nil, fmt.Errorf("xpg/cluster: select replica: %w", err) + return nil, fmt.Errorf("xpg/topology/cluster: select replica: %w", err) } if index < 0 || index >= len(c.replicas) { return nil, fmt.Errorf( - "xpg/cluster: replica selector returned invalid index %d for %d replicas", + "xpg/topology/cluster: replica selector returned invalid index %d for %d replicas", index, len(c.replicas), ) diff --git a/topology/cluster/routing_test.go b/topology/cluster/routing_test.go index bd08329..93bba97 100644 --- a/topology/cluster/routing_test.go +++ b/topology/cluster/routing_test.go @@ -45,7 +45,7 @@ func TestParseReadPolicyRejectsUnknown(t *testing.T) { t.Fatal("expected error") } - if got, want := err.Error(), `xpg/cluster: unknown read policy "nearest"`; got != want { + if got, want := err.Error(), `xpg/topology/cluster: unknown read policy "nearest"`; got != want { t.Fatalf("error = %q, want %q", got, want) } } @@ -84,7 +84,7 @@ func TestReadPoolNilCluster(t *testing.T) { t.Fatal("expected error") } - if got, want := err.Error(), "xpg/cluster: cluster is nil"; got != want { + if got, want := err.Error(), "xpg/topology/cluster: cluster is nil"; got != want { t.Fatalf("error = %q, want %q", got, want) } } @@ -330,7 +330,7 @@ func TestReadPoolRejectsSelectorIndex(t *testing.T) { } want := fmt.Sprintf( - "xpg/cluster: replica selector returned invalid index %d for 1 replicas", + "xpg/topology/cluster: replica selector returned invalid index %d for 1 replicas", test.index, ) @@ -364,7 +364,7 @@ func TestReadPoolPreservesSelectorError(t *testing.T) { t.Fatalf("ReadPool() error = %v, want wrapped selector error", err) } - if got, want := err.Error(), "xpg/cluster: select replica: boom"; got != want { + if got, want := err.Error(), "xpg/topology/cluster: select replica: boom"; got != want { t.Fatalf("error = %q, want %q", got, want) } } @@ -384,7 +384,7 @@ func TestReadPoolRejectsUnsupportedPolicy(t *testing.T) { t.Fatal("expected error") } - if got, want := err.Error(), "xpg/cluster: unsupported read policy 255"; got != want { + if got, want := err.Error(), "xpg/topology/cluster: unsupported read policy 255"; got != want { t.Fatalf("error = %q, want %q", got, want) } } diff --git a/topology/cluster/selector.go b/topology/cluster/selector.go index ab43dd6..e76f7b3 100644 --- a/topology/cluster/selector.go +++ b/topology/cluster/selector.go @@ -62,7 +62,7 @@ type ReplicaSelectorFunc func(context.Context, ReplicaSet) (int, error) // Select calls the wrapped selector function. func (selector ReplicaSelectorFunc) Select(ctx context.Context, replicas ReplicaSet) (int, error) { if selector == nil { - return -1, errors.New("xpg/cluster: replica selector function is nil") + return -1, errors.New("xpg/topology/cluster: replica selector function is nil") } return selector(ctx, replicas) diff --git a/topology/cluster/selector_test.go b/topology/cluster/selector_test.go index 348eefa..662f394 100644 --- a/topology/cluster/selector_test.go +++ b/topology/cluster/selector_test.go @@ -91,7 +91,7 @@ func TestReplicaSelectorFuncNil(t *testing.T) { t.Fatal("expected error") } - if got, want := err.Error(), "xpg/cluster: replica selector function is nil"; got != want { + if got, want := err.Error(), "xpg/topology/cluster: replica selector function is nil"; got != want { t.Fatalf("error = %q, want %q", got, want) } } diff --git a/topology/cluster/tx.go b/topology/cluster/tx.go index f2d5818..147842c 100644 --- a/topology/cluster/tx.go +++ b/topology/cluster/tx.go @@ -23,7 +23,7 @@ func (c *Cluster) InPrimaryTx( fn func(context.Context, pgx.Tx) error, ) error { if c == nil { - return errors.New("xpg/cluster: cluster is nil") + return errors.New("xpg/topology/cluster: cluster is nil") } pool, err := c.resolvePrimary() diff --git a/topology/cluster/tx_test.go b/topology/cluster/tx_test.go index 8478e22..c23f125 100644 --- a/topology/cluster/tx_test.go +++ b/topology/cluster/tx_test.go @@ -27,7 +27,7 @@ func TestInPrimaryTxNilCluster(t *testing.T) { t.Fatal("expected error") } - if got, want := err.Error(), "xpg/cluster: cluster is nil"; got != want { + if got, want := err.Error(), "xpg/topology/cluster: cluster is nil"; got != want { t.Fatalf("error = %q, want %q", got, want) } diff --git a/topology/shard/errors.go b/topology/shard/errors.go index 0ab6c86..711fe53 100644 --- a/topology/shard/errors.go +++ b/topology/shard/errors.go @@ -24,7 +24,7 @@ type UnknownShardError struct { } func (e *UnknownShardError) Error() string { - return fmt.Sprintf("xpg/shard: unknown shard %q", e.ShardID) + return fmt.Sprintf("xpg/topology/shard: unknown shard %q", e.ShardID) } func (e *UnknownShardError) Unwrap() error { @@ -41,7 +41,7 @@ type MismatchError struct { func (e *MismatchError) Error() string { return fmt.Sprintf( - "xpg/shard: key %d resolved to shard %q instead of %q", + "xpg/topology/shard: key %d resolved to shard %q instead of %q", e.Index, e.Actual, e.Expected, diff --git a/topology/shard/errors_test.go b/topology/shard/errors_test.go new file mode 100644 index 0000000..0e26a18 --- /dev/null +++ b/topology/shard/errors_test.go @@ -0,0 +1,38 @@ +package shard + +import ( + "errors" + "testing" +) + +func TestUnknownShardError(t *testing.T) { + t.Parallel() + + err := &UnknownShardError{ShardID: "missing"} + + if got, want := err.Error(), `xpg/topology/shard: unknown shard "missing"`; got != want { + t.Fatalf("Error() = %q, want %q", got, want) + } + + if !errors.Is(err, ErrUnknownShard) { + t.Fatal("errors.Is() = false, want ErrUnknownShard") + } +} + +func TestMismatchError(t *testing.T) { + t.Parallel() + + err := &MismatchError{ + Expected: "shard-a", + Actual: "shard-b", + Index: 2, + } + + if got, want := err.Error(), `xpg/topology/shard: key 2 resolved to shard "shard-b" instead of "shard-a"`; got != want { + t.Fatalf("Error() = %q, want %q", got, want) + } + + if !errors.Is(err, ErrShardMismatch) { + t.Fatal("errors.Is() = false, want ErrShardMismatch") + } +} diff --git a/topology/shard/foreach.go b/topology/shard/foreach.go index bb006e6..03b00d1 100644 --- a/topology/shard/foreach.go +++ b/topology/shard/foreach.go @@ -28,7 +28,7 @@ func (results ForEachShardResults) Err() error { errs = append( errs, fmt.Errorf( - "xpg/shard: shard %q callback: %w", + "xpg/topology/shard: shard %q callback: %w", result.ShardID, result.Err, ), @@ -51,15 +51,15 @@ func (t *Topology) ForEachShard( fn func(context.Context, Shard) error, ) (ForEachShardResults, error) { if t == nil || len(t.shards) == 0 { - return nil, errors.New("xpg/shard: topology is nil or empty") + return nil, errors.New("xpg/topology/shard: topology is nil or empty") } if concurrency <= 0 { - return nil, errors.New("xpg/shard: concurrency must be positive") + return nil, errors.New("xpg/topology/shard: concurrency must be positive") } if fn == nil { - return nil, errors.New("xpg/shard: callback is nil") + return nil, errors.New("xpg/topology/shard: callback is nil") } results := make(ForEachShardResults, len(t.shards)) diff --git a/topology/shard/foreach_test.go b/topology/shard/foreach_test.go new file mode 100644 index 0000000..a60c93f --- /dev/null +++ b/topology/shard/foreach_test.go @@ -0,0 +1,361 @@ +package shard + +import ( + "context" + "errors" + "sync/atomic" + "testing" + "time" +) + +func TestForEachShardValidatesArguments(t *testing.T) { + t.Parallel() + + topology := newTestTopology(t, "shard-a") + + tests := []struct { + name string + topology *Topology + concurrency int + fn func(context.Context, Shard) error + wantError string + }{ + { + name: "nil topology", + topology: nil, + concurrency: 1, + fn: func(context.Context, Shard) error { return nil }, + wantError: "xpg/topology/shard: topology is nil or empty", + }, + { + name: "empty topology", + topology: &Topology{}, + concurrency: 1, + fn: func(context.Context, Shard) error { return nil }, + wantError: "xpg/topology/shard: topology is nil or empty", + }, + { + name: "zero concurrency", + topology: topology, + concurrency: 0, + fn: func(context.Context, Shard) error { return nil }, + wantError: "xpg/topology/shard: concurrency must be positive", + }, + { + name: "nil callback", + topology: topology, + concurrency: 1, + fn: nil, + wantError: "xpg/topology/shard: callback is nil", + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + t.Parallel() + + _, err := test.topology.ForEachShard( + t.Context(), + test.concurrency, + test.fn, + ) + if err == nil { + t.Fatal("expected error") + } + + if got := err.Error(); got != test.wantError { + t.Fatalf("error = %q, want %q", got, test.wantError) + } + }) + } +} + +func TestForEachShardPreservesRegistrationOrder(t *testing.T) { + t.Parallel() + + topology := newTestTopology(t, "shard-c", "shard-a", "shard-b") + + results, err := topology.ForEachShard( + t.Context(), + 2, + func(context.Context, Shard) error { return nil }, + ) + if err != nil { + t.Fatalf("ForEachShard() error = %v", err) + } + + want := []ID{"shard-c", "shard-a", "shard-b"} + for index, wantID := range want { + if got := results[index].ShardID; got != wantID { + t.Fatalf("results[%d].ShardID = %q, want %q", index, got, wantID) + } + + if results[index].Err != nil { + t.Fatalf("results[%d].Err = %v", index, results[index].Err) + } + } +} + +func TestForEachShardHonorsConcurrencyLimit(t *testing.T) { + t.Parallel() + + topology := newTestTopology( + t, + "shard-a", + "shard-b", + "shard-c", + "shard-d", + "shard-e", + "shard-f", + ) + + ctx, cancel := context.WithTimeout(t.Context(), 5*time.Second) + defer cancel() + + started := make(chan struct{}, topology.Len()) + release := make(chan struct{}) + + var active atomic.Int32 + var maximum atomic.Int32 + var calls atomic.Int32 + + done := make(chan struct { + results ForEachShardResults + err error + }, 1) + + go func() { + results, err := topology.ForEachShard( + ctx, + 2, + func(ctx context.Context, _ Shard) error { + current := active.Add(1) + defer active.Add(-1) + + calls.Add(1) + + for { + observed := maximum.Load() + if current <= observed || maximum.CompareAndSwap(observed, current) { + break + } + } + + started <- struct{}{} + + select { + case <-release: + return nil + case <-ctx.Done(): + return ctx.Err() + } + }, + ) + + done <- struct { + results ForEachShardResults + err error + }{ + results: results, + err: err, + } + }() + + for range 2 { + select { + case <-started: + case <-ctx.Done(): + close(release) + t.Fatal("two callbacks did not start concurrently") + } + } + + close(release) + + var outcome struct { + results ForEachShardResults + err error + } + + select { + case outcome = <-done: + case <-ctx.Done(): + t.Fatal("ForEachShard() did not finish") + } + + if outcome.err != nil { + t.Fatalf("ForEachShard() error = %v", outcome.err) + } + + if got, want := calls.Load(), int32(topology.Len()); got != want { + t.Fatalf("callback calls = %d, want %d", got, want) + } + + if got, want := maximum.Load(), int32(2); got != want { + t.Fatalf("maximum concurrent callbacks = %d, want %d", got, want) + } + + if err := outcome.results.Err(); err != nil { + t.Fatalf("results.Err() = %v", err) + } +} + +func TestForEachShardCallbackErrorsDoNotStopOtherShards(t *testing.T) { + t.Parallel() + + topology := newTestTopology(t, "shard-a", "shard-b", "shard-c") + sentinel := errors.New("callback failed") + var calls atomic.Int32 + + results, err := topology.ForEachShard( + t.Context(), + 2, + func(_ context.Context, current Shard) error { + calls.Add(1) + if current.ID() == "shard-b" { + return sentinel + } + + return nil + }, + ) + if !errors.Is(err, sentinel) { + t.Fatalf("ForEachShard() error = %v, want wrapped sentinel", err) + } + + if got, want := calls.Load(), int32(3); got != want { + t.Fatalf("callback calls = %d, want %d", got, want) + } + + if results[0].Err != nil || !errors.Is(results[1].Err, sentinel) || results[2].Err != nil { + t.Fatalf("results = %+v", results) + } + + joined := results.Err() + if !errors.Is(joined, sentinel) { + t.Fatalf("results.Err() = %v, want wrapped sentinel", joined) + } + + if got, want := joined.Error(), `xpg/topology/shard: shard "shard-b" callback: callback failed`; got != want { + t.Fatalf("results.Err() = %q, want %q", got, want) + } + + if got, want := err.Error(), joined.Error(); got != want { + t.Fatalf("ForEachShard() error = %q, want %q", got, want) + } +} + +func TestForEachShardCanceledBeforeScheduling(t *testing.T) { + t.Parallel() + + topology := newTestTopology(t, "shard-a", "shard-b", "shard-c") + ctx, cancel := context.WithCancel(t.Context()) + cancel() + + var calls atomic.Int32 + + results, err := topology.ForEachShard( + ctx, + 2, + func(context.Context, Shard) error { + calls.Add(1) + return nil + }, + ) + if !errors.Is(err, context.Canceled) { + t.Fatalf("ForEachShard() error = %v, want context.Canceled", err) + } + + if got := calls.Load(); got != 0 { + t.Fatalf("callback calls = %d, want 0", got) + } + + for index, result := range results { + if !errors.Is(result.Err, context.Canceled) { + t.Fatalf("results[%d].Err = %v, want context.Canceled", index, result.Err) + } + } +} + +func TestForEachShardCancellationSkipsCallbacksNotStarted(t *testing.T) { + t.Parallel() + + topology := newTestTopology(t, "shard-a", "shard-b", "shard-c", "shard-d") + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + + started := make(chan struct{}) + var calls atomic.Int32 + + done := make(chan struct { + results ForEachShardResults + err error + }, 1) + + go func() { + results, err := topology.ForEachShard( + ctx, + 1, + func(ctx context.Context, _ Shard) error { + if calls.Add(1) == 1 { + close(started) + } + + <-ctx.Done() + return ctx.Err() + }, + ) + + done <- struct { + results ForEachShardResults + err error + }{results: results, err: err} + }() + + <-started + cancel() + + outcome := <-done + if !errors.Is(outcome.err, context.Canceled) { + t.Fatalf("ForEachShard() error = %v, want context.Canceled", outcome.err) + } + + if got := calls.Load(); got != 1 { + t.Fatalf("callback calls = %d, want 1", got) + } + + for index, result := range outcome.results { + if !errors.Is(result.Err, context.Canceled) { + t.Fatalf("results[%d].Err = %v, want context.Canceled", index, result.Err) + } + } +} + +func TestForEachShardResultsErr(t *testing.T) { + t.Parallel() + + first := errors.New("first") + second := errors.New("second") + + results := ForEachShardResults{ + {ShardID: "shard-a", Err: first}, + {ShardID: "shard-b"}, + {ShardID: "shard-c", Err: second}, + } + + err := results.Err() + if !errors.Is(err, first) || !errors.Is(err, second) { + t.Fatalf("Err() = %v, want both failures", err) + } + + want := "xpg/topology/shard: shard \"shard-a\" callback: first\n" + + "xpg/topology/shard: shard \"shard-c\" callback: second" + + if got := err.Error(); got != want { + t.Fatalf("Err() = %q, want %q", got, want) + } + + if err := (ForEachShardResults{{ShardID: "shard-a"}}).Err(); err != nil { + t.Fatalf("Err() = %v, want nil", err) + } +} diff --git a/topology/shard/group.go b/topology/shard/group.go index 66c4513..78b4cd8 100644 --- a/topology/shard/group.go +++ b/topology/shard/group.go @@ -9,7 +9,7 @@ import ( // shard. It returns that shard when all keys are colocated. func SameShard[K any](resolver Resolver[K], keys ...K) (Shard, error) { if resolver == nil { - return Shard{}, errors.New("xpg/shard: resolver is nil") + return Shard{}, errors.New("xpg/topology/shard: resolver is nil") } if len(keys) == 0 { @@ -18,7 +18,7 @@ func SameShard[K any](resolver Resolver[K], keys ...K) (Shard, error) { expected, err := resolver.Resolve(keys[0]) if err != nil { - return Shard{}, fmt.Errorf("xpg/shard: resolve key 0: %w", err) + return Shard{}, fmt.Errorf("xpg/topology/shard: resolve key 0: %w", err) } expectedID := expected.ID() @@ -26,7 +26,7 @@ func SameShard[K any](resolver Resolver[K], keys ...K) (Shard, error) { for index := 1; index < len(keys); index++ { actual, err := resolver.Resolve(keys[index]) if err != nil { - return Shard{}, fmt.Errorf("xpg/shard: resolve key %d: %w", index, err) + return Shard{}, fmt.Errorf("xpg/topology/shard: resolve key %d: %w", index, err) } actualID := actual.ID() @@ -55,7 +55,7 @@ type Group[K any] struct { // shard's first appearance in the input. func GroupByShard[K any](resolver Resolver[K], keys []K) ([]Group[K], error) { if resolver == nil { - return nil, errors.New("xpg/shard: resolver is nil") + return nil, errors.New("xpg/topology/shard: resolver is nil") } groups := make([]Group[K], 0) @@ -64,7 +64,7 @@ func GroupByShard[K any](resolver Resolver[K], keys []K) ([]Group[K], error) { for keyIndex, key := range keys { resolved, err := resolver.Resolve(key) if err != nil { - return nil, fmt.Errorf("xpg/shard: resolve key %d: %w", keyIndex, err) + return nil, fmt.Errorf("xpg/topology/shard: resolve key %d: %w", keyIndex, err) } id := resolved.ID() diff --git a/topology/shard/group_test.go b/topology/shard/group_test.go new file mode 100644 index 0000000..74d405f --- /dev/null +++ b/topology/shard/group_test.go @@ -0,0 +1,254 @@ +package shard + +import ( + "errors" + "slices" + "testing" +) + +func TestSameShard(t *testing.T) { + t.Parallel() + + topology := newTestTopology(t, "shard-a", "shard-b") + shardA := topology.At(0) + shardB := topology.At(1) + + resolver := testResolverFunc[int](func(key int) (Shard, error) { + if key < 100 { + return shardA, nil + } + + return shardB, nil + }) + + resolved, err := SameShard(resolver, 1, 2, 3) + if err != nil { + t.Fatalf("SameShard() error = %v", err) + } + + if got, want := resolved.ID(), ID("shard-a"); got != want { + t.Fatalf("SameShard().ID() = %q, want %q", got, want) + } +} + +func TestSameShardRejectsNilResolver(t *testing.T) { + t.Parallel() + + _, err := SameShard[int](nil, 1) + if err == nil { + t.Fatal("expected error") + } + + if got, want := err.Error(), "xpg/topology/shard: resolver is nil"; got != want { + t.Fatalf("error = %q, want %q", got, want) + } +} + +func TestSameShardRequiresKey(t *testing.T) { + t.Parallel() + + resolver := testResolverFunc[int](func(int) (Shard, error) { + t.Fatal("resolver should not be called") + return Shard{}, nil + }) + + _, err := SameShard(resolver) + if !errors.Is(err, ErrNoShard) { + t.Fatalf("error = %v, want ErrNoShard", err) + } +} + +func TestSameShardWrapsFirstResolveError(t *testing.T) { + t.Parallel() + + sentinel := errors.New("resolve failed") + resolver := testResolverFunc[int](func(int) (Shard, error) { + return Shard{}, sentinel + }) + + _, err := SameShard(resolver, 1) + if !errors.Is(err, sentinel) { + t.Fatalf("error = %v, want wrapped sentinel", err) + } + + if got, want := err.Error(), "xpg/topology/shard: resolve key 0: resolve failed"; got != want { + t.Fatalf("error = %q, want %q", got, want) + } +} + +func TestSameShardWrapsResolveErrorWithIndex(t *testing.T) { + t.Parallel() + + topology := newTestTopology(t, "shard-a") + resolved := topology.At(0) + sentinel := errors.New("resolve failed") + + resolver := testResolverFunc[int](func(key int) (Shard, error) { + if key == 2 { + return Shard{}, sentinel + } + + return resolved, nil + }) + + _, err := SameShard(resolver, 1, 2) + if !errors.Is(err, sentinel) { + t.Fatalf("error = %v, want wrapped sentinel", err) + } + + if got, want := err.Error(), "xpg/topology/shard: resolve key 1: resolve failed"; got != want { + t.Fatalf("error = %q, want %q", got, want) + } +} + +func TestSameShardReturnsMismatchDetails(t *testing.T) { + t.Parallel() + + topology := newTestTopology(t, "shard-a", "shard-b") + shardA := topology.At(0) + shardB := topology.At(1) + + resolver := testResolverFunc[int](func(key int) (Shard, error) { + if key == 3 { + return shardB, nil + } + + return shardA, nil + }) + + _, err := SameShard(resolver, 1, 2, 3) + if !errors.Is(err, ErrShardMismatch) { + t.Fatalf("error = %v, want ErrShardMismatch", err) + } + + var mismatch *MismatchError + if !errors.As(err, &mismatch) { + t.Fatalf("error = %T, want *MismatchError", err) + } + + if mismatch.Expected != "shard-a" || mismatch.Actual != "shard-b" || mismatch.Index != 2 { + t.Fatalf("mismatch = %+v", mismatch) + } +} + +func TestGroupByShardPreservesGroupAndKeyOrder(t *testing.T) { + t.Parallel() + + topology := newTestTopology(t, "shard-a", "shard-b") + shardA := topology.At(0) + shardB := topology.At(1) + + resolver := testResolverFunc[int](func(key int) (Shard, error) { + if key < 100 { + return shardA, nil + } + + return shardB, nil + }) + + groups, err := GroupByShard(resolver, []int{142, 42, 143, 43}) + if err != nil { + t.Fatalf("GroupByShard() error = %v", err) + } + + if got, want := len(groups), 2; got != want { + t.Fatalf("len(groups) = %d, want %d", got, want) + } + + if got, want := groups[0].Shard.ID(), ID("shard-b"); got != want { + t.Fatalf("groups[0].Shard.ID() = %q, want %q", got, want) + } + if got, want := groups[0].Keys, []int{142, 143}; !slices.Equal(got, want) { + t.Fatalf("groups[0].Keys = %v, want %v", got, want) + } + + if got, want := groups[1].Shard.ID(), ID("shard-a"); got != want { + t.Fatalf("groups[1].Shard.ID() = %q, want %q", got, want) + } + if got, want := groups[1].Keys, []int{42, 43}; !slices.Equal(got, want) { + t.Fatalf("groups[1].Keys = %v, want %v", got, want) + } +} + +func TestGroupByShardResolvesEachKeyOnce(t *testing.T) { + t.Parallel() + + topology := newTestTopology(t, "shard-a") + resolved := topology.At(0) + calls := 0 + + resolver := testResolverFunc[int](func(int) (Shard, error) { + calls++ + return resolved, nil + }) + + keys := []int{1, 2, 3, 4} + groups, err := GroupByShard(resolver, keys) + if err != nil { + t.Fatalf("GroupByShard() error = %v", err) + } + + if got, want := calls, len(keys); got != want { + t.Fatalf("resolve calls = %d, want %d", got, want) + } + + if got, want := len(groups), 1; got != want { + t.Fatalf("len(groups) = %d, want %d", got, want) + } +} + +func TestGroupByShardRejectsNilResolver(t *testing.T) { + t.Parallel() + + _, err := GroupByShard[int](nil, []int{1}) + if err == nil { + t.Fatal("expected error") + } + + if got, want := err.Error(), "xpg/topology/shard: resolver is nil"; got != want { + t.Fatalf("error = %q, want %q", got, want) + } +} + +func TestGroupByShardEmptyKeys(t *testing.T) { + t.Parallel() + + resolver := testResolverFunc[int](func(int) (Shard, error) { + t.Fatal("resolver should not be called") + return Shard{}, nil + }) + + groups, err := GroupByShard(resolver, nil) + if err != nil { + t.Fatalf("GroupByShard() error = %v", err) + } + + if got := len(groups); got != 0 { + t.Fatalf("len(groups) = %d, want 0", got) + } +} + +func TestGroupByShardWrapsResolveErrorWithIndex(t *testing.T) { + t.Parallel() + + topology := newTestTopology(t, "shard-a") + resolved := topology.At(0) + sentinel := errors.New("resolve failed") + + resolver := testResolverFunc[int](func(key int) (Shard, error) { + if key == 3 { + return Shard{}, sentinel + } + + return resolved, nil + }) + + _, err := GroupByShard(resolver, []int{1, 2, 3}) + if !errors.Is(err, sentinel) { + t.Fatalf("error = %v, want wrapped sentinel", err) + } + + if got, want := err.Error(), "xpg/topology/shard: resolve key 2: resolve failed"; got != want { + t.Fatalf("error = %q, want %q", got, want) + } +} diff --git a/topology/shard/helpers_test.go b/topology/shard/helpers_test.go new file mode 100644 index 0000000..63f980b --- /dev/null +++ b/topology/shard/helpers_test.go @@ -0,0 +1,70 @@ +package shard + +import ( + "testing" + + "github.com/jackc/pgx/v5/pgxpool" + "github.com/mkbeh/xpg" + "github.com/mkbeh/xpg/topology/cluster" +) + +const testDatabaseURL = "postgres://postgres@127.0.0.1:1/postgres?sslmode=disable" + +func newTestCluster(t *testing.T, id ID, labels map[string]string) *cluster.Cluster { + t.Helper() + + config, err := pgxpool.ParseConfig(testDatabaseURL) + if err != nil { + t.Fatalf("pgxpool.ParseConfig() error = %v", err) + } + + config.MinConns = 0 + config.MaxConns = 1 + + pool, err := xpg.New( + t.Context(), + config, + xpg.WithName("shard."+string(id)+".primary"), + ) + if err != nil { + t.Fatalf("xpg.New() error = %v", err) + } + + dbCluster, err := cluster.New(cluster.Config{ + ID: id, + Labels: labels, + Primary: pool, + }) + if err != nil { + pool.Close() + t.Fatalf("cluster.New() error = %v", err) + } + + t.Cleanup(dbCluster.Close) + + return dbCluster +} + +func newTestTopology(t *testing.T, ids ...ID) *Topology { + t.Helper() + + clusters := make([]*cluster.Cluster, len(ids)) + for index, id := range ids { + clusters[index] = newTestCluster(t, id, nil) + } + + topology, err := NewTopology(clusters...) + if err != nil { + t.Fatalf("NewTopology() error = %v", err) + } + + t.Cleanup(topology.Close) + + return topology +} + +type testResolverFunc[K any] func(K) (Shard, error) + +func (resolve testResolverFunc[K]) Resolve(key K) (Shard, error) { + return resolve(key) +} diff --git a/topology/shard/resolver/custom.go b/topology/shard/resolver/custom.go index 834c4cc..bd7b39c 100644 --- a/topology/shard/resolver/custom.go +++ b/topology/shard/resolver/custom.go @@ -26,7 +26,7 @@ func NewCustom[K any](topology *shard.Topology, resolve ResolveFunc[K]) (*Custom } if resolve == nil { - return nil, errors.New("xpg/shard/resolver: custom resolve function is nil") + return nil, errors.New("xpg/topology/shard/resolver: custom resolve function is nil") } return &CustomResolver[K]{ @@ -38,7 +38,7 @@ func NewCustom[K any](topology *shard.Topology, resolve ResolveFunc[K]) (*Custom // Resolve maps key to a shard and rejects IDs absent from the bound topology. func (resolver *CustomResolver[K]) Resolve(key K) (shard.Shard, error) { if resolver == nil || resolver.topology == nil || resolver.resolve == nil { - return shard.Shard{}, errors.New("xpg/shard/resolver: custom resolver is not initialized") + return shard.Shard{}, errors.New("xpg/topology/shard/resolver: custom resolver is not initialized") } id, err := resolver.resolve(key) diff --git a/topology/shard/resolver/custom_test.go b/topology/shard/resolver/custom_test.go new file mode 100644 index 0000000..04d6c53 --- /dev/null +++ b/topology/shard/resolver/custom_test.go @@ -0,0 +1,173 @@ +package resolver + +import ( + "errors" + "testing" + + "github.com/mkbeh/xpg/topology/shard" +) + +func TestNewCustomValidatesArguments(t *testing.T) { + t.Parallel() + + topology := newTestTopology(t, "shard-a") + + validResolve := ResolveFunc[int]( + func(int) (shard.ID, error) { + return "shard-a", nil + }, + ) + + tests := []struct { + name string + topology *shard.Topology + resolve ResolveFunc[int] + wantError string + }{ + { + name: "nil topology", + resolve: validResolve, + wantError: "xpg/topology/shard/resolver: topology is nil or empty", + }, + { + name: "nil resolve function", + topology: topology, + wantError: "xpg/topology/shard/resolver: custom resolve function is nil", + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + t.Parallel() + + _, err := NewCustom( + test.topology, + test.resolve, + ) + if err == nil { + t.Fatal("expected error") + } + + if got := err.Error(); got != test.wantError { + t.Fatalf("error = %q, want %q", got, test.wantError) + } + }) + } +} + +func TestCustomResolverResolve(t *testing.T) { + t.Parallel() + + topology := newTestTopology(t, "shard-a", "shard-b") + + resolver, err := NewCustom( + topology, + func(key int) (shard.ID, error) { + if key < 100 { + return "shard-a", nil + } + + return "shard-b", nil + }, + ) + if err != nil { + t.Fatalf("NewCustom() error = %v", err) + } + + resolved, err := resolver.Resolve(142) + if err != nil { + t.Fatalf("Resolve() error = %v", err) + } + + if got, want := resolved.ID(), shard.ID("shard-b"); got != want { + t.Fatalf("Resolve().ID() = %q, want %q", got, want) + } +} + +func TestCustomResolverPropagatesResolveError(t *testing.T) { + t.Parallel() + + topology := newTestTopology(t, "shard-a") + sentinel := errors.New("resolve failed") + + resolver, err := NewCustom( + topology, + func(int) (shard.ID, error) { + return "", sentinel + }, + ) + if err != nil { + t.Fatalf("NewCustom() error = %v", err) + } + + _, err = resolver.Resolve(1) + if !errors.Is(err, sentinel) { + t.Fatalf("Resolve() error = %v, want sentinel", err) + } +} + +func TestCustomResolverPropagatesErrNoShard(t *testing.T) { + t.Parallel() + + topology := newTestTopology(t, "shard-a") + + resolver, err := NewCustom( + topology, + func(int) (shard.ID, error) { + return "", shard.ErrNoShard + }, + ) + if err != nil { + t.Fatalf("NewCustom() error = %v", err) + } + + _, err = resolver.Resolve(1) + if !errors.Is(err, shard.ErrNoShard) { + t.Fatalf("Resolve() error = %v, want ErrNoShard", err) + } +} + +func TestCustomResolverRejectsUnknownShard(t *testing.T) { + t.Parallel() + + topology := newTestTopology(t, "shard-a") + + resolver, err := NewCustom( + topology, + func(int) (shard.ID, error) { + return "missing", nil + }, + ) + if err != nil { + t.Fatalf("NewCustom() error = %v", err) + } + + _, err = resolver.Resolve(1) + if !errors.Is(err, shard.ErrUnknownShard) { + t.Fatalf("Resolve() error = %v, want ErrUnknownShard", err) + } + + var unknown *shard.UnknownShardError + if !errors.As(err, &unknown) { + t.Fatalf("Resolve() error = %T, want *shard.UnknownShardError", err) + } + + if got, want := unknown.ShardID, shard.ID("missing"); got != want { + t.Fatalf("ShardID = %q, want %q", got, want) + } +} + +func TestCustomResolverUninitialized(t *testing.T) { + t.Parallel() + + var resolver *CustomResolver[int] + + _, err := resolver.Resolve(1) + if err == nil { + t.Fatal("expected error") + } + + if got, want := err.Error(), "xpg/topology/shard/resolver: custom resolver is not initialized"; got != want { + t.Fatalf("error = %q, want %q", got, want) + } +} diff --git a/topology/shard/resolver/encoder.go b/topology/shard/resolver/encoder.go index 995357e..fd98310 100644 --- a/topology/shard/resolver/encoder.go +++ b/topology/shard/resolver/encoder.go @@ -20,7 +20,7 @@ type KeyEncoderFunc[K any] func(K) ([]byte, error) // Encode calls the wrapped encoder function. func (encoder KeyEncoderFunc[K]) Encode(key K) ([]byte, error) { if encoder == nil { - return nil, errors.New("xpg/shard/resolver: key encoder function is nil") + return nil, errors.New("xpg/topology/shard/resolver: key encoder function is nil") } return encoder(key) diff --git a/topology/shard/resolver/encoder_test.go b/topology/shard/resolver/encoder_test.go new file mode 100644 index 0000000..cb78236 --- /dev/null +++ b/topology/shard/resolver/encoder_test.go @@ -0,0 +1,148 @@ +package resolver + +import ( + "bytes" + "encoding/hex" + "errors" + "testing" +) + +func TestKeyEncoderFuncNil(t *testing.T) { + t.Parallel() + + var encoder KeyEncoderFunc[int] + + _, err := encoder.Encode(1) + if err == nil { + t.Fatal("expected error") + } + + if got, want := err.Error(), "xpg/topology/shard/resolver: key encoder function is nil"; got != want { + t.Fatalf("error = %q, want %q", got, want) + } +} + +func TestKeyEncoderFunc(t *testing.T) { + t.Parallel() + + sentinel := errors.New("encode failed") + encoder := KeyEncoderFunc[int](func(key int) ([]byte, error) { + if key < 0 { + return nil, sentinel + } + + return []byte{byte(key)}, nil + }) + + encoded, err := encoder.Encode(7) + if err != nil { + t.Fatalf("Encode() error = %v", err) + } + if got, want := encoded, []byte{7}; !bytes.Equal(got, want) { + t.Fatalf("Encode() = %v, want %v", got, want) + } + + if _, err := encoder.Encode(-1); !errors.Is(err, sentinel) { + t.Fatalf("Encode() error = %v, want sentinel", err) + } +} + +func TestStringKeyEncoder(t *testing.T) { + t.Parallel() + + encoded, err := StringKeyEncoder().Encode("a\x00b") + if err != nil { + t.Fatalf("Encode() error = %v", err) + } + + if got, want := string(encoded), "a\x00b"; got != want { + t.Fatalf("Encode() = %q, want %q", got, want) + } +} + +func TestBytesKeyEncoderReturnsDefensiveCopy(t *testing.T) { + t.Parallel() + + key := []byte{1, 2, 3} + encoded, err := BytesKeyEncoder().Encode(key) + if err != nil { + t.Fatalf("Encode() error = %v", err) + } + + encoded[0] = 9 + + if got, want := key[0], byte(1); got != want { + t.Fatalf("input key changed to %d, want %d", got, want) + } +} + +func TestIntegerKeyEncoders(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + got func() ([]byte, error) + want string + }{ + { + name: "int64 positive", + got: func() ([]byte, error) { return Int64KeyEncoder().Encode(1) }, + want: "0000000000000001", + }, + { + name: "int64 negative", + got: func() ([]byte, error) { return Int64KeyEncoder().Encode(-1) }, + want: "ffffffffffffffff", + }, + { + name: "uint64", + got: func() ([]byte, error) { return Uint64KeyEncoder().Encode(0x0102030405060708) }, + want: "0102030405060708", + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + t.Parallel() + + encoded, err := test.got() + if err != nil { + t.Fatalf("Encode() error = %v", err) + } + + if got := hex.EncodeToString(encoded); got != test.want { + t.Fatalf("Encode() = %s, want %s", got, test.want) + } + }) + } +} + +func TestFixedSizeKeyEncoders(t *testing.T) { + t.Parallel() + + var key16 [16]byte + for index := range key16 { + key16[index] = byte(index) + } + + encoded16, err := Bytes16KeyEncoder().Encode(key16) + if err != nil { + t.Fatalf("Bytes16KeyEncoder.Encode() error = %v", err) + } + if got, want := encoded16, key16[:]; !bytes.Equal(got, want) { + t.Fatalf("Bytes16KeyEncoder.Encode() = %v, want %v", got, want) + } + + var key32 [32]byte + for index := range key32 { + key32[index] = byte(31 - index) + } + + encoded32, err := Bytes32KeyEncoder().Encode(key32) + if err != nil { + t.Fatalf("Bytes32KeyEncoder.Encode() error = %v", err) + } + if got, want := encoded32, key32[:]; !bytes.Equal(got, want) { + t.Fatalf("Bytes32KeyEncoder.Encode() = %v, want %v", got, want) + } +} diff --git a/topology/shard/resolver/helpers_test.go b/topology/shard/resolver/helpers_test.go new file mode 100644 index 0000000..2881ec7 --- /dev/null +++ b/topology/shard/resolver/helpers_test.go @@ -0,0 +1,67 @@ +package resolver + +import ( + "testing" + + "github.com/jackc/pgx/v5/pgxpool" + "github.com/mkbeh/xpg" + "github.com/mkbeh/xpg/topology/cluster" + "github.com/mkbeh/xpg/topology/shard" +) + +const testDatabaseURL = "postgres://postgres@127.0.0.1:1/postgres?sslmode=disable" + +func newTestTopology(t *testing.T, ids ...shard.ID) *shard.Topology { + t.Helper() + + clusters := make([]*cluster.Cluster, 0, len(ids)) + + closeClusters := func() { + for index := len(clusters) - 1; index >= 0; index-- { + clusters[index].Close() + } + } + + for _, id := range ids { + poolConfig, err := pgxpool.ParseConfig(testDatabaseURL) + if err != nil { + closeClusters() + t.Fatalf("pgxpool.ParseConfig() error = %v", err) + } + + poolConfig.MinConns = 0 + poolConfig.MaxConns = 1 + + pool, err := xpg.New( + t.Context(), + poolConfig, + xpg.WithName("shard."+string(id)+".primary"), + ) + if err != nil { + closeClusters() + t.Fatalf("xpg.New() error = %v", err) + } + + dbCluster, err := cluster.New(cluster.Config{ + ID: id, + Primary: pool, + }) + if err != nil { + pool.Close() + closeClusters() + t.Fatalf("cluster.New() error = %v", err) + } + + clusters = append(clusters, dbCluster) + } + + topology, err := shard.NewTopology(clusters...) + if err != nil { + closeClusters() + t.Fatalf("shard.NewTopology() error = %v", err) + } + + t.Cleanup(topology.Close) + + return topology +} diff --git a/topology/shard/resolver/range.go b/topology/shard/resolver/range.go index 510e2d1..9300ff9 100644 --- a/topology/shard/resolver/range.go +++ b/topology/shard/resolver/range.go @@ -44,24 +44,24 @@ func NewRange[K cmp.Ordered](topology *shard.Topology, ranges []Range[K]) (*Rang } if len(ranges) == 0 { - return nil, errors.New("xpg/shard/resolver: range resolver requires at least one range") + return nil, errors.New("xpg/topology/shard/resolver: range resolver requires at least one range") } entries := make([]rangeEntry[K], len(ranges)) for index, valueRange := range ranges { if err := requireShardID(valueRange.ShardID); err != nil { - return nil, fmt.Errorf("xpg/shard/resolver: range %d: %w", index, err) + return nil, fmt.Errorf("xpg/topology/shard/resolver: range %d: %w", index, err) } if !(valueRange.Start < valueRange.End) { //nolint:staticcheck // Negated comparison intentionally rejects NaN boundaries. - return nil, fmt.Errorf("xpg/shard/resolver: range %d must satisfy start < end", index) + return nil, fmt.Errorf("xpg/topology/shard/resolver: range %d must satisfy start < end", index) } resolved, ok := topology.Shard(valueRange.ShardID) if !ok { return nil, fmt.Errorf( - "xpg/shard/resolver: range %d: %w", + "xpg/topology/shard/resolver: range %d: %w", index, &shard.UnknownShardError{ShardID: valueRange.ShardID}, ) @@ -91,7 +91,7 @@ func NewRange[K cmp.Ordered](topology *shard.Topology, ranges []Range[K]) (*Rang } return nil, fmt.Errorf( - "xpg/shard/resolver: ranges %d and %d overlap", + "xpg/topology/shard/resolver: ranges %d and %d overlap", previous.sourceIndex, current.sourceIndex, ) @@ -108,7 +108,7 @@ func NewRange[K cmp.Ordered](topology *shard.Topology, ranges []Range[K]) (*Rang // acquire a connection, or execute a PostgreSQL query. func (resolver *RangeResolver[K]) Resolve(key K) (shard.Shard, error) { if resolver == nil || len(resolver.ranges) == 0 { - return shard.Shard{}, errors.New("xpg/shard/resolver: range resolver is not initialized") + return shard.Shard{}, errors.New("xpg/topology/shard/resolver: range resolver is not initialized") } // Non-overlap validation guarantees strictly increasing upper boundaries, diff --git a/topology/shard/resolver/range_test.go b/topology/shard/resolver/range_test.go new file mode 100644 index 0000000..9e78f28 --- /dev/null +++ b/topology/shard/resolver/range_test.go @@ -0,0 +1,275 @@ +package resolver + +import ( + "errors" + "math" + "slices" + "testing" + + "github.com/mkbeh/xpg/topology/shard" +) + +func TestNewRangeValidatesArguments(t *testing.T) { + t.Parallel() + + topology := newTestTopology(t, "shard-a") + + tests := []struct { + name string + topology *shard.Topology + ranges []Range[int] + wantError string + }{ + { + name: "nil topology", + ranges: []Range[int]{ + {Start: 0, End: 10, ShardID: "shard-a"}, + }, + wantError: "xpg/topology/shard/resolver: topology is nil or empty", + }, + { + name: "empty ranges", + topology: topology, + wantError: "xpg/topology/shard/resolver: range resolver requires at least one range", + }, + { + name: "empty shard ID", + topology: topology, + ranges: []Range[int]{ + {Start: 0, End: 10}, + }, + wantError: "xpg/topology/shard/resolver: range 0: shard ID must not be empty", + }, + { + name: "empty interval", + topology: topology, + ranges: []Range[int]{ + {Start: 10, End: 10, ShardID: "shard-a"}, + }, + wantError: "xpg/topology/shard/resolver: range 0 must satisfy start < end", + }, + { + name: "reversed interval", + topology: topology, + ranges: []Range[int]{ + {Start: 20, End: 10, ShardID: "shard-a"}, + }, + wantError: "xpg/topology/shard/resolver: range 0 must satisfy start < end", + }, + { + name: "unknown shard", + topology: topology, + ranges: []Range[int]{ + {Start: 0, End: 10, ShardID: "missing"}, + }, + wantError: `xpg/topology/shard/resolver: range 0: xpg/topology/shard: unknown shard "missing"`, + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + t.Parallel() + + _, err := NewRange( + test.topology, + test.ranges, + ) + if err == nil { + t.Fatal("expected error") + } + + if got := err.Error(); got != test.wantError { + t.Fatalf("error = %q, want %q", got, test.wantError) + } + }) + } +} + +func TestNewRangeRejectsNaNBoundaries(t *testing.T) { + t.Parallel() + + topology := newTestTopology(t, "shard-a") + + tests := []struct { + name string + valueRange Range[float64] + }{ + { + name: "NaN start", + valueRange: Range[float64]{ + Start: math.NaN(), + End: 10, + ShardID: "shard-a", + }, + }, + { + name: "NaN end", + valueRange: Range[float64]{ + Start: 0, + End: math.NaN(), + ShardID: "shard-a", + }, + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + t.Parallel() + + _, err := NewRange( + topology, + []Range[float64]{test.valueRange}, + ) + if err == nil { + t.Fatal("expected error") + } + + if got, want := err.Error(), + "xpg/topology/shard/resolver: range 0 must satisfy start < end"; got != want { + t.Fatalf("error = %q, want %q", got, want) + } + }) + } +} + +func TestNewRangeRejectsOverlapUsingSourceIndexes(t *testing.T) { + t.Parallel() + + topology := newTestTopology(t, "shard-a", "shard-b") + + _, err := NewRange(topology, []Range[int]{ + {Start: 100, End: 200, ShardID: "shard-b"}, + {Start: 50, End: 150, ShardID: "shard-a"}, + }) + if err == nil { + t.Fatal("expected overlap error") + } + + if got, want := err.Error(), "xpg/topology/shard/resolver: ranges 1 and 0 overlap"; got != want { + t.Fatalf("error = %q, want %q", got, want) + } +} + +func TestNewRangeDoesNotModifyInputOrder(t *testing.T) { + t.Parallel() + + topology := newTestTopology(t, "shard-a", "shard-b") + ranges := []Range[int]{ + {Start: 100, End: 200, ShardID: "shard-b"}, + {Start: 0, End: 100, ShardID: "shard-a"}, + } + want := slices.Clone(ranges) + + if _, err := NewRange(topology, ranges); err != nil { + t.Fatalf("NewRange() error = %v", err) + } + + if !slices.Equal(ranges, want) { + t.Fatalf("ranges = %+v, want unchanged %+v", ranges, want) + } +} + +func TestRangeResolverHalfOpenBoundariesAndGaps(t *testing.T) { + t.Parallel() + + topology := newTestTopology(t, "shard-a", "shard-b") + resolver, err := NewRange(topology, []Range[int]{ + {Start: 20, End: 30, ShardID: "shard-b"}, + {Start: 0, End: 10, ShardID: "shard-a"}, + }) + if err != nil { + t.Fatalf("NewRange() error = %v", err) + } + + tests := []struct { + key int + wantID shard.ID + wantErr error + }{ + {key: 0, wantID: "shard-a"}, + {key: 9, wantID: "shard-a"}, + {key: 10, wantErr: shard.ErrNoShard}, + {key: 19, wantErr: shard.ErrNoShard}, + {key: 20, wantID: "shard-b"}, + {key: 29, wantID: "shard-b"}, + {key: 30, wantErr: shard.ErrNoShard}, + } + + for _, test := range tests { + resolved, err := resolver.Resolve(test.key) + if test.wantErr != nil { + if !errors.Is(err, test.wantErr) { + t.Fatalf("Resolve(%d) error = %v, want %v", test.key, err, test.wantErr) + } + + continue + } + + if err != nil { + t.Fatalf("Resolve(%d) error = %v", test.key, err) + } + + if got := resolved.ID(); got != test.wantID { + t.Fatalf("Resolve(%d).ID() = %q, want %q", test.key, got, test.wantID) + } + } +} + +func TestRangeResolverAllowsAdjacentRanges(t *testing.T) { + t.Parallel() + + topology := newTestTopology(t, "shard-a", "shard-b") + resolver, err := NewRange(topology, []Range[int]{ + {Start: 0, End: 100, ShardID: "shard-a"}, + {Start: 100, End: 200, ShardID: "shard-b"}, + }) + if err != nil { + t.Fatalf("NewRange() error = %v", err) + } + + resolved, err := resolver.Resolve(100) + if err != nil { + t.Fatalf("Resolve() error = %v", err) + } + + if got, want := resolved.ID(), shard.ID("shard-b"); got != want { + t.Fatalf("Resolve().ID() = %q, want %q", got, want) + } +} + +func TestRangeResolverSupportsStrings(t *testing.T) { + t.Parallel() + + topology := newTestTopology(t, "shard-a", "shard-b") + resolver, err := NewRange(topology, []Range[string]{ + {Start: "a", End: "m", ShardID: "shard-a"}, + {Start: "m", End: "z", ShardID: "shard-b"}, + }) + if err != nil { + t.Fatalf("NewRange() error = %v", err) + } + + resolved, err := resolver.Resolve("m") + if err != nil { + t.Fatalf("Resolve() error = %v", err) + } + + if got, want := resolved.ID(), shard.ID("shard-b"); got != want { + t.Fatalf("Resolve().ID() = %q, want %q", got, want) + } +} + +func TestRangeResolverUninitialized(t *testing.T) { + t.Parallel() + + var resolver *RangeResolver[int] + + _, err := resolver.Resolve(1) + if err == nil { + t.Fatal("expected error") + } + + if got, want := err.Error(), "xpg/topology/shard/resolver: range resolver is not initialized"; got != want { + t.Fatalf("error = %q, want %q", got, want) + } +} diff --git a/topology/shard/resolver/rendezvous.go b/topology/shard/resolver/rendezvous.go index a564f0d..d34ea32 100644 --- a/topology/shard/resolver/rendezvous.go +++ b/topology/shard/resolver/rendezvous.go @@ -44,15 +44,15 @@ func NewRendezvous[K any]( } if encoder == nil { - return nil, errors.New("xpg/shard/resolver: key encoder is nil") + return nil, errors.New("xpg/topology/shard/resolver: key encoder is nil") } if namespace == "" { - return nil, errors.New("xpg/shard/resolver: rendezvous namespace must not be empty") + return nil, errors.New("xpg/topology/shard/resolver: rendezvous namespace must not be empty") } if uint64(len(namespace)) > uint64(math.MaxUint32) { - return nil, errors.New("xpg/shard/resolver: rendezvous namespace is too large") + return nil, errors.New("xpg/topology/shard/resolver: rendezvous namespace is too large") } shards := topology.Shards() @@ -62,7 +62,7 @@ func NewRendezvous[K any]( id := candidate.ID() if uint64(len(id)) > uint64(math.MaxUint32) { - return nil, errors.New("xpg/shard/resolver: shard ID is too large") + return nil, errors.New("xpg/topology/shard/resolver: shard ID is too large") } maxIDLength = max(maxIDLength, len(id)) @@ -94,16 +94,16 @@ func NewRendezvous[K any]( // Resolve maps key to a shard using rendezvous hashing. func (resolver *RendezvousResolver[K]) Resolve(key K) (shard.Shard, error) { if resolver == nil || len(resolver.shards) == 0 || resolver.encoder == nil { - return shard.Shard{}, errors.New("xpg/shard/resolver: rendezvous resolver is not initialized") + return shard.Shard{}, errors.New("xpg/topology/shard/resolver: rendezvous resolver is not initialized") } encoded, err := resolver.encoder.Encode(key) if err != nil { - return shard.Shard{}, fmt.Errorf("xpg/shard/resolver: encode rendezvous key: %w", err) + return shard.Shard{}, fmt.Errorf("xpg/topology/shard/resolver: encode rendezvous key: %w", err) } if uint64(len(encoded)) > uint64(math.MaxUint32) { - return shard.Shard{}, errors.New("xpg/shard/resolver: encoded key is too large") + return shard.Shard{}, errors.New("xpg/topology/shard/resolver: encoded key is too large") } keyLengthOffset := len(resolver.prefix) diff --git a/topology/shard/resolver/rendezvous_test.go b/topology/shard/resolver/rendezvous_test.go new file mode 100644 index 0000000..6f74dab --- /dev/null +++ b/topology/shard/resolver/rendezvous_test.go @@ -0,0 +1,255 @@ +package resolver + +import ( + "errors" + "fmt" + "testing" + + "github.com/mkbeh/xpg/topology/shard" +) + +func TestNewRendezvousValidatesArguments(t *testing.T) { + t.Parallel() + + topology := newTestTopology(t, "shard-a") + + tests := []struct { + name string + topology *shard.Topology + namespace string + encoder KeyEncoder[string] + wantError string + }{ + { + name: "nil topology", + namespace: "users", + encoder: StringKeyEncoder(), + wantError: "xpg/topology/shard/resolver: topology is nil or empty", + }, + { + name: "nil encoder", + topology: topology, + namespace: "users", + wantError: "xpg/topology/shard/resolver: key encoder is nil", + }, + { + name: "empty namespace", + topology: topology, + encoder: StringKeyEncoder(), + wantError: "xpg/topology/shard/resolver: rendezvous namespace must not be empty", + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + t.Parallel() + + _, err := NewRendezvous( + test.topology, + test.namespace, + test.encoder, + ) + if err == nil { + t.Fatal("expected error") + } + + if got := err.Error(); got != test.wantError { + t.Fatalf("error = %q, want %q", got, test.wantError) + } + }) + } +} + +func TestNewRendezvousTreatsNamespaceAsOpaqueNonEmptyString(t *testing.T) { + t.Parallel() + + topology := newTestTopology(t, "shard-a") + + for _, namespace := range []string{"users", " users ", " "} { + resolver, err := NewRendezvous(topology, namespace, StringKeyEncoder()) + if err != nil { + t.Fatalf("NewRendezvous(%q) error = %v", namespace, err) + } + + resolved, err := resolver.Resolve("alice") + if err != nil { + t.Fatalf("Resolve() error = %v", err) + } + + if got, want := resolved.ID(), shard.ID("shard-a"); got != want { + t.Fatalf("Resolve().ID() = %q, want %q", got, want) + } + } +} + +func TestRendezvousResolverStablePlacementVectors(t *testing.T) { + t.Parallel() + + topology := newTestTopology(t, "shard-a", "shard-b", "shard-c") + resolver, err := NewRendezvous(topology, "users", StringKeyEncoder()) + if err != nil { + t.Fatalf("NewRendezvous() error = %v", err) + } + + tests := []struct { + key string + want shard.ID + }{ + {key: "alice", want: "shard-a"}, + {key: "bob", want: "shard-b"}, + {key: "carol", want: "shard-b"}, + {key: "dave", want: "shard-b"}, + {key: "eve", want: "shard-c"}, + {key: "0", want: "shard-a"}, + {key: "1", want: "shard-b"}, + {key: "2", want: "shard-c"}, + } + + for _, test := range tests { + t.Run(test.key, func(t *testing.T) { + t.Parallel() + + resolved, err := resolver.Resolve(test.key) + if err != nil { + t.Fatalf("Resolve() error = %v", err) + } + + if got := resolved.ID(); got != test.want { + t.Fatalf( + "Resolve(%q).ID() = %q, want %q", + test.key, + got, + test.want, + ) + } + }) + } +} + +func TestRendezvousResolverPlacementDoesNotDependOnTopologyOrder(t *testing.T) { + t.Parallel() + + first := newTestTopology(t, "shard-a", "shard-b", "shard-c") + second := newTestTopology(t, "shard-c", "shard-a", "shard-b") + + firstResolver, err := NewRendezvous(first, "users", StringKeyEncoder()) + if err != nil { + t.Fatalf("NewRendezvous(first) error = %v", err) + } + + secondResolver, err := NewRendezvous(second, "users", StringKeyEncoder()) + if err != nil { + t.Fatalf("NewRendezvous(second) error = %v", err) + } + + for _, key := range []string{"alice", "bob", "carol", "dave", "eve", "user-123"} { + firstShard, err := firstResolver.Resolve(key) + if err != nil { + t.Fatalf("first Resolve(%q) error = %v", key, err) + } + + secondShard, err := secondResolver.Resolve(key) + if err != nil { + t.Fatalf("second Resolve(%q) error = %v", key, err) + } + + if firstShard.ID() != secondShard.ID() { + t.Fatalf( + "Resolve(%q) = %q and %q for different topology orders", + key, + firstShard.ID(), + secondShard.ID(), + ) + } + } +} + +func TestRendezvousResolverAddingShardOnlyMovesKeysToNewShard(t *testing.T) { + t.Parallel() + + before := newTestTopology(t, "shard-a", "shard-b") + after := newTestTopology(t, "shard-a", "shard-b", "shard-c") + + beforeResolver, err := NewRendezvous(before, "users", StringKeyEncoder()) + if err != nil { + t.Fatalf("NewRendezvous(before) error = %v", err) + } + + afterResolver, err := NewRendezvous(after, "users", StringKeyEncoder()) + if err != nil { + t.Fatalf("NewRendezvous(after) error = %v", err) + } + + moved := 0 + + for index := range 256 { + key := fmt.Sprintf("user-%d", index) + + previous, err := beforeResolver.Resolve(key) + if err != nil { + t.Fatalf("before Resolve(%q) error = %v", key, err) + } + + current, err := afterResolver.Resolve(key) + if err != nil { + t.Fatalf("after Resolve(%q) error = %v", key, err) + } + + if previous.ID() == current.ID() { + continue + } + + moved++ + + if current.ID() != "shard-c" { + t.Fatalf( + "Resolve(%q) moved from %q to existing shard %q", + key, + previous.ID(), + current.ID(), + ) + } + } + + if moved == 0 { + t.Fatal("expected at least one key to move to the new shard") + } +} + +func TestRendezvousResolverWrapsEncoderError(t *testing.T) { + t.Parallel() + + topology := newTestTopology(t, "shard-a") + sentinel := errors.New("encode failed") + + resolver, err := NewRendezvous( + topology, + "users", + KeyEncoderFunc[string](func(string) ([]byte, error) { + return nil, sentinel + }), + ) + if err != nil { + t.Fatalf("NewRendezvous() error = %v", err) + } + + _, err = resolver.Resolve("alice") + if !errors.Is(err, sentinel) { + t.Fatalf("Resolve() error = %v, want wrapped sentinel", err) + } +} + +func TestRendezvousResolverUninitialized(t *testing.T) { + t.Parallel() + + var resolver *RendezvousResolver[string] + + _, err := resolver.Resolve("alice") + if err == nil { + t.Fatal("expected error") + } + + if got, want := err.Error(), "xpg/topology/shard/resolver: rendezvous resolver is not initialized"; got != want { + t.Fatalf("error = %q, want %q", got, want) + } +} diff --git a/topology/shard/resolver/time_range.go b/topology/shard/resolver/time_range.go index b522401..d8d6972 100644 --- a/topology/shard/resolver/time_range.go +++ b/topology/shard/resolver/time_range.go @@ -46,27 +46,27 @@ func NewTimeRange(topology *shard.Topology, ranges []TimeRange) (*TimeRangeResol } if len(ranges) == 0 { - return nil, errors.New("xpg/shard/resolver: time range resolver requires at least one range") + return nil, errors.New("xpg/topology/shard/resolver: time range resolver requires at least one range") } entries := make([]timeRangeEntry, len(ranges)) for index, valueRange := range ranges { if err := requireShardID(valueRange.ShardID); err != nil { - return nil, fmt.Errorf("xpg/shard/resolver: time range %d: %w", index, err) + return nil, fmt.Errorf("xpg/topology/shard/resolver: time range %d: %w", index, err) } start := timeToUTC(valueRange.Start) end := timeToUTC(valueRange.End) if !start.Before(end) { - return nil, fmt.Errorf("xpg/shard/resolver: time range %d must satisfy start < end", index) + return nil, fmt.Errorf("xpg/topology/shard/resolver: time range %d must satisfy start < end", index) } resolved, ok := topology.Shard(valueRange.ShardID) if !ok { return nil, fmt.Errorf( - "xpg/shard/resolver: time range %d: %w", + "xpg/topology/shard/resolver: time range %d: %w", index, &shard.UnknownShardError{ShardID: valueRange.ShardID}, ) @@ -98,7 +98,7 @@ func NewTimeRange(topology *shard.Topology, ranges []TimeRange) (*TimeRangeResol // [00:00, 01:00) and [01:00, 02:00) if previous.end.After(current.start) { return nil, fmt.Errorf( - "xpg/shard/resolver: time ranges %d and %d overlap", + "xpg/topology/shard/resolver: time ranges %d and %d overlap", previous.sourceIndex, current.sourceIndex, ) @@ -116,7 +116,7 @@ func NewTimeRange(topology *shard.Topology, ranges []TimeRange) (*TimeRangeResol // acquire a connection, or execute a PostgreSQL query. func (resolver *TimeRangeResolver) Resolve(key time.Time) (shard.Shard, error) { if resolver == nil || len(resolver.ranges) == 0 { - return shard.Shard{}, errors.New("xpg/shard/resolver: time range resolver is not initialized") + return shard.Shard{}, errors.New("xpg/topology/shard/resolver: time range resolver is not initialized") } key = timeToUTC(key) diff --git a/topology/shard/resolver/time_range_test.go b/topology/shard/resolver/time_range_test.go new file mode 100644 index 0000000..13be909 --- /dev/null +++ b/topology/shard/resolver/time_range_test.go @@ -0,0 +1,237 @@ +package resolver + +import ( + "errors" + "slices" + "testing" + "time" + + "github.com/mkbeh/xpg/topology/shard" +) + +func TestNewTimeRangeValidatesArguments(t *testing.T) { + t.Parallel() + + topology := newTestTopology(t, "shard-a") + start := time.Date(2026, 1, 1, 0, 0, 0, 0, time.UTC) + end := start.Add(time.Hour) + + tests := []struct { + name string + topology *shard.Topology + ranges []TimeRange + wantError string + }{ + { + name: "nil topology", + ranges: []TimeRange{ + {Start: start, End: end, ShardID: "shard-a"}, + }, + wantError: "xpg/topology/shard/resolver: topology is nil or empty", + }, + { + name: "empty ranges", + topology: topology, + wantError: "xpg/topology/shard/resolver: time range resolver requires at least one range", + }, + { + name: "empty shard ID", + topology: topology, + ranges: []TimeRange{ + {Start: start, End: end}, + }, + wantError: "xpg/topology/shard/resolver: time range 0: shard ID must not be empty", + }, + { + name: "empty interval", + topology: topology, + ranges: []TimeRange{ + {Start: start, End: start, ShardID: "shard-a"}, + }, + wantError: "xpg/topology/shard/resolver: time range 0 must satisfy start < end", + }, + { + name: "reversed interval", + topology: topology, + ranges: []TimeRange{ + {Start: end, End: start, ShardID: "shard-a"}, + }, + wantError: "xpg/topology/shard/resolver: time range 0 must satisfy start < end", + }, + { + name: "unknown shard", + topology: topology, + ranges: []TimeRange{ + {Start: start, End: end, ShardID: "missing"}, + }, + wantError: `xpg/topology/shard/resolver: time range 0: xpg/topology/shard: unknown shard "missing"`, + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + t.Parallel() + + _, err := NewTimeRange( + test.topology, + test.ranges, + ) + if err == nil { + t.Fatal("expected error") + } + + if got := err.Error(); got != test.wantError { + t.Fatalf("error = %q, want %q", got, test.wantError) + } + }) + } +} + +func TestNewTimeRangeRejectsOverlapUsingSourceIndexes(t *testing.T) { + t.Parallel() + + topology := newTestTopology(t, "shard-a", "shard-b") + base := time.Date(2026, 1, 1, 0, 0, 0, 0, time.UTC) + + _, err := NewTimeRange(topology, []TimeRange{ + {Start: base.Add(2 * time.Hour), End: base.Add(4 * time.Hour), ShardID: "shard-b"}, + {Start: base.Add(time.Hour), End: base.Add(3 * time.Hour), ShardID: "shard-a"}, + }) + if err == nil { + t.Fatal("expected overlap error") + } + + if got, want := err.Error(), "xpg/topology/shard/resolver: time ranges 1 and 0 overlap"; got != want { + t.Fatalf("error = %q, want %q", got, want) + } +} + +func TestNewTimeRangeDoesNotModifyInput(t *testing.T) { + t.Parallel() + + topology := newTestTopology(t, "shard-a", "shard-b") + location := time.FixedZone("UTC+3", 3*60*60) + base := time.Date(2026, 1, 1, 0, 0, 0, 0, location) + ranges := []TimeRange{ + {Start: base.Add(time.Hour), End: base.Add(2 * time.Hour), ShardID: "shard-b"}, + {Start: base, End: base.Add(time.Hour), ShardID: "shard-a"}, + } + want := slices.Clone(ranges) + + if _, err := NewTimeRange(topology, ranges); err != nil { + t.Fatalf("NewTimeRange() error = %v", err) + } + + if !slices.Equal(ranges, want) { + t.Fatalf("ranges = %+v, want unchanged %+v", ranges, want) + } +} + +func TestTimeRangeResolverHalfOpenBoundariesAndGaps(t *testing.T) { + t.Parallel() + + topology := newTestTopology(t, "shard-a", "shard-b") + base := time.Date(2026, 1, 1, 0, 0, 0, 0, time.UTC) + resolver, err := NewTimeRange(topology, []TimeRange{ + {Start: base.Add(2 * time.Hour), End: base.Add(3 * time.Hour), ShardID: "shard-b"}, + {Start: base, End: base.Add(time.Hour), ShardID: "shard-a"}, + }) + if err != nil { + t.Fatalf("NewTimeRange() error = %v", err) + } + + tests := []struct { + key time.Time + wantID shard.ID + wantErr error + }{ + {key: base, wantID: "shard-a"}, + {key: base.Add(time.Hour - time.Nanosecond), wantID: "shard-a"}, + {key: base.Add(time.Hour), wantErr: shard.ErrNoShard}, + {key: base.Add(2 * time.Hour), wantID: "shard-b"}, + {key: base.Add(3 * time.Hour), wantErr: shard.ErrNoShard}, + } + + for _, test := range tests { + resolved, err := resolver.Resolve(test.key) + if test.wantErr != nil { + if !errors.Is(err, test.wantErr) { + t.Fatalf("Resolve(%v) error = %v, want %v", test.key, err, test.wantErr) + } + + continue + } + + if err != nil { + t.Fatalf("Resolve(%v) error = %v", test.key, err) + } + + if got := resolved.ID(); got != test.wantID { + t.Fatalf("Resolve(%v).ID() = %q, want %q", test.key, got, test.wantID) + } + } +} + +func TestTimeRangeResolverNormalizesToUTC(t *testing.T) { + t.Parallel() + + topology := newTestTopology(t, "shard-a") + start := time.Date(2026, 1, 1, 0, 0, 0, 0, time.UTC) + end := start.Add(time.Hour) + resolver, err := NewTimeRange(topology, []TimeRange{ + {Start: start, End: end, ShardID: "shard-a"}, + }) + if err != nil { + t.Fatalf("NewTimeRange() error = %v", err) + } + + location := time.FixedZone("UTC+3", 3*60*60) + key := start.Add(30 * time.Minute).In(location) + + resolved, err := resolver.Resolve(key) + if err != nil { + t.Fatalf("Resolve() error = %v", err) + } + + if got, want := resolved.ID(), shard.ID("shard-a"); got != want { + t.Fatalf("Resolve().ID() = %q, want %q", got, want) + } +} + +func TestTimeRangeResolverAllowsAdjacentRanges(t *testing.T) { + t.Parallel() + + topology := newTestTopology(t, "shard-a", "shard-b") + base := time.Date(2026, 1, 1, 0, 0, 0, 0, time.UTC) + resolver, err := NewTimeRange(topology, []TimeRange{ + {Start: base, End: base.Add(time.Hour), ShardID: "shard-a"}, + {Start: base.Add(time.Hour), End: base.Add(2 * time.Hour), ShardID: "shard-b"}, + }) + if err != nil { + t.Fatalf("NewTimeRange() error = %v", err) + } + + resolved, err := resolver.Resolve(base.Add(time.Hour)) + if err != nil { + t.Fatalf("Resolve() error = %v", err) + } + + if got, want := resolved.ID(), shard.ID("shard-b"); got != want { + t.Fatalf("Resolve().ID() = %q, want %q", got, want) + } +} + +func TestTimeRangeResolverUninitialized(t *testing.T) { + t.Parallel() + + var resolver *TimeRangeResolver + + _, err := resolver.Resolve(time.Now()) + if err == nil { + t.Fatal("expected error") + } + + if got, want := err.Error(), "xpg/topology/shard/resolver: time range resolver is not initialized"; got != want { + t.Fatalf("error = %q, want %q", got, want) + } +} diff --git a/topology/shard/resolver/validation.go b/topology/shard/resolver/validation.go index b0913a0..f81431c 100644 --- a/topology/shard/resolver/validation.go +++ b/topology/shard/resolver/validation.go @@ -8,7 +8,7 @@ import ( func requireTopology(topology *shard.Topology) error { if topology == nil || topology.Len() == 0 { - return errors.New("xpg/shard/resolver: topology is nil or empty") + return errors.New("xpg/topology/shard/resolver: topology is nil or empty") } return nil diff --git a/topology/shard/shard_test.go b/topology/shard/shard_test.go new file mode 100644 index 0000000..6afe989 --- /dev/null +++ b/topology/shard/shard_test.go @@ -0,0 +1,97 @@ +package shard + +import ( + "errors" + "testing" + + "github.com/jackc/pgx/v5" + "github.com/mkbeh/xpg/topology/cluster" +) + +func TestShardZeroValue(t *testing.T) { + t.Parallel() + + var shard Shard + + if got := shard.ID(); got != "" { + t.Fatalf("ID() = %q, want empty", got) + } + + if value, ok := shard.Label("region"); ok || value != "" { + t.Fatalf("Label() = %q, %v; want empty, false", value, ok) + } + + if labels := shard.Labels(); labels != nil { + t.Fatalf("Labels() = %#v, want nil", labels) + } + + if primary := shard.Primary(); primary != nil { + t.Fatalf("Primary() = %p, want nil", primary) + } + + if _, err := shard.ReadPool(t.Context(), cluster.ReadPrimary); !errors.Is(err, ErrNoShard) { + t.Fatalf("ReadPool() error = %v, want ErrNoShard", err) + } + + if err := shard.InPrimaryTx(t.Context(), pgx.TxOptions{}, nil); !errors.Is(err, ErrNoShard) { + t.Fatalf("InPrimaryTx() error = %v, want ErrNoShard", err) + } + + if err := shard.InReadTx( + t.Context(), + cluster.ReadPrimary, + cluster.ReadTxOptions{}, + nil, + ); !errors.Is(err, ErrNoShard) { + t.Fatalf("InReadTx() error = %v, want ErrNoShard", err) + } +} + +func TestShardDelegatesClusterMetadataAndRouting(t *testing.T) { + t.Parallel() + + dbCluster := newTestCluster(t, "shard-a", map[string]string{ + "region": "eu-west", + "role": "", + }) + + topology, err := NewTopology(dbCluster) + if err != nil { + t.Fatalf("NewTopology() error = %v", err) + } + t.Cleanup(topology.Close) + + resolved := topology.At(0) + + if got, want := resolved.ID(), ID("shard-a"); got != want { + t.Fatalf("ID() = %q, want %q", got, want) + } + + if got, ok := resolved.Label("region"); !ok || got != "eu-west" { + t.Fatalf("Label(region) = %q, %v", got, ok) + } + + if got, ok := resolved.Label("role"); !ok || got != "" { + t.Fatalf("Label(role) = %q, %v", got, ok) + } + + labels := resolved.Labels() + labels["region"] = "changed" + + if got, _ := resolved.Label("region"); got != "eu-west" { + t.Fatalf("Label(region) after mutation = %q, want eu-west", got) + } + + if resolved.Primary() != dbCluster.Primary() { + t.Fatal("Primary() did not return cluster primary") + } + + pool, err := resolved.ReadPool(t.Context(), cluster.ReadPrimary) + if err != nil { + t.Fatalf("ReadPool() error = %v", err) + } + + if pool != dbCluster.Primary() { + t.Fatal("ReadPool() did not return cluster primary") + } +} diff --git a/topology/shard/topology.go b/topology/shard/topology.go index 9fd1971..c3ea2e5 100644 --- a/topology/shard/topology.go +++ b/topology/shard/topology.go @@ -24,7 +24,7 @@ type Topology struct { // cluster registration order. Every cluster must have a unique, non-empty ID. func NewTopology(clusters ...*cluster.Cluster) (*Topology, error) { if len(clusters) == 0 { - return nil, errors.New("xpg/shard: topology must contain at least one shard") + return nil, errors.New("xpg/topology/shard: topology must contain at least one shard") } shards := make([]Shard, len(clusters)) @@ -32,16 +32,16 @@ func NewTopology(clusters ...*cluster.Cluster) (*Topology, error) { for index, candidate := range clusters { if candidate == nil { - return nil, fmt.Errorf("xpg/shard: shard %d: cluster is nil", index) + return nil, fmt.Errorf("xpg/topology/shard: shard %d: cluster is nil", index) } id := candidate.ID() if id == "" { - return nil, fmt.Errorf("xpg/shard: shard %d: cluster ID must not be empty", index) + return nil, fmt.Errorf("xpg/topology/shard: shard %d: cluster ID must not be empty", index) } if _, exists := shardsByID[id]; exists { - return nil, fmt.Errorf("xpg/shard: duplicate shard ID %q", id) + return nil, fmt.Errorf("xpg/topology/shard: duplicate shard ID %q", id) } current := Shard{cluster: candidate} diff --git a/topology/shard/topology_test.go b/topology/shard/topology_test.go new file mode 100644 index 0000000..4933b4b --- /dev/null +++ b/topology/shard/topology_test.go @@ -0,0 +1,145 @@ +package shard + +import "testing" + +func TestNewTopologyRequiresShard(t *testing.T) { + t.Parallel() + + topology, err := NewTopology() + if err == nil { + topology.Close() + t.Fatal("expected error") + } + + if got, want := err.Error(), "xpg/topology/shard: topology must contain at least one shard"; got != want { + t.Fatalf("error = %q, want %q", got, want) + } +} + +func TestNewTopologyRejectsNilCluster(t *testing.T) { + t.Parallel() + + topology, err := NewTopology(nil) + if err == nil { + topology.Close() + t.Fatal("expected error") + } + + if got, want := err.Error(), "xpg/topology/shard: shard 0: cluster is nil"; got != want { + t.Fatalf("error = %q, want %q", got, want) + } +} + +func TestNewTopologyRejectsDuplicateIDs(t *testing.T) { + t.Parallel() + + first := newTestCluster(t, "shard-a", nil) + second := newTestCluster(t, "shard-a", nil) + + topology, err := NewTopology(first, second) + if err == nil { + topology.Close() + t.Fatal("expected error") + } + + if got, want := err.Error(), `xpg/topology/shard: duplicate shard ID "shard-a"`; got != want { + t.Fatalf("error = %q, want %q", got, want) + } +} + +func TestTopologyPreservesRegistrationOrder(t *testing.T) { + t.Parallel() + + topology := newTestTopology(t, "shard-b", "shard-a", "shard-c") + + if got, want := topology.Len(), 3; got != want { + t.Fatalf("Len() = %d, want %d", got, want) + } + + want := []ID{"shard-b", "shard-a", "shard-c"} + for index, wantID := range want { + if got := topology.At(index).ID(); got != wantID { + t.Fatalf("At(%d).ID() = %q, want %q", index, got, wantID) + } + + resolved, ok := topology.Shard(wantID) + if !ok { + t.Fatalf("Shard(%q) not found", wantID) + } + + if got := resolved.ID(); got != wantID { + t.Fatalf("Shard(%q).ID() = %q, want %q", wantID, got, wantID) + } + } +} + +func TestTopologyShardsReturnsDefensiveCopy(t *testing.T) { + t.Parallel() + + topology := newTestTopology(t, "shard-a", "shard-b") + + shards := topology.Shards() + shards[0] = Shard{} + + if got, want := topology.At(0).ID(), ID("shard-a"); got != want { + t.Fatalf("At(0).ID() = %q, want %q", got, want) + } +} + +func TestTopologyShardUnknownID(t *testing.T) { + t.Parallel() + + topology := newTestTopology(t, "shard-a") + + resolved, ok := topology.Shard("missing") + if ok { + t.Fatalf("Shard() = %+v, true; want false", resolved) + } + + if got := resolved.ID(); got != "" { + t.Fatalf("Shard().ID() = %q, want empty", got) + } +} + +func TestTopologyAtPanicsOutOfRange(t *testing.T) { + t.Parallel() + + topology := newTestTopology(t, "shard-a") + + defer func() { + if recover() == nil { + t.Fatal("expected panic") + } + }() + + _ = topology.At(1) +} + +func TestTopologyNilReceiver(t *testing.T) { + t.Parallel() + + var topology *Topology + + if got := topology.Len(); got != 0 { + t.Fatalf("Len() = %d, want 0", got) + } + + if shards := topology.Shards(); shards != nil { + t.Fatalf("Shards() = %#v, want nil", shards) + } + + if resolved, ok := topology.Shard("shard-a"); ok || resolved.ID() != "" { + t.Fatalf("Shard() = %+v, %v; want zero, false", resolved, ok) + } + + topology.Close() +} + +func TestTopologyCloseIsIdempotent(t *testing.T) { + t.Parallel() + + topology := newTestTopology(t, "shard-a", "shard-b") + + topology.Close() + topology.Close() +} From 891054fd74c12f15e3a0a5d4f28532b916c27bd3 Mon Sep 17 00:00:00 2001 From: mkbeh Date: Thu, 27 Aug 2026 16:12:17 +0300 Subject: [PATCH 5/6] docs: update topology usage --- README.md | 49 ++++++++++++++++++++++++------------------------- 1 file changed, 24 insertions(+), 25 deletions(-) diff --git a/README.md b/README.md index 3c83c29..46d013b 100644 --- a/README.md +++ b/README.md @@ -26,10 +26,10 @@ connection management, routing, and common production workflows. * **Error Classification:** Classification of PostgreSQL constraint, transaction, cancellation, connection, and other common database errors. * **Advisory Locking:** Transaction-level advisory locks for coordinating concurrent database operations. -* **Primary/Replica Routing:** Explicit read policies, replica selection, primary fallback, and read-only transactions - across PostgreSQL nodes. -* **Application-Level Sharding:** Hash, range, time-based, and custom routing with colocation checks, key grouping, and - bounded parallel operations across shards. +* **Primary/Replica Routing:** Logical cluster topologies with explicit read policies, replica selection, primary + fallback, and read-only transactions across PostgreSQL nodes. +* **Application-Level Sharding:** Rendezvous, range, time-based, and custom routing with colocation checks, key + grouping, and bounded parallel operations across shards. * **Observability:** Structured logging, tracing, pool statistics, and optional OpenTelemetry metrics. ## Installation @@ -132,7 +132,7 @@ serialization failures, deadlocks, lock errors, query cancellation, and connecti ## Clustering -`xpg` groups primary and replica pools into a logical cluster with explicit read routing. +The `topology/cluster` package groups primary and replica pools into a logical cluster with explicit read routing. ```go @@ -173,61 +173,60 @@ replica is available. Replica selection is round-robin by default and can be cus ## Sharding -`xpg` provides application-level sharding with explicit key routing across an immutable shard topology. +The `topology/shard` package provides application-level sharding with explicit key routing across an immutable shard +topology. Routing strategies live under `topology/shard/resolver`. ```go -topology, err := shard.NewTopology([]shard.Config{ - {Cluster: shardA}, - {Cluster: shardB}, -}) +topology, err := shard.NewTopology(shardA, shardB) if err != nil { - panic(err) + panic(err) } defer topology.Close() // Partition user IDs into shard ranges. users, err := resolver.NewRange( - topology, - []resolver.Range[uint64]{ - {Start: 0, End: 100, ShardID: "shard-a"}, - {Start: 100, End: 200, ShardID: "shard-b"}, - }, + topology, + []resolver.Range[uint64]{ + {Start: 0, End: 100, ShardID: "shard-a"}, + {Start: 100, End: 200, ShardID: "shard-b"}, + }, ) if err != nil { - panic(err) + panic(err) } // Resolve the target shard. -shard, err := users.Resolve(userID) +targetShard, err := users.Resolve(userID) if err != nil { - panic(err) + panic(err) } // Write to the shard primary. -primaryPool := shard.Primary() +primaryPool := targetShard.Primary() _, err = primaryPool.Exec(ctx, "UPDATE users SET active = true WHERE id = $1", userID) if err != nil { - panic(err) + panic(err) } // Read from the same shard using the selected read policy. -readPool, err := shard.ReadPool(ctx, cluster.ReadReplicaPreferred) +readPool, err := targetShard.ReadPool(ctx, cluster.ReadReplicaPreferred) if err != nil { - panic(err) + panic(err) } var active bool err = readPool.QueryRow(ctx, "SELECT active FROM users WHERE id = $1", userID).Scan(&active) if err != nil { - panic(err) + panic(err) } + ``` Built-in routing strategies include rendezvous hashing, ordered ranges, time ranges, and custom resolvers. Sharding -utilities cover key colocation, grouping by shard, and parallel operations across shards. +utilities cover key colocation, grouping by shard, and bounded parallel operations across shards. ## Examples From 99e35cd0d90ef12c4f3caf612c6442b622060929 Mon Sep 17 00:00:00 2001 From: mkbeh Date: Thu, 27 Aug 2026 16:19:32 +0300 Subject: [PATCH 6/6] docs: update changelog for v0.3.0 --- CHANGELOG.md | 31 +++++++++++++++++++++++++++++++ 1 file changed, 31 insertions(+) diff --git a/CHANGELOG.md b/CHANGELOG.md index 621ec21..c15bc0e 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -2,6 +2,37 @@ All notable changes to this project will be documented in this file. +## v0.3.0 + +This release reorganizes the cluster and shard APIs under a common topology namespace and simplifies several sharding +contracts. + +### Changed + +* **Topology Package Layout:** Moved cluster and shard packages to `topology/cluster` and `topology/shard`, with shard + resolvers under `topology/shard/resolver`. +* **Cluster Identity:** Cluster IDs are now required when creating a `cluster.Cluster`. +* **Shard Topology Construction:** Simplified topology creation from `shard.NewTopology([]shard.Config{...})` to + `shard.NewTopology(clusters...)`. +* **Rendezvous Resolver:** Renamed `HashResolver` and `NewHash` to `RendezvousResolver` and `NewRendezvous`, making the + routing algorithm explicit while preserving the existing rendezvous placement contract. +* **Custom Resolvers:** Simplified custom resolver callbacks to map a key directly to `shard.ID` without receiving the + topology on every call. +* **Cross-Shard Operations:** `ForEachShard` now returns callback and cancellation failures through its function error + while preserving detailed per-shard results. + +### Fixed + +* **Rendezvous Portability:** Fixed length validation in rendezvous routing so the resolver also compiles correctly on + 32-bit architectures. + +### Removed + +* **Shard Configuration Layer:** Removed `shard.Config`; shard topologies are now created directly from clusters. +* **Generic Hash API:** Removed the `HashResolver` and `NewHash` names in favor of the explicit rendezvous API. + +--- + ## v0.2.0 Initial production release of `xpg`, built around `pgx` with PostgreSQL transaction helpers, primary/replica clustering,