Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
34 changes: 34 additions & 0 deletions src/McpAggregator.Core/Services/ConnectionManager.cs
Original file line number Diff line number Diff line change
Expand Up @@ -96,6 +96,40 @@ public async Task DisconnectAsync(string serverName)
}
}

public async Task<T> ExecuteWithRetryAsync<T>(
string serverName,
Func<McpClient, CancellationToken, Task<T>> operation,
CancellationToken ct = default)
{
var client = await GetClientAsync(serverName, ct);
try
{
return await operation(client, ct);
}
catch (Exception ex) when (ShouldRetry(ex))
{
_logger.LogWarning(ex, "Connection to '{Server}' appears broken, reconnecting", serverName);
await DisconnectAsync(serverName);
client = await GetClientAsync(serverName, ct);
return await operation(client, ct);
}
}

private static bool ShouldRetry(Exception ex)
{
if (ex is OperationCanceledException or AggregatorException)
return false;

return ex is System.IO.IOException
or System.Net.Http.HttpRequestException
or System.Net.Sockets.SocketException
or ObjectDisposedException
|| ex.InnerException is System.IO.IOException
or System.Net.Http.HttpRequestException
or System.Net.Sockets.SocketException
or ObjectDisposedException;
}

public async Task CleanupIdleConnectionsAsync(CancellationToken ct = default)
{
var cutoff = DateTimeOffset.UtcNow - _options.ConnectionIdleTimeout;
Expand Down
5 changes: 3 additions & 2 deletions src/McpAggregator.Core/Services/ToolIndex.cs
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
using McpAggregator.Core.Models;
using Microsoft.Extensions.Logging;
using Microsoft.Extensions.Options;
using ModelContextProtocol.Client;

namespace McpAggregator.Core.Services;

Expand Down Expand Up @@ -123,8 +124,8 @@ public async Task<List<ToolDetail>> GetToolsForServerAsync(string serverName, Ca
return cached.Tools;
}

var client = await _connectionManager.GetClientAsync(serverName, ct);
var mcpTools = await client.ListToolsAsync(cancellationToken: ct);
var mcpTools = await _connectionManager.ExecuteWithRetryAsync<IList<McpClientTool>>(serverName,
async (client, token) => await client.ListToolsAsync(cancellationToken: token), ct);

var tools = mcpTools.Select(t => new ToolDetail
{
Expand Down
5 changes: 3 additions & 2 deletions src/McpAggregator.Core/Tools/AdminTools.cs
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@
using System.Text.Json;
using McpAggregator.Core.Models;
using McpAggregator.Core.Services;
using ModelContextProtocol.Client;
using ModelContextProtocol.Server;

namespace McpAggregator.Core.Tools;
Expand Down Expand Up @@ -136,8 +137,8 @@ public static async Task<string> UpdateSkill(

try
{
var client = await connectionManager.GetClientAsync(server.Name, ct);
var mcpTools = await client.ListToolsAsync(cancellationToken: ct);
var mcpTools = await connectionManager.ExecuteWithRetryAsync<IList<McpClientTool>>(server.Name,
async (client, token) => await client.ListToolsAsync(cancellationToken: token), ct);

var toolSummaries = mcpTools.Select(t => new ToolSummary
{
Expand Down
5 changes: 2 additions & 3 deletions src/McpAggregator.Core/Tools/ToolProxyHandler.cs
Original file line number Diff line number Diff line change
Expand Up @@ -30,8 +30,6 @@ public async Task<string> InvokeAsync(
string? argumentsJson,
CancellationToken ct = default)
{
var client = await _connectionManager.GetClientAsync(serverName, ct);

IReadOnlyDictionary<string, object?>? args = null;
if (!string.IsNullOrWhiteSpace(argumentsJson))
{
Expand All @@ -45,7 +43,8 @@ public async Task<string> InvokeAsync(

try
{
var result = await client.CallToolAsync(toolName, args, cancellationToken: cts.Token);
var result = await _connectionManager.ExecuteWithRetryAsync<CallToolResult>(serverName,
async (client, token) => await client.CallToolAsync(toolName, args, cancellationToken: token), cts.Token);

var textContent = result.Content
.OfType<TextContentBlock>()
Expand Down
9 changes: 5 additions & 4 deletions src/McpAggregator.HttpServer/Controllers/AdminController.cs
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
using McpAggregator.Core.Models;
using McpAggregator.Core.Services;
using Microsoft.AspNetCore.Mvc;
using ModelContextProtocol.Client;

namespace McpAggregator.HttpServer.Controllers;

Expand Down Expand Up @@ -44,8 +45,8 @@ public async Task<IActionResult> RegisterServer([FromBody] RegisterServerRequest
{
try
{
var client = await _connectionManager.GetClientAsync(server.Name, ct);
var mcpTools = await client.ListToolsAsync(cancellationToken: ct);
var mcpTools = await _connectionManager.ExecuteWithRetryAsync<IList<McpClientTool>>(server.Name,
async (client, token) => await client.ListToolsAsync(cancellationToken: token), ct);

var toolSummaries = mcpTools.Select(t => new ToolSummary
{
Expand Down Expand Up @@ -80,8 +81,8 @@ public async Task<IActionResult> RegenerateSummary(string name, CancellationToke

try
{
var client = await _connectionManager.GetClientAsync(server.Name, ct);
var mcpTools = await client.ListToolsAsync(cancellationToken: ct);
var mcpTools = await _connectionManager.ExecuteWithRetryAsync<IList<McpClientTool>>(server.Name,
async (client, token) => await client.ListToolsAsync(cancellationToken: token), ct);

var toolSummaries = mcpTools.Select(t => new ToolSummary
{
Expand Down
140 changes: 140 additions & 0 deletions test/McpAggregator.Core.Tests/Services/ExecuteWithRetryTests.cs
Original file line number Diff line number Diff line change
@@ -0,0 +1,140 @@
using System.IO;
using System.Net.Sockets;
using McpAggregator.Core.Configuration;
using McpAggregator.Core.Exceptions;
using McpAggregator.Core.Models;
using McpAggregator.Core.Services;
using McpAggregator.Core.Storage;
using McpAggregator.Core.Tests.Helpers;
using Microsoft.Extensions.Logging.Abstractions;
using Rocks;

namespace McpAggregator.Core.Tests.Services;

[TestClass]
public class ExecuteWithRetryTests
{
private static (ConnectionManager Manager, ServerRegistry Registry) CreateManager(
params RegisteredServer[] servers)
{
var expectations = new IRegistryPersistenceCreateExpectations();
expectations.Setups.LoadAsync(Arg.Any<CancellationToken>())
.ReturnValue(Task.FromResult(new RegistryData { Servers = [.. servers] }));
expectations.Setups.SaveAsync(Arg.Any<RegistryData>(), Arg.Any<CancellationToken>())
.ReturnValue(Task.CompletedTask);
var persistence = expectations.Instance();

var registry = new ServerRegistry(
persistence,
TestHelpers.OptionsOf(new AggregatorOptions()),
TestHelpers.NullLoggerOf<ServerRegistry>());

var manager = new ConnectionManager(
registry,
TestHelpers.OptionsOf(new AggregatorOptions()),
NullLoggerFactory.Instance,
TestHelpers.NullLoggerOf<ConnectionManager>());

return (manager, registry);
}

[TestMethod]
public async Task ExecuteWithRetryAsync_SucceedsOnFirstTry_DoesNotRetry()
{
// We can't easily mock McpClient (concrete class), so we test ShouldRetry
// indirectly by verifying non-retryable exceptions propagate.
var (manager, registry) = CreateManager(TestHelpers.StdioServer("srv"));
await registry.EnsureLoadedAsync();

// GetClientAsync will fail with ServerUnavailableException (can't actually connect in tests)
// but that's wrapped by ConnectAsync — verify it propagates without retry
await Assert.ThrowsExceptionAsync<ServerUnavailableException>(

Check warning on line 51 in test/McpAggregator.Core.Tests/Services/ExecuteWithRetryTests.cs

View workflow job for this annotation

GitHub Actions / build-and-test

Use 'Assert.ThrowsExactly' instead of 'Assert.ThrowsException' (https://learn.microsoft.com/dotnet/core/testing/mstest-analyzers/mstest0039)
() => manager.ExecuteWithRetryAsync<string>("srv",
(client, ct) => Task.FromResult("ok")));
}

[TestMethod]
public void ShouldRetry_ReturnsFalse_ForOperationCanceledException()
{
Assert.IsFalse(InvokeShouldRetry(new OperationCanceledException()));
}

[TestMethod]
public void ShouldRetry_ReturnsFalse_ForAggregatorException()
{
Assert.IsFalse(InvokeShouldRetry(new AggregatorException("test")));
}

[TestMethod]
public void ShouldRetry_ReturnsFalse_ForServerNotFoundException()
{
Assert.IsFalse(InvokeShouldRetry(new ServerNotFoundException("srv")));
}

[TestMethod]
public void ShouldRetry_ReturnsFalse_ForToolExecutionException()
{
Assert.IsFalse(InvokeShouldRetry(new ToolExecutionException("srv", "tool", "error")));
}

[TestMethod]
public void ShouldRetry_ReturnsTrue_ForIOException()
{
Assert.IsTrue(InvokeShouldRetry(new IOException("pipe broken")));
}

[TestMethod]
public void ShouldRetry_ReturnsTrue_ForSocketException()
{
Assert.IsTrue(InvokeShouldRetry(new SocketException()));
}

[TestMethod]
public void ShouldRetry_ReturnsTrue_ForObjectDisposedException()
{
Assert.IsTrue(InvokeShouldRetry(new ObjectDisposedException("client")));
}

[TestMethod]
public void ShouldRetry_ReturnsTrue_ForHttpRequestException()
{
Assert.IsTrue(InvokeShouldRetry(new HttpRequestException("connection refused")));
}

[TestMethod]
public void ShouldRetry_ReturnsTrue_ForWrappedIOException()
{
var ex = new InvalidOperationException("outer", new IOException("inner"));
Assert.IsTrue(InvokeShouldRetry(ex));
}

[TestMethod]
public void ShouldRetry_ReturnsTrue_ForWrappedSocketException()
{
var ex = new InvalidOperationException("outer", new SocketException());
Assert.IsTrue(InvokeShouldRetry(ex));
}

[TestMethod]
public void ShouldRetry_ReturnsFalse_ForGenericException()
{
Assert.IsFalse(InvokeShouldRetry(new InvalidOperationException("generic")));
}

[TestMethod]
public void ShouldRetry_ReturnsFalse_ForArgumentException()
{
Assert.IsFalse(InvokeShouldRetry(new ArgumentException("bad arg")));
}

/// <summary>
/// Invoke the private static ShouldRetry method via reflection.
/// </summary>
private static bool InvokeShouldRetry(Exception ex)
{
var method = typeof(ConnectionManager).GetMethod("ShouldRetry",
System.Reflection.BindingFlags.NonPublic | System.Reflection.BindingFlags.Static)
?? throw new InvalidOperationException("ShouldRetry method not found");
return (bool)method.Invoke(null, [ex])!;
}
}
Loading