Skip to content
Open
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
16 changes: 8 additions & 8 deletions Dapper/SqlMapper.Async.cs
Original file line number Diff line number Diff line change
Expand Up @@ -433,17 +433,17 @@ private static async Task<IEnumerable<T>> QueryAsync<T>(this IDbConnection cnn,
if (wasClosed) await cnn.TryOpenAsync(cancel).ConfigureAwait(false);
reader = await ExecuteReaderWithFlagsFallbackAsync(cmd, wasClosed, CommandBehavior.SequentialAccess | CommandBehavior.SingleResult, cancel).ConfigureAwait(false);

var tuple = info.Deserializer;
var deserializer = info.Deserializer;
int hash = GetColumnHash(reader);
if (tuple.Func is null || tuple.Hash != hash)
if (deserializer is null || deserializer.Hash != hash)
{
if (reader.FieldCount == 0)
return Enumerable.Empty<T>();
tuple = info.Deserializer = new DeserializerState(hash, GetDeserializer(effectiveType, reader, 0, -1, false));
deserializer = info.Deserializer = new DeserializerState(hash, GetDeserializer(effectiveType, reader, 0, -1, false));
if (command.AddToCache) SetQueryCache(identity, info);
}

var func = tuple.Func;
var func = deserializer.Func;

if (command.Buffered)
{
Expand Down Expand Up @@ -1306,19 +1306,19 @@ static async IAsyncEnumerable<T> Impl(IDbConnection cnn, Type effectiveType, Com
if (wasClosed) await cnn.TryOpenAsync(cancel).ConfigureAwait(false);
reader = await ExecuteReaderWithFlagsFallbackAsync(cmd, wasClosed, CommandBehavior.SequentialAccess | CommandBehavior.SingleResult, cancel).ConfigureAwait(false);

var tuple = info.Deserializer;
var deserializer = info.Deserializer;
int hash = GetColumnHash(reader);
if (tuple.Func is null || tuple.Hash != hash)
if (deserializer is null || deserializer.Hash != hash)
{
if (reader.FieldCount == 0)
{
yield break;
}
tuple = info.Deserializer = new DeserializerState(hash, GetDeserializer(effectiveType, reader, 0, -1, false));
deserializer = info.Deserializer = new DeserializerState(hash, GetDeserializer(effectiveType, reader, 0, -1, false));
if (command.AddToCache) SetQueryCache(identity, info);
}

var func = tuple.Func;
var func = deserializer.Func;

var convertToType = Nullable.GetUnderlyingType(effectiveType) ?? effectiveType;
while (await reader.ReadAsync(cancel).ConfigureAwait(false))
Expand Down
3 changes: 1 addition & 2 deletions Dapper/SqlMapper.CacheInfo.cs
Original file line number Diff line number Diff line change
Expand Up @@ -9,8 +9,7 @@ public static partial class SqlMapper
{
private sealed class CacheInfo
{
public DeserializerState Deserializer { get; set; }
public Func<DbDataReader, object>[]? OtherDeserializers { get; set; }
public DeserializerState? Deserializer { get; set; }
public Action<IDbCommand, object?>? ParamReader { get; set; }
private int hitCount;
public int GetHitCount() { return Interlocked.CompareExchange(ref hitCount, 0, 0); }
Expand Down
11 changes: 9 additions & 2 deletions Dapper/SqlMapper.DeserializerState.cs
Original file line number Diff line number Diff line change
Expand Up @@ -6,15 +6,22 @@ namespace Dapper
{
public static partial class SqlMapper
{
private readonly struct DeserializerState
// Reference type on purpose: this is published into CacheInfo by a plain field write,
// so it must be a single atomic reference store. A multi-field struct tears, pairing a
// Hash with the Func compiled for a different result shape. OtherDeserializers lives
// here rather than in a second CacheInfo field for the same reason - two fields cannot
// be updated together, so a reader could pair one shape's Func with another's rest-set.
private sealed class DeserializerState
{
public readonly int Hash;
public readonly Func<DbDataReader, object> Func;
public readonly Func<DbDataReader, object>[]? OtherDeserializers;

public DeserializerState(int hash, Func<DbDataReader, object> func)
public DeserializerState(int hash, Func<DbDataReader, object> func, Func<DbDataReader, object>[]? otherDeserializers = null)
{
Hash = hash;
Func = func;
OtherDeserializers = otherDeserializers;
}
}
}
Expand Down
4 changes: 2 additions & 2 deletions Dapper/SqlMapper.GridReader.Async.cs
Original file line number Diff line number Diff line change
Expand Up @@ -186,7 +186,7 @@ private Func<DbDataReader, object> ValidateAndMarkConsumed(Type type, out int in
var deserializer = cache.Deserializer;

int hash = GetColumnHash(reader);
if (deserializer.Func is null || deserializer.Hash != hash)
if (deserializer is null || deserializer.Hash != hash)
{
deserializer = new DeserializerState(hash, GetDeserializer(type, reader, 0, -1, false));
cache.Deserializer = deserializer;
Expand All @@ -206,7 +206,7 @@ private async Task<T> ReadRowAsyncImpl<T>(Type type, Row row)
var deserializer = cache.Deserializer;

int hash = GetColumnHash(reader);
if (deserializer.Func is null || deserializer.Hash != hash)
if (deserializer is null || deserializer.Hash != hash)
{
deserializer = new DeserializerState(hash, GetDeserializer(type, reader, 0, -1, false));
cache.Deserializer = deserializer;
Expand Down
4 changes: 2 additions & 2 deletions Dapper/SqlMapper.GridReader.cs
Original file line number Diff line number Diff line change
Expand Up @@ -196,7 +196,7 @@ private IEnumerable<T> ReadImpl<T>(Type type, bool buffered)
var deserializer = cache.Deserializer;

int hash = GetColumnHash(reader);
if (deserializer.Func is null || deserializer.Hash != hash)
if (deserializer is null || deserializer.Hash != hash)
{
deserializer = new DeserializerState(hash, GetDeserializer(type, reader, 0, -1, false));
cache.Deserializer = deserializer;
Expand All @@ -217,7 +217,7 @@ private T ReadRow<T>(Type type, Row row)
var deserializer = cache.Deserializer;

int hash = GetColumnHash(reader);
if (deserializer.Func is null || deserializer.Hash != hash)
if (deserializer is null || deserializer.Hash != hash)
{
deserializer = new DeserializerState(hash, GetDeserializer(type, reader, 0, -1, false));
cache.Deserializer = deserializer;
Expand Down
36 changes: 16 additions & 20 deletions Dapper/SqlMapper.cs
Original file line number Diff line number Diff line change
Expand Up @@ -1219,17 +1219,17 @@ private static IEnumerable<T> QueryImpl<T>(this IDbConnection cnn, CommandDefini
// with the CloseConnection flag, so the reader will deal with the connection; we
// still need something in the "finally" to ensure that broken SQL still results
// in the connection closing itself
var tuple = info.Deserializer;
var deserializer = info.Deserializer;
int hash = GetColumnHash(reader);
if (tuple.Func is null || tuple.Hash != hash)
if (deserializer is null || deserializer.Hash != hash)
{
if (reader.FieldCount == 0) //https://code.google.com/p/dapper-dot-net/issues/detail?id=57
yield break;
tuple = info.Deserializer = new DeserializerState(hash, GetDeserializer(effectiveType, reader, 0, -1, false));
deserializer = info.Deserializer = new DeserializerState(hash, GetDeserializer(effectiveType, reader, 0, -1, false));
if (command.AddToCache) SetQueryCache(identity, info);
}

var func = tuple.Func;
var func = deserializer.Func;
var convertToType = Nullable.GetUnderlyingType(effectiveType) ?? effectiveType;
while (reader.Read())
{
Expand Down Expand Up @@ -1359,15 +1359,15 @@ private static T QueryRowImpl<T>(IDbConnection cnn, Row row, ref CommandDefiniti
[MethodImpl(MethodImplOptions.AggressiveInlining)]
private static T ReadRow<T>(CacheInfo info, Identity identity, ref CommandDefinition command, Type effectiveType, DbDataReader reader)
{
var tuple = info.Deserializer;
var deserializer = info.Deserializer;
int hash = GetColumnHash(reader);
if (tuple.Func is null || tuple.Hash != hash)
if (deserializer is null || deserializer.Hash != hash)
{
tuple = info.Deserializer = new DeserializerState(hash, GetDeserializer(effectiveType, reader, 0, -1, false));
deserializer = info.Deserializer = new DeserializerState(hash, GetDeserializer(effectiveType, reader, 0, -1, false));
if (command.AddToCache) SetQueryCache(identity, info);
}

var func = tuple.Func;
var func = deserializer.Func;
object? val = func(reader);
return GetValue<T>(reader, effectiveType, val);
}
Expand Down Expand Up @@ -1600,19 +1600,17 @@ private static IEnumerable<TReturn> MultiMapImpl<TFirst, TSecond, TThird, TFourt
ownedReader = ExecuteReaderWithFlagsFallback(ownedCommand, wasClosed, CommandBehavior.SequentialAccess | CommandBehavior.SingleResult);
reader = ownedReader;
}
var deserializer = default(DeserializerState);
Func<DbDataReader, object>[]? otherDeserializers;
var deserializer = cinfo.Deserializer;

int hash = GetColumnHash(reader);
if ((deserializer = cinfo.Deserializer).Func is null || (otherDeserializers = cinfo.OtherDeserializers) is null || hash != deserializer.Hash)
if (deserializer?.OtherDeserializers is null || hash != deserializer.Hash)
{
var deserializers = GenerateDeserializers(identity, splitOn, reader);
deserializer = cinfo.Deserializer = new DeserializerState(hash, deserializers[0]);
otherDeserializers = cinfo.OtherDeserializers = deserializers.Skip(1).ToArray();
deserializer = cinfo.Deserializer = new DeserializerState(hash, deserializers[0], deserializers.Skip(1).ToArray());
if (command.AddToCache) SetQueryCache(identity, cinfo);
}

Func<DbDataReader, TReturn> mapIt = GenerateMapper<TFirst, TSecond, TThird, TFourth, TFifth, TSixth, TSeventh, TReturn>(deserializer.Func, otherDeserializers, map);
Func<DbDataReader, TReturn> mapIt = GenerateMapper<TFirst, TSecond, TThird, TFourth, TFifth, TSixth, TSeventh, TReturn>(deserializer.Func, deserializer.OtherDeserializers!, map);

if (mapIt is not null)
{
Expand Down Expand Up @@ -1671,19 +1669,17 @@ private static IEnumerable<TReturn> MultiMapImpl<TReturn>(this IDbConnection? cn
ownedReader = ExecuteReaderWithFlagsFallback(ownedCommand, wasClosed, CommandBehavior.SequentialAccess | CommandBehavior.SingleResult);
reader = ownedReader;
}
DeserializerState deserializer;
Func<DbDataReader, object>[]? otherDeserializers;
var deserializer = cinfo.Deserializer;

int hash = GetColumnHash(reader);
if ((deserializer = cinfo.Deserializer).Func is null || (otherDeserializers = cinfo.OtherDeserializers) is null || hash != deserializer.Hash)
if (deserializer?.OtherDeserializers is null || hash != deserializer.Hash)
{
var deserializers = GenerateDeserializers(identity, splitOn, reader);
deserializer = cinfo.Deserializer = new DeserializerState(hash, deserializers[0]);
otherDeserializers = cinfo.OtherDeserializers = deserializers.Skip(1).ToArray();
deserializer = cinfo.Deserializer = new DeserializerState(hash, deserializers[0], deserializers.Skip(1).ToArray());
if (command.AddToCache) SetQueryCache(identity, cinfo);
}

Func<DbDataReader, TReturn> mapIt = GenerateMapper(types.Length, deserializer.Func, otherDeserializers, map);
Func<DbDataReader, TReturn> mapIt = GenerateMapper(types.Length, deserializer.Func, deserializer.OtherDeserializers!, map);

if (mapIt is not null)
{
Expand Down
152 changes: 152 additions & 0 deletions tests/Dapper.Tests/DeserializerCacheConcurrencyTests.cs
Original file line number Diff line number Diff line change
@@ -0,0 +1,152 @@
using System;
using System.Collections.Concurrent;
using System.Collections.Generic;
using System.Linq;
using System.Threading.Tasks;
using Microsoft.Data.Sqlite;
using Xunit;

namespace Dapper.Tests
{
/// <summary>
/// One SQL string, one connection string and one parameter type resolve to a single
/// <c>Identity</c>, so every caller shares one <c>CacheInfo</c> slot. Only the <em>value</em>
/// of <c>@mode</c> changes the result shape, which is what a branching stored procedure does
/// in production. Each test drives one of the sites that reads the cached deserializer.
/// </summary>
[Collection(NonParallelDefinition.Name)]
public class DeserializerCacheConcurrencyTests
{
// @mode picks the type of Value: text in mode 1, integer in mode 2.
private const string Sql = "select 1 as Id, case when @mode = 1 then 'abc' else 42 end as Value";

// Same idea, but the shape only varies in the second mapped type, so this exercises the
// pairing of the primary deserializer with the rest of the set.
private const string MultiMapSql = "select 1 as Id, 'fixed' as Value, 2 as Id, case when @mode = 1 then 'abc' else 42 end as Label";

private const int Threads = 16, Iterations = 50_000;

public class Row
{
public int Id { get; set; }
public string? Value { get; set; }
}

public class Tag
{
public int Id { get; set; }
public string? Label { get; set; }
}

[Fact]
public void QueryImpl_DoesNotReuseAnotherShapesDeserializer()
=> AssertShapesNeverMix((cnn, mode) => Task.FromResult(cnn.Query<Row>(Sql, new { mode }).Single().Value));

[Fact]
public void ReadRow_DoesNotReuseAnotherShapesDeserializer()
=> AssertShapesNeverMix((cnn, mode) => Task.FromResult(cnn.QuerySingle<Row>(Sql, new { mode }).Value));

[Fact]
public void QueryAsync_DoesNotReuseAnotherShapesDeserializer()
=> AssertShapesNeverMix(async (cnn, mode) => (await cnn.QueryAsync<Row>(Sql, new { mode })).Single().Value);

[Fact]
public void QueryUnbufferedAsync_DoesNotReuseAnotherShapesDeserializer()
=> AssertShapesNeverMix(async (cnn, mode) =>
{
await foreach (var row in cnn.QueryUnbufferedAsync<Row>(Sql, new { mode })) return row.Value;
throw new InvalidOperationException("no rows");
});

[Fact]
public void GridReaderReadImpl_DoesNotReuseAnotherShapesDeserializer()
=> AssertShapesNeverMix((cnn, mode) =>
{
using var grid = cnn.QueryMultiple(Sql, new { mode });
return Task.FromResult(grid.Read<Row>().Single().Value);
});

[Fact]
public void GridReaderReadRow_DoesNotReuseAnotherShapesDeserializer()
=> AssertShapesNeverMix((cnn, mode) =>
{
using var grid = cnn.QueryMultiple(Sql, new { mode });
return Task.FromResult(grid.ReadSingle<Row>().Value);
});

[Fact]
public void GridReaderReadAsyncImpl_DoesNotReuseAnotherShapesDeserializer()
=> AssertShapesNeverMix(async (cnn, mode) =>
{
using var grid = await cnn.QueryMultipleAsync(Sql, new { mode });
return (await grid.ReadAsync<Row>()).Single().Value;
});

[Fact]
public void GridReaderReadRowAsyncImpl_DoesNotReuseAnotherShapesDeserializer()
=> AssertShapesNeverMix(async (cnn, mode) =>
{
using var grid = await cnn.QueryMultipleAsync(Sql, new { mode });
return (await grid.ReadSingleAsync<Row>()).Value;
});

[Fact]
public void MultiMapImplGeneric_DoesNotPairDeserializersFromDifferentShapes()
=> AssertShapesNeverMix((cnn, mode) => Task.FromResult(
cnn.Query<Row, Tag, string?>(MultiMapSql, (_, tag) => tag.Label, new { mode }, splitOn: "Id").Single()));

[Fact]
public void MultiMapImplTypeArray_DoesNotPairDeserializersFromDifferentShapes()
=> AssertShapesNeverMix((cnn, mode) => Task.FromResult(
cnn.Query(MultiMapSql, new[] { typeof(Row), typeof(Tag) }, values => ((Tag)values[1]).Label, new { mode }, splitOn: "Id").Single()));

private static void AssertShapesNeverMix(Func<SqliteConnection, int, Task<string?>> read)
{
var failures = new ConcurrentQueue<string>();
var connections = new List<SqliteConnection>();
var workers = new List<Task>();

for (int i = 0; i < Threads; i++)
{
int mode = (i % 2) + 1;
string expected = mode == 1 ? "abc" : "42";
var connection = OpenConnection();
connections.Add(connection);
workers.Add(Task.Factory.StartNew(() =>
{
for (int n = 0; n < Iterations; n++)
{
try
{
var actual = read(connection, mode).GetAwaiter().GetResult();
if (actual != expected) failures.Enqueue($"mode {mode}: expected {expected}, got {actual}");
}
catch (Exception ex)
{
failures.Enqueue($"mode {mode}: {ex.GetBaseException().Message}");
}
}
}, TaskCreationOptions.LongRunning));
}

try
{
Task.WaitAll(workers.ToArray());
}
finally
{
foreach (var connection in connections) connection.Dispose();
}

Assert.True(failures.IsEmpty, $"{failures.Count} corrupt read(s); first few:{Environment.NewLine}"
+ string.Join(Environment.NewLine, failures.Take(5)));
}

private static SqliteConnection OpenConnection()
{
var connection = new SqliteConnection("Data Source=:memory:");
connection.Open();
return connection;
}
}
}
Loading