diff --git a/src/Microsoft.Teams.Apps/SocketMode/GeoSocket.cs b/src/Microsoft.Teams.Apps/SocketMode/GeoSocket.cs new file mode 100644 index 00000000..45b7725c --- /dev/null +++ b/src/Microsoft.Teams.Apps/SocketMode/GeoSocket.cs @@ -0,0 +1,698 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System.Diagnostics.CodeAnalysis; +using Microsoft.Extensions.Logging; + +namespace Microsoft.Teams.Apps.SocketMode; + +/// +/// Provides the retry policy, dispatch, and lifecycle reporting a needs from its transport. +/// +internal interface IGeoSocketOwner +{ + /// + /// Gets the time allowed for the initial connection to become ready. Zero permits a single unbounded attempt. + /// + TimeSpan StartupTimeout { get; } + + /// + /// Gets how long before token expiry a replacement connection is established. + /// + TimeSpan TokenRefreshMargin { get; } + + /// + /// Gets how long a replaced connection keeps dispatching after its successor is ready. + /// + TimeSpan HandoffWindow { get; } + + /// + /// Gets the server-requested retry delay for a failure, when one was provided. + /// + /// The failure that caused the retry. + TimeSpan? GetRetryAfter(Exception? error); + + /// + /// Gets the backoff delay for a zero-based retry attempt. + /// + /// The zero-based retry attempt. + TimeSpan GetBackoffDelay(int attempt); + + /// + /// Dispatches an activity received on a live connection for a geo. + /// + /// The geo that received the activity. + /// The received activity envelope. + Task DispatchAsync(string geo, SocketActivityEnvelope envelope); + + /// + /// Reports that a connection generation for a geo received SocketReady. + /// + /// The geo whose connection became ready. + /// The ready frame. + void OnGeoReady(string geo, SocketReadyFrame frame); + + /// + /// Reports that inbound delivery for a geo was unexpectedly interrupted. + /// + /// The geo that disconnected. + /// The failure that closed the connection, when available. + void OnGeoDisconnected(string geo, Exception? error); + + /// + /// Reports that inbound delivery for a geo resumed after an unexpected interruption. + /// + /// The geo that reconnected. + void OnGeoReconnected(string geo); +} + +/// +/// Keeps one geo connected: establishes the initial connection within a startup budget, reconnects after +/// unexpected closure, and rotates tokens make-before-break. +/// +/// +/// Each connection attempt is a new generation. Activities are dispatched only from the active ready generation +/// or a predecessor that is explicitly retiring during a token rotation handoff. +/// +internal sealed class GeoSocket : IAsyncDisposable +{ + private static readonly TimeSpan MinimumRefreshDelay = TimeSpan.FromSeconds(1); + + private readonly IGeoSocketOwner _owner; + private readonly Uri _negotiateUri; + private readonly ISocketConnectionFactory _connectionFactory; + private readonly ILogger _logger; + private readonly TimeProvider _timeProvider; + private readonly CancellationTokenSource _stopSource = new(); + private readonly CancellationToken _stopToken; + private readonly object _sync = new(); + private readonly HashSet _owned = []; + private readonly HashSet _retiring = []; + private readonly HashSet _retirements = []; + private readonly HashSet _releases = []; + + private long _generation; + private Generation? _active; + private ITimer? _refreshTimer; + private Task? _supervisor; + private Task? _stopTask; + private bool _disconnected; + private bool _stopping; + private int _started; + private int _disposed; + + /// + /// Initializes a supervisor for one geo. + /// + /// The transport that owns this geo. + /// The geo identifier. Empty when the negotiate URL has no geo segment. + /// The negotiate endpoint for the geo. + /// Creates one connection per generation. + /// The logger for lifecycle events. + /// The time provider for startup budgets, retries, rotation, and handoff. + internal GeoSocket( + IGeoSocketOwner owner, + string geo, + Uri negotiateUri, + ISocketConnectionFactory connectionFactory, + ILogger logger, + TimeProvider? timeProvider = null) + { + _owner = owner ?? throw new ArgumentNullException(nameof(owner)); + Geo = geo ?? throw new ArgumentNullException(nameof(geo)); + _negotiateUri = negotiateUri ?? throw new ArgumentNullException(nameof(negotiateUri)); + _connectionFactory = connectionFactory ?? throw new ArgumentNullException(nameof(connectionFactory)); + _logger = logger ?? throw new ArgumentNullException(nameof(logger)); + _timeProvider = timeProvider ?? TimeProvider.System; + _stopToken = _stopSource.Token; + } + + /// + /// Gets the geo identifier. + /// + internal string Geo { get; } + + /// + /// Establishes the initial ready connection, then supervises it in the background. + /// + /// A token for cancelling startup. + /// Thrown when no generation becomes ready within the startup budget. + /// Thrown without retrying when negotiate rejects the bot with HTTP 401 or 403. + internal async Task StartAsync(CancellationToken cancellationToken = default) + { + ObjectDisposedException.ThrowIf(Volatile.Read(ref _disposed) != 0, this); + if (Interlocked.Exchange(ref _started, 1) != 0) + { + throw new InvalidOperationException("A GeoSocket can only be started once."); + } + + Generation initial = await ConnectInitialAsync(cancellationToken).ConfigureAwait(false); + + lock (_sync) + { + if (!_stopping) + { + _supervisor = SuperviseAsync(initial); + return; + } + } + + throw new OperationCanceledException("Socket Mode geo stopped before startup completed."); + } + + /// + /// Stops supervision and stops and disposes every connection this geo owns. + /// Connection cleanup failures are logged rather than thrown. + /// + internal Task StopAsync() + { + lock (_sync) + { + _stopping = true; + return _stopTask ??= StopCoreAsync(); + } + } + + /// + public async ValueTask DisposeAsync() + { + if (Interlocked.Exchange(ref _disposed, 1) != 0) + { + return; + } + + try + { + await StopAsync().ConfigureAwait(false); + } + finally + { + _stopSource.Dispose(); + } + } + + private async Task ConnectInitialAsync(CancellationToken cancellationToken) + { + using CancellationTokenSource startSource = + CancellationTokenSource.CreateLinkedTokenSource(cancellationToken, _stopToken); + TimeSpan budget = _owner.StartupTimeout; + long startedAt = _timeProvider.GetTimestamp(); + + for (int attempt = 0; ; attempt++) + { + using CancellationTokenSource deadlineSource = new(Timeout.InfiniteTimeSpan, _timeProvider); + if (budget > TimeSpan.Zero) + { + deadlineSource.CancelAfter(RemainingBudget(budget, startedAt)); + } + + using CancellationTokenSource attemptSource = + CancellationTokenSource.CreateLinkedTokenSource(startSource.Token, deadlineSource.Token); + + try + { + return await ConnectAsync(attemptSource.Token).ConfigureAwait(false); + } + catch (Exception exception) when (!startSource.IsCancellationRequested) + { + if (SocketModeNegotiateException.IsNonRetryable(exception)) + { + throw; + } + + Exception error = deadlineSource.IsCancellationRequested + ? new TimeoutException( + $"Socket Mode geo '{Geo}' did not become ready within {budget}.", + exception) + : exception; + TimeSpan delay = GetRetryDelay(error, attempt); + + if (budget <= TimeSpan.Zero || delay >= RemainingBudget(budget, startedAt)) + { + if (ReferenceEquals(error, exception)) + { + throw; + } + + throw error; + } + + _logger.LogWarning( + exception, + "Socket Mode geo {Geo} initial connection failed; retrying in {Delay}.", + Geo, + delay); + await Task.Delay(delay, _timeProvider, startSource.Token).ConfigureAwait(false); + } + } + } + + private async Task SuperviseAsync(Generation current) + { + try + { + while (true) + { + CloseReason reason = await current.Closed.Task.WaitAsync(_stopToken).ConfigureAwait(false); + + Generation? replacement; + if (reason.Planned) + { + _logger.LogInformation("Socket Mode geo {Geo} rotating connection before token expiry.", Geo); + replacement = await ReconnectAsync(null).ConfigureAwait(false); + if (replacement is not null) + { + StartRetirement(current); + } + } + else + { + await ReleaseAsync(current.Connection).ConfigureAwait(false); + replacement = await ReconnectAsync(reason.Error).ConfigureAwait(false); + } + + if (replacement is null) + { + return; + } + + current = replacement; + + ReportReconnected(); + } + } + catch (OperationCanceledException) when (_stopToken.IsCancellationRequested) + { + } + catch (Exception exception) + { + _logger.LogError(exception, "Socket Mode geo {Geo} supervisor stopped unexpectedly.", Geo); + throw; + } + } + + /// The ready replacement, or when the geo gave up after a non-retryable failure. + private async Task ReconnectAsync(Exception? previousError) + { + Exception? error = previousError; + + for (int attempt = 1; ; attempt++) + { + await Task.Delay(GetRetryDelay(error, attempt - 1), _timeProvider, _stopToken).ConfigureAwait(false); + + try + { + return await ConnectAsync(_stopToken).ConfigureAwait(false); + } + catch (Exception exception) when (!_stopToken.IsCancellationRequested) + { + if (SocketModeNegotiateException.IsNonRetryable(exception)) + { + StopAfterNonRetryable(exception); + return null; + } + + error = exception; + _logger.LogWarning( + exception, + "Socket Mode geo {Geo} reconnect attempt {Attempt} failed.", + Geo, + attempt); + } + } + } + + /// + /// Gives up on this geo: stops every connection it owns, including one still serving during a token rotation, + /// and reports it disconnected. Other geos are unaffected. + /// + /// The non-retryable failure. + private void StopAfterNonRetryable(Exception error) + { + _logger.LogError( + error, + "Socket Mode geo {Geo} reconnect was rejected; inbound delivery for this geo has stopped until the app is restarted.", + Geo); + + // Not awaited: stopping waits for this supervisor to finish. Marking the geo as stopping happens + // synchronously, so closing connections cannot report a second disconnect. + _ = StopAsync(); + _owner.OnGeoDisconnected(Geo, error); + } + + private async Task ConnectAsync(CancellationToken cancellationToken) + { + Generation generation; + lock (_sync) + { + ThrowIfStopping(); + generation = new Generation(++_generation); + } + + generation.Connection = _connectionFactory.Create( + _negotiateUri, + new SocketConnectionHandlers( + envelope => DispatchAsync(generation, envelope), + frame => HandleReady(generation, frame), + (error, planned) => HandleClosed(generation, error, planned))); + + bool owned; + lock (_sync) + { + owned = !_stopping && _owned.Add(generation.Connection); + } + + if (!owned) + { + await StopAndDisposeAsync(generation.Connection).ConfigureAwait(false); + throw new OperationCanceledException("Socket Mode geo is stopping."); + } + + try + { + await generation.Connection.StartAsync(cancellationToken).ConfigureAwait(false); + + // StartAsync completes after SocketReady, but the ready callback may still be running. + if (!TryPromote(generation)) + { + lock (_sync) + { + ThrowIfStopping(); + } + + throw new IOException("Socket Mode connection closed before it became active."); + } + } + catch + { + await ReleaseAsync(generation.Connection).ConfigureAwait(false); + throw; + } + + ScheduleRefresh(generation); + return generation; + } + + private Task DispatchAsync(Generation generation, SocketActivityEnvelope envelope) + { + lock (_sync) + { + if (_stopping || (_active != generation && !_retiring.Contains(generation.Id))) + { + return Task.FromResult(null); + } + } + + return _owner.DispatchAsync(Geo, envelope); + } + + private void HandleReady(Generation generation, SocketReadyFrame frame) + { + if (TryPromote(generation)) + { + _owner.OnGeoReady(Geo, frame); + } + } + + private bool TryPromote(Generation generation) + { + lock (_sync) + { + if (_stopping || generation.IsClosed || generation.Id != _generation) + { + return false; + } + + if (_active is Generation previous && previous != generation) + { + _retiring.Add(previous.Id); + } + + _active = generation; + return true; + } + } + + private void HandleClosed(Generation generation, Exception? error, bool planned) + { + bool disconnected = false; + + lock (_sync) + { + generation.IsClosed = true; + _retiring.Remove(generation.Id); + + if (_active == generation) + { + _active = null; + _refreshTimer?.Dispose(); + _refreshTimer = null; + disconnected = !planned && !_stopping; + _disconnected |= disconnected; + } + } + + // Report the disconnect before waking the supervisor so a fast reconnect cannot be reported first. + try + { + if (disconnected) + { + _logger.LogWarning(error, "Socket Mode geo {Geo} disconnected; inbound delivery paused.", Geo); + _owner.OnGeoDisconnected(Geo, error); + } + } + finally + { + generation.Closed.TrySetResult(new CloseReason(error, planned)); + } + } + + private void ReportReconnected() + { + lock (_sync) + { + if (!_disconnected || _stopping) + { + return; + } + + _disconnected = false; + } + + _logger.LogInformation("Socket Mode geo {Geo} reconnected; inbound delivery resumed.", Geo); + _owner.OnGeoReconnected(Geo); + } + + private void ScheduleRefresh(Generation generation) + { + if (generation.Connection.TokenLifetime is not TimeSpan lifetime || lifetime <= TimeSpan.Zero) + { + return; + } + + TimeSpan delay = lifetime - _owner.TokenRefreshMargin; + if (delay < MinimumRefreshDelay) + { + delay = MinimumRefreshDelay; + } + + lock (_sync) + { + if (_stopping || _active != generation) + { + return; + } + + _refreshTimer?.Dispose(); + _refreshTimer = _timeProvider.CreateTimer( + _ => RequestRotation(generation), + null, + delay, + Timeout.InfiniteTimeSpan); + } + } + + private void RequestRotation(Generation generation) + { + lock (_sync) + { + if (!_stopping && _active == generation) + { + generation.Closed.TrySetResult(new CloseReason(null, Planned: true)); + } + } + } + + private void StartRetirement(Generation previous) + { + Task retirement = RetireAsync(previous); + lock (_sync) + { + if (!retirement.IsCompleted) + { + _retirements.Add(retirement); + } + } + + _ = retirement.ContinueWith( + completed => + { + lock (_sync) + { + _retirements.Remove(completed); + } + }, + CancellationToken.None, + TaskContinuationOptions.ExecuteSynchronously, + TaskScheduler.Default); + } + + private async Task RetireAsync(Generation previous) + { + try + { + TimeSpan handoff = _owner.HandoffWindow; + await Task.Delay( + handoff > TimeSpan.Zero ? handoff : TimeSpan.Zero, + _timeProvider, + _stopToken).ConfigureAwait(false); + + lock (_sync) + { + _retiring.Remove(previous.Id); + } + + await ReleaseAsync(previous.Connection).ConfigureAwait(false); + } + catch (OperationCanceledException) when (_stopToken.IsCancellationRequested) + { + } + } + + private async Task StopCoreAsync() + { + // Leave the caller's lock before cancelling and stopping connections. + await Task.Yield(); + await _stopSource.CancelAsync().ConfigureAwait(false); + + ISocketConnection[] connections; + Task[] retirements; + Task[] releases; + Task? supervisor; + lock (_sync) + { + _active = null; + _retiring.Clear(); + _refreshTimer?.Dispose(); + _refreshTimer = null; + connections = [.. _owned]; + _owned.Clear(); + retirements = [.. _retirements]; + releases = [.. _releases]; + supervisor = _supervisor; + } + + try + { + await Task.WhenAll(connections.Select(StopAndDisposeAsync)).ConfigureAwait(false); + } + finally + { + await Task.WhenAll(retirements).ConfigureAwait(false); + await Task.WhenAll(releases).ConfigureAwait(false); + if (supervisor is not null) + { + await supervisor.ConfigureAwait(ConfigureAwaitOptions.SuppressThrowing); + } + } + } + + private async Task ReleaseAsync(ISocketConnection connection) + { + TaskCompletionSource released = new(TaskCreationOptions.RunContinuationsAsynchronously); + lock (_sync) + { + if (!_owned.Remove(connection)) + { + return; + } + + // Registered with the ownership claim so StopAsync can await cleanup it no longer owns. + _releases.Add(released.Task); + } + + try + { + await StopAndDisposeAsync(connection).ConfigureAwait(false); + } + finally + { + lock (_sync) + { + _releases.Remove(released.Task); + } + + released.TrySetResult(); + } + } + + [SuppressMessage( + "Design", + "CA1031:Do not catch general exception types", + Justification = "Connection cleanup is best-effort; failures are logged so they cannot mask the caller's outcome.")] + private async Task StopAndDisposeAsync(ISocketConnection connection) + { + try + { + await connection.StopAsync(CancellationToken.None).ConfigureAwait(false); + } + catch (Exception exception) + { + _logger.LogWarning(exception, "Socket Mode geo {Geo} failed to stop a connection.", Geo); + } + + try + { + await connection.DisposeAsync().ConfigureAwait(false); + } + catch (Exception exception) + { + _logger.LogWarning(exception, "Socket Mode geo {Geo} failed to dispose a connection.", Geo); + } + } + + private TimeSpan GetRetryDelay(Exception? error, int attempt) + { + TimeSpan delay = _owner.GetRetryAfter(error) ?? _owner.GetBackoffDelay(attempt); + return delay > TimeSpan.Zero ? delay : TimeSpan.Zero; + } + + private TimeSpan RemainingBudget(TimeSpan budget, long startedAt) + { + TimeSpan remaining = budget - _timeProvider.GetElapsedTime(startedAt); + return remaining > TimeSpan.Zero ? remaining : TimeSpan.Zero; + } + + private void ThrowIfStopping() + { + if (_stopping) + { + throw new OperationCanceledException("Socket Mode geo is stopping."); + } + } + + private sealed class Generation(long id) + { + internal long Id { get; } = id; + + internal ISocketConnection Connection { get; set; } = null!; + + internal TaskCompletionSource Closed { get; } = + new(TaskCreationOptions.RunContinuationsAsynchronously); + + internal bool IsClosed { get; set; } + } + + private sealed record CloseReason(Exception? Error, bool Planned); +} + diff --git a/src/Microsoft.Teams.Apps/SocketMode/SocketModeNegotiator.cs b/src/Microsoft.Teams.Apps/SocketMode/SocketModeNegotiator.cs index 20a19203..cb15f6ee 100644 --- a/src/Microsoft.Teams.Apps/SocketMode/SocketModeNegotiator.cs +++ b/src/Microsoft.Teams.Apps/SocketMode/SocketModeNegotiator.cs @@ -240,7 +240,7 @@ internal sealed class SocketModeNegotiateException : Exception internal SocketModeNegotiateException( HttpStatusCode statusCode, TimeSpan? retryAfter) - : base($"Socket Mode negotiate failed with HTTP {(int)statusCode}.") + : base(CreateMessage(statusCode)) { StatusCode = statusCode; RetryAfter = retryAfter; @@ -255,4 +255,32 @@ internal SocketModeNegotiateException( /// Gets the server-provided retry delay, when available. /// internal TimeSpan? RetryAfter { get; } + + /// + /// Gets whether the service rejected the bot (HTTP 401 or 403), which retrying cannot fix. + /// + internal bool IsAuthError => StatusCode is HttpStatusCode.Unauthorized or HttpStatusCode.Forbidden; + + /// + /// Determines whether a connection failure must not be retried. + /// + /// The connection failure. + /// when negotiate rejected the bot with HTTP 401 or 403. + internal static bool IsNonRetryable(Exception? error) + => error is SocketModeNegotiateException { IsAuthError: true }; + + private static string CreateMessage(HttpStatusCode statusCode) + { + string failure = $"Socket Mode negotiate failed with HTTP {(int)statusCode}."; + return statusCode switch + { + HttpStatusCode.Unauthorized => + $"{failure} The bot could not be authenticated: verify the bot credentials (client ID and secret, " + + "certificate, or managed identity) and restart the app after correcting them.", + HttpStatusCode.Forbidden => + $"{failure} This bot is not authorized to use Socket Mode: verify the bot registration and Socket Mode " + + "access for this environment, then restart the app.", + _ => failure, + }; + } } diff --git a/src/Microsoft.Teams.Apps/SocketMode/SocketModeTransport.cs b/src/Microsoft.Teams.Apps/SocketMode/SocketModeTransport.cs new file mode 100644 index 00000000..518abf6b --- /dev/null +++ b/src/Microsoft.Teams.Apps/SocketMode/SocketModeTransport.cs @@ -0,0 +1,479 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System.Diagnostics.CodeAnalysis; +using System.Runtime.ExceptionServices; +using Microsoft.Extensions.Logging; +using Microsoft.Teams.Apps.Schema; +using Microsoft.Teams.Core.Schema; + +namespace Microsoft.Teams.Apps.SocketMode; + +/// +/// Describes the lifecycle state of a Socket Mode transport or one of its geos. +/// +internal enum SocketModeStatus +{ + /// + /// The transport has not started. + /// + Idle, + + /// + /// A connection is being established. + /// + Connecting, + + /// + /// Every connection is ready to receive activities. + /// + Ready, + + /// + /// A connection dropped unexpectedly and is recovering. + /// + Disconnected, + + /// + /// The transport has stopped. + /// + Stopped, +} + +/// +/// Configures a . +/// +internal sealed class SocketModeTransportOptions +{ + /// + /// Gets the base URL used to negotiate each geo connection. + /// + internal Uri NegotiateBaseUri { get; init; } = new(SocketModeProtocol.DefaultNegotiateBaseUrl); + + /// + /// Gets the geos to connect. An empty string connects to the base URL without a geo segment. + /// + internal IReadOnlyList Geos { get; init; } = SocketModeProtocol.DefaultGeos; + + /// + /// Gets the time each geo has to establish its initial connection. + /// + internal TimeSpan StartupTimeout { get; init; } = TimeSpan.FromSeconds(30); + + /// + /// Gets an explicit reconnect delay schedule. The last delay repeats. When omitted, capped exponential + /// backoff with full jitter is used. + /// + internal IReadOnlyList? ReconnectDelays { get; init; } +} + +/// +/// Connects one per configured geo and routes their activities to the app. +/// +/// +/// Startup is all-or-nothing: every geo must become ready or the transport stops and startup fails. Once +/// started, each geo is supervised independently. +/// +internal sealed class SocketModeTransport : IGeoSocketOwner, IAsyncDisposable +{ + private static readonly TimeSpan TokenRefreshMargin = TimeSpan.FromSeconds(60); + private static readonly TimeSpan HandoffWindow = TimeSpan.FromSeconds(5); + private static readonly TimeSpan ReconnectInitialDelay = TimeSpan.FromSeconds(1); + private static readonly TimeSpan ReconnectMaxDelay = TimeSpan.FromSeconds(15); + + private readonly SocketModeTransportOptions _options; + private readonly IReadOnlyList<(string Geo, Uri NegotiateUri)> _geos; + private readonly ISocketConnectionFactory _connectionFactory; + private readonly Func> _dispatch; + private readonly Func? _onError; + private readonly string? _botKey; + private readonly ILogger _logger; + private readonly TimeProvider _timeProvider; + private readonly Random _random; + private readonly object _sync = new(); + private readonly Dictionary _geoStatuses = []; + + private GeoSocket[] _geoSockets = []; + private SocketModeStatus _lifecycle = SocketModeStatus.Idle; + private Task? _stopTask; + + /// + /// Initializes a Socket Mode transport and resolves its geos. + /// + /// The transport options. + /// Creates one connection per geo generation. + /// Processes an inbound activity and returns its status and optional body. + /// The logger for transport lifecycle and dispatch failures. + /// The bot key echoed on reply frames, when available. + /// Observes activity processing failures before the failure reply is returned. + /// The time provider for supervision timing and reply timestamps. + /// The random source for reconnect jitter. + /// Thrown when the geos are empty, null, or duplicated. + /// Thrown when the startup timeout is negative. + internal SocketModeTransport( + SocketModeTransportOptions options, + ISocketConnectionFactory connectionFactory, + Func> dispatch, + ILogger logger, + string? botKey = null, + Func? onError = null, + TimeProvider? timeProvider = null, + Random? random = null) + { + _options = options ?? throw new ArgumentNullException(nameof(options)); + _connectionFactory = connectionFactory ?? throw new ArgumentNullException(nameof(connectionFactory)); + _dispatch = dispatch ?? throw new ArgumentNullException(nameof(dispatch)); + _logger = logger ?? throw new ArgumentNullException(nameof(logger)); + _botKey = botKey; + _onError = onError; + _timeProvider = timeProvider ?? TimeProvider.System; + _random = random ?? Random.Shared; + + ArgumentNullException.ThrowIfNull(options.NegotiateBaseUri, nameof(options)); + ArgumentOutOfRangeException.ThrowIfLessThan(options.StartupTimeout, TimeSpan.Zero, nameof(options)); + _geos = ResolveGeos(options.NegotiateBaseUri, options.Geos); + } + + /// + /// Gets the aggregate status: ready when every geo is ready, disconnected when any geo is recovering from a + /// drop, and connecting otherwise. + /// + internal SocketModeStatus Status + { + get + { + lock (_sync) + { + if (_lifecycle is SocketModeStatus.Idle or SocketModeStatus.Stopped) + { + return _lifecycle; + } + + if (_geoStatuses.Count > 0 && _geoStatuses.Values.All(status => status == SocketModeStatus.Ready)) + { + return SocketModeStatus.Ready; + } + + return _geoStatuses.ContainsValue(SocketModeStatus.Disconnected) + ? SocketModeStatus.Disconnected + : SocketModeStatus.Connecting; + } + } + } + + /// + /// Gets a snapshot of each geo's status. + /// + internal IReadOnlyDictionary GeoStatuses + { + get + { + lock (_sync) + { + return new Dictionary(_geoStatuses); + } + } + } + + /// + TimeSpan IGeoSocketOwner.StartupTimeout => _options.StartupTimeout; + + /// + TimeSpan IGeoSocketOwner.TokenRefreshMargin => TokenRefreshMargin; + + /// + TimeSpan IGeoSocketOwner.HandoffWindow => HandoffWindow; + + /// + /// Connects every geo. Completes once all are ready; otherwise stops the transport and throws the first + /// geo failure. + /// + /// A token for cancelling startup. + internal async Task StartAsync(CancellationToken cancellationToken = default) + { + GeoSocket[] geoSockets; + lock (_sync) + { + if (_lifecycle != SocketModeStatus.Idle) + { + throw new InvalidOperationException("A Socket Mode transport can only be started once."); + } + + _lifecycle = SocketModeStatus.Connecting; + geoSockets = [.. _geos.Select(geo => new GeoSocket( + this, + geo.Geo, + geo.NegotiateUri, + _connectionFactory, + _logger, + _timeProvider))]; + foreach (GeoSocket geoSocket in geoSockets) + { + _geoStatuses[geoSocket.Geo] = SocketModeStatus.Connecting; + } + + _geoSockets = geoSockets; + } + + _logger.LogInformation( + "Socket Mode connecting {Count} geo(s): {Geos}.", + geoSockets.Length, + string.Join(", ", geoSockets.Select(geoSocket => geoSocket.Geo))); + + using CancellationTokenSource startSource = CancellationTokenSource.CreateLinkedTokenSource(cancellationToken); + Exception? firstFailure = null; + + async Task StartGeoAsync(GeoSocket geoSocket) + { + try + { + await geoSocket.StartAsync(startSource.Token).ConfigureAwait(false); + } + catch (Exception exception) + { + if (Interlocked.CompareExchange(ref firstFailure, exception, null) is null) + { + await startSource.CancelAsync().ConfigureAwait(false); + } + + throw; + } + } + + try + { + await Task.WhenAll(geoSockets.Select(StartGeoAsync)).ConfigureAwait(false); + } + catch (Exception exception) + { + bool stoppedExternally; + lock (_sync) + { + stoppedExternally = _stopTask is not null; + } + + await StopAsync().ConfigureAwait(false); + + if (stoppedExternally) + { + throw new OperationCanceledException("Socket Mode transport stopped before startup completed.", exception); + } + + ExceptionDispatchInfo.Throw(firstFailure ?? exception); + } + + lock (_sync) + { + if (_stopTask is not null) + { + throw new OperationCanceledException("Socket Mode transport stopped before startup completed."); + } + } + + _logger.LogInformation("Socket Mode ready across {Count} geo(s).", geoSockets.Length); + } + + /// + /// Stops every geo and disposes their connections. Idempotent; connection cleanup failures are logged rather than thrown. + /// + internal Task StopAsync() + { + lock (_sync) + { + if (_stopTask is not null) + { + return _stopTask; + } + + _lifecycle = SocketModeStatus.Stopped; + foreach (string geo in _geoStatuses.Keys.ToArray()) + { + _geoStatuses[geo] = SocketModeStatus.Stopped; + } + _stopTask = Task.WhenAll(_geoSockets.Select(geoSocket => geoSocket.DisposeAsync().AsTask())); + return _stopTask; + } + } + + /// + public ValueTask DisposeAsync() => new(StopAsync()); + + /// + TimeSpan? IGeoSocketOwner.GetRetryAfter(Exception? error) + => (error as SocketModeNegotiateException)?.RetryAfter; + + /// + [SuppressMessage( + "Security", + "CA5394:Do not use insecure randomness", + Justification = "Jitter only spreads reconnect attempts; it has no security purpose.")] + TimeSpan IGeoSocketOwner.GetBackoffDelay(int attempt) + { + if (_options.ReconnectDelays is { Count: > 0 } schedule) + { + return schedule[Math.Min(attempt, schedule.Count - 1)]; + } + + double capSeconds = Math.Min( + ReconnectInitialDelay.TotalSeconds * Math.Pow(2, Math.Min(attempt, 30)), + ReconnectMaxDelay.TotalSeconds); + return TimeSpan.FromSeconds(capSeconds * _random.NextDouble()); + } + + /// + Task IGeoSocketOwner.DispatchAsync(string geo, SocketActivityEnvelope envelope) + => HandleEnvelopeAsync(envelope); + + /// + void IGeoSocketOwner.OnGeoReady(string geo, SocketReadyFrame frame) + { + SetGeoStatus(geo, SocketModeStatus.Ready); + _logger.LogDebug("Socket Mode geo {Geo} ready with connection {ConnectionId}.", geo, frame.ConnectionId); + } + + /// + void IGeoSocketOwner.OnGeoDisconnected(string geo, Exception? error) + => SetGeoStatus(geo, SocketModeStatus.Disconnected); + + /// + void IGeoSocketOwner.OnGeoReconnected(string geo) + => SetGeoStatus(geo, SocketModeStatus.Ready); + + [SuppressMessage( + "Design", + "CA1031:Do not catch general exception types", + Justification = "Any activity processing failure must be logged and answered with a 500 reply frame.")] + private async Task HandleEnvelopeAsync(SocketActivityEnvelope envelope) + { + long receivedAt = NowUnixMilliseconds(); + + if (envelope.ProtocolVersion > SocketModeProtocol.CurrentVersion) + { + _logger.LogWarning( + "Socket Mode rejecting envelope {EnvelopeId} with unsupported protocol version {Version}.", + envelope.EnvelopeId, + envelope.ProtocolVersion); + return SocketModeEnvelope.CreateInvokeReply( + envelope, + _botKey, + new SocketDispatchResult(400, new { error = $"unsupported protocolVersion {envelope.ProtocolVersion}" }), + receivedAt, + NowUnixMilliseconds()); + } + + if (!SocketModeEnvelope.TryReadActivity(envelope, out CoreActivity? activity) || activity is null) + { + _logger.LogWarning("Socket Mode envelope {EnvelopeId} had no activity payload; dropping.", envelope.EnvelopeId); + return null; + } + + bool invoke = IsInvoke(envelope, activity); + + try + { + SocketDispatchResult result = await _dispatch(activity).ConfigureAwait(false); + return invoke + ? SocketModeEnvelope.CreateInvokeReply(envelope, _botKey, result, receivedAt, NowUnixMilliseconds()) + : SocketModeEnvelope.CreateAcknowledgement( + envelope, + _botKey, + receivedAt, + NowUnixMilliseconds(), + result.Status); + } + catch (Exception exception) + { + _logger.LogError( + exception, + "Socket Mode failed to process activity {ActivityType} in envelope {EnvelopeId}.", + activity.Type, + envelope.EnvelopeId); + await ReportErrorAsync(exception).ConfigureAwait(false); + + return invoke + ? SocketModeEnvelope.CreateInvokeReply( + envelope, + _botKey, + new SocketDispatchResult(500, new { error = "bot handler error" }), + receivedAt, + NowUnixMilliseconds()) + : SocketModeEnvelope.CreateAcknowledgement(envelope, _botKey, receivedAt, NowUnixMilliseconds(), 500); + } + } + + [SuppressMessage( + "Design", + "CA1031:Do not catch general exception types", + Justification = "A failing error observer is logged so it cannot prevent the failure reply.")] + private async Task ReportErrorAsync(Exception exception) + { + if (_onError is null) + { + return; + } + + try + { + await _onError(exception).ConfigureAwait(false); + } + catch (Exception hookException) + { + _logger.LogWarning(hookException, "Socket Mode error observer failed."); + } + } + + private void SetGeoStatus(string geo, SocketModeStatus status) + { + lock (_sync) + { + if (_lifecycle != SocketModeStatus.Stopped) + { + _geoStatuses[geo] = status; + } + } + } + + private long NowUnixMilliseconds() => _timeProvider.GetUtcNow().ToUnixTimeMilliseconds(); + + // An envelope is an invoke when its type says so, falling back to the activity type when the envelope has none. + private static bool IsInvoke(SocketActivityEnvelope envelope, CoreActivity activity) + => string.Equals( + string.IsNullOrEmpty(envelope.Type) ? activity.Type : envelope.Type, + TeamsActivityTypes.Invoke, + StringComparison.OrdinalIgnoreCase); + + private static List<(string Geo, Uri NegotiateUri)> ResolveGeos(Uri baseUri, IReadOnlyList geos) + { + if (geos is null || geos.Count == 0) + { + throw new ArgumentException( + "Socket Mode requires at least one geo. Use an empty string to connect without a geo segment.", + nameof(geos)); + } + + string baseUrl = baseUri.AbsoluteUri.TrimEnd('/'); + HashSet seen = new(StringComparer.OrdinalIgnoreCase); + List<(string Geo, Uri NegotiateUri)> resolved = []; + + foreach (string geo in geos) + { + if (geo is null) + { + throw new ArgumentException("Socket Mode geos cannot contain null.", nameof(geos)); + } + + string segment = geo.Trim().Trim('/'); + if (!seen.Add(segment)) + { + throw new ArgumentException($"Socket Mode geo '{geo}' is listed more than once.", nameof(geos)); + } + + string path = segment.Length > 0 + ? $"/{Uri.EscapeDataString(segment)}{SocketModeProtocol.NegotiatePath}" + : SocketModeProtocol.NegotiatePath; + resolved.Add((segment, new Uri(baseUrl + path, UriKind.Absolute))); + } + + return resolved; + } +} diff --git a/test/Microsoft.Teams.Apps.UnitTests/SocketMode/GeoSocketTests.cs b/test/Microsoft.Teams.Apps.UnitTests/SocketMode/GeoSocketTests.cs new file mode 100644 index 00000000..b092dfe4 --- /dev/null +++ b/test/Microsoft.Teams.Apps.UnitTests/SocketMode/GeoSocketTests.cs @@ -0,0 +1,943 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System.Diagnostics; +using System.Net; +using Microsoft.Extensions.Logging; +using Microsoft.Teams.Apps.SocketMode; + +namespace Microsoft.Teams.Apps.UnitTests.SocketMode; + +public class GeoSocketTests +{ + private static readonly Uri NegotiateUri = new("https://botapi.skype.com/amer/v3/websockets/connect"); + + [Fact] + public async Task StartAsync_CompletesOnlyAfterSocketReady() + { + Harness harness = new(); + FakeConnection connection = harness.Factory.Enqueue(); + + Task start = harness.Socket.StartAsync(); + await connection.Started.Task; + Assert.False(start.IsCompleted); + + connection.Ready("initial"); + await start; + + Assert.NotNull(await connection.Activity("a1")); + Assert.Equal(["initial"], harness.Owner.ReadyIds); + Assert.Equal(["a1"], harness.Owner.Dispatched); + Assert.Empty(harness.Owner.Disconnections); + } + + [Fact] + public async Task StartAsync_RetriesWithBackoffWithinBudget() + { + Harness harness = new() { Backoff = _ => TimeSpan.FromSeconds(2) }; + FakeConnection failed = harness.Factory.Enqueue(startError: new IOException("first")); + FakeConnection retry = harness.Factory.Enqueue(); + + Task start = harness.Socket.StartAsync(); + await harness.Time.WaitForTimerAsync(TimeSpan.FromSeconds(2)); + harness.Time.Advance(TimeSpan.FromSeconds(2)); + await retry.Started.Task; + retry.Ready("retry"); + await start; + + Assert.Equal(1, failed.StopCount); + Assert.Equal(1, failed.DisposeCount); + Assert.Null(await failed.Activity("stale")); + Assert.Equal(["retry"], harness.Owner.ReadyIds); + } + + [Fact] + public async Task StartAsync_PrefersRetryAfterOverBackoff() + { + Harness harness = new() { Backoff = _ => TimeSpan.FromSeconds(2) }; + harness.Factory.Enqueue(startError: new SocketModeNegotiateException( + HttpStatusCode.TooManyRequests, + TimeSpan.FromSeconds(7))); + FakeConnection retry = harness.Factory.Enqueue(); + + Task start = harness.Socket.StartAsync(); + await harness.Time.WaitForTimerAsync(TimeSpan.FromSeconds(7)); + Assert.False(harness.Time.HasTimerDueIn(TimeSpan.FromSeconds(2))); + + harness.Time.Advance(TimeSpan.FromSeconds(7)); + await retry.Started.Task; + retry.Ready("retry"); + await start; + } + + [Fact] + public async Task StartAsync_FailsWhenPendingAttemptExceedsBudget() + { + Harness harness = new() { StartupTimeout = TimeSpan.FromSeconds(5) }; + FakeConnection connection = harness.Factory.Enqueue(); + + Task start = harness.Socket.StartAsync(); + await connection.Started.Task; + await harness.Time.WaitForTimerAsync(TimeSpan.FromSeconds(5)); + harness.Time.Advance(TimeSpan.FromSeconds(5)); + + await Assert.ThrowsAsync(() => start); + Assert.Equal(1, harness.Factory.CreateCount); + Assert.Equal(1, connection.DisposeCount); + } + + [Fact] + public async Task StartAsync_ZeroBudgetMakesOneAttemptWithoutRetry() + { + Harness harness = new() { StartupTimeout = TimeSpan.Zero }; + harness.Factory.Enqueue(startError: new IOException("failed")); + + IOException error = await Assert.ThrowsAsync(() => harness.Socket.StartAsync()); + + Assert.Equal("failed", error.Message); + Assert.Equal(1, harness.Factory.CreateCount); + } + + [Theory] + [InlineData(HttpStatusCode.Unauthorized)] + [InlineData(HttpStatusCode.Forbidden)] + public async Task StartAsync_AuthRejectionFailsWithoutRetry(HttpStatusCode status) + { + Harness harness = new(); + SocketModeNegotiateException rejected = new(status, retryAfter: null); + FakeConnection connection = harness.Factory.Enqueue(startError: rejected); + + SocketModeNegotiateException error = + await Assert.ThrowsAsync(() => harness.Socket.StartAsync()); + + Assert.Same(rejected, error); + Assert.Equal(1, harness.Factory.CreateCount); + Assert.Equal(0, harness.Time.ScheduledTimerCount); + Assert.Equal(1, connection.DisposeCount); + Assert.Empty(harness.Logger.Warnings); + } + + [Fact] + public async Task Reconnect_AuthRejectionStopsTheGeo() + { + Harness harness = new() { Backoff = _ => TimeSpan.FromSeconds(3) }; + FakeConnection initial = harness.Factory.Enqueue(); + SocketModeNegotiateException rejected = new(HttpStatusCode.Unauthorized, retryAfter: null); + FakeConnection rejectedAttempt = harness.Factory.Enqueue(startError: rejected); + await harness.StartReadyAsync(initial); + + IOException dropped = new("dropped"); + initial.Close(dropped); + await harness.Time.WaitForTimerAsync(TimeSpan.FromSeconds(3)); + harness.Time.Advance(TimeSpan.FromSeconds(3)); + await WaitUntilAsync(() => harness.Owner.Disconnections.Count == 2); + await harness.Socket.StopAsync(); + + Assert.Equal([dropped, rejected], harness.Owner.Disconnections); + Assert.Same(rejected, Assert.Single(harness.Logger.Errors)); + Assert.Equal(2, harness.Factory.CreateCount); + Assert.Equal(0, harness.Time.ScheduledTimerCount); + Assert.Equal(1, rejectedAttempt.DisposeCount); + Assert.Equal(0, harness.Owner.Reconnections); + } + + [Fact] + public async Task Rotation_AuthRejectionStopsTheConnectionStillServing() + { + Harness harness = new(); + FakeConnection initial = harness.Factory.Enqueue(TimeSpan.FromSeconds(10)); + SocketModeNegotiateException rejected = new(HttpStatusCode.Forbidden, retryAfter: null); + FakeConnection rejectedAttempt = harness.Factory.Enqueue(startError: rejected); + await harness.StartReadyAsync(initial); + + harness.Time.Advance(TimeSpan.FromSeconds(5)); + await harness.AdvanceThroughBackoffAsync(rejectedAttempt); + await initial.Disposed.Task; + await WaitUntilAsync(() => harness.Owner.Disconnections.Count == 1); + + Assert.Same(rejected, Assert.Single(harness.Owner.Disconnections)); + Assert.Equal(1, initial.StopCount); + Assert.Null(await initial.Activity("after-rejection")); + Assert.DoesNotContain("after-rejection", harness.Owner.Dispatched); + Assert.Equal(2, harness.Factory.CreateCount); + Assert.DoesNotContain(harness.Logger.Warnings, warning => warning.Message.Contains("paused", StringComparison.Ordinal)); + } + + [Fact] + public async Task UnexpectedClose_ReportsDisconnectedThenReconnected() + { + Harness harness = new() { Backoff = _ => TimeSpan.FromSeconds(3) }; + FakeConnection initial = harness.Factory.Enqueue(); + FakeConnection replacement = harness.Factory.Enqueue(); + await harness.StartReadyAsync(initial); + + IOException dropped = new("dropped"); + initial.Close(dropped); + + Assert.Same(dropped, Assert.Single(harness.Owner.Disconnections)); + Assert.Null(await initial.Activity("after-close")); + await harness.Time.WaitForTimerAsync(TimeSpan.FromSeconds(3)); + harness.Time.Advance(TimeSpan.FromSeconds(3)); + await replacement.Started.Task; + replacement.Ready("replacement"); + await WaitUntilAsync(() => harness.Owner.Reconnections == 1); + + Assert.Equal(1, initial.DisposeCount); + Assert.NotNull(await replacement.Activity("resumed")); + Assert.DoesNotContain("after-close", harness.Owner.Dispatched); + } + + [Fact] + public async Task UnexpectedClose_ReportsDisconnectedBeforeSupervisorReconnects() + { + Harness harness = new() { Backoff = _ => TimeSpan.Zero }; + FakeConnection initial = harness.Factory.Enqueue(); + FakeConnection replacement = harness.Factory.Enqueue(); + await harness.StartReadyAsync(initial); + + using ManualResetEventSlim entered = new(); + using ManualResetEventSlim release = new(); + harness.Owner.Disconnecting = () => + { + entered.Set(); + release.Wait(TimeSpan.FromSeconds(5)); + }; + + Task closing = Task.Run(() => initial.Close(new IOException("dropped"))); + Assert.True(entered.Wait(TimeSpan.FromSeconds(5))); + + Assert.False(SpinWait.SpinUntil(() => initial.DisposeCount > 0, TimeSpan.FromMilliseconds(200))); + Assert.Equal(1, harness.Factory.CreateCount); + + release.Set(); + await closing; + await replacement.Started.Task; + replacement.Ready("replacement"); + await WaitUntilAsync(() => harness.Owner.Reconnections == 1); + + Assert.Single(harness.Owner.Disconnections); + Assert.Equal(1, initial.DisposeCount); + } + + [Fact] + public async Task Rotation_IsMakeBeforeBreakWithHandoff() + { + Harness harness = new(); + FakeConnection initial = harness.Factory.Enqueue(TimeSpan.FromSeconds(10)); + FakeConnection replacement = harness.Factory.Enqueue(); + await harness.StartReadyAsync(initial); + + harness.Time.Advance(TimeSpan.FromSeconds(5)); + await harness.AdvanceThroughBackoffAsync(replacement); + Assert.NotNull(await initial.Activity("during-replacement-startup")); + Assert.Equal(0, initial.StopCount); + + replacement.Ready("replacement"); + await harness.Time.WaitForTimerAsync(harness.Owner.HandoffWindow); + Assert.NotNull(await initial.Activity("retiring")); + Assert.NotNull(await replacement.Activity("active")); + Assert.Equal(0, initial.StopCount); + + harness.Time.Advance(harness.Owner.HandoffWindow); + await initial.Disposed.Task; + + Assert.Equal(1, initial.StopCount); + Assert.Null(await initial.Activity("retired")); + Assert.DoesNotContain("retired", harness.Owner.Dispatched); + Assert.Empty(harness.Owner.Disconnections); + Assert.Equal(0, harness.Owner.Reconnections); + Assert.Equal(2, harness.Owner.ReadyIds.Count); + } + + [Fact] + public async Task Rotation_UsesOneSecondMinimumDelay() + { + Harness harness = new() { TokenRefreshMargin = TimeSpan.FromSeconds(30) }; + FakeConnection initial = harness.Factory.Enqueue(TimeSpan.FromSeconds(10)); + FakeConnection replacement = harness.Factory.Enqueue(); + await harness.StartReadyAsync(initial); + + Assert.True(harness.Time.HasTimerDueIn(TimeSpan.FromSeconds(1))); + harness.Time.Advance(TimeSpan.FromSeconds(1)); + await harness.AdvanceThroughBackoffAsync(replacement); + } + + [Fact] + public async Task Rotation_WaitsForFirstBackoffDelayBeforeConnectingReplacement() + { + List attempts = []; + Harness harness = new() + { + Backoff = attempt => + { + lock (attempts) + { + attempts.Add(attempt); + } + + return TimeSpan.FromMilliseconds(700); + }, + }; + FakeConnection initial = harness.Factory.Enqueue(TimeSpan.FromSeconds(10)); + FakeConnection replacement = harness.Factory.Enqueue(); + await harness.StartReadyAsync(initial); + + harness.Time.Advance(TimeSpan.FromSeconds(5)); + await harness.Time.WaitForTimerAsync(TimeSpan.FromMilliseconds(700)); + + Assert.Equal(1, harness.Factory.CreateCount); + lock (attempts) + { + Assert.Equal([0], attempts); + } + + Assert.NotNull(await initial.Activity("during-backoff")); + + harness.Time.Advance(TimeSpan.FromMilliseconds(700)); + await replacement.Started.Task; + + Assert.Equal(2, harness.Factory.CreateCount); + } + + [Fact] + public async Task TokenLifetimeAbsent_DoesNotScheduleRotation() + { + Harness harness = new(); + FakeConnection connection = harness.Factory.Enqueue(tokenLifetime: null); + await harness.StartReadyAsync(connection); + + Assert.Equal(0, harness.Time.ScheduledTimerCount); + harness.Time.Advance(TimeSpan.FromDays(1)); + + Assert.Equal(1, harness.Factory.CreateCount); + } + + [Fact] + public async Task PredecessorDropDuringRotation_ReportsOutageUntilReplacementReady() + { + Harness harness = new(); + FakeConnection initial = harness.Factory.Enqueue(TimeSpan.FromSeconds(10)); + FakeConnection replacement = harness.Factory.Enqueue(); + await harness.StartReadyAsync(initial); + + harness.Time.Advance(TimeSpan.FromSeconds(5)); + await harness.AdvanceThroughBackoffAsync(replacement); + initial.Close(new IOException("dropped")); + + Assert.Single(harness.Owner.Disconnections); + Assert.Null(await initial.Activity("dropped")); + + replacement.Ready("replacement"); + await WaitUntilAsync(() => harness.Owner.Reconnections == 1); + await harness.Time.WaitForTimerAsync(harness.Owner.HandoffWindow); + harness.Time.Advance(harness.Owner.HandoffWindow); + await initial.Disposed.Task; + + Assert.Equal(1, initial.DisposeCount); + } + + [Fact] + public async Task RepeatedRotation_ReleasesEveryPredecessorOnce() + { + Harness harness = new() { HandoffWindow = TimeSpan.Zero }; + FakeConnection[] connections = + [ + harness.Factory.Enqueue(TimeSpan.FromSeconds(10)), + harness.Factory.Enqueue(TimeSpan.FromSeconds(10)), + harness.Factory.Enqueue(TimeSpan.FromSeconds(10)), + harness.Factory.Enqueue(TimeSpan.FromSeconds(10)), + ]; + await harness.StartReadyAsync(connections[0]); + + for (int i = 1; i < connections.Length; i++) + { + await harness.Time.WaitForTimerAsync(TimeSpan.FromSeconds(5)); + harness.Time.Advance(TimeSpan.FromSeconds(5)); + await harness.AdvanceThroughBackoffAsync(connections[i]); + connections[i].Ready($"gen-{i}"); + await connections[i - 1].Disposed.Task; + } + + Assert.All(connections[..^1], connection => Assert.Equal(1, connection.DisposeCount)); + Assert.Equal(0, connections[^1].DisposeCount); + Assert.Empty(harness.Owner.Disconnections); + + await harness.Socket.StopAsync(); + Assert.All(connections, connection => Assert.Equal(1, connection.DisposeCount)); + } + + [Fact] + public async Task StopAsync_DuringStartup_CancelsAndDisposesOnce() + { + Harness harness = new(); + FakeConnection connection = harness.Factory.Enqueue(); + + Task start = harness.Socket.StartAsync(); + await connection.Started.Task; + await harness.Socket.StopAsync(); + + await Assert.ThrowsAnyAsync(() => start); + await harness.Socket.StopAsync(); + await harness.Socket.DisposeAsync(); + await harness.Socket.DisposeAsync(); + + Assert.Equal(1, connection.StopCount); + Assert.Equal(1, connection.DisposeCount); + Assert.Empty(harness.Owner.Disconnections); + } + + [Fact] + public async Task StartAsync_CallerCancellationInterruptsRetryDelay() + { + Harness harness = new() { Backoff = _ => TimeSpan.FromMinutes(1), StartupTimeout = TimeSpan.FromMinutes(5) }; + FakeConnection failed = harness.Factory.Enqueue(startError: new IOException("failed")); + using CancellationTokenSource cancellation = new(); + + Task start = harness.Socket.StartAsync(cancellation.Token); + await harness.Time.WaitForTimerAsync(TimeSpan.FromMinutes(1)); + await cancellation.CancelAsync(); + + await Assert.ThrowsAnyAsync(() => start); + Assert.Equal(1, harness.Factory.CreateCount); + Assert.Equal(1, failed.DisposeCount); + } + + [Fact] + public async Task StopAsync_CompletesWhenSupervisorHasFaulted() + { + Harness harness = new() { Backoff = _ => TimeSpan.Zero }; + FakeConnection initial = harness.Factory.Enqueue(); + FakeConnection replacement = harness.Factory.Enqueue(); + harness.Owner.ReconnectedError = new InvalidOperationException("observer failed"); + await harness.StartReadyAsync(initial); + + initial.Close(new IOException("dropped")); + await replacement.Started.Task; + replacement.Ready("replacement"); + await WaitUntilAsync(() => harness.Owner.Reconnections == 1); + + await harness.Socket.StopAsync(); + + Assert.Equal(1, replacement.StopCount); + Assert.Equal(1, replacement.DisposeCount); + } + + [Fact] + public async Task StopAsync_DuringReconnectDelay_CancelsWithoutNewGeneration() + { + Harness harness = new() { Backoff = _ => TimeSpan.FromMinutes(1) }; + FakeConnection initial = harness.Factory.Enqueue(); + await harness.StartReadyAsync(initial); + + initial.Close(new IOException("dropped")); + await harness.Time.WaitForTimerAsync(TimeSpan.FromMinutes(1)); + await harness.Socket.StopAsync(); + harness.Time.Advance(TimeSpan.FromMinutes(1)); + + Assert.Equal(1, harness.Factory.CreateCount); + Assert.Equal(1, initial.DisposeCount); + Assert.Equal(0, harness.Owner.Reconnections); + } + + [Fact] + public async Task StopAsync_DuringRotation_DisposesActiveAndConnectingOnce() + { + Harness harness = new(); + FakeConnection active = harness.Factory.Enqueue(TimeSpan.FromSeconds(10)); + FakeConnection connecting = harness.Factory.Enqueue(); + await harness.StartReadyAsync(active); + + harness.Time.Advance(TimeSpan.FromSeconds(5)); + await harness.AdvanceThroughBackoffAsync(connecting); + await harness.Socket.StopAsync(); + + Assert.Equal(1, active.StopCount); + Assert.Equal(1, active.DisposeCount); + Assert.Equal(1, connecting.StopCount); + Assert.Equal(1, connecting.DisposeCount); + Assert.Empty(harness.Owner.Disconnections); + Assert.Null(await active.Activity("stopped")); + } + + [Fact] + public async Task StopAsync_DuringHandoff_DisposesActiveAndRetiringOnce() + { + Harness harness = new() { HandoffWindow = TimeSpan.FromMinutes(1) }; + FakeConnection retiring = harness.Factory.Enqueue(TimeSpan.FromSeconds(10)); + FakeConnection active = harness.Factory.Enqueue(); + await harness.StartReadyAsync(retiring); + + harness.Time.Advance(TimeSpan.FromSeconds(5)); + await harness.AdvanceThroughBackoffAsync(active); + active.Ready("replacement"); + await harness.Time.WaitForTimerAsync(TimeSpan.FromMinutes(1)); + await harness.Socket.StopAsync(); + harness.Time.Advance(TimeSpan.FromMinutes(1)); + + Assert.Equal(1, retiring.StopCount); + Assert.Equal(1, retiring.DisposeCount); + Assert.Equal(1, active.StopCount); + Assert.Equal(1, active.DisposeCount); + Assert.Empty(harness.Owner.Disconnections); + } + + [Fact] + public async Task StopAsync_WaitsForFailedStartupAttemptCleanup() + { + Harness harness = new(); + TaskCompletionSource disposeGate = new(TaskCreationOptions.RunContinuationsAsynchronously); + FakeConnection failed = harness.Factory.Enqueue(startError: new IOException("connect failed")); + failed.DisposeGate = disposeGate.Task; + + Task start = harness.Socket.StartAsync(); + await failed.Disposed.Task; + + Task stop = harness.Socket.StopAsync(); + Assert.NotSame(stop, await Task.WhenAny(stop, Task.Delay(TimeSpan.FromMilliseconds(200)))); + + disposeGate.SetResult(); + await stop; + await Assert.ThrowsAsync(() => start); + + Assert.Equal(1, failed.StopCount); + Assert.Equal(1, failed.DisposeCount); + Assert.Equal(1, harness.Factory.CreateCount); + } + + [Fact] + public async Task StopAsync_LogsConnectionCleanupFailuresAndStillDisposes() + { + Harness harness = new(); + FakeConnection connection = harness.Factory.Enqueue(); + connection.StopError = new IOException("stop failed"); + connection.DisposeError = new IOException("dispose failed"); + await harness.StartReadyAsync(connection); + + await harness.Socket.StopAsync(); + + Assert.Equal(1, connection.StopCount); + Assert.Equal(1, connection.DisposeCount); + Assert.Equal( + [connection.StopError, connection.DisposeError], + harness.Logger.Warnings.Select(warning => warning.Exception)); + } + + [Fact] + public async Task Retirement_LogsDisposeFailureAndKeepsReplacementActive() + { + Harness harness = new() { HandoffWindow = TimeSpan.FromSeconds(5) }; + FakeConnection retiring = harness.Factory.Enqueue(TimeSpan.FromSeconds(10)); + FakeConnection active = harness.Factory.Enqueue(); + retiring.DisposeError = new IOException("dispose failed"); + await harness.StartReadyAsync(retiring); + + harness.Time.Advance(TimeSpan.FromSeconds(5)); + await harness.AdvanceThroughBackoffAsync(active); + active.Ready("replacement"); + await harness.Time.WaitForTimerAsync(TimeSpan.FromSeconds(5)); + harness.Time.Advance(TimeSpan.FromSeconds(5)); + await WaitUntilAsync(() => harness.Logger.Warnings.Length > 0); + + Assert.Same(retiring.DisposeError, Assert.Single(harness.Logger.Warnings).Exception); + Assert.NotNull(await active.Activity("after-retirement")); + + await harness.Socket.StopAsync(); + + Assert.Equal(1, retiring.DisposeCount); + Assert.Equal(1, active.DisposeCount); + } + + private static async Task WaitUntilAsync(Func condition) + { + Stopwatch elapsed = Stopwatch.StartNew(); + while (!condition()) + { + if (elapsed.Elapsed > TimeSpan.FromSeconds(5)) + { + throw new TimeoutException("The asynchronous test condition was not met."); + } + + await Task.Yield(); + } + } + + private sealed class Harness + { + private GeoSocket? _socket; + + internal ManualTimeProvider Time { get; } = new(); + + internal FakeOwner Owner { get; } = new(); + + internal FakeConnectionFactory Factory { get; } = new(); + + internal RecordingLogger Logger { get; } = new(); + + internal TimeSpan StartupTimeout { init => Owner.StartupTimeout = value; } + + internal TimeSpan TokenRefreshMargin { init => Owner.TokenRefreshMargin = value; } + + internal TimeSpan HandoffWindow { init => Owner.HandoffWindow = value; } + + internal Func Backoff { init => Owner.Backoff = value; } + + internal GeoSocket Socket => + _socket ??= new GeoSocket(Owner, "amer", NegotiateUri, Factory, Logger, Time); + + internal async Task StartReadyAsync(FakeConnection connection) + { + Task start = Socket.StartAsync(); + await connection.Started.Task; + connection.Ready("initial"); + await start; + } + + internal async Task AdvanceThroughBackoffAsync(FakeConnection next) + { + TimeSpan delay = Owner.Backoff(0); + await Time.WaitForTimerAsync(delay); + Time.Advance(delay); + await next.Started.Task; + } + } + + private sealed class FakeOwner : IGeoSocketOwner + { + private readonly object _sync = new(); + private readonly List _dispatched = []; + private readonly List _readyIds = []; + private readonly List _disconnections = []; + private int _reconnections; + + public TimeSpan StartupTimeout { get; set; } = TimeSpan.FromSeconds(30); + + public TimeSpan TokenRefreshMargin { get; set; } = TimeSpan.FromSeconds(5); + + public TimeSpan HandoffWindow { get; set; } = TimeSpan.FromSeconds(5); + + internal Func Backoff { get; set; } = _ => TimeSpan.FromSeconds(1); + + internal List Dispatched { get { lock (_sync) { return [.. _dispatched]; } } } + + internal List ReadyIds { get { lock (_sync) { return [.. _readyIds]; } } } + + internal List Disconnections { get { lock (_sync) { return [.. _disconnections]; } } } + + internal int Reconnections => Volatile.Read(ref _reconnections); + + public TimeSpan? GetRetryAfter(Exception? error) + => (error as SocketModeNegotiateException)?.RetryAfter; + + public TimeSpan GetBackoffDelay(int attempt) => Backoff(attempt); + + public Task DispatchAsync(string geo, SocketActivityEnvelope envelope) + { + lock (_sync) + { + _dispatched.Add(envelope.EnvelopeId!); + } + + return Task.FromResult(new SocketReplyFrame { Status = 200 }); + } + + public void OnGeoReady(string geo, SocketReadyFrame frame) + { + lock (_sync) + { + _readyIds.Add(frame.ConnectionId!); + } + } + + internal Action? Disconnecting { get; set; } + + public void OnGeoDisconnected(string geo, Exception? error) + { + Disconnecting?.Invoke(); + lock (_sync) + { + _disconnections.Add(error); + } + } + + internal Exception? ReconnectedError { get; set; } + + public void OnGeoReconnected(string geo) + { + Interlocked.Increment(ref _reconnections); + if (ReconnectedError is not null) + { + throw ReconnectedError; + } + } + } + + private sealed class FakeConnectionFactory : ISocketConnectionFactory + { + private readonly Queue _pending = []; + private int _createCount; + + internal int CreateCount => Volatile.Read(ref _createCount); + + internal FakeConnection Enqueue(TimeSpan? tokenLifetime = null, Exception? startError = null) + { + FakeConnection connection = new(tokenLifetime, startError); + lock (_pending) + { + _pending.Enqueue(connection); + } + + return connection; + } + + public ISocketConnection Create(Uri negotiateUri, SocketConnectionHandlers handlers) + { + Interlocked.Increment(ref _createCount); + FakeConnection connection; + lock (_pending) + { + connection = _pending.Dequeue(); + } + + connection.Handlers = handlers; + return connection; + } + } + + private sealed class FakeConnection(TimeSpan? tokenLifetime, Exception? startError) : ISocketConnection + { + private readonly TaskCompletionSource _ready = new(TaskCreationOptions.RunContinuationsAsynchronously); + private int _closed; + private int _stopCount; + private int _disposeCount; + + public TimeSpan? TokenLifetime { get; } = tokenLifetime; + + internal SocketConnectionHandlers Handlers { get; set; } = null!; + + internal TaskCompletionSource Started { get; } = new(TaskCreationOptions.RunContinuationsAsynchronously); + + internal TaskCompletionSource Disposed { get; } = new(TaskCreationOptions.RunContinuationsAsynchronously); + + internal int StopCount => Volatile.Read(ref _stopCount); + + internal int DisposeCount => Volatile.Read(ref _disposeCount); + + internal Exception? StopError { get; set; } + + internal Exception? DisposeError { get; set; } + + public async Task StartAsync(CancellationToken cancellationToken) + { + Started.TrySetResult(); + if (startError is not null) + { + throw startError; + } + + await _ready.Task.WaitAsync(cancellationToken); + } + + public Task StopAsync(CancellationToken cancellationToken = default) + { + Interlocked.Increment(ref _stopCount); + _ready.TrySetCanceled(CancellationToken.None); + RaiseClosed(null, planned: true); + return StopError is null ? Task.CompletedTask : Task.FromException(StopError); + } + + internal Task? DisposeGate { get; set; } + + public async ValueTask DisposeAsync() + { + Interlocked.Increment(ref _disposeCount); + Disposed.TrySetResult(); + if (DisposeGate is not null) + { + await DisposeGate; + } + + if (DisposeError is not null) + { + throw DisposeError; + } + } + + internal void Ready(string connectionId) + { + Handlers.OnReady(new SocketReadyFrame { ConnectionId = connectionId }); + _ready.TrySetResult(); + } + + internal void Close(Exception error) => RaiseClosed(error, planned: false); + + internal Task Activity(string envelopeId) + => Handlers.OnActivity(new SocketActivityEnvelope { EnvelopeId = envelopeId }); + + private void RaiseClosed(Exception? error, bool planned) + { + if (Interlocked.Exchange(ref _closed, 1) == 0) + { + Handlers.OnClosed(error, planned); + } + } + } + + private sealed class RecordingLogger : ILogger + { + private readonly List<(string Message, Exception? Exception)> _warnings = []; + private readonly List _errors = []; + + internal Exception?[] Errors + { + get + { + lock (_errors) + { + return [.. _errors]; + } + } + } + + internal (string Message, Exception? Exception)[] Warnings + { + get + { + lock (_warnings) + { + return [.. _warnings]; + } + } + } + + public IDisposable? BeginScope(TState state) + where TState : notnull => null; + + public bool IsEnabled(LogLevel logLevel) => true; + + public void Log( + LogLevel logLevel, + EventId eventId, + TState state, + Exception? exception, + Func formatter) + { + if (logLevel == LogLevel.Warning) + { + lock (_warnings) + { + _warnings.Add((formatter(state, exception), exception)); + } + } + else if (logLevel == LogLevel.Error) + { + lock (_errors) + { + _errors.Add(exception); + } + } + } + } + + private sealed class ManualTimeProvider : TimeProvider + { + private readonly object _sync = new(); + private readonly List _timers = []; + private long _now; + + public override long TimestampFrequency => TimeSpan.TicksPerSecond; + + internal int ScheduledTimerCount + { + get + { + lock (_sync) + { + return _timers.Count(timer => timer.Due is not null); + } + } + } + + public override long GetTimestamp() + { + lock (_sync) + { + return _now; + } + } + + public override DateTimeOffset GetUtcNow() => DateTimeOffset.UnixEpoch.AddTicks(GetTimestamp()); + + public override ITimer CreateTimer(TimerCallback callback, object? state, TimeSpan dueTime, TimeSpan period) + { + ManualTimer timer = new(this, callback, state); + timer.Change(dueTime, period); + lock (_sync) + { + _timers.Add(timer); + } + + return timer; + } + + internal bool HasTimerDueIn(TimeSpan dueIn) + { + lock (_sync) + { + return _timers.Any(timer => timer.Due == _now + dueIn.Ticks); + } + } + + internal Task WaitForTimerAsync(TimeSpan dueIn) => WaitUntilAsync(() => HasTimerDueIn(dueIn)); + + internal void Advance(TimeSpan amount) + { + List due; + lock (_sync) + { + _now += amount.Ticks; + due = [.. _timers.Where(timer => timer.Due <= _now)]; + foreach (ManualTimer timer in due) + { + timer.Due = null; + } + } + + foreach (ManualTimer timer in due) + { + timer.Fire(); + } + } + + private sealed class ManualTimer(ManualTimeProvider owner, TimerCallback callback, object? state) : ITimer + { + internal long? Due { get; set; } + + public bool Change(TimeSpan dueTime, TimeSpan period) + { + lock (owner._sync) + { + Due = dueTime == Timeout.InfiniteTimeSpan ? null : owner._now + dueTime.Ticks; + } + + return true; + } + + public void Dispose() + { + lock (owner._sync) + { + Due = null; + owner._timers.Remove(this); + } + } + + public ValueTask DisposeAsync() + { + Dispose(); + return ValueTask.CompletedTask; + } + + internal void Fire() => callback(state); + } + } +} diff --git a/test/Microsoft.Teams.Apps.UnitTests/SocketMode/SocketModeEndToEndTests.cs b/test/Microsoft.Teams.Apps.UnitTests/SocketMode/SocketModeEndToEndTests.cs new file mode 100644 index 00000000..eff08935 --- /dev/null +++ b/test/Microsoft.Teams.Apps.UnitTests/SocketMode/SocketModeEndToEndTests.cs @@ -0,0 +1,287 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System.Collections.Concurrent; +using System.Net; +using System.Net.Http.Headers; +using System.Text; +using System.Text.Json; +using Microsoft.Extensions.Logging.Abstractions; +using Microsoft.Teams.Apps.SocketMode; +using Microsoft.Teams.Core.Schema; + +namespace Microsoft.Teams.Apps.UnitTests.SocketMode; + +/// +/// Runs Socket Mode end to end in process: transport, geo supervisors, negotiator and SignalR connection factory are real; only HTTP and the SignalR client are faked. +/// +public class SocketModeEndToEndTests +{ + [Fact] + public async Task Transport_ComposesNegotiatorAndSignalRFactoryAcrossGeos() + { + EndToEnd stack = new(); + + Task start = stack.Transport.StartAsync(); + FakeSignalRClient[] clients = await stack.SignalR.WaitForAsync(3); + foreach (FakeSignalRClient client in clients) + { + client.Ready(); + } + + await start; + + Assert.Equal(SocketModeStatus.Ready, stack.Transport.Status); + Assert.Equal( + [ + "https://botapi.skype.com/amer/v3/websockets/connect", + "https://botapi.skype.com/apac/v3/websockets/connect", + "https://botapi.skype.com/emea/v3/websockets/connect", + ], + stack.Http.Requests.Select(request => request.AbsoluteUri).Order(StringComparer.Ordinal)); + Assert.All(stack.Http.Authorizations, authorization => Assert.Equal("Bearer bot-token", authorization)); + Assert.Equal( + ["https://signalr.test/amer", "https://signalr.test/apac", "https://signalr.test/emea"], + clients.Select(client => client.Url.AbsoluteUri).Order(StringComparer.Ordinal)); + + FakeSignalRClient emea = clients.Single(client => client.Url.AbsolutePath == "/emea"); + SocketReplyFrame? reply = await emea.ActivityAsync(new SocketActivityEnvelope + { + ProtocolVersion = SocketModeProtocol.CurrentVersion, + EnvelopeId = "env-1", + Type = "invoke", + Payload = JsonDocument.Parse("""{"type":"invoke","id":"a1"}""").RootElement.Clone(), + }); + + Assert.NotNull(reply); + Assert.Equal("env-1", reply.EnvelopeId); + Assert.Equal("bot-id", reply.BotKey); + Assert.Equal(200, reply.Status); + Assert.Equal("a1", Assert.Single(stack.Dispatched).Id); + + await stack.Transport.StopAsync(); + + Assert.Equal(SocketModeStatus.Stopped, stack.Transport.Status); + Assert.All(clients, client => + { + Assert.Equal(1, client.StopCount); + Assert.Equal(1, client.DisposeCount); + }); + } + + [Fact] + public async Task Transport_HonorsNegotiatorRetryAfterDuringStartup() + { + EndToEnd stack = new(new SocketModeTransportOptions { Geos = ["amer"] }); + stack.Http.ThrottleOnce(TimeSpan.Zero); + + Task start = stack.Transport.StartAsync(); + FakeSignalRClient client = (await stack.SignalR.WaitForAsync(1))[0]; + client.Ready(); + await start; + + Assert.Equal(2, stack.Http.Requests.Count); + Assert.Equal(SocketModeStatus.Ready, stack.Transport.Status); + + await stack.Transport.DisposeAsync(); + } + + [Fact] + public async Task Transport_FailsStartupWithNegotiatorErrorWhenBudgetIsExhausted() + { + EndToEnd stack = new(new SocketModeTransportOptions { Geos = ["amer"], StartupTimeout = TimeSpan.Zero }); + stack.Http.ThrottleOnce(TimeSpan.Zero); + + SocketModeNegotiateException error = + await Assert.ThrowsAsync(() => stack.Transport.StartAsync()); + + Assert.Equal(HttpStatusCode.TooManyRequests, error.StatusCode); + Assert.Equal(SocketModeStatus.Stopped, stack.Transport.Status); + Assert.Empty(stack.SignalR.Clients); + } + + private sealed class EndToEnd + { + internal EndToEnd(SocketModeTransportOptions? options = null) + { + SocketModeNegotiator negotiator = new( + new HttpClient(Http), + _ => Task.FromResult("bot-token")); + SignalRSocketConnectionFactory factory = new( + negotiator, + TimeSpan.FromSeconds(30), + TimeSpan.FromSeconds(15), + TimeSpan.FromSeconds(30), + NullLogger.Instance, + SignalR.Create); + Transport = new SocketModeTransport( + options ?? new SocketModeTransportOptions(), + factory, + activity => + { + Dispatched.Enqueue(activity); + return Task.FromResult(new SocketDispatchResult(200, null)); + }, + NullLogger.Instance, + botKey: "bot-id"); + } + + internal NegotiateHandler Http { get; } = new(); + internal FakeSignalRClientFactory SignalR { get; } = new(); + internal ConcurrentQueue Dispatched { get; } = new(); + internal SocketModeTransport Transport { get; } + } + + private sealed class NegotiateHandler : HttpMessageHandler + { + private readonly object _gate = new(); + private TimeSpan? _throttle; + + internal ConcurrentQueue Requests { get; } = new(); + internal ConcurrentQueue Authorizations { get; } = new(); + + internal void ThrottleOnce(TimeSpan retryAfter) + { + lock (_gate) + { + _throttle = retryAfter; + } + } + + protected override Task SendAsync( + HttpRequestMessage request, + CancellationToken cancellationToken) + { + Uri uri = request.RequestUri ?? throw new InvalidOperationException("Missing request URI."); + Requests.Enqueue(uri); + Authorizations.Enqueue(request.Headers.Authorization?.ToString()); + + TimeSpan? throttle; + lock (_gate) + { + throttle = _throttle; + _throttle = null; + } + + if (throttle is TimeSpan retryAfter) + { + HttpResponseMessage throttled = new(HttpStatusCode.TooManyRequests); + throttled.Headers.RetryAfter = new RetryConditionHeaderValue(retryAfter); + return Task.FromResult(throttled); + } + + string geo = uri.Segments[1].TrimEnd('/'); + string json = $$"""{"url":"https://signalr.test/{{geo}}","accessToken":"token-{{geo}}","expiresIn":3600}"""; + return Task.FromResult(new HttpResponseMessage(HttpStatusCode.OK) + { + Content = new StringContent(json, Encoding.UTF8, "application/json"), + }); + } + } + + private sealed class FakeSignalRClientFactory + { + private readonly object _gate = new(); + private TaskCompletionSource _changed = NewSignal(); + + internal List Clients { get; } = []; + + internal ISignalRClientConnection Create( + Uri url, + string accessToken, + TimeSpan keepAliveInterval, + TimeSpan serverTimeout) + { + FakeSignalRClient client = new(url, Changed); + lock (_gate) + { + Clients.Add(client); + } + + Changed(); + return client; + } + + internal async Task WaitForAsync(int count) + { + using CancellationTokenSource timeout = new(TimeSpan.FromSeconds(10)); + while (true) + { + Task changed; + lock (_gate) + { + if (Clients.Count >= count && Clients.Take(count).All(client => client.StartCount > 0)) + { + return [.. Clients.Take(count)]; + } + + changed = _changed.Task; + } + + await changed.WaitAsync(timeout.Token); + } + } + + private void Changed() + { + TaskCompletionSource previous; + lock (_gate) + { + previous = _changed; + _changed = NewSignal(); + } + + previous.TrySetResult(); + } + + private static TaskCompletionSource NewSignal() => new(TaskCreationOptions.RunContinuationsAsynchronously); + } + + private sealed class FakeSignalRClient(Uri url, Action changed) : ISignalRClientConnection + { + private Func>? _onActivity; + private Action? _onReady; + private Action? _onClosed; + private int _startCount; + private int _stopCount; + private int _disposeCount; + + internal Uri Url { get; } = url; + internal int StartCount => Volatile.Read(ref _startCount); + internal int StopCount => Volatile.Read(ref _stopCount); + internal int DisposeCount => Volatile.Read(ref _disposeCount); + + public void OnActivity(Func> handler) => _onActivity = handler; + + public void OnReady(Action handler) => _onReady = handler; + + public void OnClosed(Action handler) => _onClosed = handler; + + public Task StartAsync(CancellationToken cancellationToken) + { + cancellationToken.ThrowIfCancellationRequested(); + Interlocked.Increment(ref _startCount); + changed(); + return Task.CompletedTask; + } + + public Task StopAsync(CancellationToken cancellationToken) + { + Interlocked.Increment(ref _stopCount); + _onClosed?.Invoke(null); + return Task.CompletedTask; + } + + public ValueTask DisposeAsync() + { + Interlocked.Increment(ref _disposeCount); + return ValueTask.CompletedTask; + } + + internal void Ready() + => (_onReady ?? throw new InvalidOperationException())(new SocketReadyFrame { BotKey = "bot-id" }); + + internal Task ActivityAsync(SocketActivityEnvelope envelope) + => (_onActivity ?? throw new InvalidOperationException())(envelope); + } +} diff --git a/test/Microsoft.Teams.Apps.UnitTests/SocketMode/SocketModeNegotiatorTests.cs b/test/Microsoft.Teams.Apps.UnitTests/SocketMode/SocketModeNegotiatorTests.cs index cd017376..5be6b3e3 100644 --- a/test/Microsoft.Teams.Apps.UnitTests/SocketMode/SocketModeNegotiatorTests.cs +++ b/test/Microsoft.Teams.Apps.UnitTests/SocketMode/SocketModeNegotiatorTests.cs @@ -216,6 +216,36 @@ await Assert.ThrowsAsync( Assert.Null(exception.RetryAfter); } + [Theory] + [InlineData(HttpStatusCode.Unauthorized, "verify the bot credentials")] + [InlineData(HttpStatusCode.Forbidden, "not authorized to use Socket Mode")] + public async Task NegotiateAsync_AuthFailureIsNonRetryableWithActionableMessage(HttpStatusCode status, string guidance) + { + RecordingHandler handler = new((_, _) => Task.FromResult(new HttpResponseMessage(status))); + SocketModeNegotiator negotiator = CreateNegotiator(handler); + + SocketModeNegotiateException exception = + await Assert.ThrowsAsync( + () => negotiator.NegotiateAsync(NegotiateUri)); + + Assert.True(exception.IsAuthError); + Assert.True(SocketModeNegotiateException.IsNonRetryable(exception)); + Assert.Contains($"HTTP {(int)status}", exception.Message, StringComparison.Ordinal); + Assert.Contains(guidance, exception.Message, StringComparison.Ordinal); + } + + [Theory] + [InlineData(HttpStatusCode.TooManyRequests)] + [InlineData(HttpStatusCode.ServiceUnavailable)] + public void OtherNegotiateFailuresRemainRetryable(HttpStatusCode status) + { + SocketModeNegotiateException exception = new(status, retryAfter: null); + + Assert.False(exception.IsAuthError); + Assert.False(SocketModeNegotiateException.IsNonRetryable(exception)); + Assert.False(SocketModeNegotiateException.IsNonRetryable(new IOException("dropped"))); + } + [Fact] public async Task NegotiateAsync_HonorsCallerCancellation() { diff --git a/test/Microsoft.Teams.Apps.UnitTests/SocketMode/SocketModeTransportTests.cs b/test/Microsoft.Teams.Apps.UnitTests/SocketMode/SocketModeTransportTests.cs new file mode 100644 index 00000000..49b86699 --- /dev/null +++ b/test/Microsoft.Teams.Apps.UnitTests/SocketMode/SocketModeTransportTests.cs @@ -0,0 +1,516 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +using System.Net; +using System.Text.Json; +using Microsoft.Extensions.Logging.Abstractions; +using Microsoft.Teams.Apps.SocketMode; +using Microsoft.Teams.Core.Schema; + +namespace Microsoft.Teams.Apps.UnitTests.SocketMode; + +public class SocketModeTransportTests +{ + private static readonly DateTimeOffset Now = DateTimeOffset.FromUnixTimeMilliseconds(1_700_000_000_000); + + [Fact] + public async Task StartAsync_ConnectsEveryDefaultGeo() + { + Harness harness = new(); + + Task start = harness.Transport.StartAsync(); + FakeConnection[] connections = await harness.Factory.WaitForAsync(3); + + Assert.Equal( + [ + "https://botapi.skype.com/amer/v3/websockets/connect", + "https://botapi.skype.com/apac/v3/websockets/connect", + "https://botapi.skype.com/emea/v3/websockets/connect", + ], + SortedUris(connections)); + Assert.False(start.IsCompleted); + Assert.Equal(SocketModeStatus.Connecting, harness.Transport.Status); + + foreach (FakeConnection connection in connections) + { + connection.Ready(); + } + + await start; + Assert.Equal(SocketModeStatus.Ready, harness.Transport.Status); + Assert.All(harness.Transport.GeoStatuses.Values, status => Assert.Equal(SocketModeStatus.Ready, status)); + } + + [Fact] + public async Task StartAsync_BuildsGeoUrisFromCustomBaseAndEmptyGeo() + { + Harness harness = new(new SocketModeTransportOptions + { + NegotiateBaseUri = new Uri("http://localhost:3978/"), + Geos = ["", " /eu/ "], + }); + + Task start = harness.Transport.StartAsync(); + FakeConnection[] connections = await harness.Factory.WaitForAsync(2); + + Assert.Equal( + [ + "http://localhost:3978/eu/v3/websockets/connect", + "http://localhost:3978/v3/websockets/connect", + ], + SortedUris(connections)); + foreach (FakeConnection connection in connections) + { + connection.Ready(); + } + + await start; + } + + [Theory] + [InlineData(new object[] { new string[0] })] + [InlineData(new object[] { new[] { "amer", "AMER" } })] + public void Constructor_RejectsInvalidGeos(string[] geos) + { + Assert.Throws(() => new Harness(new SocketModeTransportOptions { Geos = geos }).Transport); + } + + [Fact] + public void Constructor_RejectsNegativeStartupTimeout() + { + Assert.Throws(() => new Harness(new SocketModeTransportOptions + { + StartupTimeout = TimeSpan.FromSeconds(-1), + }).Transport); + } + + [Fact] + public async Task StartAsync_WhenAnyGeoFails_StopsEveryGeoAndThrowsThatFailure() + { + IOException failure = new("emea failed"); + Harness harness = new(new SocketModeTransportOptions { StartupTimeout = TimeSpan.Zero }); + harness.Factory.FailStart("emea", failure); + + IOException thrown = await Assert.ThrowsAsync(() => harness.Transport.StartAsync()); + + Assert.Same(failure, thrown); + Assert.Equal(SocketModeStatus.Stopped, harness.Transport.Status); + Assert.All(harness.Factory.Connections, connection => Assert.Equal(1, connection.DisposeCount)); + } + + [Fact] + public async Task StartAsync_WhenCleanupFails_StillThrowsTheStartupFailure() + { + IOException failure = new("emea failed"); + Harness harness = new(new SocketModeTransportOptions { StartupTimeout = TimeSpan.Zero }); + harness.Factory.FailDispose("amer", new IOException("amer dispose failed")); + harness.Factory.FailDispose("apac", new IOException("apac dispose failed")); + + Task start = harness.Transport.StartAsync(); + FakeConnection[] connections = await harness.Factory.WaitForAsync(3); + connections.Single(connection => connection.Geo == "amer").Ready(); + connections.Single(connection => connection.Geo == "apac").Ready(); + connections.Single(connection => connection.Geo == "emea").Fail(failure); + + IOException thrown = await Assert.ThrowsAsync(() => start); + + Assert.Same(failure, thrown); + Assert.Equal(SocketModeStatus.Stopped, harness.Transport.Status); + Assert.All(harness.Factory.Connections, connection => Assert.Equal(1, connection.DisposeCount)); + } + + [Fact] + public async Task StopAsync_DuringStartup_CancelsStartAndDisposesConnections() + { + Harness harness = new(); + Task start = harness.Transport.StartAsync(); + FakeConnection[] connections = await harness.Factory.WaitForAsync(3); + + await harness.Transport.StopAsync(); + + await Assert.ThrowsAnyAsync(() => start); + Assert.All(connections, connection => Assert.Equal(1, connection.DisposeCount)); + Assert.Equal(SocketModeStatus.Stopped, harness.Transport.Status); + } + + [Fact] + public async Task StopAsync_IsIdempotentAndDisposesEachConnectionOnce() + { + Harness harness = new(); + FakeConnection[] connections = await harness.StartReadyAsync(); + + await harness.Transport.StopAsync(); + await harness.Transport.StopAsync(); + await harness.Transport.DisposeAsync(); + + Assert.All(connections, connection => + { + Assert.Equal(1, connection.StopCount); + Assert.Equal(1, connection.DisposeCount); + }); + Assert.Equal(SocketModeStatus.Stopped, harness.Transport.Status); + await Assert.ThrowsAsync(() => harness.Transport.StartAsync()); + } + + [Fact] + public async Task Status_ReportsDisconnectedUntilDroppedGeoRecovers() + { + Harness harness = new(new SocketModeTransportOptions { ReconnectDelays = [TimeSpan.Zero] }); + FakeConnection[] connections = await harness.StartReadyAsync(); + FakeConnection amer = connections.Single(connection => connection.Geo == "amer"); + + amer.Close(new IOException("dropped")); + + Assert.Equal(SocketModeStatus.Disconnected, harness.Transport.Status); + Assert.Equal(SocketModeStatus.Disconnected, harness.Transport.GeoStatuses["amer"]); + Assert.Equal(SocketModeStatus.Ready, harness.Transport.GeoStatuses["emea"]); + + FakeConnection replacement = (await harness.Factory.WaitForAsync(4))[3]; + replacement.Ready(); + await WaitUntilAsync(() => harness.Transport.Status == SocketModeStatus.Ready); + } + + [Fact] + public async Task Dispatch_InvokeReturnsHandlerStatusAndBody() + { + Harness harness = new() { Dispatch = _ => Task.FromResult(new SocketDispatchResult(201, "created")) }; + FakeConnection connection = (await harness.StartReadyAsync())[0]; + + SocketReplyFrame? reply = await connection.ActivityAsync(Envelope("invoke", """{"type":"invoke","id":"a1"}""")); + + Assert.NotNull(reply); + Assert.Equal("env-1", reply.EnvelopeId); + Assert.Equal("bot-id", reply.BotKey); + Assert.Equal(201, reply.Status); + Assert.Equal("created", reply.Body); + Assert.Equal(Now.ToUnixTimeMilliseconds(), reply.ReceivedAtUnixMilliseconds); + Assert.Equal("a1", Assert.Single(harness.Dispatched).Id); + } + + [Fact] + public async Task Dispatch_MessageReturnsAcknowledgementWithHandlerStatus() + { + Harness harness = new() { Dispatch = _ => Task.FromResult(new SocketDispatchResult(202, "ignored")) }; + FakeConnection connection = (await harness.StartReadyAsync())[0]; + + SocketReplyFrame? reply = await connection.ActivityAsync(Envelope("message", """{"type":"message"}""")); + + Assert.NotNull(reply); + Assert.Equal(202, reply.Status); + Assert.Null(reply.Body); + } + + [Fact] + public async Task Dispatch_ClassifiesInvokeFromActivityWhenEnvelopeTypeIsAbsent() + { + Harness harness = new() { Dispatch = _ => Task.FromResult(new SocketDispatchResult(200, "result")) }; + FakeConnection connection = (await harness.StartReadyAsync())[0]; + + SocketReplyFrame? reply = await connection.ActivityAsync(Envelope(null, """{"type":"invoke"}""")); + + Assert.Equal("result", reply?.Body); + } + + [Theory] + [InlineData("invoke", """{"type":"invoke"}""", true)] + [InlineData("message", """{"type":"message"}""", false)] + public async Task Dispatch_HandlerFailureReturns500AndReportsError(string type, string payload, bool hasBody) + { + InvalidOperationException failure = new("handler failed"); + List reported = []; + Harness harness = new() + { + Dispatch = _ => throw failure, + OnError = exception => + { + reported.Add(exception); + return Task.CompletedTask; + }, + }; + FakeConnection connection = (await harness.StartReadyAsync())[0]; + + SocketReplyFrame? reply = await connection.ActivityAsync(Envelope(type, payload)); + + Assert.NotNull(reply); + Assert.Equal(500, reply.Status); + Assert.Equal(hasBody, reply.Body is not null); + Assert.Same(failure, Assert.Single(reported)); + } + + [Fact] + public async Task Dispatch_FailingErrorObserverStillReturns500() + { + Harness harness = new() + { + Dispatch = _ => Task.FromException(new InvalidOperationException("handler failed")), + OnError = _ => throw new InvalidOperationException("observer failed"), + }; + FakeConnection connection = (await harness.StartReadyAsync())[0]; + + SocketReplyFrame? reply = await connection.ActivityAsync(Envelope("invoke", """{"type":"invoke"}""")); + + Assert.Equal(500, reply?.Status); + } + + [Fact] + public async Task Dispatch_RejectsUnsupportedProtocolVersionWithoutDispatching() + { + Harness harness = new(); + FakeConnection connection = (await harness.StartReadyAsync())[0]; + + SocketReplyFrame? reply = await connection.ActivityAsync( + Envelope("message", """{"type":"message"}""", protocolVersion: SocketModeProtocol.CurrentVersion + 1)); + + Assert.Equal(400, reply?.Status); + Assert.Empty(harness.Dispatched); + } + + [Fact] + public async Task Dispatch_DropsEnvelopeWithoutActivity() + { + Harness harness = new(); + FakeConnection connection = (await harness.StartReadyAsync())[0]; + + SocketReplyFrame? reply = await connection.ActivityAsync(Envelope("message", """{"text":"no type"}""")); + + Assert.Null(reply); + Assert.Empty(harness.Dispatched); + } + + [Fact] + public void RetryPolicy_UsesScheduleThenRepeatsLastDelay() + { + IGeoSocketOwner owner = new Harness(new SocketModeTransportOptions + { + ReconnectDelays = [TimeSpan.FromSeconds(1), TimeSpan.FromSeconds(4)], + }).Transport; + + Assert.Equal(TimeSpan.FromSeconds(1), owner.GetBackoffDelay(0)); + Assert.Equal(TimeSpan.FromSeconds(4), owner.GetBackoffDelay(1)); + Assert.Equal(TimeSpan.FromSeconds(4), owner.GetBackoffDelay(9)); + } + + [Fact] + public void RetryPolicy_JittersCappedExponentialBackoff() + { + IGeoSocketOwner owner = new Harness { Random = new FixedRandom(0.5) }.Transport; + + Assert.Equal(TimeSpan.FromSeconds(0.5), owner.GetBackoffDelay(0)); + Assert.Equal(TimeSpan.FromSeconds(2), owner.GetBackoffDelay(2)); + Assert.Equal(TimeSpan.FromSeconds(7.5), owner.GetBackoffDelay(100)); + } + + [Fact] + public void RetryPolicy_ReadsRetryAfterFromNegotiateFailure() + { + IGeoSocketOwner owner = new Harness().Transport; + + Assert.Equal( + TimeSpan.FromSeconds(7), + owner.GetRetryAfter(new SocketModeNegotiateException(HttpStatusCode.TooManyRequests, TimeSpan.FromSeconds(7)))); + Assert.Null(owner.GetRetryAfter(new IOException())); + Assert.Null(owner.GetRetryAfter(null)); + } + + private static SocketActivityEnvelope Envelope(string? type, string payload, int protocolVersion = 1) + => new() + { + ProtocolVersion = protocolVersion, + EnvelopeId = "env-1", + Type = type, + Payload = JsonDocument.Parse(payload).RootElement.Clone(), + }; + + private static string[] SortedUris(FakeConnection[] connections) + => [.. connections.Select(connection => connection.NegotiateUri.AbsoluteUri).Order(StringComparer.Ordinal)]; + + private static async Task WaitUntilAsync(Func condition) + { + using CancellationTokenSource timeout = new(TimeSpan.FromSeconds(5)); + while (!condition()) + { + await Task.Delay(1, timeout.Token); + } + } + + private sealed class Harness + { + private readonly SocketModeTransportOptions _options; + private SocketModeTransport? _transport; + + internal Harness(SocketModeTransportOptions? options = null) + { + _options = options ?? new SocketModeTransportOptions(); + } + + internal FakeConnectionFactory Factory { get; } = new(); + + internal List Dispatched { get; } = []; + + internal Func> Dispatch { get; init; } = + _ => Task.FromResult(new SocketDispatchResult(200)); + + internal Func? OnError { get; init; } + + internal Random? Random { get; init; } + + internal SocketModeTransport Transport => _transport ??= new SocketModeTransport( + _options, + Factory, + activity => + { + lock (Dispatched) + { + Dispatched.Add(activity); + } + + return Dispatch(activity); + }, + NullLogger.Instance, + "bot-id", + OnError, + new FixedTimeProvider(), + Random); + + internal async Task StartReadyAsync() + { + Task start = Transport.StartAsync(); + FakeConnection[] connections = await Factory.WaitForAsync(_options.Geos.Count); + foreach (FakeConnection connection in connections) + { + connection.Ready(); + } + + await start; + return connections; + } + } + + private sealed class FakeConnectionFactory : ISocketConnectionFactory + { + private readonly Dictionary _startFailures = []; + private readonly Dictionary _disposeFailures = []; + private readonly List _connections = []; + + internal FakeConnection[] Connections + { + get + { + lock (_connections) + { + return [.. _connections]; + } + } + } + + internal void FailStart(string geo, Exception error) => _startFailures[geo] = error; + + internal void FailDispose(string geo, Exception error) => _disposeFailures[geo] = error; + + internal async Task WaitForAsync(int count) + { + await WaitUntilAsync(() => Connections.Length >= count); + FakeConnection[] connections = Connections; + await Task.WhenAll(connections.Select(connection => connection.Started.Task)); + return connections; + } + + public ISocketConnection Create(Uri negotiateUri, SocketConnectionHandlers handlers) + { + string[] segments = negotiateUri.AbsolutePath.Split('/', StringSplitOptions.RemoveEmptyEntries); + string geo = segments.Length > 3 ? segments[0] : string.Empty; + FakeConnection connection = new( + negotiateUri, + geo, + handlers, + _startFailures.GetValueOrDefault(geo), + _disposeFailures.GetValueOrDefault(geo)); + lock (_connections) + { + _connections.Add(connection); + } + + return connection; + } + } + + private sealed class FakeConnection( + Uri negotiateUri, + string geo, + SocketConnectionHandlers handlers, + Exception? startError, + Exception? disposeError) : ISocketConnection + { + private readonly TaskCompletionSource _ready = new(TaskCreationOptions.RunContinuationsAsynchronously); + private int _closed; + private int _stopCount; + private int _disposeCount; + + public TimeSpan? TokenLifetime => null; + + internal Uri NegotiateUri { get; } = negotiateUri; + + internal string Geo { get; } = geo; + + internal TaskCompletionSource Started { get; } = new(TaskCreationOptions.RunContinuationsAsynchronously); + + internal int StopCount => Volatile.Read(ref _stopCount); + + internal int DisposeCount => Volatile.Read(ref _disposeCount); + + public async Task StartAsync(CancellationToken cancellationToken) + { + Started.TrySetResult(); + if (startError is not null) + { + throw startError; + } + + await _ready.Task.WaitAsync(cancellationToken); + } + + public Task StopAsync(CancellationToken cancellationToken = default) + { + Interlocked.Increment(ref _stopCount); + _ready.TrySetCanceled(CancellationToken.None); + RaiseClosed(null, planned: true); + return Task.CompletedTask; + } + + public ValueTask DisposeAsync() + { + Interlocked.Increment(ref _disposeCount); + return disposeError is null ? ValueTask.CompletedTask : ValueTask.FromException(disposeError); + } + + internal void Ready() + { + handlers.OnReady(new SocketReadyFrame { ConnectionId = NegotiateUri.AbsolutePath }); + _ready.TrySetResult(); + } + + internal void Fail(Exception error) => _ready.TrySetException(error); + + internal void Close(Exception error) => RaiseClosed(error, planned: false); + + internal Task ActivityAsync(SocketActivityEnvelope envelope) => handlers.OnActivity(envelope); + + private void RaiseClosed(Exception? error, bool planned) + { + if (Interlocked.Exchange(ref _closed, 1) == 0) + { + handlers.OnClosed(error, planned); + } + } + } + + private sealed class FixedTimeProvider : TimeProvider + { + public override DateTimeOffset GetUtcNow() => Now; + } + + private sealed class FixedRandom(double value) : Random + { + public override double NextDouble() => value; + } +}