Skip to content
Closed
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
29 changes: 27 additions & 2 deletions src/Azure.DataApiBuilder.Mcp/Core/McpStdioServer.cs
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
using Azure.DataApiBuilder.Config.ObjectModel;
using Azure.DataApiBuilder.Core.AuthenticationHelpers.AuthenticationSimulator;
using Azure.DataApiBuilder.Core.Configurations;
using Azure.DataApiBuilder.Core.Services.MetadataProviders;
using Azure.DataApiBuilder.Core.Telemetry;
using Azure.DataApiBuilder.Mcp.Model;
using Azure.DataApiBuilder.Mcp.Telemetry;
Expand All @@ -30,6 +31,8 @@ public class McpStdioServer : IMcpStdioServer
private readonly McpStdoutWriter _stdoutWriter;
private readonly TextReader? _inputReader;
private readonly string _protocolVersion;
private readonly object _initializationLock = new();
private Task? _initializationTask;

private const int MAX_LINE_LENGTH = 1024 * 1024; // 1 MB limit for incoming JSON-RPC requests

Expand Down Expand Up @@ -134,7 +137,7 @@ public async Task RunAsync(CancellationToken cancellationToken)
break;

case "tools/list":
HandleListTools(id);
await HandleListToolsAsync(id);
break;

case "tools/call":
Expand Down Expand Up @@ -284,8 +287,10 @@ private void HandleInitialize(JsonElement? id, JsonElement root)
/// <param name="id">
/// The request identifier extracted from the incoming JSON-RPC request. Used to correlate the response with the request.
/// </param>
private void HandleListTools(JsonElement? id)
private async Task HandleListToolsAsync(JsonElement? id)
{
await EnsureToolsInitializedAsync();

List<object> toolsWire = new();
int count = 0;

Expand All @@ -308,6 +313,24 @@ private void HandleListTools(JsonElement? id)
WriteResult(id, new { tools = toolsWire });
}

private Task EnsureToolsInitializedAsync()
{
lock (_initializationLock)
{
return _initializationTask ??= InitializeToolsAsync();
}
}

private async Task InitializeToolsAsync()
{
IMetadataProviderFactory metadataProviderFactory =
_serviceProvider.GetRequiredService<IMetadataProviderFactory>();
await metadataProviderFactory.InitializeAsync();

IEnumerable<IMcpTool> tools = _serviceProvider.GetServices<IMcpTool>();
McpToolRegistry.InitializeAndRegisterTools(tools, _toolRegistry, _serviceProvider);
}

/// <summary>
/// Handles the "logging/setLevel" JSON-RPC method by updating the runtime log level.
/// </summary>
Expand Down Expand Up @@ -455,6 +478,8 @@ private async Task HandleCallToolAsync(JsonElement? id, JsonElement root, Cancel
return;
}

await EnsureToolsInitializedAsync();

if (!_toolRegistry.TryGetTool(toolName!, out IMcpTool? tool) || tool is null)
{
WriteError(id, McpStdioJsonRpcErrorCodes.INVALID_PARAMS, $"Tool not found: {toolName}");
Expand Down
147 changes: 147 additions & 0 deletions src/Service.Tests/UnitTests/McpStdioServerRunAsyncTests.cs
Original file line number Diff line number Diff line change
Expand Up @@ -4,10 +4,18 @@
#nullable enable

using System;
using System.Collections.Generic;
using System.Diagnostics.CodeAnalysis;
using System.IO;
using System.Text.Json;
using System.Threading;
using System.Threading.Tasks;
using Azure.DataApiBuilder.Config;
using Azure.DataApiBuilder.Config.DatabasePrimitives;
using Azure.DataApiBuilder.Config.ObjectModel;
using Azure.DataApiBuilder.Core.Configurations;
using Azure.DataApiBuilder.Core.Services;
using Azure.DataApiBuilder.Core.Services.MetadataProviders;
using Azure.DataApiBuilder.Mcp.Core;
using Azure.DataApiBuilder.Mcp.Model;
using Microsoft.Extensions.DependencyInjection;
Expand Down Expand Up @@ -91,6 +99,63 @@ public async Task RunAsync_OutOfRangeNumericId_PreservesIdAndContinuesProcessing
Assert.IsTrue(shutdownResponse.RootElement.GetProperty("result").GetProperty("ok").GetBoolean());
}

