diff --git a/base/bucket.go b/base/bucket.go index abdaaec8ea..6ed2675823 100644 --- a/base/bucket.go +++ b/base/bucket.go @@ -18,6 +18,7 @@ import ( "io" "net/http" "net/url" + "regexp" "strconv" "strings" "time" @@ -541,9 +542,9 @@ func GetSourceID(ctx context.Context, bucket Bucket) (string, error) { } gocbBucket, err := AsGocbV2Bucket(bucket) if err != nil { - // for rosmar bucket and testing, use the bucket name as the source ID to make it easier to identify the source + // for rosmar bucket and testing, base the source ID on the bucket name to make it easier to identify the source if underGoTest() { - return bucket.GetName(), nil + return testSourceID(bucket.GetName(), bucketUUID) } serverUUID := "" return CreateEncodedSourceID(bucketUUID, serverUUID) @@ -557,6 +558,30 @@ func GetSourceID(ctx context.Context, bucket Bucket) (string, error) { return CreateEncodedSourceID(bucketUUID, serverUUID) } +// encodedSourceIDLength is the length of a source ID in its encoded form: 16 bytes of unpadded +// base64. Couchbase Lite rejects a version whose source ID is any other length, or any string that +// is not the canonical encoding of 16 bytes. +const encodedSourceIDLength = 22 + +// testSourceIDName matches a bucket name testSourceID can embed: letters and digits only, so it never +// contains the '+' padding, and short enough to leave room for the final 'A'. +var testSourceIDName = regexp.MustCompile(`^[A-Za-z0-9]{0,21}$`) + +// testSourceID returns a source ID for a test bucket that still reads as the bucket's name, so that +// a real Couchbase Lite in a test accepts the versions Sync Gateway writes. The name is padded with +// '+', which is visually quiet and never part of a name, so distinct names always give distinct IDs: +// "rosmar1" becomes "rosmar1++++++++++++++A". A name that cannot be embedded that way gets the same +// encoded ID production would. +func testSourceID(bucketName, bucketUUID string) (string, error) { + if !testSourceIDName.MatchString(bucketName) { + return CreateEncodedSourceID(bucketUUID, "") + } + // The final character carries only the last 2 of the 128 bits, and its other 4 bits must be zero, + // so it cannot be another '+'. 'A' is the zero digit: it is valid there and sets no bits of its own, + // so it reads as an end marker rather than as part of the ID. + return bucketName + strings.Repeat("+", encodedSourceIDLength-1-len(bucketName)) + "A", nil +} + // CreateEncodedSourceID will hash the bucket UUID and cluster UUID using md5 hash function then will base64 encode it // This function is in sync with xdcr implementation of UUIDstoDocumentSource https://github.com/couchbase/goxdcr/blob/dfba7a5b4251d93db46e2b0b4b55ea014218931b/hlv/hlv.go#L51 func CreateEncodedSourceID(bucketUUID, clusterUUID string) (string, error) { diff --git a/base/bucket_test.go b/base/bucket_test.go index 5c6d97a3d1..d4cff1598f 100644 --- a/base/bucket_test.go +++ b/base/bucket_test.go @@ -16,6 +16,7 @@ import ( "crypto/rand" "crypto/x509" "crypto/x509/pkix" + "encoding/base64" "encoding/pem" "math/big" "os" @@ -453,6 +454,42 @@ func TestTLSConfig(t *testing.T) { assert.Empty(t, conf) } +func TestTestSourceID(t *testing.T) { + const bucketUUID = "a1b2c3d4e5f60718293a4b5c6d7e8f90" + encoded, err := CreateEncodedSourceID(bucketUUID, "") + require.NoError(t, err) + + tests := []struct { + bucketName string + expected string + }{ + {bucketName: "rosmar1", expected: "rosmar1++++++++++++++A"}, + {bucketName: "rosmar12", expected: "rosmar12+++++++++++++A"}, + {bucketName: "rosmar1A", expected: "rosmar1A+++++++++++++A"}, + {bucketName: "abcdefghijklmnopqrstu", expected: "abcdefghijklmnopqrstuA"}, + {bucketName: "", expected: "+++++++++++++++++++++A"}, + // No room left for the final 'A'. + {bucketName: "abcdefghijklmnopqrstuv", expected: encoded}, + {bucketName: "sg_int_0", expected: encoded}, + {bucketName: "rosmar+1", expected: encoded}, + } + for _, test := range tests { + t.Run(test.bucketName, func(t *testing.T) { + sourceID, err := testSourceID(test.bucketName, bucketUUID) + require.NoError(t, err) + assert.Equal(t, test.expected, sourceID) + + // Couchbase Lite only accepts the canonical, unpadded encoding of 16 bytes, and reserves + // the all-zero ID for itself. + require.Len(t, sourceID, encodedSourceIDLength) + decoded, err := base64.RawStdEncoding.Strict().DecodeString(sourceID) + require.NoError(t, err) + require.Len(t, decoded, 16) + assert.NotEqual(t, make([]byte, 16), decoded) + }) + } +} + func TestBaseBucket(t *testing.T) { tests := []struct { diff --git a/db/blip_handler.go b/db/blip_handler.go index 3e13448b48..e39240777b 100644 --- a/db/blip_handler.go +++ b/db/blip_handler.go @@ -1755,5 +1755,5 @@ func GetHLVFromRevMessage(msg *blip.Message) (*HybridLogicalVector, []string, er versionVectorStr += ";" + historyStr } } - return extractHLVFromBlipString(versionVectorStr) + return ExtractHLVFromBlipString(versionVectorStr) } diff --git a/db/blip_sync_context.go b/db/blip_sync_context.go index b6cace07d6..d1a5149a39 100644 --- a/db/blip_sync_context.go +++ b/db/blip_sync_context.go @@ -970,7 +970,7 @@ func (bsc *BlipSyncContext) getKnownRevs(ctx context.Context, docID string, know if revID, ok := knownRevsArray[0].(string); ok { if bsc.useHLV() && !base.IsRevTreeID(revID) { // extract cv from the known revs array - msgHLV, _, deltaSrcErr := extractHLVFromBlipString(revID) + msgHLV, _, deltaSrcErr := ExtractHLVFromBlipString(revID) if deltaSrcErr != nil { base.DebugfCtx(ctx, base.KeySync, "Invalid known rev format for hlv on doc: %s falling back to full body replication. Err: %v KnownRev: %s", base.UD(docID), deltaSrcErr, revID) deltaSrcRev = "" // will force falling back to full body replication below @@ -987,7 +987,7 @@ func (bsc *BlipSyncContext) getKnownRevs(ctx context.Context, docID string, know for _, rev := range knownRevsArray { if revID, ok := rev.(string); ok { // extract cv from the known revs array - msgHLV, _, err := extractHLVFromBlipString(revID) + msgHLV, _, err := ExtractHLVFromBlipString(revID) if err != nil { // assume we have received legacy rev if the following conditions are met: // - we cannot parse cv from known revs diff --git a/db/crud.go b/db/crud.go index 996580c383..2353d809af 100644 --- a/db/crud.go +++ b/db/crud.go @@ -4314,7 +4314,7 @@ func (db *DatabaseCollectionWithUser) CheckProposedVersion(ctx context.Context, // Temporary (CBG-4466): check the full HLV that's being sent by CBL with proposeChanges messages. // If the current server cv is dominated by the incoming HLV (i.e. the incoming HLV has an entry for the same source // with a version that's greater than or equal to the server's cv), then we can accept the proposed version. - proposedHLV, _, err := extractHLVFromBlipString(proposedHLVString) + proposedHLV, _, err := ExtractHLVFromBlipString(proposedHLVString) if err != nil { base.WarnfCtx(ctx, "CheckProposedVersion for doc %s unable to extract proposedHLV from rev message, will be treated as conflict: %v", base.UD(docid), err) } else if proposedHLV.DominatesSource(localDocCV) { diff --git a/db/hybrid_logical_vector.go b/db/hybrid_logical_vector.go index 1d6554e452..66274e4e86 100644 --- a/db/hybrid_logical_vector.go +++ b/db/hybrid_logical_vector.go @@ -590,7 +590,7 @@ func (hlv *HybridLogicalVector) toHistoryForHLV(sortFunc func(HLVVersions) iter. return s.String() } -// extractHLVFromBlipMessage extracts the full HLV a string in the format seen over Blip +// ExtractHLVFromBlipString extracts the full HLV from a string in the format seen over Blip // blip string may be the following formats // 1. cv only: cv // 2. cv and pv: cv;pv @@ -607,7 +607,7 @@ func (hlv *HybridLogicalVector) toHistoryForHLV(sortFunc func(HLVVersions) iter. // // Function will return list of revIDs if legacy rev ID was found in the HLV history section (PV) // TODO: CBG-3662 - Optimise once we've settled on and tested the format with CBL -func extractHLVFromBlipString(versionVectorStr string) (*HybridLogicalVector, []string, error) { +func ExtractHLVFromBlipString(versionVectorStr string) (*HybridLogicalVector, []string, error) { hlv := &HybridLogicalVector{} vectorFields := strings.Split(versionVectorStr, ";") diff --git a/db/hybrid_logical_vector_test.go b/db/hybrid_logical_vector_test.go index c2b045336c..bdf7be44ea 100644 --- a/db/hybrid_logical_vector_test.go +++ b/db/hybrid_logical_vector_test.go @@ -259,7 +259,7 @@ func createHLVForTest(tb *testing.T, input string) *HybridLogicalVector { if input == "" { return NewHybridLogicalVector() } - hlv, _, err := extractHLVFromBlipString(input) + hlv, _, err := ExtractHLVFromBlipString(input) require.NoError(tb, err) return hlv } @@ -838,7 +838,7 @@ func TestInvalidHLVInBlipMessageForm(t *testing.T) { for _, testCase := range testCases { t.Run(testCase.name, func(t *testing.T) { require.NotEmpty(t, testCase.errMsg) // make sure err msg is specified - hlv, legacyRevs, err := extractHLVFromBlipString(testCase.hlv) + hlv, legacyRevs, err := ExtractHLVFromBlipString(testCase.hlv) require.ErrorContains(t, err, testCase.errMsg, "expected err for %s", testCase.hlv) require.Nil(t, hlv) require.Nil(t, legacyRevs) @@ -1100,12 +1100,12 @@ func getHLVTestCases(t testing.TB) []extractHLVFromBlipMsgBMarkCases { } // TestExtractHLVFromChangesMessage: -// - Each test case gets run through extractHLVFromBlipString and assert that the resulting HLV +// - Each test case gets run through ExtractHLVFromBlipString and assert that the resulting HLV // is correct to what is expected func TestExtractHLVFromChangesMessage(t *testing.T) { for _, test := range getHLVTestCases(t) { t.Run(test.name, func(t *testing.T) { - hlv, legacyRevs, err := extractHLVFromBlipString(test.hlvString) + hlv, legacyRevs, err := ExtractHLVFromBlipString(test.hlvString) require.NoError(t, err) require.Equal(t, test.expectedHLV, *hlv, "HLV not parsed correctly for %s", test.hlvString) @@ -1176,7 +1176,7 @@ func BenchmarkExtractHLVFromBlipMessage(b *testing.B) { for _, bm := range getHLVTestCases(b) { b.Run(bm.name, func(b *testing.B) { for b.Loop() { - _, _, _ = extractHLVFromBlipString(bm.hlvString) + _, _, _ = ExtractHLVFromBlipString(bm.hlvString) } }) } @@ -1497,7 +1497,7 @@ func TestHLVAddVersion(t *testing.T) { for _, tc := range testCases { t.Run(tc.name, func(t *testing.T) { - hlv, _, err := extractHLVFromBlipString(tc.initialHLV) + hlv, _, err := ExtractHLVFromBlipString(tc.initialHLV) require.NoError(t, err, "unable to parse initialHLV") newVersion, err := ParseVersion(tc.newVersion) require.NoError(t, err) @@ -1505,7 +1505,7 @@ func TestHLVAddVersion(t *testing.T) { err = hlv.AddVersion(newVersion) require.NoError(t, err) - expectedHLV, _, err := extractHLVFromBlipString(tc.expectedHLV) + expectedHLV, _, err := ExtractHLVFromBlipString(tc.expectedHLV) require.NoError(t, err) require.True(t, hlv.Equal(expectedHLV), "expected %#v does not match actual %#v", expectedHLV, hlv) @@ -1625,9 +1625,9 @@ func TestHLVIsInConflict(t *testing.T) { } for _, tc := range testCases { t.Run(tc.name, func(t *testing.T) { - localHLV, _, err := extractHLVFromBlipString(tc.localHLV) + localHLV, _, err := ExtractHLVFromBlipString(tc.localHLV) require.NoError(t, err) - incomingHLV, _, err := extractHLVFromBlipString(tc.incomingHLV) + incomingHLV, _, err := ExtractHLVFromBlipString(tc.incomingHLV) require.NoError(t, err) require.Equal(t, tc.conflict, IsInConflict(t.Context(), localHLV, incomingHLV))