From fc85fe62d64861485685fb7525665e9d1ed65c59 Mon Sep 17 00:00:00 2001 From: Emil Lenngren Date: Sun, 11 Apr 2021 15:31:09 +0200 Subject: [PATCH] Recover the buffer position when type handler throws, for sequential read --- src/Npgsql/Internal/NpgsqlReadBuffer.cs | 7 +++++++ src/Npgsql/NpgsqlDataReader.cs | 22 +++++++++++----------- test/Npgsql.Tests/ReaderTests.cs | 12 ++++++++++++ 3 files changed, 30 insertions(+), 11 deletions(-) diff --git a/src/Npgsql/Internal/NpgsqlReadBuffer.cs b/src/Npgsql/Internal/NpgsqlReadBuffer.cs index a796556ae2..bc723eb8e9 100644 --- a/src/Npgsql/Internal/NpgsqlReadBuffer.cs +++ b/src/Npgsql/Internal/NpgsqlReadBuffer.cs @@ -77,6 +77,9 @@ internal TimeSpan Timeout internal int ReadPosition { get; set; } internal int ReadBytesLeft => FilledBytes - ReadPosition; + int _flushedBytes; // modulo 2^32 + internal int CumulativeReadPosition => unchecked(_flushedBytes + ReadPosition); + internal readonly byte[] Buffer; internal int FilledBytes; @@ -161,6 +164,7 @@ static async Task EnsureLong( { Array.Copy(buffer.Buffer, buffer.ReadPosition, buffer.Buffer, 0, buffer.ReadBytesLeft); buffer.FilledBytes = buffer.ReadBytesLeft; + buffer._flushedBytes = unchecked(buffer._flushedBytes + buffer.ReadPosition); buffer.ReadPosition = 0; } @@ -494,6 +498,7 @@ public int Read(Span output) var read = Underlying.Read(output); if (read == 0) throw new EndOfStreamException(); + _flushedBytes = unchecked(_flushedBytes + read); return read; } catch (Exception e) @@ -526,6 +531,7 @@ static async ValueTask ReadAsyncLong(NpgsqlReadBuffer buffer, Memory var read = await buffer.Underlying.ReadAsync(output, cancellationToken); if (read == 0) throw new EndOfStreamException(); + buffer._flushedBytes = unchecked(buffer._flushedBytes + read); return read; } catch (Exception e) @@ -637,6 +643,7 @@ public void Dispose() internal void Clear() { + _flushedBytes = unchecked(_flushedBytes + FilledBytes); ReadPosition = 0; FilledBytes = 0; } diff --git a/src/Npgsql/NpgsqlDataReader.cs b/src/Npgsql/NpgsqlDataReader.cs index c0b7609797..2abe3fa053 100644 --- a/src/Npgsql/NpgsqlDataReader.cs +++ b/src/Npgsql/NpgsqlDataReader.cs @@ -1520,8 +1520,8 @@ public override T GetFieldValue(int ordinal) { if (Connector.State != ConnectorState.Broken) { - var writtenBytes = Buffer.ReadPosition - position; - var remainingBytes = ColumnLen - writtenBytes; + var readBytes = Buffer.ReadPosition - position; + var remainingBytes = ColumnLen - readBytes; if (remainingBytes > 0) Buffer.Skip(remainingBytes, false).GetAwaiter().GetResult(); } @@ -1554,7 +1554,7 @@ async ValueTask GetFieldValueSequential(int column, bool async, Cancellati ThrowHelper.ThrowInvalidCastException_NoValue(field); } - var position = Buffer.ReadPosition; + var position = Buffer.CumulativeReadPosition; try { return NullableHandler.Exists @@ -1573,8 +1573,8 @@ async ValueTask GetFieldValueSequential(int column, bool async, Cancellati { if (Connector.State != ConnectorState.Broken) { - var writtenBytes = Buffer.ReadPosition - position; - var remainingBytes = ColumnLen - writtenBytes; + var readBytes = unchecked(Buffer.CumulativeReadPosition - position); + var remainingBytes = ColumnLen - readBytes; if (remainingBytes > 0) await Buffer.Skip(remainingBytes, async); } @@ -1610,7 +1610,7 @@ public override object GetValue(int ordinal) return DBNull.Value; object result; - var position = Buffer.ReadPosition; + var position = Buffer.CumulativeReadPosition; try { result = _isSequential @@ -1621,8 +1621,8 @@ public override object GetValue(int ordinal) { if (Connector.State != ConnectorState.Broken) { - var writtenBytes = Buffer.ReadPosition - position; - var remainingBytes = ColumnLen - writtenBytes; + var readBytes = unchecked(Buffer.CumulativeReadPosition - position); + var remainingBytes = ColumnLen - readBytes; if (remainingBytes > 0) Buffer.Skip(remainingBytes, false).GetAwaiter().GetResult(); } @@ -1666,7 +1666,7 @@ public override object GetProviderSpecificValue(int ordinal) if (ColumnLen == -1) return DBNull.Value; - var position = Buffer.ReadPosition; + var position = Buffer.CumulativeReadPosition; try { return _isSequential @@ -1677,8 +1677,8 @@ public override object GetProviderSpecificValue(int ordinal) { if (Connector.State != ConnectorState.Broken) { - var writtenBytes = Buffer.ReadPosition - position; - var remainingBytes = ColumnLen - writtenBytes; + var readBytes = unchecked(Buffer.CumulativeReadPosition - position); + var remainingBytes = ColumnLen - readBytes; if (remainingBytes > 0) Buffer.Skip(remainingBytes, false).GetAwaiter().GetResult(); } diff --git a/test/Npgsql.Tests/ReaderTests.cs b/test/Npgsql.Tests/ReaderTests.cs index e6d72adb17..e34651e1fd 100644 --- a/test/Npgsql.Tests/ReaderTests.cs +++ b/test/Npgsql.Tests/ReaderTests.cs @@ -1224,6 +1224,18 @@ await pgMock Assert.DoesNotThrowAsync(async () => await reader.DisposeAsync()); } + [Test] + [IssueLink("https://github.com/npgsql/npgsql/issues/3648")] + public async Task RecoversFromExceptionInLongColumn() + { + using var conn = await OpenConnectionAsync(); + using var cmd = new NpgsqlCommand(@"SELECT array_cat(array_fill(NULL::numeric, ARRAY[2048]), ARRAY['NaN'::numeric]), 9", conn); + using var reader = await cmd.ExecuteReaderAsync(Behavior); + await reader.ReadAsync(); + Assert.Throws(() => reader.GetFieldValue(0)); + Assert.That(reader.GetInt32(1), Is.EqualTo(9)); + } + #region GetBytes / GetStream [Test]