diff --git a/Dapper/SqlDataRecordListTVPParameter.cs b/Dapper/SqlDataRecordListTVPParameter.cs index 28a6ab27c..fdb522de3 100644 --- a/Dapper/SqlDataRecordListTVPParameter.cs +++ b/Dapper/SqlDataRecordListTVPParameter.cs @@ -2,7 +2,6 @@ using System.Collections; using System.Collections.Generic; using System.Data; -using System.Linq; using System.Reflection; using System.Reflection.Emit; @@ -37,7 +36,15 @@ void SqlMapper.ICustomQueryParameter.AddParameter(IDbCommand command, string nam internal static void Set(IDbDataParameter parameter, IEnumerable? data, string? typeName) { - parameter.Value = data is not null && data.Any() ? data : null; + // don't enumerate "data" here - it may be a single-pass source (an open reader, a + // streaming iterator, etc); only short-circuit to null when we can tell for free + // that it is empty, since providers reject an empty TVP enumerable + parameter.Value = data switch + { + null => null, + IReadOnlyCollection { Count: 0 } => null, + _ => data, + }; StructuredHelper.ConfigureTVP(parameter, typeName); } } diff --git a/tests/Dapper.Tests/ParameterTests.cs b/tests/Dapper.Tests/ParameterTests.cs index 5eb455c65..9fee6a434 100644 --- a/tests/Dapper.Tests/ParameterTests.cs +++ b/tests/Dapper.Tests/ParameterTests.cs @@ -131,6 +131,25 @@ private static IEnumerable CreateSqlDataRecordList(IDbConnection co return number_list; } + // wraps a sequence to prove it is only ever enumerated once, guarding against + // https://github.com/DapperLib/Dapper/issues/2064 + private class SingleEnumerationEnumerable : IEnumerable + { + private readonly IEnumerable data; + + public SingleEnumerationEnumerable(IEnumerable data) => this.data = data; + + public bool WasEnumerated { get; private set; } + + public IEnumerator GetEnumerator() + { + Assert.False(WasEnumerated, "Source was enumerated more than once."); + WasEnumerated = true; + return data.GetEnumerator(); + } + + IEnumerator IEnumerable.GetEnumerator() => GetEnumerator(); + } private class IntDynamicParam : SqlMapper.IDynamicParameters { @@ -486,6 +505,80 @@ public void TestSqlDataRecordListParametersWithAsTableValuedParameter() } } + [Fact] + public void TestSqlDataRecordListParametersWithAsTableValuedParameterSinglePassSource() + { + try + { + connection.Execute("CREATE TYPE int_list_type AS TABLE (n int NOT NULL PRIMARY KEY)"); + connection.Execute("CREATE PROC get_ints @integers int_list_type READONLY AS select * from @integers"); + + // SingleEnumerationEnumerable has to be instantiated against the provider's + // concrete SqlDataRecord type here, not the shared IDataRecord interface. SqlClient + // recognizes a structured/TVP parameter value via a hard "value is + // IEnumerable" check (see SqlParameter.CoerceValue); that check + // inspects the object's actual runtime interface implementations, and generic + // interface implementations aren't covariant the way assignments are, so a wrapper + // built as IEnumerable never satisfies it, even though every element it + // yields really is a SqlDataRecord. Wrapping the provider-specific list directly + // (matching the pattern already used in TestSqlDataRecordListParametersWithTypeHandlers + // below) keeps that concrete typing intact while still proving single enumeration. + SqlMapper.ICustomQueryParameter tvp; +#pragma warning disable CS0618 // Type or member is obsolete + if (connection is System.Data.SqlClient.SqlConnection) + { + var records = new SingleEnumerationEnumerable(CreateSqlDataRecordList_SD(new int[] { 1, 2, 3 })); + tvp = records.AsTableValuedParameter(); + } +#pragma warning restore CS0618 // Type or member is obsolete + else if (connection is Microsoft.Data.SqlClient.SqlConnection) + { + var records = new SingleEnumerationEnumerable(CreateSqlDataRecordList_MD(new int[] { 1, 2, 3 })); + tvp = records.AsTableValuedParameter(); + } + else + { + throw new ArgumentException(nameof(connection)); + } + + var nums = connection.Query("get_ints", new { integers = tvp }, commandType: CommandType.StoredProcedure).ToList(); + Assert.Equal(new int[] { 1, 2, 3 }, nums); + } + finally + { + try + { + connection.Execute("DROP PROC get_ints"); + } + finally + { + connection.Execute("DROP TYPE int_list_type"); + } + } + } + + [Fact] + public void AsTableValuedParameterDoesNotEnumerateNonCollectionSource() + { + var records = new SingleEnumerationEnumerable(Enumerable.Empty()); + var parameter = Provider.CreateRawParameter("integers", DBNull.Value); + + SqlDataRecordListTVPParameter.Set(parameter, records, "int_list_type"); + + Assert.False(records.WasEnumerated); + Assert.Same(records, parameter.Value); + } + + [Fact] + public void AsTableValuedParameterNullsOutEmptyCollectionSource() + { + var parameter = Provider.CreateRawParameter("integers", DBNull.Value); + + SqlDataRecordListTVPParameter.Set(parameter, Array.Empty(), "int_list_type"); + + Assert.Null(parameter.Value); + } + [Fact] public void TestEmptySqlDataRecordListParametersWithAsTableValuedParameter() {