| File: StreamTracker.cs | Web Access |
| Project: src\aspnetcore\src\SignalR\server\Core\src\Microsoft.AspNetCore.SignalR.Core.csproj (Microsoft.AspNetCore.SignalR.Core) |
// 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.Concurrent; using System.Diagnostics; using System.Diagnostics.CodeAnalysis; using System.Linq; using System.Reflection; using System.Runtime.CompilerServices; using System.Threading.Channels; using Microsoft.AspNetCore.SignalR.Protocol; namespace Microsoft.AspNetCore.SignalR; internal sealed class StreamTracker { private static readonly MethodInfo _buildConverterMethod = typeof(StreamTracker).GetMethods(BindingFlags.NonPublic | BindingFlags.Static).Single(m => m.Name.Equals(nameof(BuildStream))); private readonly object[] _streamConverterArgs; private readonly ConcurrentDictionary<string, (long Owner, IStreamConverter Converter)> _lookup = new(); private long _nextStreamOwner; public StreamTracker(int streamBufferCapacity) { _streamConverterArgs = new object[] { streamBufferCapacity }; } public long GetNextStreamOwner() { return Interlocked.Increment(ref _nextStreamOwner); } /// <summary> /// Creates a new stream and returns the ChannelReader for it as an object. /// </summary> [UnconditionalSuppressMessage("Trimming", "IL2060:MakeGenericMethod", Justification = "BuildStream doesn't have trimming annotations.")] [UnconditionalSuppressMessage("AOT", "IL3050:RequiresDynamicCode", Justification = "HubMethodDescriptor checks for ValueType streaming item types when PublishAot=true. Developers will get an exception in this situation before publishing.")] public object AddStream(string streamId, Type itemType, Type targetType, long streamOwner) { Debug.Assert(RuntimeFeature.IsDynamicCodeSupported || !itemType.IsValueType, "HubMethodDescriptor ensures itemType is not a ValueType when PublishAot=true."); var newConverter = (IStreamConverter)_buildConverterMethod.MakeGenericMethod(itemType).Invoke(null, _streamConverterArgs)!; if (!_lookup.TryAdd(streamId, (streamOwner, newConverter))) { throw new HubException($"Stream ID '{streamId}' is already in use."); } return newConverter.GetReaderAsObject(targetType); } private bool TryGetRegistration(string streamId, out (long Owner, IStreamConverter Converter) registration) { if (_lookup.TryGetValue(streamId, out registration)) { return true; } return false; } public bool TryProcessItem(StreamItemMessage message, [NotNullWhen(true)] out Task? task) { if (TryGetRegistration(message.InvocationId!, out var registration)) { task = registration.Converter.WriteToStream(message.Item); return true; } task = default; return false; } public Type GetStreamItemType(string streamId) { if (TryGetRegistration(streamId, out var registration)) { return registration.Converter.GetItemType(); } throw new KeyNotFoundException($"No stream with id '{streamId}' could be found."); } public bool TryComplete(CompletionMessage message) { if (!_lookup.TryRemove(message.InvocationId!, out var registration)) { return false; } registration.Converter.TryComplete(message.HasResult || message.Error == null ? null : new HubException(message.Error)); return true; } public bool TryComplete(string streamId, long streamOwner) { if (!_lookup.TryGetValue(streamId, out var registration) || registration.Owner != streamOwner) { return false; } if (!_lookup.TryRemove(KeyValuePair.Create(streamId, registration))) { return false; } registration.Converter.TryComplete(null); return true; } public void CompleteAll(Exception ex) { foreach (var converter in _lookup) { converter.Value.Converter.TryComplete(ex); } } private static IStreamConverter BuildStream<T>(int streamBufferCapacity) { return new ChannelConverter<T>(streamBufferCapacity); } private interface IStreamConverter { Type GetItemType(); object GetReaderAsObject(Type type); Task WriteToStream(object? item); void TryComplete(Exception? ex); } private sealed class ChannelConverter<T> : IStreamConverter { private readonly Channel<T?> _channel; public ChannelConverter(int streamBufferCapacity) { _channel = Channel.CreateBounded<T?>(streamBufferCapacity); } public Type GetItemType() { return typeof(T); } public object GetReaderAsObject(Type type) { if (ReflectionHelper.IsIAsyncEnumerable(type)) { return _channel.Reader.ReadAllAsync(); } else { return _channel.Reader; } } public Task WriteToStream(object? o) { return _channel.Writer.WriteAsync((T?)o).AsTask(); } public void TryComplete(Exception? ex) { _channel.Writer.TryComplete(ex); } } }