diff --git a/src/OpenRiaServices.Client.DomainClients.Http/Framework/MessagePack/Converters/ObjectConverterFactory.cs b/src/OpenRiaServices.Client.DomainClients.Http/Framework/MessagePack/Converters/ObjectConverterFactory.cs index 005919934..c9284d684 100644 --- a/src/OpenRiaServices.Client.DomainClients.Http/Framework/MessagePack/Converters/ObjectConverterFactory.cs +++ b/src/OpenRiaServices.Client.DomainClients.Http/Framework/MessagePack/Converters/ObjectConverterFactory.cs @@ -49,10 +49,6 @@ public ObjectConverterFactory(IEnumerable entityTypes) if (knownTypes.Count <= 1) continue; - // Skip base types, it has KnownType attributes and is handled correctly - if (item.Key.BaseType == typeof(Entity)) - continue; - converters.Add(item.Key, (MessagePackConverter)Activator.CreateInstance(typeof(ObjectConverter<>).MakeGenericType(item.Key), [knownTypes])!); } converters.Add(typeof(object), new ObjectConverter(allTypes)); @@ -60,6 +56,15 @@ public ObjectConverterFactory(IEnumerable entityTypes) _converters = converters.ToFrozenDictionary(); } + /// + /// Generate DerivedTypeUnion required for the ObjectConverterFactory to work. + /// + /// + /// In order to prevent Nerdabank.MessagePack from creating additional discriminators we must disable the built in inheritance support for the types. + /// + public IEnumerable GetDerivedTypeUnions() + => _converters.Keys.Select(DerivedTypeUnion.CreateDisabled); + MessagePackConverter? IMessagePackConverterFactory.CreateConverter(Type type, ITypeShape? shape, in ConverterContext context) { // Only return type shapes for specified converter, this allows nerdbank default converters to be generated for other providers diff --git a/src/OpenRiaServices.Client.DomainClients.Http/Framework/MessagePackHttpDomainClientFactory.cs b/src/OpenRiaServices.Client.DomainClients.Http/Framework/MessagePackHttpDomainClientFactory.cs index e92ac3cc2..f0b85d842 100644 --- a/src/OpenRiaServices.Client.DomainClients.Http/Framework/MessagePackHttpDomainClientFactory.cs +++ b/src/OpenRiaServices.Client.DomainClients.Http/Framework/MessagePackHttpDomainClientFactory.cs @@ -45,9 +45,11 @@ internal MessagePackSerializer GetSerializer(Type service, IEnumerable kno { return _serializerCache.GetOrAdd(service, static (_, args) => { + var converterFactory = new ObjectConverterFactory(args.knownTypes); return args.Item1.BaseSerializerSerializer with { - ConverterFactories = [new ObjectConverterFactory(args.knownTypes), .. args.Item1.BaseSerializerSerializer.ConverterFactories] + ConverterFactories = [converterFactory, .. args.Item1.BaseSerializerSerializer.ConverterFactories], + DerivedTypeUnions = [..converterFactory.GetDerivedTypeUnions(), .. args.Item1.BaseSerializerSerializer.DerivedTypeUnions] }; }, (this, knownTypes)); diff --git a/src/Test/OpenRiaservices.EndToEnd.Wcf.Test/Data/InheritanceTests.cs b/src/Test/OpenRiaservices.EndToEnd.Wcf.Test/Data/InheritanceTests.cs index 03f5af9a4..6cf1e9fcf 100644 --- a/src/Test/OpenRiaservices.EndToEnd.Wcf.Test/Data/InheritanceTests.cs +++ b/src/Test/OpenRiaservices.EndToEnd.Wcf.Test/Data/InheritanceTests.cs @@ -5,12 +5,10 @@ using System.Globalization; using System.Linq; using System.Reflection; +using System.Threading.Tasks; using Cities; using Microsoft.Silverlight.Testing; -using Microsoft.VisualStudio.TestTools.UnitTesting; using OpenRiaServices.Silverlight.Testing; -using TestDomainServices; -using TestDomainServices.NamedUpdates; namespace OpenRiaServices.Client.Test { @@ -394,6 +392,27 @@ public void Inherit_Run_Call_Derived_Custom_Method_On_Abstract_Base_() EnqueueTestComplete(); } + [TestMethod] + [Description("Ensure querying works for 3 levels of inheritance")] + public async Task Inherit_Queries() + { + // Inheritance is City <-- CityWithEditHistory <-- CityWithInfo + // Ensure we can methods retrieve entities at all levels of the hierarchy + CityDomainContext baseContext = new CityDomainContext(); + CityDomainContext middleContext = new CityDomainContext(); + CityDomainContext mostDerivedContext = new CityDomainContext(); + + var baseResult = await baseContext.LoadAsync(baseContext.GetCitiesQuery()); + var editResult = await middleContext.LoadAsync(middleContext.GetCitiesWithEditHistoryQuery()); + var mostDerivedResult = await mostDerivedContext.LoadAsync(mostDerivedContext.GetCitiesWithInfoQuery()); + + // CityWithEditHistory is abstract so should be return exactly the same number of entities as CityWithInfo + Assert.HasCount(2, mostDerivedResult.Entities); + Assert.HasCount(2, editResult.Entities); + Assert.HasCount(2, baseResult.Entities.OfType()); + Assert.HasCount(11, baseResult.Entities); + } + #endregion //Custom methods #region CUD