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