Skip to content

Commit 0217aa5

Browse files
committed
Clean MCP sessions and deduplicate tool notifications
1 parent bd72241 commit 0217aa5

3 files changed

Lines changed: 157 additions & 12 deletions

File tree

Server/src/services/tools/__init__.py

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -167,7 +167,9 @@ async def sync_tool_visibility_from_unity(
167167
len(enabled_tools), len(tools),
168168
)
169169

170-
PluginHub._sync_server_tool_visibility(enabled_tools)
170+
tool_list_changed = PluginHub._tool_list_changed(enabled_tools)
171+
if tool_list_changed:
172+
PluginHub._sync_server_tool_visibility(enabled_tools)
171173

172174
# Register custom (non-built-in) tools via CustomToolService.
173175
# The extended get_tool_states response includes is_built_in,
@@ -235,7 +237,7 @@ async def sync_tool_visibility_from_unity(
235237
"Update MCPForUnity to enable custom tool sync in stdio mode."
236238
)
237239

238-
if notify:
240+
if notify and tool_list_changed:
239241
await PluginHub._notify_mcp_tool_list_changed()
240242

241243
# Build summary

Server/src/transport/plugin_hub.py

Lines changed: 63 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,8 @@
33
from __future__ import annotations
44

55
import asyncio
6+
import hashlib
7+
import json
68
import logging
79
import os
810
import time
@@ -62,23 +64,39 @@ def _read_bounded_wait_env(name: str, default_s: float, max_s: float) -> float:
6264
_session_tracking_installed = False
6365

6466

67+
def _track_mcp_session(session: Any) -> None:
68+
_active_mcp_sessions.add(session)
69+
70+
71+
def _untrack_mcp_session(session: Any) -> None:
72+
_active_mcp_sessions.discard(session)
73+
74+
6575
def _install_session_tracking() -> None:
6676
"""Patch *MiddlewareServerSession* to track active MCP client sessions."""
6777
global _session_tracking_installed
6878
if _session_tracking_installed:
6979
return
70-
_session_tracking_installed = True
7180

7281
from fastmcp.server.low_level import MiddlewareServerSession
7382

7483
_original_aenter = MiddlewareServerSession.__aenter__
84+
_original_aexit = MiddlewareServerSession.__aexit__
7585

7686
async def _tracking_aenter(self): # type: ignore[override]
7787
result = await _original_aenter(self)
78-
_active_mcp_sessions.add(self)
88+
_track_mcp_session(self)
7989
return result
8090

91+
async def _tracking_aexit(self, exc_type, exc_value, traceback): # type: ignore[override]
92+
try:
93+
return await _original_aexit(self, exc_type, exc_value, traceback)
94+
finally:
95+
_untrack_mcp_session(self)
96+
8197
MiddlewareServerSession.__aenter__ = _tracking_aenter # type: ignore[assignment]
98+
MiddlewareServerSession.__aexit__ = _tracking_aexit # type: ignore[assignment]
99+
_session_tracking_installed = True
82100

83101

84102
class PluginDisconnectedError(RuntimeError):
@@ -147,6 +165,7 @@ class PluginHub(WebSocketEndpoint):
147165
_last_pong: ClassVar[dict[str, float]] = {}
148166
# session_id -> ping task
149167
_ping_tasks: ClassVar[dict[str, asyncio.Task]] = {}
168+
_published_tool_fingerprint: ClassVar[str | None] = None
150169

