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
7 changes: 7 additions & 0 deletions .github/copilot-instructions.md
Original file line number Diff line number Diff line change
Expand Up @@ -63,3 +63,10 @@
- Fix root causes instead of patching symptoms.
- Preserve public behavior unless the task explicitly changes the CLI or output contract.
- If behavior changes, update tests and docs in the same change.

## Commits

- When creating Git commits, do not add Copilot or AI `Co-authored-by` trailers.
- Ensure commit messages are clear, concise, and follow the project's existing style.
- Use the imperative mood in commit messages (e.g., "Add feature" instead of "Added feature").
- Reference relevant issues or pull requests in the commit message when applicable.
69 changes: 69 additions & 0 deletions CosmosDBShell.Tests/McpConfirmationTests.cs
Original file line number Diff line number Diff line change
@@ -0,0 +1,69 @@
// ------------------------------------------------------------
// Copyright (c) Microsoft Corporation. All rights reserved.
// ------------------------------------------------------------

namespace CosmosShell.Tests;

using System.Text.Json;
using Azure.Data.Cosmos.Shell.Mcp;
using Microsoft.AspNetCore.Hosting.Server;
using Microsoft.AspNetCore.Hosting.Server.Features;
using Microsoft.Extensions.DependencyInjection;
using Microsoft.Extensions.Hosting;
using ModelContextProtocol.Client;
using ModelContextProtocol.Protocol;

[Collection(CosmosShell.Tests.Shell.ThemeStateTestCollection.Name)]
public class McpConfirmationTests
{
[Theory]
[InlineData(null, "2026-07-28")] // Stateless request; confirmation uses native multi-round-trip requests.
[InlineData("2025-11-25", "2025-11-25")] // Initialize handshake; confirmation is sent over the session.
public async Task DestructiveCommand_OverHttp_AsksClientAndHonorsDecline(string? requestedVersion, string negotiatedVersion)
{
using var timeout = CancellationTokenSource.CreateLinkedTokenSource(TestContext.Current.CancellationToken);
timeout.CancelAfter(TimeSpan.FromSeconds(10));
using var host = McpServer.CreateHost(new Program.CosmosShellOptions { McpPort = 0 });
await host.StartAsync(timeout.Token);
try
{
var prompts = new List<string?>();
var address = host.Services.GetRequiredService<IServer>().Features.Get<IServerAddressesFeature>()!.Addresses.Single();
var transport = new HttpClientTransport(new HttpClientTransportOptions
{
Endpoint = new Uri(address.TrimEnd('/') + "/"),
});
await using var client = await McpClient.CreateAsync(
transport,
new McpClientOptions
{
ProtocolVersion = requestedVersion,
Handlers = new McpClientHandlers
{
ElicitationHandler = (request, _) =>
{
prompts.Add(request?.Message);
return ValueTask.FromResult(new ElicitResult { Action = "decline" });
},
},
},
cancellationToken: timeout.Token);
Assert.Equal(negotiatedVersion, client.NegotiatedProtocolVersion);

var result = await client.CallToolAsync(
"rmdb",
new Dictionary<string, object?> { ["name"] = "McpConfirmationTestDb" },
cancellationToken: timeout.Token);

var prompt = Assert.Single(prompts);
Assert.Contains("rmdb \"McpConfirmationTestDb\"", prompt);
Assert.True(result.IsError);
using var document = JsonDocument.Parse(Assert.IsType<TextContentBlock>(Assert.Single(result.Content)).Text);
Assert.Contains("was not approved by the user", document.RootElement.GetProperty("error").GetString());
}
finally
{
await host.StopAsync(TestContext.Current.CancellationToken);
}
}
}
194 changes: 190 additions & 4 deletions CosmosDBShell.Tests/McpLocationSubscriptionTests.cs
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@ namespace CosmosShell.Tests;
using Microsoft.Extensions.Hosting;
using ModelContextProtocol;
using ModelContextProtocol.Client;
using ModelContextProtocol.Protocol;

[Collection(CosmosShell.Tests.Shell.ThemeStateTestCollection.Name)]
public class McpLocationSubscriptionTests
Expand Down Expand Up @@ -105,25 +106,210 @@ public async Task EndedSession_RemovesLocationSubscription()
}

