diff --git a/src/libraries/System.Private.DataContractSerialization/src/System.Private.DataContractSerialization.csproj b/src/libraries/System.Private.DataContractSerialization/src/System.Private.DataContractSerialization.csproj index 8ee00c3c19581b..bd9a7cc8028184 100644 --- a/src/libraries/System.Private.DataContractSerialization/src/System.Private.DataContractSerialization.csproj +++ b/src/libraries/System.Private.DataContractSerialization/src/System.Private.DataContractSerialization.csproj @@ -145,6 +145,7 @@ + @@ -163,6 +164,7 @@ + diff --git a/src/libraries/System.Private.DataContractSerialization/src/System/Runtime/Serialization/ContextAware.cs b/src/libraries/System.Private.DataContractSerialization/src/System/Runtime/Serialization/ContextAware.cs new file mode 100644 index 00000000000000..3e49158213c412 --- /dev/null +++ b/src/libraries/System.Private.DataContractSerialization/src/System/Runtime/Serialization/ContextAware.cs @@ -0,0 +1,103 @@ +// Licensed to the .NET Foundation under one or more agreements. +// The .NET Foundation licenses this file to you under the MIT license. + +using System.Collections; +using System.Collections.Concurrent; +using System.Diagnostics.CodeAnalysis; +using System.Reflection; +using System.Runtime.CompilerServices; +using System.Runtime.Loader; +using System.Runtime.Serialization.DataContracts; + +namespace System.Runtime.Serialization +{ + internal sealed class ContextAwareDataContractIndex + { + private (DataContract? strong, WeakReference? weak)[] _contracts; + private ConditionalWeakTable _keepAlive; + + public int Length => _contracts.Length; + + public ContextAwareDataContractIndex(int size) + { + _contracts = new (DataContract?, WeakReference?)[size]; + _keepAlive = new ConditionalWeakTable(); + } + + public DataContract? GetItem(int index) => _contracts[index].strong ?? (_contracts[index].weak?.TryGetTarget(out DataContract? ret) == true ? ret : null); + + public void SetItem(int index, DataContract dataContract) + { + // Check for unloadability to decide how to store the value + AssemblyLoadContext? alc = AssemblyLoadContext.GetLoadContext(dataContract.UnderlyingType.Assembly); + if (alc == null || !alc.IsCollectible) + { + _contracts[index].strong = dataContract; + } + else + { + _contracts[index].weak = new WeakReference(dataContract); + _keepAlive.Add(dataContract.UnderlyingType, dataContract); + } + } + + public void Resize(int newSize) + { + Array.Resize<(DataContract?, WeakReference?)>(ref _contracts, newSize); + } + } + + internal sealed class ContextAwareDictionary + where TKey : Type + where TValue : class? + { + private readonly ConcurrentDictionary _fastDictionary = new(); + private readonly ConditionalWeakTable _collectibleTable = new(); + + + internal TValue GetOrAdd(TKey t, Func f) + { + TValue? ret; + + // The fast and most common default case + if (_fastDictionary.TryGetValue(t, out ret)) + return ret; + + // Common case for collectible contexts + if (_collectibleTable.TryGetValue(t, out ret)) + return ret; + + // Not found. Do the slower work of creating the value in the correct collection. + AssemblyLoadContext? alc = AssemblyLoadContext.GetLoadContext(t.Assembly); + + // Null and non-collectible load contexts use the default table + if (alc == null || !alc.IsCollectible) + { + // The create delegate here could be quite expensive. ConcurrentDictionary semantics would let use + // do this without a lock and not corrupt the dictionary, but the delegate still might be called multiple + // times. So we use a lock to ensure the delegate is only called once. Keep the lock off the hot path. + if (!_fastDictionary.TryGetValue(t, out ret)) + { + lock (_fastDictionary) + { + return _fastDictionary.GetOrAdd(t, f); + } + } + } + + // Collectible load contexts should use the ConditionalWeakTable so they can be unloaded + else + { + if (!_collectibleTable.TryGetValue(t, out ret)) + { + lock (_collectibleTable) + { + return _collectibleTable.GetValue(t, k => f(k)); + } + } + } + + return ret; + } + } +} diff --git a/src/libraries/System.Private.DataContractSerialization/src/System/Runtime/Serialization/DataContract.cs b/src/libraries/System.Private.DataContractSerialization/src/System/Runtime/Serialization/DataContract.cs index 78fe30328b8c74..1028c02ff4c1a0 100644 --- a/src/libraries/System.Private.DataContractSerialization/src/System/Runtime/Serialization/DataContract.cs +++ b/src/libraries/System.Private.DataContractSerialization/src/System/Runtime/Serialization/DataContract.cs @@ -299,9 +299,9 @@ internal virtual bool IsValidContract() internal class DataContractCriticalHelper { private static readonly ConcurrentDictionary s_typeToIDCache = new(); - private static DataContract[] s_dataContractCache = new DataContract[32]; + private static readonly ContextAwareDataContractIndex s_dataContractCache = new(32); private static int s_dataContractID; - private static readonly ConcurrentDictionary s_typeToBuiltInContract = new(); + private static readonly ContextAwareDictionary s_typeToBuiltInContract = new(); private static Dictionary? s_nameToBuiltInContract; private static Dictionary? s_typeNameToBuiltInContract; private static readonly Hashtable s_namespaces = new Hashtable(); @@ -335,7 +335,7 @@ internal class DataContractCriticalHelper [RequiresUnreferencedCode(DataContract.SerializerTrimmerWarning)] internal static DataContract GetDataContractSkipValidation(int id, RuntimeTypeHandle typeHandle, Type? type) { - DataContract dataContract = s_dataContractCache[id]; + DataContract? dataContract = s_dataContractCache.GetItem(id); if (dataContract == null) { dataContract = CreateDataContract(id, typeHandle, type); @@ -351,13 +351,13 @@ internal static DataContract GetDataContractSkipValidation(int id, RuntimeTypeHa [RequiresUnreferencedCode(DataContract.SerializerTrimmerWarning)] internal static DataContract GetGetOnlyCollectionDataContractSkipValidation(int id, RuntimeTypeHandle typeHandle, Type? type) { - DataContract dataContract = s_dataContractCache[id] ?? CreateGetOnlyCollectionDataContract(id, typeHandle, type); + DataContract dataContract = s_dataContractCache.GetItem(id) ?? CreateGetOnlyCollectionDataContract(id, typeHandle, type); return dataContract; } internal static DataContract GetDataContractForInitialization(int id) { - DataContract dataContract = s_dataContractCache[id]; + DataContract? dataContract = s_dataContractCache.GetItem(id); if (dataContract == null) { throw new SerializationException(SR.DataContractCacheOverflow); @@ -368,14 +368,14 @@ internal static DataContract GetDataContractForInitialization(int id) internal static int GetIdForInitialization(ClassDataContract classContract) { int id = DataContract.GetId(classContract.TypeForInitialization.TypeHandle); - if (id < s_dataContractCache.Length && ContractMatches(classContract, s_dataContractCache[id])) + if (id < s_dataContractCache.Length && ContractMatches(classContract, s_dataContractCache.GetItem(id))) { return id; } int currentDataContractId = DataContractCriticalHelper.s_dataContractID; for (int i = 0; i < currentDataContractId; i++) { - if (ContractMatches(classContract, s_dataContractCache[i])) + if (ContractMatches(classContract, s_dataContractCache.GetItem(id))) { return i; } @@ -383,7 +383,7 @@ internal static int GetIdForInitialization(ClassDataContract classContract) throw new SerializationException(SR.DataContractCacheOverflow); } - private static bool ContractMatches(DataContract contract, DataContract cachedContract) + private static bool ContractMatches(DataContract contract, DataContract? cachedContract) { return (cachedContract != null && cachedContract.UnderlyingType == contract.UnderlyingType); } @@ -410,7 +410,7 @@ internal static int GetId(RuntimeTypeHandle typeHandle) Debug.Fail("DataContract cache overflow"); throw new SerializationException(SR.DataContractCacheOverflow); } - Array.Resize(ref s_dataContractCache, newSize); + s_dataContractCache.Resize(newSize); } return nextId; }); @@ -427,12 +427,12 @@ internal static int GetId(RuntimeTypeHandle typeHandle) [RequiresUnreferencedCode(DataContract.SerializerTrimmerWarning)] private static DataContract CreateDataContract(int id, RuntimeTypeHandle typeHandle, Type? type) { - DataContract? dataContract = s_dataContractCache[id]; + DataContract? dataContract = s_dataContractCache.GetItem(id); if (dataContract == null) { lock (s_createDataContractLock) { - dataContract = s_dataContractCache[id]; + dataContract = s_dataContractCache.GetItem(id); if (dataContract == null) { type ??= Type.GetTypeFromHandle(typeHandle)!; @@ -499,7 +499,7 @@ private static void AssignDataContractToId(DataContract dataContract, int id) { lock (s_cacheLock) { - s_dataContractCache[id] = dataContract; + s_dataContractCache.SetItem(id, dataContract); } } @@ -510,7 +510,7 @@ private static DataContract CreateGetOnlyCollectionDataContract(int id, RuntimeT DataContract? dataContract = null; lock (s_createDataContractLock) { - dataContract = s_dataContractCache[id]; + dataContract = s_dataContractCache.GetItem(id); if (dataContract == null) { type ??= Type.GetTypeFromHandle(typeHandle)!; diff --git a/src/libraries/System.Runtime.Serialization.Xml/tests/DataContractSerializer.cs b/src/libraries/System.Runtime.Serialization.Xml/tests/DataContractSerializer.cs index 8484775ea230a0..0ae22abf361777 100644 --- a/src/libraries/System.Runtime.Serialization.Xml/tests/DataContractSerializer.cs +++ b/src/libraries/System.Runtime.Serialization.Xml/tests/DataContractSerializer.cs @@ -19,7 +19,8 @@ using System.Xml.Schema; using Xunit; using System.Runtime.Serialization.Tests; - +using System.Runtime.CompilerServices; +using System.Runtime.Loader; public static partial class DataContractSerializerTests { @@ -1115,6 +1116,88 @@ public static void DCS_TypeNamesWithSpecialCharacters() Assert.Equal(x.PropertyNameWithSpecialCharacters\u6F22\u00F1, y.PropertyNameWithSpecialCharacters\u6F22\u00F1); } + [Fact] +#if XMLSERIALIZERGENERATORTESTS + // Lack of AssemblyDependencyResolver results in assemblies that are not loaded by path to get + // loaded in the default ALC, which causes problems for this test. + [SkipOnPlatform(TestPlatforms.Browser, "AssemblyDependencyResolver not supported in wasm")] +#endif + [ActiveIssue("34072", TestRuntimes.Mono)] + public static void DCS_TypeInCollectibleALC() + { + ExecuteAndUnload("SerializableAssembly.dll", "SerializationTypes.SimpleType", makeCollection: false, out var weakRef); + + for (int i = 0; weakRef.IsAlive && i < 10; i++) + { + GC.Collect(); + GC.WaitForPendingFinalizers(); + } + Assert.True(!weakRef.IsAlive); + } + + [Fact] +#if XMLSERIALIZERGENERATORTESTS + // Lack of AssemblyDependencyResolver results in assemblies that are not loaded by path to get + // loaded in the default ALC, which causes problems for this test. + [SkipOnPlatform(TestPlatforms.Browser, "AssemblyDependencyResolver not supported in wasm")] +#endif + [ActiveIssue("34072", TestRuntimes.Mono)] + public static void DCS_CollectionTypeInCollectibleALC() + { + ExecuteAndUnload("SerializableAssembly.dll", "SerializationTypes.SimpleType", makeCollection: true, out var weakRef); + + for (int i = 0; weakRef.IsAlive && i < 10; i++) + { + GC.Collect(); + GC.WaitForPendingFinalizers(); + } + Assert.True(!weakRef.IsAlive); + } + + [MethodImpl(MethodImplOptions.NoInlining)] + private static void ExecuteAndUnload(string assemblyfile, string typename, bool makeCollection, out WeakReference wref) + { + var fullPath = Path.GetFullPath(assemblyfile); + var alc = new TestAssemblyLoadContext("DataContractSerializerTests", true, fullPath); + object obj; + wref = new WeakReference(alc); + + // Load assembly by path. By name, and it gets loaded in the default ALC. + var asm = alc.LoadFromAssemblyPath(fullPath); + + // Ensure the type loaded in the intended non-Default ALC + var type = asm.GetType(typename); + Assert.Equal(AssemblyLoadContext.GetLoadContext(type.Assembly), alc); + Assert.NotEqual(alc, AssemblyLoadContext.Default); + + if (makeCollection) + { + int arrayLength = 3; + var array = Array.CreateInstance(type, arrayLength); + for (int i = 0; i < arrayLength; i++) + { + array.SetValue(Activator.CreateInstance(type), i); + } + type = array.GetType(); + obj = array; + } + else + { + obj = Activator.CreateInstance(type); + } + + // Round-Trip the instance + var dcs = new DataContractSerializer(type); + var rtobj = DataContractSerializerHelper.SerializeAndDeserialize(obj, null, null, () => dcs, true, false); + Assert.NotNull(rtobj); + if (makeCollection) + Assert.Equal(obj, rtobj); + else + Assert.True(rtobj.Equals(obj)); + + alc.Unload(); + } + [Fact] public static void DCS_JaggedArrayAsRoot() { diff --git a/src/libraries/System.Runtime.Serialization.Xml/tests/ReflectionOnly/System.Runtime.Serialization.Xml.ReflectionOnly.Tests.csproj b/src/libraries/System.Runtime.Serialization.Xml/tests/ReflectionOnly/System.Runtime.Serialization.Xml.ReflectionOnly.Tests.csproj index bdc0d2c86c22c0..d128172dd2acc6 100644 --- a/src/libraries/System.Runtime.Serialization.Xml/tests/ReflectionOnly/System.Runtime.Serialization.Xml.ReflectionOnly.Tests.csproj +++ b/src/libraries/System.Runtime.Serialization.Xml/tests/ReflectionOnly/System.Runtime.Serialization.Xml.ReflectionOnly.Tests.csproj @@ -3,10 +3,14 @@ $(DefineConstants);ReflectionOnly $(NetCoreAppCurrent) + + + + - + diff --git a/src/libraries/System.Runtime.Serialization.Xml/tests/System.Runtime.Serialization.Xml.Tests.csproj b/src/libraries/System.Runtime.Serialization.Xml/tests/System.Runtime.Serialization.Xml.Tests.csproj index 1628ec4ae6ebd1..c299e17295d0a0 100644 --- a/src/libraries/System.Runtime.Serialization.Xml/tests/System.Runtime.Serialization.Xml.Tests.csproj +++ b/src/libraries/System.Runtime.Serialization.Xml/tests/System.Runtime.Serialization.Xml.Tests.csproj @@ -2,9 +2,13 @@ $(NetCoreAppCurrent) + + + + - +