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);
        }
    }
}