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
2 changes: 1 addition & 1 deletion libs/client/RespReadResponseUtils.cs
Original file line number Diff line number Diff line change
Expand Up @@ -120,7 +120,7 @@ public static bool TryReadStringWithLengthHeader(MemoryPool<byte> pool, out Memo
static bool TryReadPtrWithSignedLengthHeader(ref byte* keyPtr, ref int length, ref byte* ptr, byte* end)
{
// Parse RESP string header
if (!RespReadUtils.TryReadSignedLengthHeader(out length, ref ptr, end))
if (!RespReadUtils.TryReadSignedLengthHeader(out length, ref ptr, end) || length > RespReadUtils.MaxArgumentLengthBytes)
{
return false;
}
Expand Down
17 changes: 12 additions & 5 deletions libs/common/RespReadUtils.cs
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,13 @@ namespace Garnet.common
/// </summary>
public static unsafe class RespReadUtils
{
/// <summary>
/// Maximum size a single RESP element can be.
///
/// This matches the default Redis max string size of 512MB.
/// </summary>
public const int MaxArgumentLengthBytes = 512 * 1_024 * 1_024;

Comment thread
kevin-montrose marked this conversation as resolved.
/// <summary>
/// Tries to read the leading sign of the given ASCII-encoded number.
/// </summary>
Expand Down Expand Up @@ -690,7 +697,7 @@ public static bool TryReadUInt64WithLengthHeader(out ulong number, ref byte* ptr
public static bool TrySkipByteArrayWithLengthHeader(ref byte* ptr, byte* end)
{
// Parse RESP string header
if (!TryReadUnsignedLengthHeader(out var length, ref ptr, end))
if (!TryReadUnsignedLengthHeader(out var length, ref ptr, end) || length > RespReadUtils.MaxArgumentLengthBytes)
return false;

// Advance read pointer to the end of the array (including terminator)
Expand Down Expand Up @@ -753,7 +760,7 @@ public static bool TrySliceWithLengthHeader(out ReadOnlySpan<byte> result, scope
result = default;

// Parse RESP string header
if (!TryReadUnsignedLengthHeader(out var length, ref ptr, end))
if (!TryReadUnsignedLengthHeader(out var length, ref ptr, end) || length > RespReadUtils.MaxArgumentLengthBytes)
return false;

// Advance read pointer to the end of the array (including terminator)
Expand Down Expand Up @@ -860,7 +867,7 @@ public static bool TryReadSpanWithLengthHeader(out ReadOnlySpan<byte> result, re
return false;

// Parse RESP string header
if (!TryReadUnsignedLengthHeader(out var length, ref ptr, end))
if (!TryReadUnsignedLengthHeader(out var length, ref ptr, end) || length > RespReadUtils.MaxArgumentLengthBytes)
return false;

// Extract string content + '\r\n' terminator
Expand Down Expand Up @@ -925,7 +932,7 @@ public static bool TryReadStringResponseWithLengthHeader(out string result, ref
public static bool TryReadPtrWithSignedLengthHeader(ref byte* stringPtr, ref int length, ref byte* ptr, byte* end)
{
// Parse RESP string header
if (!TryReadSignedLengthHeader(out length, ref ptr, end))
if (!TryReadSignedLengthHeader(out length, ref ptr, end) || length > RespReadUtils.MaxArgumentLengthBytes)
{
return false;
}
Expand Down Expand Up @@ -1148,7 +1155,7 @@ public static bool TryReadDoubleWithLengthHeader(out double result, out bool par
public static bool TryReadPtrWithLengthHeader(ref byte* result, ref int len, ref byte* ptr, byte* end)
{
// Parse RESP string header
if (!TryReadUnsignedLengthHeader(out len, ref ptr, end))
if (!TryReadUnsignedLengthHeader(out len, ref ptr, end) || len > RespReadUtils.MaxArgumentLengthBytes)
{
return false;
}
Expand Down
2 changes: 1 addition & 1 deletion libs/server/Resp/Parser/SessionParseState.cs
Original file line number Diff line number Diff line change
Expand Up @@ -347,7 +347,7 @@ public readonly bool Read(int i, ref byte* ptr, byte* end)
ref var slice = ref Unsafe.AsRef<PinnedSpanByte>(bufferPtr + i);

// Parse RESP string header
if (!RespReadUtils.TryReadUnsignedLengthHeader(out var length, ref ptr, end))
if (!RespReadUtils.TryReadUnsignedLengthHeader(out var length, ref ptr, end) || length > RespReadUtils.MaxArgumentLengthBytes)
return false;
slice.Set(ptr, length);

Expand Down
106 changes: 106 additions & 0 deletions test/standalone/Garnet.test/Resp/RespReadUtilsTests.cs
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
// Copyright (c) Microsoft Corporation.
// Licensed under the MIT license.

using System;
using System.Text;
using Garnet.common;
using Garnet.common.Parsing;
Expand Down Expand Up @@ -380,5 +381,110 @@ public static unsafe void GetSerializedRecordSpanInsufficientHeaderTest()
ClassicAssert.AreEqual(0, recordSpan.Length);
}
}

[Test]
public static unsafe void RejectOversizedLengths()
{
var bigBuffer = GC.AllocateUninitializedArray<byte>(RespReadUtils.MaxArgumentLengthBytes + 1, pinned: true);

var smallToCopy = Encoding.ASCII.GetBytes("$1234\r\n" + new string('a', 1234) + "\r\n");
var bigToCopy = "$536870913\r\n"u8;

Comment thread
kevin-montrose marked this conversation as resolved.
fixed (byte* startPtr = bigBuffer)
{
var endPtr = startPtr + bigBuffer.Length;

// TryReadPtrWithSignedLengthHeader
{
// Pass in bounds
var ptr = startPtr;
smallToCopy.CopyTo(new Span<byte>(startPtr, bigBuffer.Length));

var len = 0;
byte* outPtr = null;
ClassicAssert.True(RespReadUtils.TryReadPtrWithSignedLengthHeader(ref outPtr, ref len, ref ptr, endPtr));
ClassicAssert.True(outPtr == (startPtr + 7));
ClassicAssert.AreEqual(1234, len);
ClassicAssert.True(ptr == (outPtr + len + 2));

// Fail too big
ptr = startPtr;

len = 0;
outPtr = null;
bigToCopy.CopyTo(new Span<byte>(startPtr, bigBuffer.Length));
ClassicAssert.False(RespReadUtils.TryReadPtrWithSignedLengthHeader(ref outPtr, ref len, ref ptr, endPtr));
}

// TrySkipByteArrayWithLengthHeader
{
// Pass in bounds
var ptr = startPtr;
smallToCopy.CopyTo(new Span<byte>(startPtr, bigBuffer.Length));

ClassicAssert.True(RespReadUtils.TrySkipByteArrayWithLengthHeader(ref ptr, endPtr));
ClassicAssert.True(ptr == (startPtr + 7 + 1234 + 2));

// Fail too big
ptr = startPtr;
bigToCopy.CopyTo(new Span<byte>(startPtr, bigBuffer.Length));
ClassicAssert.False(RespReadUtils.TrySkipByteArrayWithLengthHeader(ref ptr, endPtr));
}

// TrySliceWithLengthHeader
{
// Pass in bounds
var ptr = startPtr;
smallToCopy.CopyTo(new Span<byte>(startPtr, bigBuffer.Length));

ClassicAssert.True(RespReadUtils.TrySliceWithLengthHeader(out var slice, ref ptr, endPtr));
ClassicAssert.AreEqual(1234, slice.Length);
ClassicAssert.True(ptr == (startPtr + 7 + slice.Length + 2));

// Fail too big
ptr = startPtr;
bigToCopy.CopyTo(new Span<byte>(startPtr, bigBuffer.Length));
ClassicAssert.False(RespReadUtils.TrySliceWithLengthHeader(out slice, ref ptr, endPtr));
}

// TryReadSpanWithLengthHeader
{
// Pass in bounds
var ptr = startPtr;
smallToCopy.CopyTo(new Span<byte>(startPtr, bigBuffer.Length));

ClassicAssert.True(RespReadUtils.TryReadSpanWithLengthHeader(out var slice, ref ptr, endPtr));
ClassicAssert.AreEqual(1234, slice.Length);
ClassicAssert.True(ptr == (startPtr + 7 + slice.Length + 2));

// Fail too big
ptr = startPtr;
bigToCopy.CopyTo(new Span<byte>(startPtr, bigBuffer.Length));
ClassicAssert.False(RespReadUtils.TryReadSpanWithLengthHeader(out slice, ref ptr, endPtr));
}

// TryReadPtrWithLengthHeader
{
// Pass in bounds
var ptr = startPtr;
smallToCopy.CopyTo(new Span<byte>(startPtr, bigBuffer.Length));

var len = 0;
byte* resultPtr = null;
ClassicAssert.True(RespReadUtils.TryReadPtrWithLengthHeader(ref resultPtr, ref len, ref ptr, endPtr));
ClassicAssert.True(resultPtr == (startPtr + 7));
ClassicAssert.AreEqual(1234, len);
ClassicAssert.True(ptr == (startPtr + 7 + len + 2));

// Fail too big
ptr = startPtr;

len = 0;
resultPtr = null;
bigToCopy.CopyTo(new Span<byte>(startPtr, bigBuffer.Length));
ClassicAssert.False(RespReadUtils.TryReadPtrWithLengthHeader(ref resultPtr, ref len, ref ptr, endPtr));
}
}
}
}
}
Loading