Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
47 changes: 46 additions & 1 deletion src/libraries/System.Memory/tests/Span/CommonPrefixLength.T.cs
Original file line number Diff line number Diff line change
Expand Up @@ -62,7 +62,7 @@ private static void ValidateWithDefaultValues<T>(int length1, int length2, IEqua
}

[Fact]
public static void PartialEquals_ReturnsPrefixLength_ValueType()
public static void PartialEquals_ReturnsPrefixLength_Byte()
{
byte[] arr1 = new byte[] { 1, 2, 3, 4, 5 };
byte[] arr2 = new byte[] { 1, 2, 3, 6, 7 };
Expand All @@ -76,6 +76,51 @@ public static void PartialEquals_ReturnsPrefixLength_ValueType()
Assert.Equal(3, MemoryExtensions.CommonPrefixLength((Span<byte>)arr1, arr2, null));
Assert.Equal(3, MemoryExtensions.CommonPrefixLength((Span<byte>)arr1, arr2, EqualityComparer<byte>.Default));
Assert.Equal(3, MemoryExtensions.CommonPrefixLength((Span<byte>)arr1, arr2, NonDefaultEqualityComparer<byte>.Instance));

// Vectorized code path
arr1 = new byte[] { 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17 };
arr2 = new byte[] { 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 42, 15, 16, 17 };

Assert.Equal(13, MemoryExtensions.CommonPrefixLength((ReadOnlySpan<byte>)arr1, arr2));
Assert.Equal(13, MemoryExtensions.CommonPrefixLength((ReadOnlySpan<byte>)arr1, arr2, null));
Assert.Equal(13, MemoryExtensions.CommonPrefixLength((ReadOnlySpan<byte>)arr1, arr2, EqualityComparer<byte>.Default));
Assert.Equal(13, MemoryExtensions.CommonPrefixLength((ReadOnlySpan<byte>)arr1, arr2, NonDefaultEqualityComparer<byte>.Instance));

Assert.Equal(13, MemoryExtensions.CommonPrefixLength((Span<byte>)arr1, arr2));
Assert.Equal(13, MemoryExtensions.CommonPrefixLength((Span<byte>)arr1, arr2, null));
Assert.Equal(13, MemoryExtensions.CommonPrefixLength((Span<byte>)arr1, arr2, EqualityComparer<byte>.Default));
Assert.Equal(13, MemoryExtensions.CommonPrefixLength((Span<byte>)arr1, arr2, NonDefaultEqualityComparer<byte>.Instance));
}

[Fact]
public static void PartialEquals_ReturnsPrefixLength_ValueType()
{
int[] arr1 = new int[] { 1, 2, 3 };
int[] arr2 = new int[] { 1, 2, 6 };

Assert.Equal(2, MemoryExtensions.CommonPrefixLength((ReadOnlySpan<int>)arr1, arr2));
Assert.Equal(2, MemoryExtensions.CommonPrefixLength((ReadOnlySpan<int>)arr1, arr2, null));
Assert.Equal(2, MemoryExtensions.CommonPrefixLength((ReadOnlySpan<int>)arr1, arr2, EqualityComparer<int>.Default));
Assert.Equal(2, MemoryExtensions.CommonPrefixLength((ReadOnlySpan<int>)arr1, arr2, NonDefaultEqualityComparer<int>.Instance));

Assert.Equal(2, MemoryExtensions.CommonPrefixLength((Span<int>)arr1, arr2));
Assert.Equal(2, MemoryExtensions.CommonPrefixLength((Span<int>)arr1, arr2, null));
Assert.Equal(2, MemoryExtensions.CommonPrefixLength((Span<int>)arr1, arr2, EqualityComparer<int>.Default));
Assert.Equal(2, MemoryExtensions.CommonPrefixLength((Span<int>)arr1, arr2, NonDefaultEqualityComparer<int>.Instance));

// Vectorized code path
arr1 = new int[] { 1, 2, 3, 4, 5 };
arr2 = new int[] { 1, 2, 3, 6, 7 };

Assert.Equal(3, MemoryExtensions.CommonPrefixLength((ReadOnlySpan<int>)arr1, arr2));
Assert.Equal(3, MemoryExtensions.CommonPrefixLength((ReadOnlySpan<int>)arr1, arr2, null));
Assert.Equal(3, MemoryExtensions.CommonPrefixLength((ReadOnlySpan<int>)arr1, arr2, EqualityComparer<int>.Default));
Assert.Equal(3, MemoryExtensions.CommonPrefixLength((ReadOnlySpan<int>)arr1, arr2, NonDefaultEqualityComparer<int>.Instance));

Assert.Equal(3, MemoryExtensions.CommonPrefixLength((Span<int>)arr1, arr2));
Assert.Equal(3, MemoryExtensions.CommonPrefixLength((Span<int>)arr1, arr2, null));
Assert.Equal(3, MemoryExtensions.CommonPrefixLength((Span<int>)arr1, arr2, EqualityComparer<int>.Default));
Assert.Equal(3, MemoryExtensions.CommonPrefixLength((Span<int>)arr1, arr2, NonDefaultEqualityComparer<int>.Instance));
}

