Skip to content
Merged
Show file tree
Hide file tree
Changes from 3 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
10 changes: 0 additions & 10 deletions MCPForUnity/Editor/Services/ToolDiscoveryService.cs
Original file line number Diff line number Diff line change
Expand Up @@ -226,16 +226,6 @@ private void EnsurePreferenceInitialized(ToolMetadata metadata)
{
bool defaultValue = metadata.AutoRegister || metadata.IsBuiltIn;
EditorPrefs.SetBool(key, defaultValue);
return;
}

if (metadata.IsBuiltIn && !metadata.AutoRegister)
{
bool currentValue = EditorPrefs.GetBool(key, metadata.AutoRegister);
if (currentValue == metadata.AutoRegister)
{
EditorPrefs.SetBool(key, true);
}
}
}

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -14,5 +14,6 @@ public interface IMcpTransportClient
Task<bool> StartAsync();
Task StopAsync();
Task<bool> VerifyAsync();
Task ReregisterToolsAsync();
}
}
24 changes: 14 additions & 10 deletions MCPForUnity/Editor/Services/Transport/TransportManager.cs
Original file line number Diff line number Diff line change
Expand Up @@ -42,16 +42,6 @@ private IMcpTransportClient GetOrCreateClient(TransportMode mode)
};
}

private IMcpTransportClient GetClient(TransportMode mode)
{
return mode switch
{
TransportMode.Http => _httpClient,
TransportMode.Stdio => _stdioClient,
_ => throw new ArgumentOutOfRangeException(nameof(mode), mode, "Unsupported transport mode"),
};
}

public async Task<bool> StartAsync(TransportMode mode)
{
IMcpTransportClient client = GetOrCreateClient(mode);
Expand Down Expand Up @@ -128,6 +118,20 @@ public TransportState GetState(TransportMode mode)

public bool IsRunning(TransportMode mode) => GetState(mode).IsConnected;

/// <summary>
/// Gets the active transport client for the specified mode.
/// Returns null if the client hasn't been created yet.
/// </summary>
public IMcpTransportClient GetClient(TransportMode mode)
{
return mode switch
{
TransportMode.Http => _httpClient,
TransportMode.Stdio => _stdioClient,
_ => throw new ArgumentOutOfRangeException(nameof(mode), mode, "Unsupported transport mode"),
};
}
Comment thread
whatevertogo marked this conversation as resolved.

private void UpdateState(TransportMode mode, TransportState state)
{
switch (mode)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -46,5 +46,12 @@ public Task<bool> VerifyAsync()
return Task.FromResult(running);
}

public Task ReregisterToolsAsync()
{
// Stdio transport doesn't support dynamic tool reregistration
// Tools are registered at server startup
return Task.CompletedTask;
}

}
}
Original file line number Diff line number Diff line change
Expand Up @@ -506,6 +506,29 @@ private async Task SendRegisterToolsAsync(CancellationToken token)
McpLog.Info($"[WebSocket] Sent {tools.Count} tools registration", false);
}

public async Task ReregisterToolsAsync()
{
if (!IsConnected || _lifecycleCts == null)
{
McpLog.Warn("[WebSocket] Cannot reregister tools: not connected");
return;
}

try
{
await SendRegisterToolsAsync(_lifecycleCts.Token).ConfigureAwait(false);
McpLog.Info("[WebSocket] Tool reregistration completed", false);
}
catch (System.OperationCanceledException)
{
McpLog.Warn("[WebSocket] Tool reregistration cancelled");
}
catch (System.Exception ex)
{
McpLog.Error($"[WebSocket] Tool reregistration failed: {ex.Message}");
}
}

