// Licensed to the .NET Foundation under one or more agreements.
// The .NET Foundation licenses this file to you under the MIT license.
using System.Globalization;
using System.Net.Sockets;
using System.Runtime.CompilerServices;
using Aspire.Cli.Resources;
using Aspire.Cli.Telemetry;
using Aspire.Hosting;
using Microsoft.Extensions.Logging;
using StreamJsonRpc;
namespace Aspire.Cli.Backchannel;
internal interface IAppHostCliBackchannel
{
Task RequestStopAsync(CancellationToken cancellationToken);
Task NotifyAppHostReadyAsync(CancellationToken cancellationToken);
Task<DashboardUrlsState> GetDashboardUrlsAsync(CancellationToken cancellationToken);
IAsyncEnumerable<BackchannelLogEntry> GetAppHostLogEntriesAsync(CancellationToken cancellationToken);
IAsyncEnumerable<RpcResourceState> GetResourceStatesAsync(CancellationToken cancellationToken);
Task WaitForDisconnectAsync(CancellationToken cancellationToken);
Task ConnectAsync(string socketPath, int retryCount, CancellationToken cancellationToken);
Task ConnectAsync(string socketPath, bool autoReconnect, int retryCount, CancellationToken cancellationToken);
IAsyncEnumerable<PublishingActivity> GetPublishingActivitiesAsync(CancellationToken cancellationToken);
Task<string[]> GetCapabilitiesAsync(CancellationToken cancellationToken);
Task CompletePromptResponseAsync(string promptId, PublishingPromptInputAnswer[] answers, CancellationToken cancellationToken);
Task UpdatePromptResponseAsync(string promptId, PublishingPromptInputAnswer[] answers, CancellationToken cancellationToken);
Task<GetPipelineStepsResponse> GetPipelineStepsAsync(string? step, CancellationToken cancellationToken);
Task<UploadFileResponse> UploadFileAsync(string filePath, string fileName, int interactionId, string inputName, CancellationToken cancellationToken);
}
internal sealed class AppHostCliBackchannel(
ILogger<AppHostCliBackchannel> logger,
IEnvironment environment,
AspireCliTelemetry telemetry,
ProfilingTelemetry profilingTelemetry) : IAppHostCliBackchannel
{
private const string BaselineCapability = "baseline.v2";
private TaskCompletionSource<JsonRpc> _rpcTaskCompletionSource = new();
private TaskCompletionSource _disconnectTaskCompletionSource = new(TaskCreationOptions.RunContinuationsAsynchronously);
private string? _socketPath;
private bool _autoReconnect;
private CancellationToken _cancellationToken;
private readonly object _lock = new();
private volatile bool _isReconnecting;
/// <summary>
/// Gets the current RPC task in a thread-safe manner.
/// </summary>
private Task<JsonRpc> GetRpcTaskAsync()
{
lock (_lock)
{
return _rpcTaskCompletionSource.Task;
}
}
public async Task WaitForDisconnectAsync(CancellationToken cancellationToken)
{
Task disconnectTask;
lock (_lock)
{
disconnectTask = _disconnectTaskCompletionSource.Task;
}
await disconnectTask.WaitAsync(cancellationToken).ConfigureAwait(false);
}
public async Task RequestStopAsync(CancellationToken cancellationToken)
{
// This RPC call is required to allow the CLI to trigger a clean shutdown
// of the AppHost process. The AppHost process will then trigger the shutdown
// which will allow the CLI to await the pending run.
using var activity = telemetry.StartDiagnosticActivity();
var rpc = await GetRpcTaskAsync().WaitAsync(cancellationToken).ConfigureAwait(false);
logger.LogDebug("Requesting stop");
await rpc.InvokeWithProfilingAsync(
profilingTelemetry,
"apphost",
"RequestStopAsync",
[],
cancellationToken);
}
public async Task NotifyAppHostReadyAsync(CancellationToken cancellationToken)
{
var rpc = await GetRpcTaskAsync().WaitAsync(cancellationToken).ConfigureAwait(false);
logger.LogDebug("Notifying AppHost startup readiness");
try
{
await rpc.InvokeWithProfilingAsync(
profilingTelemetry,
"apphost",
"NotifyAppHostReadyAsync",
[],
cancellationToken);
}
catch (RemoteMethodNotFoundException ex)
{
logger.LogDebug(ex, "NotifyAppHostReadyAsync RPC method not available on the remote AppHost. The AppHost may be running an older version.");
}
}
public async Task<DashboardUrlsState> GetDashboardUrlsAsync(CancellationToken cancellationToken)
{
using var activity = profilingTelemetry.StartBackchannelGetDashboardUrls();
activity.AddBackchannelWaitForRpcEvent();
var rpc = await GetRpcTaskAsync().WaitAsync(cancellationToken).ConfigureAwait(false);
activity.AddBackchannelRpcReadyEvent();
logger.LogDebug("Requesting dashboard URL");
activity.AddBackchannelGetDashboardUrlsInvokeEvent();
var state = await rpc.InvokeWithProfilingAsync<DashboardUrlsState>(
profilingTelemetry,
"apphost",
"GetDashboardUrlsAsync",
[],
cancellationToken);
activity.SetAppHostDashboardUrls(state);
activity.AddBackchannelGetDashboardUrlsResponseEvent();
return state;
}
public IAsyncEnumerable<BackchannelLogEntry> GetAppHostLogEntriesAsync(CancellationToken cancellationToken)
{
return InvokeStreamingRpcAsync<BackchannelLogEntry>(
(rpc, ct) => rpc.InvokeStreamingWithProfilingAsync<BackchannelLogEntry>(
profilingTelemetry, "apphost", "GetAppHostLogEntriesAsync", [], ct, ProfilingJsonRpcExtensions.StreamingSpanLifetime.FirstItem),
"AppHost log entries",
cancellationToken);
}
public IAsyncEnumerable<RpcResourceState> GetResourceStatesAsync(CancellationToken cancellationToken)
{
return InvokeStreamingRpcAsync<RpcResourceState>(
(rpc, ct) => rpc.InvokeStreamingWithProfilingAsync<RpcResourceState>(
profilingTelemetry, "apphost", "GetResourceStatesAsync", [], ct),
"resource states",
cancellationToken);
}
/// <summary>
/// Invokes a streaming RPC method, handling reconnection when auto-reconnect is enabled.
/// </summary>
private async IAsyncEnumerable<T> InvokeStreamingRpcAsync<T>(
Func<JsonRpc, CancellationToken, Task<IAsyncEnumerable<T>>> startStream,
string operationName,
[EnumeratorCancellation] CancellationToken cancellationToken)
{
while (!cancellationToken.IsCancellationRequested)
{
IAsyncEnumerable<T>? items = null;
try
{
using var activity = telemetry.StartDiagnosticActivity();
var rpc = await GetRpcTaskAsync().WaitAsync(cancellationToken).ConfigureAwait(false);
logger.LogDebug("Requesting {OperationName}", operationName);
items = await startStream(rpc, cancellationToken).ConfigureAwait(false);
logger.LogDebug("Received {OperationName} async enumerable", operationName);
}
catch (Exception ex) when (_autoReconnect && !cancellationToken.IsCancellationRequested && IsConnectionLostException(ex))
{
logger.LogDebug("Connection lost while getting {OperationName}, waiting for reconnect...", operationName);
await WaitForReconnectionAsync(cancellationToken).ConfigureAwait(false);
continue;
}
var reportingEnumerable = new ReportingAsyncEnumerable<T>(items);
await foreach (var item in EnumerateWithReconnect(reportingEnumerable, cancellationToken))
{
yield return item;
}
// If we exit the enumeration loop because of a connection loss, the reporting enumerable will indicate that we should retry.
// If not then the enumerable ended with no more data. We can exit the method.
if (!reportingEnumerable.RetryBecauseConnectionLost)
{
yield break;
}
}
}
private sealed class ReportingAsyncEnumerable<T>(IAsyncEnumerable<T> source) : IAsyncEnumerable<T>
{
public bool RetryBecauseConnectionLost { get; set; }
public IAsyncEnumerator<T> GetAsyncEnumerator(CancellationToken cancellationToken = default)
{
return source.GetAsyncEnumerator(cancellationToken);
}
}
private async IAsyncEnumerable<T> EnumerateWithReconnect<T>(ReportingAsyncEnumerable<T> source, [EnumeratorCancellation] CancellationToken cancellationToken)
{
var enumerator = source.GetAsyncEnumerator(cancellationToken);
try
{
while (true)
{
bool hasNext;
T current;
try
{
hasNext = await enumerator.MoveNextAsync().ConfigureAwait(false);
if (!hasNext)
{
yield break;
}
current = enumerator.Current;
}
catch (Exception ex) when (_autoReconnect && !cancellationToken.IsCancellationRequested && IsConnectionLostException(ex))
{
source.RetryBecauseConnectionLost = true;
logger.LogDebug("Connection lost during enumeration, will restart after reconnect");
yield break; // Exit this enumeration, outer loop will restart
}
yield return current;
}
}
finally
{
// Disposing a dead connection's enumerator may throw - suppress it
try
{
await enumerator.DisposeAsync().ConfigureAwait(false);
}
catch (Exception ex) when (IsConnectionLostException(ex))
{
logger.LogDebug("Ignoring connection lost exception during enumerator disposal");
}
}
}
private static bool IsConnectionLostException(Exception ex)
{
return ex is ConnectionLostException
|| ex is ObjectDisposedException
|| (ex is OperationCanceledException && ex.InnerException is ConnectionLostException);
}
private async Task WaitForReconnectionAsync(CancellationToken cancellationToken)
{
// Wait for the TCS to be reset and then completed again
var startTime = DateTime.UtcNow;
var maxWait = TimeSpan.FromSeconds(60);
// First, wait for the reconnection to start (TCS to be reset)
// This handles the race where we catch the exception before OnDisconnected fires
Task<JsonRpc>? initialTask = null;
while (!cancellationToken.IsCancellationRequested && DateTime.UtcNow - startTime < maxWait)
{
var currentTask = GetRpcTaskAsync();
// If this is a new TCS (different from what we had), reconnection has started
if (initialTask is not null && !ReferenceEquals(currentTask, initialTask))
{
break;
}
// If we haven't captured the initial task yet, do so
initialTask ??= currentTask;
// If the current task is not completed, reconnection has started (TCS was reset)
if (!currentTask.IsCompleted)
{
break;
}
await Task.Delay(100, cancellationToken).ConfigureAwait(false);
}
// Now wait for the reconnection to complete
while (!cancellationToken.IsCancellationRequested && DateTime.UtcNow - startTime < maxWait)
{
var rpcTask = GetRpcTaskAsync();
if (rpcTask.IsCompletedSuccessfully)
{
logger.LogDebug("Reconnection completed successfully");
return;
}
await Task.Delay(500, cancellationToken).ConfigureAwait(false);
}
logger.LogWarning("Timed out waiting for backchannel reconnection");
}
public Task ConnectAsync(string socketPath, int retryCount, CancellationToken cancellationToken)
=> ConnectAsync(socketPath, autoReconnect: false, retryCount: retryCount, cancellationToken);
public async Task ConnectAsync(string socketPath, bool autoReconnect, int retryCount, CancellationToken cancellationToken)
{
try
{
using var activity = profilingTelemetry.StartBackchannelConnect(socketPath, autoReconnect, retryCount);
lock (_lock)
{
if (_rpcTaskCompletionSource.Task.IsCompleted && !_rpcTaskCompletionSource.Task.IsFaulted)
{
throw new InvalidOperationException(ErrorStrings.AlreadyConnectedToBackchannel);
}
}
_socketPath = socketPath;
_autoReconnect = autoReconnect;
_cancellationToken = cancellationToken;
lock (_lock)
{
_disconnectTaskCompletionSource = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously);
}
var connectingLogLevel = retryCount % 10 == 0 ? LogLevel.Debug : LogLevel.Trace;
logger.Log(connectingLogLevel, "Connecting to AppHost backchannel at {SocketPath} (autoReconnect={AutoReconnect}, retryCount={RetryCount})", socketPath, autoReconnect, retryCount);
var socket = new Socket(AddressFamily.Unix, SocketType.Stream, ProtocolType.Unspecified);
var endpoint = new UnixDomainSocketEndPoint(socketPath);
activity.AddBackchannelSocketConnectStartEvent();
await socket.ConnectAsync(endpoint, cancellationToken);
activity.AddBackchannelSocketConnectedEvent();
logger.LogDebug("Connected to AppHost backchannel at {SocketPath} (retryCount={RetryCount})", socketPath, retryCount);
var stream = new NetworkStream(socket, true);
JsonRpc? rpc = null;
try
{
rpc = new JsonRpc(new HeaderDelimitedMessageHandler(stream, stream, BackchannelJsonSerializerContext.CreateRpcMessageFormatter()))
{
ActivityTracingStrategy = new ActivityTracingStrategy()
};
rpc.StartListening();
activity.AddBackchannelRpcListeningEvent();
activity.AddBackchannelGetCapabilitiesStartEvent();
var capabilities = await rpc.InvokeWithProfilingAsync<string[]>(
profilingTelemetry,
"apphost",
"GetCapabilitiesAsync",
[],
cancellationToken);
activity.SetBackchannelCapabilitySummary(capabilities, BaselineCapability);
activity.AddBackchannelGetCapabilitiesResponseEvent();
if (!capabilities.Any(s => s == BaselineCapability))
{
throw new AppHostIncompatibleException(
string.Format(CultureInfo.CurrentCulture, ErrorStrings.AppHostIncompatibleWithCli, BaselineCapability),
BaselineCapability
);
}
rpc.Disconnected += OnRpcDisconnected;
// Set up auto-reconnect if enabled
if (autoReconnect)
{
rpc.Disconnected += OnDisconnected;
}
}
catch
{
rpc?.Dispose();
throw;
}
lock (_lock)
{
_rpcTaskCompletionSource.SetResult(rpc);
}
}
catch (RemoteMethodNotFoundException ex)
{
logger.LogError(ex, "Failed to connect to AppHost backchannel. The AppHost must be updated to a version that supports the {BaselineCapability} capability.", BaselineCapability);
throw new AppHostIncompatibleException(
string.Format(CultureInfo.CurrentCulture, ErrorStrings.AppHostIncompatibleWithCli, BaselineCapability),
BaselineCapability
);
}
}
private void OnDisconnected(object? sender, JsonRpcDisconnectedEventArgs args)
{
// Prevent concurrent reconnection attempts
lock (_lock)
{
if (_isReconnecting)
{
logger.LogDebug("Backchannel disconnected but reconnection already in progress, ignoring.");
return;
}
_isReconnecting = true;
}
logger.LogInformation("Backchannel disconnected: {Reason}. Attempting to reconnect...", args.Reason);
_ = Task.Run(async () =>
{
try
{
await ReconnectInternalAsync().ConfigureAwait(false);
}
catch (Exception ex)
{
logger.LogWarning(ex, "Failed to reconnect backchannel");
}
finally
{
lock (_lock)
{
_isReconnecting = false;
}
}
});
}
private void OnRpcDisconnected(object? sender, JsonRpcDisconnectedEventArgs args)
{
logger.LogDebug("Backchannel disconnected: {Reason}", args.Reason);
lock (_lock)
{
_disconnectTaskCompletionSource.TrySetResult();
}
}
private void ResetForReconnection()
{
lock (_lock)
{
logger.LogDebug("Resetting backchannel for reconnection");
_rpcTaskCompletionSource = new TaskCompletionSource<JsonRpc>();
}
}
private async Task ReconnectInternalAsync()
{
if (_socketPath is null)
{
throw new InvalidOperationException("Cannot reconnect: no previous connection.");
}
ResetForReconnection();
// Wait for the new socket to appear (the new DistributedApplication needs to start)
var startTime = DateTime.UtcNow;
var maxWait = TimeSpan.FromSeconds(30);
var retryCount = 0;
while (!_cancellationToken.IsCancellationRequested)
{
try
{
await ConnectAsync(_socketPath, _autoReconnect, retryCount, _cancellationToken).ConfigureAwait(false);
logger.LogInformation("Successfully reconnected to backchannel");
return;
}
catch (SocketException) when (DateTime.UtcNow - startTime < maxWait)
{
retryCount++;
// Socket not ready yet, wait and retry
await Task.Delay(500, _cancellationToken).ConfigureAwait(false);
}
catch (SocketException)
{
// Timeout exceeded — fall through to warning
break;
}
}
logger.LogWarning("Timed out waiting for backchannel reconnection");
}
public IAsyncEnumerable<PublishingActivity> GetPublishingActivitiesAsync(CancellationToken cancellationToken)
{
return InvokeStreamingRpcAsync<PublishingActivity>(
(rpc, ct) => rpc.InvokeStreamingWithProfilingAsync<PublishingActivity>(
profilingTelemetry, "apphost", "GetPublishingActivitiesAsync", [], ct),
"publishing activities",
cancellationToken);
}
public async Task<string[]> GetCapabilitiesAsync(CancellationToken cancellationToken)
{
using var activity = telemetry.StartDiagnosticActivity();
var rpc = await GetRpcTaskAsync().WaitAsync(cancellationToken).ConfigureAwait(false);
logger.LogDebug("Requesting capabilities");
var capabilities = await rpc.InvokeWithProfilingAsync<string[]>(
profilingTelemetry,
"apphost",
"GetCapabilitiesAsync",
[],
cancellationToken).ConfigureAwait(false);
return capabilities;
}
public async Task CompletePromptResponseAsync(string promptId, PublishingPromptInputAnswer[] answers, CancellationToken cancellationToken)
{
using var activity = telemetry.StartDiagnosticActivity();
var rpc = await GetRpcTaskAsync().WaitAsync(cancellationToken).ConfigureAwait(false);
logger.LogDebug("Providing prompt responses for prompt ID {PromptId}", promptId);
await rpc.InvokeWithProfilingAsync(
profilingTelemetry,
"apphost",
"CompletePromptResponseAsync",
[promptId, answers],
cancellationToken).ConfigureAwait(false);
}
public async Task UpdatePromptResponseAsync(string promptId, PublishingPromptInputAnswer[] answers, CancellationToken cancellationToken)
{
using var activity = telemetry.StartDiagnosticActivity();
var rpc = await GetRpcTaskAsync().WaitAsync(cancellationToken).ConfigureAwait(false);
logger.LogDebug("Providing prompt responses for prompt ID {PromptId}", promptId);
await rpc.InvokeWithProfilingAsync(
profilingTelemetry,
"apphost",
"UpdatePromptResponseAsync",
[promptId, answers],
cancellationToken).ConfigureAwait(false);
}
public async Task<GetPipelineStepsResponse> GetPipelineStepsAsync(string? step, CancellationToken cancellationToken)
{
using var activity = telemetry.StartDiagnosticActivity();
var rpc = await GetRpcTaskAsync().WaitAsync(cancellationToken).ConfigureAwait(false);
logger.LogDebug("Requesting pipeline steps.");
var response = await rpc.InvokeWithProfilingAsync<GetPipelineStepsResponse>(
profilingTelemetry,
"apphost",
"GetPipelineStepsAsync",
[new GetPipelineStepsRequest { Step = step }],
cancellationToken).ConfigureAwait(false);
logger.LogDebug("Received {StepCount} pipeline steps.", response.Steps.Length);
return response;
}
public async Task<UploadFileResponse> UploadFileAsync(string filePath, string fileName, int interactionId, string inputName, CancellationToken cancellationToken)
{
using var activity = telemetry.StartDiagnosticActivity();
logger.LogDebug("Uploading file {FileName} from {FilePath}", fileName, filePath);
// Enforce the server-side upload limit on the client before reading the file into memory,
// preventing unbounded memory allocation for very large files. The server also checks this
// limit, so this is a client-side guard to avoid OOM before the rejection arrives.
var maxUploadFileSize = long.TryParse(environment.GetEnvironmentVariable(KnownConfigNames.MaxFileUploadSize), out var parsed)
? parsed
: FileUploadHelpers.DefaultMaxFileUploadSize;
var fileInfo = new FileInfo(filePath);
if (fileInfo.Length > maxUploadFileSize)
{
throw new InvalidOperationException(
$"File '{fileName}' ({fileInfo.Length} bytes) exceeds the maximum upload size of {maxUploadFileSize} bytes. " +
$"To increase the limit, set the {KnownConfigNames.MaxFileUploadSize} environment variable.");
}
var rpc = await GetRpcTaskAsync().WaitAsync(cancellationToken).ConfigureAwait(false);
// Known limitation: the entire file is loaded into memory because StreamJsonRpc does not
// support streaming byte payloads. The server-side upload limit (default 100 MB) bounds
// worst-case memory usage. The Dashboard path uses gRPC streaming and avoids this.
var data = await File.ReadAllBytesAsync(filePath, cancellationToken).ConfigureAwait(false);
var response = await rpc.InvokeWithProfilingAsync<UploadFileResponse>(
profilingTelemetry,
"apphost",
"UploadFileAsync",
[new UploadFileRequest { Data = data, FileName = fileName, InteractionId = interactionId, InputName = inputName }],
cancellationToken).ConfigureAwait(false);
logger.LogDebug("File uploaded with ID {FileId}", response.FileId);
return response;
}
}