diff --git a/crdt.go b/crdt.go index 4559dfd..a1737ef 100644 --- a/crdt.go +++ b/crdt.go @@ -4,6 +4,7 @@ import ( "bytes" "context" "encoding/gob" + "errors" "fmt" "maps" "strconv" @@ -209,8 +210,8 @@ func (c *CRDT) getEntry(ctx context.Context, reader datastore.Read, key string) } func (c *CRDT) buildSetEntry( - key string, value []byte, oldEntry KeyEntry, peerSeq uint64, -) (newEntry KeyEntry, wasPresent bool, newSeq uint64) { + key string, value []byte, oldEntry KeyEntry, peerNextSeq uint64, +) (newEntry KeyEntry, wasPresent bool) { wasPresent = oldEntry.Meta.CausalLength%2 == 1 meta := oldEntry.Meta @@ -221,20 +222,19 @@ func (c *CRDT) buildSetEntry( meta.ValueVersion++ } - newSeq = peerSeq + 1 meta.PeerID = c.PeerID - meta.PeerSeq = newSeq - return KeyEntry{Key: key, Value: value, Meta: meta}, wasPresent, newSeq + meta.PeerSeq = peerNextSeq + return KeyEntry{Key: key, Value: value, Meta: meta}, wasPresent } -func (c *CRDT) buildDeleteEntry(key string, oldEntry KeyEntry, peerSeq uint64) (newEntry KeyEntry, newSeq uint64) { + +func (c *CRDT) buildDeleteEntry(key string, oldEntry KeyEntry, peerNextSeq uint64) (newEntry KeyEntry) { meta := oldEntry.Meta meta.CausalLength++ meta.ValueVersion = 0 - newSeq = peerSeq + 1 meta.PeerID = c.PeerID - meta.PeerSeq = newSeq - return KeyEntry{Key: key, Value: []byte{}, Meta: meta}, newSeq + meta.PeerSeq = peerNextSeq + return KeyEntry{Key: key, Value: []byte{}, Meta: meta} } func (c *CRDT) applyEntry( @@ -307,8 +307,8 @@ func (c *CRDT) setWithTransaction(ctx context.Context, txnDs datastore.TxnDatast if bytes.Equal(entry.Value, value) { return 0, nil // No change needed } - - newEntry, wasPresent, newSeq := c.buildSetEntry(key, value, entry, peerSeq) + peerNextSeq := peerSeq + 1 + newEntry, wasPresent := c.buildSetEntry(key, value, entry, peerNextSeq) err = c.applyEntry(ctx, txn, entry, newEntry) if err != nil { @@ -336,12 +336,12 @@ func (c *CRDT) setWithTransaction(ctx context.Context, txnDs datastore.TxnDatast } } - return newSeq, nil + return peerNextSeq, nil }) } func (c *CRDT) setWithBatch(ctx context.Context, batchDs datastore.Batching, key string, value []byte, peerSeq uint64, afterCommit *func()) error { - // For batching, we need to read first, then batch writes + // For batching, we need to read first, then Batch writes entry, err := c.getEntry(ctx, c.ds, key) if err != nil && err != datastore.ErrNotFound { return err @@ -351,16 +351,17 @@ func (c *CRDT) setWithBatch(ctx context.Context, batchDs datastore.Batching, key return nil } - newEntry, wasPresent, newSeq := c.buildSetEntry(key, value, entry, peerSeq) + peerNextSeq := peerSeq + 1 + newEntry, wasPresent := c.buildSetEntry(key, value, entry, peerNextSeq) // Without transaction support, claim the sequence number first to prevent reuse - if err := c.saveTrackedPeer(ctx, c.ds, c.PeerID, newSeq); err != nil { + if err := c.saveTrackedPeer(ctx, c.ds, c.PeerID, peerNextSeq); err != nil { return err } batch, err := batchDs.Batch(ctx) if err != nil { - return fmt.Errorf("failed to create batch: %w", err) + return fmt.Errorf("failed to create Batch: %w", err) } err = c.applyEntry(ctx, batch, entry, newEntry) @@ -380,12 +381,12 @@ func (c *CRDT) setWithBatch(ctx context.Context, batchDs datastore.Batching, key } if err := batch.Commit(ctx); err != nil { - return fmt.Errorf("batch commit error: %w", err) + return fmt.Errorf("Batch commit error: %w", err) } - // If batch succeeded, update in-memory state - c.PeerSeq = newSeq - c.trackedPeers[c.PeerID] = newSeq + // If Batch succeeded, update in-memory state + c.PeerSeq = peerNextSeq + c.trackedPeers[c.PeerID] = peerNextSeq return nil } @@ -401,10 +402,11 @@ func (c *CRDT) setDirect(ctx context.Context, ds datastore.Datastore, key string return nil } - newEntry, wasPresent, newSeq := c.buildSetEntry(key, value, entry, peerSeq) + peerNextSeq := peerSeq + 1 + newEntry, wasPresent := c.buildSetEntry(key, value, entry, peerNextSeq) // Without transaction support, claim the sequence number first to prevent reuse - if err := c.saveTrackedPeer(ctx, c.ds, c.PeerID, newSeq); err != nil { + if err := c.saveTrackedPeer(ctx, c.ds, c.PeerID, peerNextSeq); err != nil { return err } @@ -425,8 +427,8 @@ func (c *CRDT) setDirect(ctx context.Context, ds datastore.Datastore, key string } // Update in-memory state - c.PeerSeq = peerSeq - c.trackedPeers[c.PeerID] = peerSeq + c.PeerSeq = peerNextSeq + c.trackedPeers[c.PeerID] = peerNextSeq return nil } @@ -477,7 +479,8 @@ func (c *CRDT) deleteWithTransaction(ctx context.Context, txnDs datastore.TxnDat meta := entry.Meta if meta.CausalLength%2 == 1 { - deletedEntry, newSeq := c.buildDeleteEntry(key, entry, peerSeq) + peerNextSeq := peerSeq + 1 + deletedEntry := c.buildDeleteEntry(key, entry, peerNextSeq) err := c.applyEntry(ctx, txn, entry, deletedEntry) if err != nil { @@ -492,11 +495,11 @@ func (c *CRDT) deleteWithTransaction(ctx context.Context, txnDs datastore.TxnDat } // post-commit hook - if hooks := c.getDeleteHooks(); len(c.deleteHooks) > 0 { + if hooks := c.getDeleteHooks(); len(hooks) > 0 { *afterCommit = func() { c.runDeleteHooks(key, entry.Value, entry.Meta, hooks) } } - return newSeq, nil + return peerNextSeq, nil } return 0, nil // No deletion needed }) @@ -510,16 +513,17 @@ func (c *CRDT) deleteWithBatch(ctx context.Context, batchDs datastore.Batching, meta := entry.Meta if meta.CausalLength%2 == 1 { - deletedEntry, newSeq := c.buildDeleteEntry(key, entry, peerSeq) + peerNextSeq := peerSeq + 1 + deletedEntry := c.buildDeleteEntry(key, entry, peerNextSeq) // Without transaction support, claim the sequence number first to prevent reuse - if err := c.saveTrackedPeer(ctx, c.ds, c.PeerID, newSeq); err != nil { + if err := c.saveTrackedPeer(ctx, c.ds, c.PeerID, peerNextSeq); err != nil { return err } batch, err := batchDs.Batch(ctx) if err != nil { - return fmt.Errorf("failed to create batch: %w", err) + return fmt.Errorf("failed to create `Batch`: %w", err) } err = c.applyEntry(ctx, batch, entry, deletedEntry) @@ -528,7 +532,7 @@ func (c *CRDT) deleteWithBatch(ctx context.Context, batchDs datastore.Batching, } if err := batch.Commit(ctx); err != nil { - return fmt.Errorf("batch commit error: %w", err) + return fmt.Errorf("Batch commit error: %w", err) } // post-commit hook @@ -536,9 +540,9 @@ func (c *CRDT) deleteWithBatch(ctx context.Context, batchDs datastore.Batching, *afterCommit = func() { c.runDeleteHooks(key, entry.Value, entry.Meta, hooks) } } - // If batch succeeded, update in-memory state. - c.PeerSeq = newSeq - c.trackedPeers[c.PeerID] = newSeq + // If Batch succeeded, update in-memory state. + c.PeerSeq = peerNextSeq + c.trackedPeers[c.PeerID] = peerNextSeq } return nil } @@ -551,10 +555,11 @@ func (c *CRDT) deleteDirect(ctx context.Context, ds datastore.Datastore, key str meta := entry.Meta if meta.CausalLength%2 == 1 { - deletedEntry, newSeq := c.buildDeleteEntry(key, entry, peerSeq) + peerNextSeq := peerSeq + 1 + deletedEntry := c.buildDeleteEntry(key, entry, peerNextSeq) // Without transaction support, claim the sequence number first to prevent reuse - if err := c.saveTrackedPeer(ctx, c.ds, c.PeerID, newSeq); err != nil { + if err := c.saveTrackedPeer(ctx, c.ds, c.PeerID, peerNextSeq); err != nil { return err } @@ -569,8 +574,8 @@ func (c *CRDT) deleteDirect(ctx context.Context, ds datastore.Datastore, key str } // Update in-memory state. - c.PeerSeq = newSeq - c.trackedPeers[c.PeerID] = newSeq + c.PeerSeq = peerNextSeq + c.trackedPeers[c.PeerID] = peerNextSeq } return nil } @@ -869,3 +874,192 @@ func (c *CRDT) MergeChanges(fromPeerID string, changes []KeyEntry, tracked map[s return nil } + +func (c *CRDT) Batch(_ context.Context) (batch *Batch, err error) { + switch c.ds.(type) { + case datastore.TxnDatastore, datastore.Batching: + return &Batch{parent: c}, nil + default: + return nil, errors.New("batch not implemented for this datastore type") + } +} + +func (c *CRDT) handleBatchCommit(ctx context.Context, batch *Batch) error { + c.mu.Lock() + defer c.mu.Unlock() + + var dsBatch datastore.Batch + + switch ds := c.ds.(type) { + case datastore.TxnDatastore: + txn, err := ds.NewTransaction(context.Background(), false) + if err != nil { + return fmt.Errorf("failed to create transaction: %w", err) + } + defer txn.Discard(ctx) + dsBatch = txn + case datastore.Batching: + dsBatchImpl, err := ds.Batch(ctx) + if err != nil { + return fmt.Errorf("failed to create batch: %w", err) + } + dsBatch = dsBatchImpl + default: + return errors.New("Batch not implemented for this datastore type") + } + + peerSeq, err := batch.commit(ctx, c.ds, dsBatch, c.PeerSeq) + if err != nil { + return fmt.Errorf("failed to execute batch function: %w", err) + } + + c.PeerSeq = peerSeq + c.trackedPeers[c.PeerID] = peerSeq + + return nil +} + +// Batch implements Batch operations using a transactional datastore. +type Batch struct { + // These fields are set initially. + parent *CRDT + inactive bool + + // These fields accumulate operations to be performed during Commit. + operations []func() error + afterCommit []func() + + // Those fields are not set initially, but lazily when batch is to be commited. + // They are used to provide a consistent view during batch operations. + // Prior to setting those fields, parent CRDT needs to be locked. + reader datastore.Read + batch datastore.Batch + lastPeerSeq uint64 +} + +func (b *Batch) Set(ctx context.Context, key string, value []byte) error { + b.operations = append(b.operations, func() error { + return b.set(ctx, key, value) + }) + return nil +} + +func (b *Batch) set(ctx context.Context, key string, value []byte) error { + // Read existing entry. + entry, err := b.parent.getEntry(ctx, b.reader, key) + if err != nil && err != datastore.ErrNotFound { + return fmt.Errorf("transactional Batch set: get entry error: %w", err) + } + + // Check if value is unchanged. + if bytes.Equal(entry.Value, value) { + return nil // No change needed + } + + // Build new entry. + b.lastPeerSeq++ + newEntry, wasPresent := b.parent.buildSetEntry(key, value, entry, b.lastPeerSeq) + + // Apply new entry to batch. + err = b.parent.applyEntry(ctx, b.batch, entry, newEntry) + if err != nil { + return fmt.Errorf("transactional Batch set: apply entry error: %w", err) + } + + // In case batch is provided by a transactional datastore, call transactional hooks. + if b.batch.(datastore.Txn) != nil { + if !wasPresent && b.parent.insertTxnHook != nil { + if err := b.parent.insertTxnHook(ctx, b.batch, key, value, newEntry.Meta); err != nil { + return fmt.Errorf("transactional Batch set: insert batchingDS hook error: %w", err) + } + } else if wasPresent && b.parent.updateTxnHook != nil { + if err := b.parent.updateTxnHook(ctx, b.batch, key, entry.Value, entry.Meta, value, newEntry.Meta); err != nil { + return fmt.Errorf("transactional Batch set: update batchingDS hook error: %w", err) + } + } + } + + // post-commit hook + if !wasPresent { + if hooks := b.parent.getInsertHooks(); len(hooks) > 0 { + afterCommit := func() { b.parent.runInsertHooks(key, value, newEntry.Meta, hooks) } + b.afterCommit = append(b.afterCommit, afterCommit) + } + } else { + if hooks := b.parent.getUpdateHooks(); len(hooks) > 0 { + afterCommit := func() { b.parent.runUpdateHooks(key, entry.Value, entry.Meta, value, newEntry.Meta, hooks) } + b.afterCommit = append(b.afterCommit, afterCommit) + } + } + + return nil +} + +func (b *Batch) Delete(ctx context.Context, key string) error { + b.operations = append(b.operations, func() error { + return b.delete(ctx, key) + }) + return nil +} + +func (b *Batch) delete(ctx context.Context, key string) error { + entry, err := b.parent.getEntry(ctx, b.reader, key) + if err != nil && err != datastore.ErrNotFound { + return fmt.Errorf("transactional Batch delete: get entry error: %w", err) + } + + meta := entry.Meta + if meta.CausalLength%2 == 1 { + b.lastPeerSeq++ + deletedEntry := b.parent.buildDeleteEntry(key, entry, b.lastPeerSeq) + + err := b.parent.applyEntry(ctx, b.batch, entry, deletedEntry) + if err != nil { + return fmt.Errorf("transactional Batch delete: apply entry error: %w", err) + } + + // transactional hook + if b.parent.deleteTxnHook != nil { + if err := b.parent.deleteTxnHook(ctx, b.batch, key, entry.Value, entry.Meta); err != nil { + return fmt.Errorf("transactional Batch delete: delete batchingDS hook error: %w", err) + } + } + + // post-commit hook + if hooks := b.parent.getDeleteHooks(); len(hooks) > 0 { + afterCommit := func() { b.parent.runDeleteHooks(key, entry.Value, entry.Meta, hooks) } + b.afterCommit = append(b.afterCommit, afterCommit) + } + + return nil + } + + return nil // No deletion needed +} + +func (b *Batch) Commit(ctx context.Context) error { + return b.parent.handleBatchCommit(ctx, b) +} + +func (b *Batch) commit(ctx context.Context, reader datastore.Read, batch datastore.Batch, peerSeq uint64) (newPeerSeq uint64, err error) { + if b.inactive { + return 0, errors.New("batch already commited or discarded") + } + b.inactive = true + b.reader = reader + b.batch = batch + b.lastPeerSeq = peerSeq + + for _, op := range b.operations { + if err := op(); err != nil { + return 0, fmt.Errorf("batch operation error: %w", err) + } + } + b.operations = nil + + if err := b.batch.Commit(ctx); err != nil { + return 0, fmt.Errorf("batch commit error: %w", err) + } + + return b.lastPeerSeq, nil +} diff --git a/crdt_test.go b/crdt_test.go index d1a90e0..f7781d9 100644 --- a/crdt_test.go +++ b/crdt_test.go @@ -1,6 +1,7 @@ package clset_test import ( + "context" "fmt" "testing" @@ -489,3 +490,163 @@ func TestCRDT_RemoteHooks_OnMerge(t *testing.T) { require.NoError(t, crdt2.MergeChanges("p1", changes, tracked)) assert.Contains(t, inserts, "foo:v3") } + +type batchOpType int + +const ( + OpSet batchOpType = iota + OpDelete +) + +type batchOperation struct { + Op batchOpType + Key string + Value []byte +} + +type testCase struct { + initial map[string][]byte + operations []batchOperation + expectToExist map[string][]byte + expectToNotExist []string + expectErr bool +} + +func TestBatch(t *testing.T) { + tests := map[string]testCase{ + "put multiple keys": { + operations: []batchOperation{ + {Op: OpSet, Key: "a", Value: []byte("A")}, + {Op: OpSet, Key: "b", Value: []byte("B")}, + {Op: OpSet, Key: "c", Value: []byte("C")}, + }, + expectToExist: map[string][]byte{ + "a": []byte("A"), + "b": []byte("B"), + "c": []byte("C"), + }, + }, + "delete multiple keys": { + initial: map[string][]byte{ + "x": []byte("X"), + "y": []byte("Y"), + }, + operations: []batchOperation{ + {Op: OpDelete, Key: "x"}, + {Op: OpDelete, Key: "y"}, + }, + expectToNotExist: []string{"x", "y"}, + }, + "overwrite existing key": { + initial: map[string][]byte{ + "k": []byte("old"), + }, + operations: []batchOperation{ + {Op: OpSet, Key: "k", Value: []byte("new")}, + }, + expectToExist: map[string][]byte{ + "k": []byte("new"), + }, + }, + "mixed put and delete": { + initial: map[string][]byte{ + "m": []byte("M0"), + "n": []byte("N0"), + }, + operations: []batchOperation{ + {Op: OpSet, Key: "x", Value: []byte("X1")}, + {Op: OpDelete, Key: "m"}, + {Op: OpSet, Key: "n", Value: []byte("N1")}, + }, + expectToExist: map[string][]byte{ + "x": []byte("X1"), + "n": []byte("N1"), + }, + expectToNotExist: []string{"m"}, + }, + "multiple operations same key": { + initial: map[string][]byte{ + "a": []byte("init"), + }, + operations: []batchOperation{ + {Op: OpSet, Key: "a", Value: []byte("v1")}, + {Op: OpDelete, Key: "a"}, + {Op: OpSet, Key: "a", Value: []byte("v2")}, + }, + expectToExist: map[string][]byte{ + "a": []byte("v2"), + }, + }, + "empty Batch operations": { + operations: []batchOperation{}, + expectToExist: map[string][]byte{}, + expectToNotExist: []string{}, + }, + "delete non-existing key": { + initial: map[string][]byte{ + "exists": []byte("ok"), + }, + operations: []batchOperation{ + {Op: OpDelete, Key: "missing"}, + }, + expectToExist: map[string][]byte{ + "exists": []byte("ok"), + }, + expectToNotExist: []string{"missing"}, + }, + } + + for name, tc := range tests { + t.Run(name, func(t *testing.T) { + ds := createTestDatastore(t) + crdt := clset.New("peer1", ds) + ctx := context.Background() + + // Load initial state + if tc.initial != nil { + for key, val := range tc.initial { + require.NoError(t, crdt.Set(key, val)) + } + } + + // Create Batch + batch, err := crdt.Batch(ctx) + require.NoError(t, err) + + // Apply operations + for _, op := range tc.operations { + switch op.Op { + case OpSet: + err = batch.Set(ctx, op.Key, op.Value) + require.NoError(t, err) + case OpDelete: + err = batch.Delete(ctx, op.Key) + require.NoError(t, err) + } + } + + // Commit + err = batch.Commit(ctx) + if tc.expectErr { + require.Error(t, err) + return + } + require.NoError(t, err) + + // Expected to exist + for key, expected := range tc.expectToExist { + got, exist, err := crdt.Get(key) + require.NoError(t, err) + require.True(t, exist) + require.Equal(t, expected, got) + } + + // Expected NOT to exist + for _, key := range tc.expectToNotExist { + _, exists, err := crdt.Get(key) + require.NoError(t, err) + require.False(t, exists) + } + }) + } +} diff --git a/go.mod b/go.mod index 62aea7e..536b1b3 100644 --- a/go.mod +++ b/go.mod @@ -34,8 +34,6 @@ require ( github.com/huin/goupnp v1.3.0 // indirect github.com/inconshreveable/mousetrap v1.1.0 // indirect github.com/ipfs/go-cid v0.5.0 // indirect - github.com/ipfs/go-detect-race v0.0.1 // indirect - github.com/ipfs/go-ipfs-delay v0.0.1 // indirect github.com/ipfs/go-log/v2 v2.6.0 // indirect github.com/jackpal/go-nat-pmp v1.0.2 // indirect github.com/jbenet/go-temp-err-catcher v0.1.0 // indirect @@ -117,7 +115,7 @@ require ( golang.org/x/mod v0.25.0 // indirect golang.org/x/net v0.41.0 // indirect golang.org/x/sync v0.15.0 // indirect - golang.org/x/sys v0.34.0 // indirect + golang.org/x/sys v0.36.0 // indirect golang.org/x/text v0.26.0 // indirect golang.org/x/time v0.12.0 // indirect golang.org/x/tools v0.34.0 // indirect diff --git a/go.sum b/go.sum index b707687..29b3e34 100644 --- a/go.sum +++ b/go.sum @@ -102,8 +102,6 @@ github.com/ipfs/go-detect-race v0.0.1 h1:qX/xay2W3E4Q1U7d9lNs1sU9nvguX0a7319XbyQ github.com/ipfs/go-detect-race v0.0.1/go.mod h1:8BNT7shDZPo99Q74BpGMK+4D8Mn4j46UU0LZ723meps= github.com/ipfs/go-ds-badger4 v0.1.8 h1:frNczf5CjCVm62RJ5mW5tD/oLQY/9IKAUpKviRV9QAI= github.com/ipfs/go-ds-badger4 v0.1.8/go.mod h1:FdqSLA5TMsyqooENB/Hf4xzYE/iH0z/ErLD6ogtfMrA= -github.com/ipfs/go-ipfs-delay v0.0.1 h1:r/UXYyRcddO6thwOnhiznIAiSvxMECGgtv35Xs1IeRQ= -github.com/ipfs/go-ipfs-delay v0.0.1/go.mod h1:8SP1YXK1M1kXuc4KJZINY3TQQ03J2rwBG9QfXmbRPrw= github.com/ipfs/go-log/v2 v2.6.0 h1:2Nu1KKQQ2ayonKp4MPo6pXCjqw1ULc9iohRqWV5EYqg= github.com/ipfs/go-log/v2 v2.6.0/go.mod h1:p+Efr3qaY5YXpx9TX7MoLCSEZX5boSWj9wh86P5HJa8= github.com/jackpal/go-nat-pmp v1.0.2 h1:KzKSgb7qkJvOUTqYl9/Hg/me3pWgBmERKrTGD7BdWus= @@ -448,8 +446,8 @@ golang.org/x/sys v0.12.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.16.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA= golang.org/x/sys v0.17.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA= golang.org/x/sys v0.20.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA= -golang.org/x/sys v0.34.0 h1:H5Y5sJ2L2JRdyv7ROF1he/lPdvFsd0mJHFw2ThKHxLA= -golang.org/x/sys v0.34.0/go.mod h1:BJP2sWEmIv4KK5OTEluFJCKSidICx8ciO85XgH3Ak8k= +golang.org/x/sys v0.36.0 h1:KVRy2GtZBrk1cBYA7MKu5bEZFxQk4NIDV6RLVcC8o0k= +golang.org/x/sys v0.36.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks= golang.org/x/telemetry v0.0.0-20240228155512-f48c80bd79b2/go.mod h1:TeRTkGYfJXctD9OcfyVLyj2J3IxLnKwHJR8f4D8a3YE= golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo= golang.org/x/term v0.0.0-20210927222741-03fcf44c2211/go.mod h1:jbD1KX2456YbFQfuXm/mYQcufACuNUgVhRMnK/tPxf8= diff --git a/p2p_sync_test.go b/p2p_sync_test.go index f5f5b42..d030c05 100644 --- a/p2p_sync_test.go +++ b/p2p_sync_test.go @@ -1,9 +1,11 @@ package clset_test import ( + "context" "crypto/rand" "fmt" "log" + "sync" "testing" "time" @@ -142,3 +144,96 @@ func TestPeerTracking(t *testing.T) { assert.Contains(t, crdt2.GetTrackedPeers(), "p1") assert.True(t, crdt2.GetTrackedPeers()["p1"] > 0) } + +func TestSync_AfterBatch(t *testing.T) { + crdt1 := createTestCRDT(t, "batchPeer1") + crdt2 := createTestCRDT(t, "batchPeer2") + + p2p1 := createPeer(t, crdt1, 15001) + defer p2p1.Close() + + p2p2 := createPeer(t, crdt2, 15002) + defer p2p2.Close() + + time.Sleep(1 * time.Second) + connectAddr := p2p1.Host.Addrs()[0].String() + "/p2p/" + p2p1.Host.ID().String() + require.NoError(t, p2p2.ManualConnect(connectAddr)) + + // Perform batch updates on peer1 + ctx1 := context.Background() + batch, err := crdt1.Batch(ctx1) + require.NoError(t, err) + for i := 0; i < 10; i++ { + require.NoError(t, batch.Set(ctx1, fmt.Sprintf("key-%d", i), []byte(fmt.Sprintf("value-%d", i)))) + } + require.NoError(t, batch.Commit(ctx1)) + + p2p2.SyncNow() + time.Sleep(2 * time.Second) + + // Verify all keys are synced to peer2 + for i := 0; i < 10; i++ { + val, exists, err := crdt2.Get(fmt.Sprintf("key-%d", i)) + require.NoError(t, err) + assert.True(t, exists) + assert.Equal(t, []byte(fmt.Sprintf("value-%d", i)), val) + } +} + +func TestSync_ConflictingBatch(t *testing.T) { + crdt1 := createTestCRDT(t, "batchPeer1") + crdt2 := createTestCRDT(t, "batchPeer2") + + p2p1 := createPeer(t, crdt1, 15001) + defer p2p1.Close() + + p2p2 := createPeer(t, crdt2, 15002) + defer p2p2.Close() + + time.Sleep(1 * time.Second) + connectAddr := p2p1.Host.Addrs()[0].String() + "/p2p/" + p2p1.Host.ID().String() + require.NoError(t, p2p2.ManualConnect(connectAddr)) + + // Perform conflicting batch updates concurrently + var wg sync.WaitGroup + wg.Add(2) + + concurrentBatchFn := func(crdt *clset.CRDT, peerName string) { + defer wg.Done() + ctx := context.Background() + batch, err := crdt1.Batch(ctx) + require.NoError(t, err) + for j := 0; j < 10; j++ { + key := fmt.Sprintf("key-%d", j) + value := []byte(fmt.Sprintf("%s_value-%d", peerName, j)) + require.NoError(t, batch.Set(ctx, key, value)) + } + require.NoError(t, batch.Commit(ctx), "Peer %s failed to commit batch", peerName) + } + go concurrentBatchFn(crdt1, "peer1") + go concurrentBatchFn(crdt2, "peer2") + wg.Wait() + + p2p1.SyncNow() + time.Sleep(2 * time.Second) + + // Verify both peers resolved to the same value for the conflicting key + peer1KV := make(map[string][]byte, 10) + peer2KV := make(map[string][]byte, 10) + + for i := 0; i < 10; i++ { + key := fmt.Sprintf("key-%d", i) + + val1, exists, err := crdt1.Get(key) + require.NoError(t, err) + require.True(t, exists) + peer1KV[key] = val1 + + val2, exists, err := crdt2.Get(key) + require.NoError(t, err) + require.True(t, exists) + peer2KV[key] = val2 + } + + require.Equal(t, peer1KV, peer2KV, "Both peers should have the same resolved values after sync") +}