private async Task HandleExecuteAsync(JObject payload, CancellationToken token)
{
string commandId = payload.Value<string>("id");
Expand Down
53 changes: 50 additions & 3 deletions MCPForUnity/Editor/Windows/Components/Tools/McpToolsSection.cs
Original file line number Diff line number Diff line change
@@ -1,9 +1,11 @@
using System;
using System.Collections.Generic;
using System.Linq;
using System.Threading;
using MCPForUnity.Editor.Constants;
using MCPForUnity.Editor.Helpers;
using MCPForUnity.Editor.Services;
using MCPForUnity.Editor.Services.Transport;
using MCPForUnity.Editor.Tools;
using UnityEditor;
using UnityEngine.UIElements;
Expand Down Expand Up @@ -223,23 +225,61 @@ private VisualElement CreateToolRow(ToolMetadata tool)
return row;
}

private void HandleToggleChange(ToolMetadata tool, bool enabled, bool updateSummary = true)
private void HandleToggleChange(
ToolMetadata tool,
bool enabled,
bool updateSummary = true,
bool reregisterTools = true)
{
MCPServiceLocator.ToolDiscovery.SetToolEnabled(tool.Name, enabled);

if (updateSummary)
{
UpdateSummary();
}

if (reregisterTools)
{
// Trigger tool reregistration with connected MCP server
ReregisterToolsAsync();
}
}

private void ReregisterToolsAsync()
{
// Fire and forget - don't block UI
ThreadPool.QueueUserWorkItem(_ =>
{
Comment thread
whatevertogo marked this conversation as resolved.
try
{
var transportManager = MCPServiceLocator.TransportManager;
var client = transportManager.GetClient(TransportMode.Http);
if (client != null && client.IsConnected)
{
client.ReregisterToolsAsync().Wait();
}
}
catch (Exception ex)
{
McpLog.Warn($"Failed to reregister tools: {ex.Message}");
Comment thread
whatevertogo marked this conversation as resolved.
Outdated
}
});
Comment thread
whatevertogo marked this conversation as resolved.
Outdated
}
Comment thread
sourcery-ai[bot] marked this conversation as resolved.

private void SetAllToolsState(bool enabled)
{
bool hasChanges = false;

foreach (var tool in allTools)
{
if (!toolToggleMap.TryGetValue(tool.Name, out var toggle))
{
MCPServiceLocator.ToolDiscovery.SetToolEnabled(tool.Name, enabled);
bool currentEnabled = MCPServiceLocator.ToolDiscovery.IsToolEnabled(tool.Name);
if (currentEnabled != enabled)
{
MCPServiceLocator.ToolDiscovery.SetToolEnabled(tool.Name, enabled);
hasChanges = true;
}
continue;
}

Expand All @@ -249,10 +289,17 @@ private void SetAllToolsState(bool enabled)
}

toggle.SetValueWithoutNotify(enabled);
HandleToggleChange(tool, enabled, updateSummary: false);
HandleToggleChange(tool, enabled, updateSummary: false, reregisterTools: false);
hasChanges = true;
}

UpdateSummary();

if (hasChanges)
{
// Trigger a single reregistration after bulk change
ReregisterToolsAsync();
}
}

private void UpdateSummary()
Expand Down
142 changes: 142 additions & 0 deletions Server/src/transport/unity_instance_middleware.py
Original file line number Diff line number Diff line change
Expand Up @@ -55,6 +55,43 @@ def __init__(self):
super().__init__()
self._active_by_key: dict[str, str] = {}
self._lock = RLock()
self._unity_managed_tool_names = {
"batch_execute",
"execute_menu_item",
"find_gameobjects",
"get_test_job",
"manage_asset",
"manage_components",
"manage_editor",
"manage_gameobject",
"manage_material",
"manage_prefabs",
"manage_scene",
"manage_script",
"manage_scriptable_object",
"manage_shader",
"manage_texture",
"manage_vfx",
"read_console",
"refresh_unity",
"run_tests",
}
Comment thread
whatevertogo marked this conversation as resolved.
Outdated
self._tool_alias_to_unity_target = {
# Server-side script helpers route to Unity's manage_script command.
"apply_text_edits": "manage_script",
"create_script": "manage_script",
"delete_script": "manage_script",
"find_in_file": "manage_script",
"get_sha": "manage_script",
"script_apply_edits": "manage_script",
"validate_script": "manage_script",
}
self._server_only_tool_names = {
"debug_request_context",
"execute_custom_tool",
"manage_script_capabilities",
"set_active_instance",
}

