Skip to content
Open
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
77 changes: 77 additions & 0 deletions CosmosDBShell.Tests/McpLocationSubscriptionTests.cs
Original file line number Diff line number Diff line change
@@ -0,0 +1,77 @@
// ------------------------------------------------------------
// Copyright (c) Microsoft Corporation. All rights reserved.
// ------------------------------------------------------------

namespace CosmosShell.Tests;

using System.Net;
using System.Net.Sockets;
using System.Text.Json;
using Azure.Data.Cosmos.Shell.Core;
using Azure.Data.Cosmos.Shell.Mcp;
using Azure.Data.Cosmos.Shell.States;
using ModelContextProtocol;
using ModelContextProtocol.Client;

[Collection(CosmosShell.Tests.Shell.ThemeStateTestCollection.Name)]
public class McpLocationSubscriptionTests
{
[Fact]
public async Task SubscribedClient_ReceivesInteractiveLocationChange()
{
using var timeout = CancellationTokenSource.CreateLinkedTokenSource(TestContext.Current.CancellationToken);
timeout.CancelAfter(TimeSpan.FromSeconds(10));
using var listener = new TcpListener(IPAddress.Loopback, 0);
listener.Start();
var port = ((IPEndPoint)listener.LocalEndpoint).Port;
listener.Stop();

using var host = McpServer.CreateHost(new Program.CosmosShellOptions { McpPort = port });
await host.StartAsync(timeout.Token);
try
{
var transport = new HttpClientTransport(new HttpClientTransportOptions
{
Endpoint = new Uri($"http://127.0.0.1:{port}/"),
});
await using var client = await McpClient.CreateAsync(transport, cancellationToken: timeout.Token);
var resources = await client.ListResourcesAsync(cancellationToken: timeout.Token);
Assert.Contains(resources, resource => resource.Uri == ResourceOperations.CurrentLocationUri);

var invalid = await Assert.ThrowsAsync<McpProtocolException>(
() => client.SubscribeToResourceAsync("cosmos://docs/scripting", cancellationToken: timeout.Token));
Assert.Equal(McpErrorCode.InvalidParams, invalid.ErrorCode);
Assert.Contains(ResourceOperations.CurrentLocationUri, invalid.Message);

var updated = new TaskCompletionSource<string>(TaskCreationOptions.RunContinuationsAsynchronously);
await using var subscription = await client.SubscribeToResourceAsync(
ResourceOperations.CurrentLocationUri,
(notification, _) =>
{
updated.TrySetResult(notification.Uri);
return ValueTask.CompletedTask;
},
cancellationToken: timeout.Token);

var originalState = ShellInterpreter.Instance.State;
try
{
ShellInterpreter.Instance.State = new DatabaseState("McpNotificationTest", null!);
Assert.Equal(ResourceOperations.CurrentLocationUri, await updated.Task.WaitAsync(timeout.Token));

var resource = await client.ReadResourceAsync(ResourceOperations.CurrentLocationUri, cancellationToken: timeout.Token);
var content = Assert.Single(resource.Contents);
using var json = JsonDocument.Parse(Assert.IsType<ModelContextProtocol.Protocol.TextResourceContents>(content).Text);
Assert.Equal("/McpNotificationTest", json.RootElement.GetProperty("currentLocation").GetString());
}
finally
{
ShellInterpreter.Instance.State = originalState;
}
}
finally
{
await host.StopAsync(TestContext.Current.CancellationToken);
}
}
}
22 changes: 22 additions & 0 deletions CosmosDBShell.Tests/ResourceOperationsTests.cs
Original file line number Diff line number Diff line change
Expand Up @@ -4,10 +4,32 @@

namespace CosmosShell.Tests;

using System.Text.Json;
using Azure.Data.Cosmos.Shell.Mcp;
using Azure.Data.Cosmos.Shell.States;

