diff --git a/Dapper/SqlMapper.Async.cs b/Dapper/SqlMapper.Async.cs index eade08cb2..89d615341 100644 --- a/Dapper/SqlMapper.Async.cs +++ b/Dapper/SqlMapper.Async.cs @@ -433,17 +433,17 @@ private static async Task> QueryAsync(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(); - 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) { @@ -1306,19 +1306,19 @@ static async IAsyncEnumerable 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)) diff --git a/Dapper/SqlMapper.CacheInfo.cs b/Dapper/SqlMapper.CacheInfo.cs index 69edc4eea..7fe757fc0 100644 --- a/Dapper/SqlMapper.CacheInfo.cs +++ b/Dapper/SqlMapper.CacheInfo.cs @@ -9,8 +9,7 @@ public static partial class SqlMapper { private sealed class CacheInfo { - public DeserializerState Deserializer { get; set; } - public Func[]? OtherDeserializers { get; set; } + public DeserializerState? Deserializer { get; set; } public Action? ParamReader { get; set; } private int hitCount; public int GetHitCount() { return Interlocked.CompareExchange(ref hitCount, 0, 0); } diff --git a/Dapper/SqlMapper.DeserializerState.cs b/Dapper/SqlMapper.DeserializerState.cs index 4b594e0f5..77e33bf7b 100644 --- a/Dapper/SqlMapper.DeserializerState.cs +++ b/Dapper/SqlMapper.DeserializerState.cs @@ -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 Func; + public readonly Func[]? OtherDeserializers; - public DeserializerState(int hash, Func func) + public DeserializerState(int hash, Func func, Func[]? otherDeserializers = null) { Hash = hash; Func = func; + OtherDeserializers = otherDeserializers; } } } diff --git a/Dapper/SqlMapper.GridReader.Async.cs b/Dapper/SqlMapper.GridReader.Async.cs index a32d53124..ae106aae5 100644 --- a/Dapper/SqlMapper.GridReader.Async.cs +++ b/Dapper/SqlMapper.GridReader.Async.cs @@ -186,7 +186,7 @@ private Func 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; @@ -206,7 +206,7 @@ private async Task ReadRowAsyncImpl(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; diff --git a/Dapper/SqlMapper.GridReader.cs b/Dapper/SqlMapper.GridReader.cs index 1370f7cc9..3bdb7bd5b 100644 --- a/Dapper/SqlMapper.GridReader.cs +++ b/Dapper/SqlMapper.GridReader.cs @@ -196,7 +196,7 @@ private IEnumerable ReadImpl(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; @@ -217,7 +217,7 @@ private T ReadRow(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; diff --git a/Dapper/SqlMapper.cs b/Dapper/SqlMapper.cs index 2fa0e72b7..bd9ced44c 100644 --- a/Dapper/SqlMapper.cs +++ b/Dapper/SqlMapper.cs @@ -1219,17 +1219,17 @@ private static IEnumerable QueryImpl(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()) { @@ -1359,15 +1359,15 @@ private static T QueryRowImpl(IDbConnection cnn, Row row, ref CommandDefiniti [MethodImpl(MethodImplOptions.AggressiveInlining)] private static T ReadRow(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(reader, effectiveType, val); } @@ -1600,19 +1600,17 @@ private static IEnumerable MultiMapImpl[]? 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 mapIt = GenerateMapper(deserializer.Func, otherDeserializers, map); + Func mapIt = GenerateMapper(deserializer.Func, deserializer.OtherDeserializers!, map); if (mapIt is not null) { @@ -1671,19 +1669,17 @@ private static IEnumerable MultiMapImpl(this IDbConnection? cn ownedReader = ExecuteReaderWithFlagsFallback(ownedCommand, wasClosed, CommandBehavior.SequentialAccess | CommandBehavior.SingleResult); reader = ownedReader; } - DeserializerState deserializer; - Func[]? 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 mapIt = GenerateMapper(types.Length, deserializer.Func, otherDeserializers, map); + Func mapIt = GenerateMapper(types.Length, deserializer.Func, deserializer.OtherDeserializers!, map); if (mapIt is not null) { diff --git a/tests/Dapper.Tests/DeserializerCacheConcurrencyTests.cs b/tests/Dapper.Tests/DeserializerCacheConcurrencyTests.cs new file mode 100644 index 000000000..46ab41f18 --- /dev/null +++ b/tests/Dapper.Tests/DeserializerCacheConcurrencyTests.cs @@ -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 +{ + /// + /// One SQL string, one connection string and one parameter type resolve to a single + /// Identity, so every caller shares one CacheInfo slot. Only the value + /// of @mode 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. + /// + [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(Sql, new { mode }).Single().Value)); + + [Fact] + public void ReadRow_DoesNotReuseAnotherShapesDeserializer() + => AssertShapesNeverMix((cnn, mode) => Task.FromResult(cnn.QuerySingle(Sql, new { mode }).Value)); + + [Fact] + public void QueryAsync_DoesNotReuseAnotherShapesDeserializer() + => AssertShapesNeverMix(async (cnn, mode) => (await cnn.QueryAsync(Sql, new { mode })).Single().Value); + + [Fact] + public void QueryUnbufferedAsync_DoesNotReuseAnotherShapesDeserializer() + => AssertShapesNeverMix(async (cnn, mode) => + { + await foreach (var row in cnn.QueryUnbufferedAsync(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().Single().Value); + }); + + [Fact] + public void GridReaderReadRow_DoesNotReuseAnotherShapesDeserializer() + => AssertShapesNeverMix((cnn, mode) => + { + using var grid = cnn.QueryMultiple(Sql, new { mode }); + return Task.FromResult(grid.ReadSingle().Value); + }); + + [Fact] + public void GridReaderReadAsyncImpl_DoesNotReuseAnotherShapesDeserializer() + => AssertShapesNeverMix(async (cnn, mode) => + { + using var grid = await cnn.QueryMultipleAsync(Sql, new { mode }); + return (await grid.ReadAsync()).Single().Value; + }); + + [Fact] + public void GridReaderReadRowAsyncImpl_DoesNotReuseAnotherShapesDeserializer() + => AssertShapesNeverMix(async (cnn, mode) => + { + using var grid = await cnn.QueryMultipleAsync(Sql, new { mode }); + return (await grid.ReadSingleAsync()).Value; + }); + + [Fact] + public void MultiMapImplGeneric_DoesNotPairDeserializersFromDifferentShapes() + => AssertShapesNeverMix((cnn, mode) => Task.FromResult( + cnn.Query(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> read) + { + var failures = new ConcurrentQueue(); + var connections = new List(); + var workers = new List(); + + 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; + } + } +}