def get_session_key(self, ctx) -> str:
"""
Expand Down Expand Up @@ -260,3 +297,108 @@ async def on_read_resource(self, context: MiddlewareContext, call_next):
"""Inject active Unity instance into resource context if available."""
await self._inject_unity_instance(context)
return await call_next(context)

async def on_list_tools(self, context: MiddlewareContext, call_next):
"""Filter MCP tool listing to the Unity-enabled set when session data is available."""
await self._inject_unity_instance(context)
tools = await call_next(context)

if not self._should_filter_tool_listing():
return tools

enabled_tool_names = await self._resolve_enabled_tool_names_for_context(context)
if not enabled_tool_names:
return tools
Comment thread
whatevertogo marked this conversation as resolved.
Outdated

filtered = []
for tool in tools:
tool_name = getattr(tool, "name", None)
if self._is_tool_visible(tool_name, enabled_tool_names):
filtered.append(tool)

return filtered

def _should_filter_tool_listing(self) -> bool:
transport = (config.transport_mode or "stdio").lower()
return transport == "http" and PluginHub.is_configured()

async def _resolve_enabled_tool_names_for_context(
self,
context: MiddlewareContext,
) -> set[str] | None:
ctx = context.fastmcp_context
user_id = ctx.get_state("user_id") if config.http_remote_hosted else None
active_instance = ctx.get_state("unity_instance")

project_hashes = self._resolve_candidate_project_hashes(active_instance)
if not project_hashes:
try:
sessions_data = await PluginHub.get_sessions(user_id=user_id)
sessions = sessions_data.sessions if sessions_data else {}
except Exception:
return None
Comment thread
whatevertogo marked this conversation as resolved.
Outdated

Comment thread
whatevertogo marked this conversation as resolved.
Outdated
if not sessions:
return None

if len(sessions) == 1:
only_session = next(iter(sessions.values()))
only_hash = getattr(only_session, "hash", None)
if only_hash:
project_hashes = [only_hash]
else:
# Multiple sessions without explicit selection: use a union so we don't
# hide tools that are valid in at least one visible Unity instance.
project_hashes = [
session.hash
for session in sessions.values()
if getattr(session, "hash", None)
]

if not project_hashes:
return None

enabled_tool_names: set[str] = set()
for project_hash in project_hashes:
try:
registered_tools = await PluginHub.get_tools_for_project(project_hash)
Comment thread
whatevertogo marked this conversation as resolved.
Outdated
except Exception:
continue
Comment thread
whatevertogo marked this conversation as resolved.

for tool in registered_tools:
tool_name = getattr(tool, "name", None)
if isinstance(tool_name, str) and tool_name:
enabled_tool_names.add(tool_name)

return enabled_tool_names or None

@staticmethod
def _resolve_candidate_project_hashes(active_instance: str | None) -> list[str]:
if not active_instance:
return []

if "@" in active_instance:
_, _, suffix = active_instance.rpartition("@")
return [suffix] if suffix else []

return [active_instance]

def _is_tool_visible(self, tool_name: str | None, enabled_tool_names: set[str]) -> bool:
if not isinstance(tool_name, str) or not tool_name:
return True

if tool_name in self._server_only_tool_names:
return True

if tool_name in enabled_tool_names:
return True

unity_target = self._tool_alias_to_unity_target.get(tool_name)
if unity_target:
return unity_target in enabled_tool_names

# Keep unknown tools visible for forward compatibility.
if tool_name not in self._unity_managed_tool_names:
return True

return False
Loading
Loading