public class ResourceOperationsTests
{
[Theory]
[InlineData(null, null, null)]
[InlineData("", null, "/")]
[InlineData("db", null, "/db")]
[InlineData("db", "container", "/db/container")]
public void GetCurrentLocation_ReturnsSharedLocationAsJson(string? database, string? container, string? expected)
{
var state = database is null
? (State)new DisconnectedState()
: database.Length == 0
? new ConnectedState(null!)
: container is null
? new DatabaseState(database, null!)
: new ContainerState(container, database, null!);

using var document = JsonDocument.Parse(ResourceOperations.GetCurrentLocation(state));
var location = document.RootElement.GetProperty("currentLocation");
Assert.Equal(expected, location.ValueKind == JsonValueKind.Null ? null : location.GetString());
}

[Fact]
public void GetScriptingGuide_ReturnsEmbeddedProgrammingMarkdown()
{
Expand Down
79 changes: 79 additions & 0 deletions CosmosDBShell.Tests/ShellLocationChangedTests.cs
Original file line number Diff line number Diff line change
@@ -0,0 +1,79 @@
// ------------------------------------------------------------
// Copyright (c) Microsoft Corporation. All rights reserved.
// ------------------------------------------------------------

namespace CosmosShell.Tests;

using System.Runtime.CompilerServices;
using Azure.Data.Cosmos.Shell.Core;
using Azure.Data.Cosmos.Shell.States;
using Microsoft.Azure.Cosmos;

public class ShellLocationChangedTests
{
[Fact]
public void State_NotifiesOnlyWhenLocationChanges()
{
using var shell = new ShellInterpreter();
var changes = 0;
shell.LocationChanged += () => changes++;

shell.State = new DisconnectedState();
Assert.Equal(0, changes);

shell.State = new ConnectedState(null!);
Assert.Equal(1, changes);

shell.State = new DatabaseState("db", null!);
Assert.Equal(2, changes);

shell.State = new DatabaseState("db", null!);
Assert.Equal(2, changes);

shell.State = new ContainerState("container", "db", null!);
Assert.Equal(3, changes);

shell.State = new DisconnectedState();
Assert.Equal(4, changes);
}

[Fact]
public void State_NotifiesWhenClientChangesAtSameLocation()
{
using var shell = new ShellInterpreter();
using var firstClient = CreateTestClient();
using var secondClient = CreateTestClient();
shell.State = new DatabaseState("db", firstClient);
var changes = 0;
shell.LocationChanged += () => changes++;

shell.State = new DatabaseState("db", secondClient);
Assert.Equal(1, changes);

shell.State = new DisconnectedState();
}

[Fact]
public void State_DoesNotNotifyWhenOnlyArmContextChanges()
{
using var shell = new ShellInterpreter();
using var client = CreateTestClient();
shell.State = new ConnectedState(client);
var changes = 0;
shell.LocationChanged += () => changes++;

var armContext = (ArmCosmosContext)RuntimeHelpers.GetUninitializedObject(typeof(ArmCosmosContext));
shell.State = new ConnectedState(client, armContext);
Assert.Equal(0, changes);

shell.State = new DisconnectedState();
}

private static CosmosClient CreateTestClient()
{
return new CosmosClient(
"https://localhost:8081",
Convert.ToBase64String(new byte[64]),
new CosmosClientOptions { ConnectionMode = ConnectionMode.Gateway });
}
}
14 changes: 11 additions & 3 deletions CosmosDBShell.Tests/ToolOperationsCallToolTests.cs
Original file line number Diff line number Diff line change
Expand Up @@ -28,11 +28,19 @@ namespace CosmosShell.Tests;
// success path (which writes the highlighted command line through AnsiConsole) does
// not race with other tests that swap the global console or theme.
[Collection(CosmosShell.Tests.Shell.ThemeStateTestCollection.Name)]
public class ToolOperationsCallToolTests
public class ToolOperationsCallToolTests : IDisposable
{
private static ToolOperations CreateToolOperations()
private readonly LocationResourceSubscriptions locationSubscriptions =
new(NullLogger<LocationResourceSubscriptions>.Instance);

public void Dispose()
{
this.locationSubscriptions.Dispose();
}

private ToolOperations CreateToolOperations()
{
return new ToolOperations(NullLogger<ToolOperations>.Instance);
return new ToolOperations(NullLogger<ToolOperations>.Instance, this.locationSubscriptions);
}

private static RequestContext<CallToolRequestParams> CallContext(string? name, Dictionary<string, JsonElement>? arguments = null)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -133,6 +133,8 @@ internal ShellInterpreter(string? configPath = null)
this.editorCancelTokenSource = new CancellationTokenSource();
}

internal event Action? LocationChanged;

/// <summary>
/// Gets the line editor instance used by the shell, or <c>null</c> if not available.
/// </summary>
Expand Down Expand Up @@ -288,8 +290,15 @@ internal State State
get;
set
{
var oldState = field;
field = value;
Interlocked.Increment(ref this.stateVersion);
if (oldState != null
&& (ShellLocation.GetCurrentLocation(oldState) != ShellLocation.GetCurrentLocation(value)
|| (oldState as ConnectedState)?.Client != (value as ConnectedState)?.Client))
Comment thread
mkrueger marked this conversation as resolved.
{
this.LocationChanged?.Invoke();
}
}
}

Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,119 @@
// ------------------------------------------------------------
// Copyright (c) Microsoft Corporation. All rights reserved.
// ------------------------------------------------------------

namespace Azure.Data.Cosmos.Shell.Mcp;

using System.Threading.Channels;
using Azure.Data.Cosmos.Shell.Core;
using Microsoft.Extensions.Hosting;
using Microsoft.Extensions.Logging;
using ModelContextProtocol;
using ModelContextProtocol.Protocol;
using ModelContextProtocol.Server;

internal sealed class LocationResourceSubscriptions : BackgroundService
{
private readonly object sync = new();

private readonly List<WeakReference<ModelContextProtocol.Server.McpServer>> subscribers = [];

// Notifications carry only the URI, so pending changes are coalesced into one.
private readonly Channel<bool> changes = Channel.CreateBounded<bool>(
new BoundedChannelOptions(1) { FullMode = BoundedChannelFullMode.DropWrite, SingleReader = true });

private readonly ILogger<LocationResourceSubscriptions> logger;

public LocationResourceSubscriptions(ILogger<LocationResourceSubscriptions> logger)
{
this.logger = logger;
ShellInterpreter.Instance.LocationChanged += this.OnLocationChanged;
}

public void Subscribe(ModelContextProtocol.Server.McpServer server, string uri)
{
ValidateUri(uri);
lock (this.sync)
{
this.PruneSubscribers();
if (!this.subscribers.Any(reference => reference.TryGetTarget(out var target) && ReferenceEquals(target, server)))
{
this.subscribers.Add(new WeakReference<ModelContextProtocol.Server.McpServer>(server));
}
}
}

public void Unsubscribe(ModelContextProtocol.Server.McpServer server, string uri)
{
ValidateUri(uri);
lock (this.sync)
{
this.subscribers.RemoveAll(reference => !reference.TryGetTarget(out var target) || ReferenceEquals(target, server));
}
}

protected override async Task ExecuteAsync(CancellationToken stoppingToken)
{
await foreach (var change in this.changes.Reader.ReadAllAsync(stoppingToken))
{
ModelContextProtocol.Server.McpServer[] servers;
lock (this.sync)
{
this.PruneSubscribers();
servers = this.subscribers
.Select(reference => reference.TryGetTarget(out var server) ? server : null)
.OfType<ModelContextProtocol.Server.McpServer>()
.ToArray();
}

foreach (var server in servers)
{
try
{
await server.SendNotificationAsync(
NotificationMethods.ResourceUpdatedNotification,
new ResourceUpdatedNotificationParams { Uri = ResourceOperations.CurrentLocationUri },
cancellationToken: stoppingToken);
}
catch (OperationCanceledException) when (stoppingToken.IsCancellationRequested)
{
return;
}
catch (Exception ex) when (!stoppingToken.IsCancellationRequested)
{
this.logger.LogWarning(ex, "Could not notify an MCP client about the shell location change.");
lock (this.sync)
{
this.subscribers.RemoveAll(reference => !reference.TryGetTarget(out var target) || ReferenceEquals(target, server));
}
}
}
}
}

public override void Dispose()
{
ShellInterpreter.Instance.LocationChanged -= this.OnLocationChanged;
base.Dispose();
}

internal static void ValidateUri(string uri)
{
if (!string.Equals(uri, ResourceOperations.CurrentLocationUri, StringComparison.Ordinal))
{
throw new McpProtocolException(
$"Resource '{uri}' does not support subscriptions. Only '{ResourceOperations.CurrentLocationUri}' can be subscribed to.",
McpErrorCode.InvalidParams);
}
}

private void OnLocationChanged()
{
this.changes.Writer.TryWrite(true);
}

private void PruneSubscribers()
{
this.subscribers.RemoveAll(reference => !reference.TryGetTarget(out _));
}
}
6 changes: 5 additions & 1 deletion CosmosDBShell/Azure.Data.Cosmos.Shell.Mcp/McpServer.cs
Original file line number Diff line number Diff line change
Expand Up @@ -50,6 +50,8 @@ public static IHost CreateHost(CosmosShellOptions serverArguments)
private static void ConfigureMcpServer(IServiceCollection services)
{
services.AddSingleton<ToolOperations>();
services.AddSingleton<LocationResourceSubscriptions>();
services.AddHostedService(services => services.GetRequiredService<LocationResourceSubscriptions>());
services.AddOptions<McpServerOptions>()
.Configure<ToolOperations>((mcpServerOptions, toolOperations) =>
{
Expand All @@ -66,13 +68,15 @@ private static void ConfigureMcpServer(IServiceCollection services)
mcpServerOptions.Capabilities = new ServerCapabilities
{
Tools = new ToolsCapability(),
Resources = new ResourcesCapability(),
Resources = new ResourcesCapability { Subscribe = true },
};

mcpServerOptions.Handlers = new McpServerHandlers
{
CallToolHandler = toolOperations.CallToolHandler,
ListToolsHandler = toolOperations.ListToolsHandler,
SubscribeToResourcesHandler = toolOperations.SubscribeToResourcesHandler,
UnsubscribeFromResourcesHandler = toolOperations.UnsubscribeFromResourcesHandler,
};

mcpServerOptions.ServerInstructions = LoadServerInstructions();
Expand Down
Loading
Loading