[Fact]
public void HttpTransport_UsesStatefulSessionsWithBoundedIdleTimeout()
public async Task ListeningClient_ReceivesLocationChangeOnListenStream()
{
using var timeout = CancellationTokenSource.CreateLinkedTokenSource(TestContext.Current.CancellationToken);
timeout.CancelAfter(TimeSpan.FromSeconds(10));
using var host = McpServer.CreateHost(new Program.CosmosShellOptions { McpPort = 0 });
await host.StartAsync(timeout.Token);
try
{
var subscriptions = host.Services.GetRequiredService<LocationResourceSubscriptions>();
await using var client = await ConnectAsync(host, timeout.Token, protocolVersion: null);
Assert.Equal("2026-07-28", client.NegotiatedProtocolVersion);

var acknowledged = new TaskCompletionSource<JsonRpcNotification>(TaskCreationOptions.RunContinuationsAsynchronously);
var updated = new TaskCompletionSource<JsonRpcNotification>(TaskCreationOptions.RunContinuationsAsynchronously);
await using var acknowledgedHandler = client.RegisterNotificationHandler(
NotificationMethods.SubscriptionsAcknowledgedNotification,
(notification, _) =>
{
acknowledged.TrySetResult(notification);
return ValueTask.CompletedTask;
});
await using var updatedHandler = client.RegisterNotificationHandler(
NotificationMethods.ResourceUpdatedNotification,
(notification, _) =>
{
updated.TrySetResult(notification);
return ValueTask.CompletedTask;
});

using var listenCancellation = CancellationTokenSource.CreateLinkedTokenSource(timeout.Token);
var listen = client.SendRequestAsync(
new JsonRpcRequest
{
Id = new RequestId("location-listen"),
Method = RequestMethods.SubscriptionsListen,
Params = JsonSerializer.SerializeToNode(new SubscriptionsListenRequestParams
{
Notifications = new SubscriptionsListenNotifications
{
ResourceSubscriptions = [ResourceOperations.CurrentLocationUri, "cosmos://docs/scripting"],
ToolsListChanged = true,
},
}),
},
listenCancellation.Token);

// Only the current-location resource is honored.
var acknowledgement = await acknowledged.Task.WaitAsync(timeout.Token);
var granted = acknowledgement.Params!["notifications"]!.AsObject();
Assert.Equal(ResourceOperations.CurrentLocationUri, Assert.Single(granted["resourceSubscriptions"]!.AsArray())!.GetValue<string>());
Assert.False(granted.ContainsKey("toolsListChanged"));
Assert.Equal("location-listen", acknowledgement.Params["_meta"]![MetaKeys.SubscriptionId]!.GetValue<string>());
while (subscriptions.ListenerCount != 1)
{
await Task.Delay(20, timeout.Token);
}

var originalState = ShellInterpreter.Instance.State;
using var cosmosClient = new CosmosClient(
"https://localhost:8081",
Convert.ToBase64String(new byte[64]),
new CosmosClientOptions { ConnectionMode = ConnectionMode.Gateway });
try
{
ShellInterpreter.Instance.State = new DatabaseState("McpListenTest", cosmosClient);
var notification = await updated.Task.WaitAsync(timeout.Token);
Assert.Equal(ResourceOperations.CurrentLocationUri, notification.Params!["uri"]!.GetValue<string>());
Assert.Equal("location-listen", notification.Params["_meta"]![MetaKeys.SubscriptionId]!.GetValue<string>());
}
finally
{
ShellInterpreter.Instance.State = originalState;
}

// Cancelling the listen request ends the stream and releases the listener.
await listenCancellation.CancelAsync();
try
{
await listen;
Assert.Fail("The listen request should end only when cancelled.");
}
catch (OperationCanceledException)
{
Assert.True(listenCancellation.IsCancellationRequested);
}

while (subscriptions.ListenerCount != 0)
{
await Task.Delay(20, timeout.Token);
}
}
finally
{
await host.StopAsync(TestContext.Current.CancellationToken);
}
}

