using ErsatzTV.Controllers; using ErsatzTV.Core.Domain; using ErsatzTV.Core.Interfaces.Streaming; using ErsatzTV.Core.Streaming; using Microsoft.AspNetCore.Http; using Microsoft.AspNetCore.Mvc; using Microsoft.AspNetCore.Routing; using Microsoft.Extensions.DependencyInjection; using NSubstitute; using NUnit.Framework; using Shouldly; namespace ErsatzTV.Tests.Controllers; [TestFixture] public class TrackedFileStreamResultTests { [Test] public async Task Should_Track_Session_Only_While_Result_Executes() { var tracker = new DirectStreamSessionTracker(); var stream = new BlockingReadStream(); TrackedFileStreamResult result = new( stream, "video/mp2t", tracker, "1", StreamingMode.TransportStream); ActionContext context = GetActionContext(); tracker.IsActive("1").ShouldBeFalse(); Task execute = result.ExecuteResultAsync(context); await stream.WaitForRead(); tracker.IsActive("1").ShouldBeTrue(); tracker.GetViewerCount("1").ShouldBe(1); stream.Complete(); await execute; tracker.IsActive("1").ShouldBeFalse(); tracker.GetViewerCount("1").ShouldBe(0); } [Test] public async Task Should_Remove_Session_When_Response_Stream_Fails() { var tracker = new DirectStreamSessionTracker(); TrackedFileStreamResult result = new( new ThrowingReadStream(new IOException("stream failed")), "video/mp2t", tracker, "1", StreamingMode.HttpLiveStreamingDirect); Func execute = () => result.ExecuteResultAsync(GetActionContext()); await execute.ShouldThrowAsync(); tracker.IsActive("1").ShouldBeFalse(); tracker.GetViewerCount("1").ShouldBe(0); } [Test] public async Task Should_Remove_Session_When_Response_Stream_Is_Aborted() { var tracker = new DirectStreamSessionTracker(); TrackedFileStreamResult result = new( new ThrowingReadStream(new OperationCanceledException("client aborted")), "video/mp2t", tracker, "1", StreamingMode.HttpLiveStreamingDirect); await result.ExecuteResultAsync(GetActionContext()); tracker.IsActive("1").ShouldBeFalse(); tracker.GetViewerCount("1").ShouldBe(0); } [Test] public async Task Should_Not_Track_Session_For_Head_Request() { IDirectStreamSessionTracker tracker = Substitute.For(); var stream = new BlockingReadStream(); TrackedFileStreamResult result = new( stream, "video/mp2t", tracker, "1", StreamingMode.TransportStream); await result.ExecuteResultAsync(GetActionContext(HttpMethods.Head)); tracker.DidNotReceive().Register(Arg.Any(), Arg.Any()); } private static ActionContext GetActionContext(string method = "GET") { var httpContext = new DefaultHttpContext(); httpContext.Request.Method = method; httpContext.RequestServices = new ServiceCollection() .AddLogging() .AddControllers() .Services .BuildServiceProvider(); httpContext.Response.Body = new MemoryStream(); return new ActionContext(httpContext, new RouteData(), new Microsoft.AspNetCore.Mvc.Abstractions.ActionDescriptor()); } private sealed class BlockingReadStream : Stream { private readonly TaskCompletionSource _continue = new(TaskCreationOptions.RunContinuationsAsynchronously); private readonly TaskCompletionSource _readStarted = new(TaskCreationOptions.RunContinuationsAsynchronously); public override bool CanRead => true; public override bool CanSeek => false; public override bool CanWrite => false; public override long Length => 0; public override long Position { get; set; } public Task WaitForRead() => _readStarted.Task.WaitAsync(TimeSpan.FromSeconds(5)); public void Complete() => _continue.SetResult(); public override ValueTask ReadAsync(Memory buffer, CancellationToken cancellationToken = default) { _readStarted.TrySetResult(); return ReadAfterContinue(); } public override Task ReadAsync(byte[] buffer, int offset, int count, CancellationToken cancellationToken) => ReadAsync(buffer.AsMemory(offset, count), cancellationToken).AsTask(); public override int Read(byte[] buffer, int offset, int count) => throw new NotSupportedException(); public override void Flush() { } public override long Seek(long offset, SeekOrigin origin) => throw new NotSupportedException(); public override void SetLength(long value) => throw new NotSupportedException(); public override void Write(byte[] buffer, int offset, int count) => throw new NotSupportedException(); private async ValueTask ReadAfterContinue() { await _continue.Task; return 0; } } private sealed class ThrowingReadStream(Exception exception) : Stream { public override bool CanRead => true; public override bool CanSeek => false; public override bool CanWrite => false; public override long Length => 0; public override long Position { get; set; } public override ValueTask ReadAsync(Memory buffer, CancellationToken cancellationToken = default) => ValueTask.FromException(exception); public override Task ReadAsync(byte[] buffer, int offset, int count, CancellationToken cancellationToken) => ReadAsync(buffer.AsMemory(offset, count), cancellationToken).AsTask(); public override int Read(byte[] buffer, int offset, int count) => throw exception; public override void Flush() { } public override long Seek(long offset, SeekOrigin origin) => throw new NotSupportedException(); public override void SetLength(long value) => throw new NotSupportedException(); public override void Write(byte[] buffer, int offset, int count) => throw new NotSupportedException(); } }