[TestMethod]
public async Task RunAsync_InitializeRespondsBeforeMetadataInferenceCompletes()
{
const string INPUT =
"{\"jsonrpc\":\"2.0\",\"id\":1,\"method\":\"initialize\",\"params\":{}}\n" +
"{\"jsonrpc\":\"2.0\",\"method\":\"notifications/initialized\",\"params\":{}}\n" +
"{\"jsonrpc\":\"2.0\",\"id\":2,\"method\":\"tools/list\",\"params\":{}}\n" +
"{\"jsonrpc\":\"2.0\",\"id\":3,\"method\":\"tools/list\",\"params\":{}}\n" +
"{\"jsonrpc\":\"2.0\",\"id\":4,\"method\":\"shutdown\"}\n";

SignalingTextWriter stdoutCapture = new();
BlockingMetadataProviderFactory metadataProviderFactory = new();
RuntimeConfig runtimeConfig = new(
Schema: RuntimeConfig.DEFAULT_CONFIG_SCHEMA_LINK,
DataSource: null,
Entities: new RuntimeEntities(new Dictionary<string, Entity>()),
Runtime: new RuntimeOptions(
Rest: null,
GraphQL: null,
Mcp: new McpRuntimeOptions(),
Host: null));
RuntimeConfigProvider runtimeConfigProvider = new StubRuntimeConfigProvider(runtimeConfig);

ServiceProvider serviceProvider = new ServiceCollection()
.AddSingleton(new McpStdoutWriter(stdoutCapture))
.AddSingleton<McpToolRegistry>()
.AddSingleton<IMetadataProviderFactory>(metadataProviderFactory)
.AddSingleton(runtimeConfigProvider)
.BuildServiceProvider();
McpStdioServer server = new(
serviceProvider.GetRequiredService<McpToolRegistry>(),
serviceProvider,
new StringReader(INPUT));

Task runTask = server.RunAsync(CancellationToken.None);

string initializeResponse = await stdoutCapture.FirstLineWritten.Task.WaitAsync(TimeSpan.FromSeconds(5));
await metadataProviderFactory.InitializeStarted.Task.WaitAsync(TimeSpan.FromSeconds(5));

using (JsonDocument response = JsonDocument.Parse(initializeResponse))
{
Assert.AreEqual(1, response.RootElement.GetProperty("id").GetInt32(),
"Initialize must respond before metadata inference completes.");
}

Assert.AreEqual(1, stdoutCapture.LineCount,
"tools/list must wait until metadata inference and tool registration complete.");

metadataProviderFactory.CompleteInitialization();
await runTask.WaitAsync(TimeSpan.FromSeconds(5));

Assert.AreEqual(4, stdoutCapture.LineCount,
"Expected initialize, two tools/list, and shutdown responses.");
Assert.AreEqual(1, metadataProviderFactory.InitializeAsyncCallCount,
"Metadata inference must run exactly once.");
}

private static (McpStdioServer server, StringWriter stdoutCapture) CreateServerWithCapturedOutput(TextReader inputReader)
{
StringWriter stdoutCapture = new();
Expand All @@ -108,5 +173,87 @@ private static (McpStdioServer server, StringWriter stdoutCapture) CreateServerW

return (server, stdoutCapture);
}

private sealed class SignalingTextWriter : StringWriter
{
public TaskCompletionSource<string> FirstLineWritten { get; } =
new(TaskCreationOptions.RunContinuationsAsynchronously);

public int LineCount { get; private set; }

public override void WriteLine(string? value)
{
base.WriteLine(value);
LineCount++;
FirstLineWritten.TrySetResult(value ?? string.Empty);
}
}

private sealed class BlockingMetadataProviderFactory : IMetadataProviderFactory
{
private readonly TaskCompletionSource _completeInitialization =
new(TaskCreationOptions.RunContinuationsAsynchronously);

public TaskCompletionSource InitializeStarted { get; } = new(TaskCreationOptions.RunContinuationsAsynchronously);

public int InitializeAsyncCallCount { get; private set; }

public async Task InitializeAsync()
{
InitializeAsyncCallCount++;
InitializeStarted.TrySetResult();
await _completeInitialization.Task;
}

public void CompleteInitialization()
{
_completeInitialization.TrySetResult();
}

public ISqlMetadataProvider GetMetadataProvider(string dataSourceName)
=> throw new NotImplementedException();

public IEnumerable<ISqlMetadataProvider> ListMetadataProviders()
=> Array.Empty<ISqlMetadataProvider>();

public List<Exception> GetAllMetadataExceptions()
=> new();

public void InitializeAsync(
Dictionary<string, Dictionary<string, DatabaseObject>> entityToDatabaseObjectMap,
Dictionary<string, Dictionary<string, string>> graphQLStoredProcedureExposedNameToEntityNameMap)
{
InitializeAsyncCallCount++;
}
}

private sealed class StubRuntimeConfigProvider : RuntimeConfigProvider
{
private readonly RuntimeConfig _runtimeConfig;

public StubRuntimeConfigProvider(RuntimeConfig runtimeConfig) : base(new StubRuntimeConfigLoader())
{
_runtimeConfig = runtimeConfig;
}

public override RuntimeConfig GetConfig()
{
return _runtimeConfig;
}
}

private sealed class StubRuntimeConfigLoader : RuntimeConfigLoader
{
public override bool TryLoadKnownConfig([NotNullWhen(true)] out RuntimeConfig? config, bool replaceEnvVar = false)
{
config = null;
return false;
}

public override string GetPublishedDraftSchemaLink()
{
return RuntimeConfig.DEFAULT_CONFIG_SCHEMA_LINK;
}
}
}
}
7 changes: 0 additions & 7 deletions src/Service/Utilities/McpStdioHelper.cs
Original file line number Diff line number Diff line change
Expand Up @@ -78,13 +78,6 @@ public static bool RunMcpStdioHost(IHost host)
{
try
{
Mcp.Core.McpToolRegistry registry =
host.Services.GetRequiredService<Mcp.Core.McpToolRegistry>();
IEnumerable<Mcp.Model.IMcpTool> tools =
host.Services.GetServices<Mcp.Model.IMcpTool>();

Mcp.Core.McpToolRegistry.InitializeAndRegisterTools(tools, registry, host.Services);

IHostApplicationLifetime lifetime =
host.Services.GetRequiredService<IHostApplicationLifetime>();
Mcp.Core.IMcpStdioServer stdio =
Expand Down
Loading