[Fact]
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -2033,6 +2033,29 @@ public static int CommonPrefixLength<T>(this Span<T> span, ReadOnlySpan<T> other
/// <returns>The length of the common prefix shared by the two spans. If there's no shared prefix, 0 is returned.</returns>
public static int CommonPrefixLength<T>(this ReadOnlySpan<T> span, ReadOnlySpan<T> other)
{
if (RuntimeHelpers.IsBitwiseEquatable<T>())
{
nuint length = Math.Min((nuint)(uint)span.Length, (nuint)(uint)other.Length);

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I was going to suggest nuint.Min but then I remembered that IntPtr and UIntPtr still have them explicitly implemented until the work around dotnet/csharplang#6031 goes in.

nuint size = (uint)Unsafe.SizeOf<T>();
nuint index = SpanHelpers.CommonPrefixLength(
ref Unsafe.As<T, byte>(ref MemoryMarshal.GetReference(span)),
ref Unsafe.As<T, byte>(ref MemoryMarshal.GetReference(other)),
length * size);

// A byte-wise comparison in CommonPrefixLength can be used for multi-byte types,
// that are bitwise-equatable, too. In order to get the correct index in terms of type T
// of the first mismatch, integer division by the size of T is used.
//
// Example for short:
// index (byte-based): b-1, b, b+1, b+2, b+3
// index (short-based): s-1, s, s+1
// byte sequence 1: { ..., [0x42, 0x43], [0x37, 0x38], ... }
// byte sequence 2: { ..., [0x42, 0x43], [0x37, 0xAB], ... }
// So the mismatch is a byte-index b+3, which gives integer divided by the size of short:
// 3 / 2 = 1, thus the expected index short-based.
return (int)(index / size);

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The use of byte and then dividing to get rid of partials is clever... but deserves a big comment :)

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Too verbose now?

}

// Shrink one of the spans if necessary to ensure they're both the same length. We can then iterate until
// the Length of one of them and at least have bounds checks removed from that one.
SliceLongerSpanToMatchShorterLength(ref span, ref other);
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -2111,6 +2111,92 @@ public static unsafe int SequenceCompareTo(ref byte first, int firstLength, ref
return firstLength - secondLength;
}

public static nuint CommonPrefixLength(ref byte first, ref byte second, nuint length)
{
nuint i;

// It is ordered this way to match the default branch predictor rules, to don't have too much
// overhead for short input-lengths.
if (!Vector128.IsHardwareAccelerated || length < (nuint)Vector128<byte>.Count)
Comment thread
gfoidl marked this conversation as resolved.
{
// To have kind of fast path for small inputs, we handle as much elements needed
// so that either we are done or can use the unrolled loop below.
i = length % 4;

if (i > 0)
{
if (first != second)
{
return 0;
}

if (i > 1)
{
if (Unsafe.Add(ref first, 1) != Unsafe.Add(ref second, 1))
{
return 1;
}

if (i > 2 && Unsafe.Add(ref first, 2) != Unsafe.Add(ref second, 2))
{
return 2;
}
}
}

for (; (nint)i <= (nint)length - 4; i += 4)
{
if (Unsafe.Add(ref first, i + 0) != Unsafe.Add(ref second, i + 0)) return i + 0;
if (Unsafe.Add(ref first, i + 1) != Unsafe.Add(ref second, i + 1)) return i + 1;
if (Unsafe.Add(ref first, i + 2) != Unsafe.Add(ref second, i + 2)) return i + 2;
if (Unsafe.Add(ref first, i + 3) != Unsafe.Add(ref second, i + 3)) return i + 3;
}

return length;
}

Debug.Assert(length >= (uint)Vector128<byte>.Count);

uint mask;
nuint lengthToExamine = length - (nuint)Vector128<byte>.Count;

Vector128<byte> maskVec;
i = 0;

while (i < lengthToExamine)
{
maskVec = Vector128.Equals(
Vector128.LoadUnsafe(ref first, i),
Vector128.LoadUnsafe(ref second, i));

mask = maskVec.ExtractMostSignificantBits();
if (mask != 0xFFFF)
{
goto Found;
}

i += (nuint)Vector128<byte>.Count;
}

// Do final compare as Vector128<byte>.Count from end rather than start
i = lengthToExamine;
maskVec = Vector128.Equals(
Vector128.LoadUnsafe(ref first, i),
Vector128.LoadUnsafe(ref second, i));

mask = maskVec.ExtractMostSignificantBits();
if (mask != 0xFFFF)
{
goto Found;
}

return length;

Found:
mask = ~mask;
return i + uint.TrailingZeroCount(mask);
}

// Vector sub-search adapted from https://github.com/aspnet/KestrelHttpServer/pull/1138
[MethodImpl(MethodImplOptions.AggressiveInlining)]
private static int LocateLastFoundByte(Vector<byte> match)
Expand Down