diff --git a/src/Avalonia.Base/Data/Core/Plugins/ObservableStreamPlugin.cs b/src/Avalonia.Base/Data/Core/Plugins/ObservableStreamPlugin.cs index 14ca8ee79e..c41097c274 100644 --- a/src/Avalonia.Base/Data/Core/Plugins/ObservableStreamPlugin.cs +++ b/src/Avalonia.Base/Data/Core/Plugins/ObservableStreamPlugin.cs @@ -2,6 +2,9 @@ // Licensed under the MIT license. See licence.md file in the project root for full license information. using System; +using System.Linq; +using System.Reactive.Linq; +using System.Reflection; namespace Avalonia.Data.Core.Plugins { @@ -10,12 +13,19 @@ namespace Avalonia.Data.Core.Plugins /// public class ObservableStreamPlugin : IStreamPlugin { + static MethodInfo observableSelect; + /// /// Checks whether this plugin handles the specified value. /// /// A weak reference to the value. /// True if the plugin can handle the value; otherwise false. - public virtual bool Match(WeakReference reference) => reference.Target is IObservable; + public virtual bool Match(WeakReference reference) + { + return reference.Target.GetType().GetInterfaces().Any(x => + x.IsGenericType && + x.GetGenericTypeDefinition() == typeof(IObservable<>)); + } /// /// Starts producing output based on the specified value. @@ -26,7 +36,69 @@ namespace Avalonia.Data.Core.Plugins /// public virtual IObservable Start(WeakReference reference) { - return reference.Target as IObservable; + var target = reference.Target; + + // If the observable returns a reference type then we can cast it. + if (target is IObservable result) + { + return result; + }; + + // If the observable returns a value type then we need to call Observable.Select on it. + // First get the type of T in `IObservable`. + var sourceType = reference.Target.GetType().GetInterfaces().First(x => + x.IsGenericType && + x.GetGenericTypeDefinition() == typeof(IObservable<>)).GetGenericArguments()[0]; + + // Get the Observable.Select method. + var select = GetObservableSelect(sourceType); + + // Make a Box<> delegate of the correct type. + var funcType = typeof(Func<,>).MakeGenericType(sourceType, typeof(object)); + var box = GetType().GetMethod(nameof(Box), BindingFlags.Static | BindingFlags.NonPublic) + .MakeGenericMethod(sourceType) + .CreateDelegate(funcType); + + // Call Observable.Select(target, box); + return (IObservable)select.Invoke( + null, + new object[] { target, box }); + } + + private static MethodInfo GetObservableSelect(Type source) + { + return GetObservableSelect().MakeGenericMethod(source, typeof(object)); } + + private static MethodInfo GetObservableSelect() + { + if (observableSelect == null) + { + observableSelect = typeof(Observable).GetRuntimeMethods().First(x => + { + if (x.Name == nameof(Observable.Select) && + x.ContainsGenericParameters && + x.GetGenericArguments().Length == 2) + { + var parameters = x.GetParameters(); + + if (parameters.Length == 2 && + parameters[0].ParameterType.IsConstructedGenericType && + parameters[0].ParameterType.GetGenericTypeDefinition() == typeof(IObservable<>) && + parameters[1].ParameterType.IsConstructedGenericType && + parameters[1].ParameterType.GetGenericTypeDefinition() == typeof(Func<,>)) + { + return true; + } + } + + return false; + }); + } + + return observableSelect; + } + + private static object Box(T value) => (object)value; } } diff --git a/tests/Avalonia.Base.UnitTests/Data/Core/ExpressionObserverTests_Observable.cs b/tests/Avalonia.Base.UnitTests/Data/Core/ExpressionObserverTests_Observable.cs index 701fdbce9c..4585181ab7 100644 --- a/tests/Avalonia.Base.UnitTests/Data/Core/ExpressionObserverTests_Observable.cs +++ b/tests/Avalonia.Base.UnitTests/Data/Core/ExpressionObserverTests_Observable.cs @@ -150,6 +150,26 @@ namespace Avalonia.Base.UnitTests.Data.Core } } + [Fact] + public void Should_Work_With_Value_Type() + { + using (var sync = UnitTestSynchronizationContext.Begin()) + { + var source = new BehaviorSubject(1); + var data = new { Foo = source }; + var target = ExpressionObserver.Create(data, o => o.Foo.StreamBinding()); + var result = new List(); + + var sub = target.Subscribe(x => result.Add((int)x)); + source.OnNext(42); + sync.ExecutePostedCallbacks(); + + Assert.Equal(new[] { 1, 42 }, result); + + GC.KeepAlive(data); + } + } + private class Class1 : NotifyingBase { public Subject Next { get; } = new Subject();