151170
@classmethod
152171
def configure(
@@ -160,6 +179,7 @@ def configure(
160179
cls._loop = loop or asyncio.get_running_loop()
161180
# Ensure coordination primitives are bound to the configured loop
162181
cls._lock = asyncio.Lock()
182+
cls._published_tool_fingerprint = None
163183
# Start tracking MCP client sessions for tool-change notifications
164184
if mcp is not None:
165185
_install_session_tracking()
@@ -530,13 +550,16 @@ async def _handle_register_tools(self, websocket: WebSocket, payload: RegisterTo
530550
logger.info(
531551
f"Registered {len(payload.tools)} tools for session {session_id}")
532552

533-
# Sync server-level FastMCP visibility so new MCP client sessions
534-
# (e.g. new Claude Code conversations) see the correct tool set.
535-
self._sync_server_tool_visibility(payload.tools)
553+
if cls._tool_list_changed(payload.tools):
554+
# Sync server-level FastMCP visibility so new MCP client sessions
555+
# (e.g. new Claude Code conversations) see the correct tool set.
556+
self._sync_server_tool_visibility(payload.tools)
536557

537-
# Notify any already-connected MCP clients (e.g. CC over stdio) that
538-
# the tool list has changed so they re-fetch.
539-
await cls._notify_mcp_tool_list_changed()
558+
# Notify any already-connected MCP clients that the published tool
559+
# schema changed so they re-fetch it.
560+
await cls._notify_mcp_tool_list_changed()
561+
else:
562+
logger.debug("Unity tool schema unchanged; skipping tools/list_changed")
540563

541564
try:
542565
from services.custom_tool_service import CustomToolService
@@ -555,6 +578,33 @@ async def _handle_register_tools(self, websocket: WebSocket, payload: RegisterTo
555578
exc_info=exc,
556579
)
557580

581+
@classmethod
582+
def _tool_list_changed(cls, registered_tools: list) -> bool:
583+
serialized: list[str] = []
584+
for tool in registered_tools:
585+
if isinstance(tool, dict):
586+
payload = tool
587+
elif hasattr(tool, "model_dump"):
588+
try:
589+
payload = tool.model_dump(mode="json")
590+
except TypeError:
591+
payload = tool.model_dump()
592+
else:
593+
payload = vars(tool)
594+
595+
serialized.append(
596+
json.dumps(payload, sort_keys=True, separators=(",", ":"), default=str)
597+
)
598+
599+
digest = hashlib.sha256(
600+
"\n".join(sorted(serialized)).encode("utf-8")
601+
).hexdigest()
602+
if digest == cls._published_tool_fingerprint:
603+
return False
604+
605+
cls._published_tool_fingerprint = digest
606+
return True
607+
558608
@classmethod
559609
def _sync_server_tool_visibility(cls, registered_tools: list) -> None:
560610
"""Sync FastMCP server-level tool group visibility to match Unity's state.
@@ -641,17 +691,20 @@ async def _notify_mcp_tool_list_changed(cls) -> None:
641691
sessions = list(_active_mcp_sessions)
642692
if not sessions:
643693
return
694+
notified = 0
644695
for session in sessions:
645696
try:
646697
await session.send_tool_list_changed()
698+
notified += 1
647699
except Exception:
700+
_untrack_mcp_session(session)
648701
logger.debug(
649-
"Failed to notify MCP session of tool list change",
702+
"Failed to notify MCP session of tool list change; removed stale session",
650703
exc_info=True,
651704
)
652705
logger.info(
653706
"Sent tools/list_changed notification to %d MCP session(s)",
654-
len(sessions),
707+
notified,
655708
)
656709

657710
async def _handle_command_result(self, payload: CommandResultMessage) -> None:
Lines changed: 90 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,90 @@
1+
from __future__ import annotations
2+
3+
import pytest
4+
5+
from transport import plugin_hub
6+
from transport.plugin_hub import PluginHub
7+
8+
9+
class _FakeSession:
10+
def __init__(self, fail: bool = False) -> None:
11+
self.fail = fail
12+
self.notifications = 0
13+
14+
async def send_tool_list_changed(self) -> None:
15+
if self.fail:
16+
raise RuntimeError("closed")
17+
self.notifications += 1
18+
19+
20+
@pytest.fixture(autouse=True)
21+
def _reset_session_tracking() -> None:
22+
plugin_hub._active_mcp_sessions.clear()
23+
PluginHub._published_tool_fingerprint = None
24+
yield
25+
plugin_hub._active_mcp_sessions.clear()
26+
PluginHub._published_tool_fingerprint = None
27+
28+
29+
def test_session_tracking_removes_exited_session() -> None:
30+
session = _FakeSession()
31+
32+
plugin_hub._track_mcp_session(session)
33+
assert session in plugin_hub._active_mcp_sessions
34+
35+
plugin_hub._untrack_mcp_session(session)
36+
assert session not in plugin_hub._active_mcp_sessions
37+
38+
39+
def test_twenty_connect_disconnect_cycles_return_to_baseline() -> None:
40+
sessions = [_FakeSession() for _ in range(20)]
41+
42+
for session in sessions:
43+
plugin_hub._track_mcp_session(session)
44+
plugin_hub._untrack_mcp_session(session)
45+
46+
assert list(plugin_hub._active_mcp_sessions) == []
47+
48+
49+
@pytest.mark.asyncio
50+
async def test_notification_prunes_closed_sessions() -> None:
51+
active = _FakeSession()
52+
closed = _FakeSession(fail=True)
53+
plugin_hub._track_mcp_session(active)
54+
plugin_hub._track_mcp_session(closed)
55+
56+
await PluginHub._notify_mcp_tool_list_changed()
57+
58+
assert active.notifications == 1
59+
assert active in plugin_hub._active_mcp_sessions
60+
assert closed not in plugin_hub._active_mcp_sessions
61+
62+
63+
@pytest.mark.asyncio
64+
async def test_unchanged_tools_do_not_rebroadcast() -> None:
65+
sessions = [_FakeSession(), _FakeSession()]
66+
for session in sessions:
67+
plugin_hub._track_mcp_session(session)
68+
69+
tools = [{"name": "read_console", "description": "Console"}]
70+
assert PluginHub._tool_list_changed(tools) is True
71+
await PluginHub._notify_mcp_tool_list_changed()
72+
assert PluginHub._tool_list_changed(tools) is False
73+
74+
assert [session.notifications for session in sessions] == [1, 1]
75+
76+
77+
def test_tool_fingerprint_deduplicates_reordered_payload() -> None:
78+
first = [
79+
{"name": "manage_scene", "description": "Scene"},
80+
{"name": "read_console", "description": "Console"},
81+
]
82+
reordered = list(reversed(first))
83+
changed_schema = [
84+
{"name": "manage_scene", "description": "Scene changed"},
85+
{"name": "read_console", "description": "Console"},
86+
]
87+
88+
assert PluginHub._tool_list_changed(first) is True
89+
assert PluginHub._tool_list_changed(reordered) is False
90+
assert PluginHub._tool_list_changed(changed_schema) is True

0 commit comments

Comments
 (0)