From 0dca42fd9522a2061c864b6fb9d68da1918471fe Mon Sep 17 00:00:00 2001 From: souvikghosh04 Date: Tue, 8 Sep 2026 14:19:53 +0530 Subject: [PATCH] [MCP] Respond to stdio initialize before metadata inference Defer metadata provider initialization and tool registration until the first tools/list or tools/call request so the MCP handshake is not blocked by schema introspection. Fixes #3430 --- .../Core/McpStdioServer.cs | 29 +++- .../UnitTests/McpStdioServerRunAsyncTests.cs | 147 ++++++++++++++++++ src/Service/Utilities/McpStdioHelper.cs | 7 - 3 files changed, 174 insertions(+), 9 deletions(-) diff --git a/src/Azure.DataApiBuilder.Mcp/Core/McpStdioServer.cs b/src/Azure.DataApiBuilder.Mcp/Core/McpStdioServer.cs index aef3b63dcc..49974296be 100644 --- a/src/Azure.DataApiBuilder.Mcp/Core/McpStdioServer.cs +++ b/src/Azure.DataApiBuilder.Mcp/Core/McpStdioServer.cs @@ -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; @@ -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 @@ -134,7 +137,7 @@ public async Task RunAsync(CancellationToken cancellationToken) break; case "tools/list": - HandleListTools(id); + await HandleListToolsAsync(id); break; case "tools/call": @@ -284,8 +287,10 @@ private void HandleInitialize(JsonElement? id, JsonElement root) /// /// The request identifier extracted from the incoming JSON-RPC request. Used to correlate the response with the request. /// - private void HandleListTools(JsonElement? id) + private async Task HandleListToolsAsync(JsonElement? id) { + await EnsureToolsInitializedAsync(); + List toolsWire = new(); int count = 0; @@ -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(); + await metadataProviderFactory.InitializeAsync(); + + IEnumerable tools = _serviceProvider.GetServices(); + McpToolRegistry.InitializeAndRegisterTools(tools, _toolRegistry, _serviceProvider); + } + /// /// Handles the "logging/setLevel" JSON-RPC method by updating the runtime log level. /// @@ -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}"); diff --git a/src/Service.Tests/UnitTests/McpStdioServerRunAsyncTests.cs b/src/Service.Tests/UnitTests/McpStdioServerRunAsyncTests.cs index 224d534158..fd5c955a25 100644 --- a/src/Service.Tests/UnitTests/McpStdioServerRunAsyncTests.cs +++ b/src/Service.Tests/UnitTests/McpStdioServerRunAsyncTests.cs @@ -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; @@ -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()), + 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() + .AddSingleton(metadataProviderFactory) + .AddSingleton(runtimeConfigProvider) + .BuildServiceProvider(); + McpStdioServer server = new( + serviceProvider.GetRequiredService(), + 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(); @@ -108,5 +173,87 @@ private static (McpStdioServer server, StringWriter stdoutCapture) CreateServerW return (server, stdoutCapture); } + + private sealed class SignalingTextWriter : StringWriter + { + public TaskCompletionSource 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 ListMetadataProviders() + => Array.Empty(); + + public List GetAllMetadataExceptions() + => new(); + + public void InitializeAsync( + Dictionary> entityToDatabaseObjectMap, + Dictionary> 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; + } + } } } diff --git a/src/Service/Utilities/McpStdioHelper.cs b/src/Service/Utilities/McpStdioHelper.cs index 4ee403b98e..169fa87982 100644 --- a/src/Service/Utilities/McpStdioHelper.cs +++ b/src/Service/Utilities/McpStdioHelper.cs @@ -78,13 +78,6 @@ public static bool RunMcpStdioHost(IHost host) { try { - Mcp.Core.McpToolRegistry registry = - host.Services.GetRequiredService(); - IEnumerable tools = - host.Services.GetServices(); - - Mcp.Core.McpToolRegistry.InitializeAndRegisterTools(tools, registry, host.Services); - IHostApplicationLifetime lifetime = host.Services.GetRequiredService(); Mcp.Core.IMcpStdioServer stdio =