[Fact]
public async Task ListenWithOnlyUnsupportedFilters_AcknowledgesNothingAndCompletes()
{
using var timeout = CancellationTokenSource.CreateLinkedTokenSource(TestContext.Current.CancellationToken);
timeout.CancelAfter(TimeSpan.FromSeconds(10));
using var host = McpServer.CreateHost(new Program.CosmosShellOptions { McpPort = 0 });
await host.StartAsync(timeout.Token);
try
{
var subscriptions = host.Services.GetRequiredService<LocationResourceSubscriptions>();
await using var client = await ConnectAsync(host, timeout.Token, protocolVersion: null);

var acknowledged = new TaskCompletionSource<JsonRpcNotification>(TaskCreationOptions.RunContinuationsAsynchronously);
await using var acknowledgedHandler = client.RegisterNotificationHandler(
NotificationMethods.SubscriptionsAcknowledgedNotification,
(notification, _) =>
{
acknowledged.TrySetResult(notification);
return ValueTask.CompletedTask;
});

var response = await client.SendRequestAsync(
new JsonRpcRequest
{
Id = new RequestId("unsupported-listen"),
Method = RequestMethods.SubscriptionsListen,
Params = JsonSerializer.SerializeToNode(new SubscriptionsListenRequestParams
{
Notifications = new SubscriptionsListenNotifications
{
ResourceSubscriptions = ["cosmos://docs/scripting"],
ToolsListChanged = true,
},
}),
},
timeout.Token);

Assert.Equal("unsupported-listen", response.Id.ToString());
var acknowledgement = await acknowledged.Task.WaitAsync(timeout.Token);
var granted = acknowledgement.Params!["notifications"]!.AsObject();
Assert.False(granted.ContainsKey("resourceSubscriptions"));
Assert.False(granted.ContainsKey("toolsListChanged"));
Assert.Equal("unsupported-listen", acknowledgement.Params["_meta"]![MetaKeys.SubscriptionId]!.GetValue<string>());
Assert.Equal(0, subscriptions.ListenerCount);
}
finally
{
await host.StopAsync(TestContext.Current.CancellationToken);
}
}

[Fact]
public async Task HostStop_EndsOpenListenStreamWithoutWaitingForShutdownTimeout()
{
using var timeout = CancellationTokenSource.CreateLinkedTokenSource(TestContext.Current.CancellationToken);
timeout.CancelAfter(TimeSpan.FromSeconds(10));
using var host = McpServer.CreateHost(new Program.CosmosShellOptions { McpPort = 0 });
await host.StartAsync(timeout.Token);
var subscriptions = host.Services.GetRequiredService<LocationResourceSubscriptions>();
await using var client = await ConnectAsync(host, timeout.Token, protocolVersion: null);

_ = client.SendRequestAsync(
new JsonRpcRequest
{
Id = new RequestId("shutdown-listen"),
Method = RequestMethods.SubscriptionsListen,
Params = JsonSerializer.SerializeToNode(new SubscriptionsListenRequestParams
{
Notifications = new SubscriptionsListenNotifications
{
ResourceSubscriptions = [ResourceOperations.CurrentLocationUri],
},
}),
},
timeout.Token);
while (subscriptions.ListenerCount != 1)
{
await Task.Delay(20, timeout.Token);
}

// The web server stops first and waits for open requests until the host shutdown timeout (30 s by default).
await host.StopAsync(CancellationToken.None).WaitAsync(TimeSpan.FromSeconds(5), TestContext.Current.CancellationToken);

Assert.Equal(0, subscriptions.ListenerCount);
}

[Fact]
public void HttpTransport_UsesSessionsOnlyForInitializeClientsWithBoundedIdleTimeout()
{
var options = new ModelContextProtocol.AspNetCore.HttpServerTransportOptions();
McpServer.ConfigureHttpTransport(options);

Assert.Equal(ModelContextProtocol.AspNetCore.HttpServerSessionMode.Stateful, options.SessionMode);
Assert.Equal(ModelContextProtocol.AspNetCore.HttpServerSessionMode.StatefulForInitializeClients, options.SessionMode);
#pragma warning disable MCP9006, MCPEXP002
Assert.Equal(McpServer.SessionIdleTimeout, options.IdleTimeout);
Assert.NotNull(options.RunSessionHandler);
#pragma warning restore MCP9006, MCPEXP002
}

private static async Task<McpClient> ConnectAsync(IHost host, CancellationToken cancellationToken)
// resources/subscribe and session lifetimes exist only for clients that use the initialize handshake.
private static async Task<McpClient> ConnectAsync(IHost host, CancellationToken cancellationToken, string? protocolVersion = "2025-11-25")
{
var address = host.Services.GetRequiredService<IServer>().Features.Get<IServerAddressesFeature>()!.Addresses.Single();
var transport = new HttpClientTransport(new HttpClientTransportOptions
{
Endpoint = new Uri(address.TrimEnd('/') + "/"),
});
return await McpClient.CreateAsync(transport, cancellationToken: cancellationToken);
return await McpClient.CreateAsync(transport, new McpClientOptions { ProtocolVersion = protocolVersion }, cancellationToken: cancellationToken);
}
}
Loading
Loading