Skip to content

Commit a4d7002

Browse files
liuzqkgithub-actions[bot]
authored andcommitted
Clean MCP sessions and deduplicate tool notifications
1 parent bd72241 commit a4d7002

3 files changed

Lines changed: 395 additions & 22 deletions

File tree

Server/src/services/tools/__init__.py

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

170-
PluginHub._sync_server_tool_visibility(enabled_tools)
170+
visibility_synced = PluginHub._sync_server_tool_visibility(enabled_tools)
171171

172172
# Register custom (non-built-in) tools via CustomToolService.
173173
# The extended get_tool_states response includes is_built_in,
@@ -235,8 +235,11 @@ async def sync_tool_visibility_from_unity(
235235
"Update MCPForUnity to enable custom tool sync in stdio mode."
236236
)
237237

238-
if notify:
239-
await PluginHub._notify_mcp_tool_list_changed()
238+
if visibility_synced:
239+
await PluginHub._record_and_notify_tool_list_change(
240+
enabled_tools,
241+
notify=notify,
242+
)
240243

241244
# Build summary
242245
from services.registry import get_group_tool_names

Server/src/transport/plugin_hub.py

Lines changed: 178 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -3,13 +3,16 @@
33
from __future__ import annotations
44

55
import asyncio
6+
import hashlib
7+
import json
68
import logging
79
import os
810
import time
911
import uuid
1012
import weakref
1113
from typing import TYPE_CHECKING, Any, ClassVar
1214

15+
import anyio
1316
from starlette.endpoints import WebSocketEndpoint
1417
from starlette.websockets import WebSocket, WebSocketState
1518

@@ -60,25 +63,57 @@ def _read_bounded_wait_env(name: str, default_s: float, max_s: float) -> float:
6063
# session so we can send ``tools/list_changed`` notifications later.
6164
_active_mcp_sessions: weakref.WeakSet = weakref.WeakSet()
6265
_session_tracking_installed = False
66+
_SESSION_TRACKING_STATE_ATTRIBUTE = "_mcpforunity_active_sessions"
67+
68+
69+
def _track_mcp_session(session: Any) -> None:
70+
_active_mcp_sessions.add(session)
71+
72+
73+
def _untrack_mcp_session(session: Any) -> None:
74+
_active_mcp_sessions.discard(session)
6375

6476

6577
def _install_session_tracking() -> None:
6678
"""Patch *MiddlewareServerSession* to track active MCP client sessions."""
67-
global _session_tracking_installed
68-
if _session_tracking_installed:
69-
return
70-
_session_tracking_installed = True
79+
global _active_mcp_sessions, _session_tracking_installed
7180

7281
from fastmcp.server.low_level import MiddlewareServerSession
7382

83+
existing_sessions = getattr(
84+
MiddlewareServerSession,
85+
_SESSION_TRACKING_STATE_ATTRIBUTE,
86+
None,
87+
)
88+
if isinstance(existing_sessions, weakref.WeakSet):
89+
_active_mcp_sessions = existing_sessions
90+
_session_tracking_installed = True
91+
return
92+
if _session_tracking_installed:
93+
return
94+
7495
_original_aenter = MiddlewareServerSession.__aenter__
96+
_original_aexit = MiddlewareServerSession.__aexit__
7597

7698
async def _tracking_aenter(self): # type: ignore[override]
7799
result = await _original_aenter(self)
78-
_active_mcp_sessions.add(self)
100+
_track_mcp_session(self)
79101
return result
80102

103+
async def _tracking_aexit(self, exc_type, exc_value, traceback): # type: ignore[override]
104+
try:
105+
return await _original_aexit(self, exc_type, exc_value, traceback)
106+
finally:
107+
_untrack_mcp_session(self)
108+
81109
MiddlewareServerSession.__aenter__ = _tracking_aenter # type: ignore[assignment]
110+
MiddlewareServerSession.__aexit__ = _tracking_aexit # type: ignore[assignment]
111+
setattr(
112+
MiddlewareServerSession,
113+
_SESSION_TRACKING_STATE_ATTRIBUTE,
114+
_active_mcp_sessions,
115+
)
116+
_session_tracking_installed = True
82117

83118

84119
class PluginDisconnectedError(RuntimeError):
@@ -147,6 +182,9 @@ class PluginHub(WebSocketEndpoint):
147182
_last_pong: ClassVar[dict[str, float]] = {}
148183
# session_id -> ping task
149184
_ping_tasks: ClassVar[dict[str, asyncio.Task]] = {}
185+
_published_tool_fingerprint: ClassVar[str | None] = None
186+
_pending_tool_list_notifications: ClassVar[weakref.WeakSet] = weakref.WeakSet()
187+
_TOOL_LIST_NOTIFY_TIMEOUT_SECONDS = 1.0
150188

151189
@classmethod
152190
def configure(
@@ -160,6 +198,8 @@ def configure(
160198
cls._loop = loop or asyncio.get_running_loop()
161199
# Ensure coordination primitives are bound to the configured loop
162200
cls._lock = asyncio.Lock()
201+
cls._published_tool_fingerprint = None
202+
cls._pending_tool_list_notifications.clear()
163203
# Start tracking MCP client sessions for tool-change notifications
164204
if mcp is not None:
165205
_install_session_tracking()
@@ -532,11 +572,7 @@ async def _handle_register_tools(self, websocket: WebSocket, payload: RegisterTo
532572

533573
# Sync server-level FastMCP visibility so new MCP client sessions
534574
# (e.g. new Claude Code conversations) see the correct tool set.
535-
self._sync_server_tool_visibility(payload.tools)
536-
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()
575+
visibility_synced = self._sync_server_tool_visibility(payload.tools)
540576

541577
try:
542578
from services.custom_tool_service import CustomToolService
@@ -555,8 +591,100 @@ async def _handle_register_tools(self, websocket: WebSocket, payload: RegisterTo
555591
exc_info=exc,
556592
)
557593

594+
if visibility_synced:
595+
await cls._record_and_notify_tool_list_change(payload.tools)
596+
597+
@classmethod
598+
def _tool_fingerprint(cls, tools: list) -> str:
599+
serialized: list[str] = []
600+
for tool in tools:
601+
if hasattr(tool, "to_mcp_tool"):
602+
tool = tool.to_mcp_tool()
603+
604+
if isinstance(tool, dict):
605+
payload = tool
606+
elif hasattr(tool, "model_dump"):
607+
try:
608+
payload = tool.model_dump(mode="json")
609+
except TypeError:
610+
payload = tool.model_dump()
611+
else:
612+
payload = vars(tool)
613+
614+
serialized.append(
615+
json.dumps(
616+
payload,
617+
sort_keys=True,
618+
separators=(",", ":"),
619+
default=str,
620+
)
621+
)
622+
623+
return hashlib.sha256(
624+
"\n".join(sorted(serialized)).encode("utf-8")
625+
).hexdigest()
626+
627+
@classmethod
628+
async def _record_published_tool_list(cls, registered_tools: list) -> bool:
629+
"""Record a stable fingerprint after the publication transaction succeeds."""
630+
published_tools = registered_tools
631+
if cls._mcp is not None:
632+
try:
633+
published_tools = list(await cls._mcp.list_tools())
634+
except Exception:
635+
logger.warning(
636+
"Failed to inspect the published FastMCP tool list; "
637+
"leaving its fingerprint unchanged for retry",
638+
exc_info=True,
639+
)
640+
return False
641+
642+
try:
643+
digest = cls._tool_fingerprint(published_tools)
644+
except Exception:
645+
logger.warning(
646+
"Failed to fingerprint the published FastMCP tool list; "
647+
"leaving its fingerprint unchanged for retry",
648+
exc_info=True,
649+
)
650+
return False
651+
if digest == cls._published_tool_fingerprint:
652+
return False
653+
654+
cls._published_tool_fingerprint = digest
655+
return True
656+
657+
@classmethod
658+
async def _record_and_notify_tool_list_change(
659+
cls,
660+
registered_tools: list,
661+
notify: bool = True,
662+
) -> bool:
663+
changed = await cls._record_published_tool_list(registered_tools)
664+
if not notify:
665+
return changed
666+
667+
if changed:
668+
cls._pending_tool_list_notifications.update(
669+
list(_active_mcp_sessions)
670+
)
671+
672+
if cls._pending_tool_list_notifications:
673+
# Retry only sessions that have not yet observed the latest
674+
# published schema, avoiding duplicate notifications to peers that
675+
# already succeeded.
676+
await cls._notify_mcp_tool_list_changed(
677+
list(cls._pending_tool_list_notifications)
678+
)
679+
else:
680+
logger.debug(
681+
"Published tool schema unchanged; skipping tools/list_changed"
682+
)
683+
684+
return changed
685+
558686
@classmethod
559-
def _sync_server_tool_visibility(cls, registered_tools: list) -> None:
687+
def _sync_server_tool_visibility(cls, registered_tools: list) -> bool:
560688
"""Sync FastMCP server-level tool group visibility to match Unity's state.
561689
562690
When Unity sends ``register_tools``, some groups may have been toggled
@@ -573,7 +701,7 @@ def _sync_server_tool_visibility(cls, registered_tools: list) -> None:
573701
"""
574702
mcp = cls._mcp
575703
if mcp is None:
576-
return
704+
return True
577705

578706
try:
579707
from services.registry import get_group_tool_names, TOOL_GROUPS
@@ -621,14 +749,19 @@ def _sync_server_tool_visibility(cls, registered_tools: list) -> None:
621749
len(mcp._transforms),
622750
cls._unity_transform_start or 0,
623751
)
752+
return True
624753
except Exception:
625754
logger.debug(
626755
"Failed to sync server-level tool visibility",
627756
exc_info=True,
628757
)
758+
return False
629759

630760
@classmethod
631-
async def _notify_mcp_tool_list_changed(cls) -> None:
761+
async def _notify_mcp_tool_list_changed(
762+
cls,
763+
sessions: list[Any] | None = None,
764+
) -> None:
632765
"""Send ``tools/list_changed`` to every connected MCP client session.
633766
634767
After server-level tool visibility is updated (e.g. when Unity reports
@@ -638,20 +771,46 @@ async def _notify_mcp_tool_list_changed(cls) -> None:
638771
transforms but do **not** push notifications to already-connected
639772
sessions — we do that here.
640773
"""
641-
sessions = list(_active_mcp_sessions)
774+
sessions = list(_active_mcp_sessions) if sessions is None else sessions
642775
if not sessions:
643776
return
644-
for session in sessions:
777+
778+
async def notify_session(session: Any) -> bool:
645779
try:
646-
await session.send_tool_list_changed()
780+
await asyncio.wait_for(
781+
session.send_tool_list_changed(),
782+
timeout=cls._TOOL_LIST_NOTIFY_TIMEOUT_SECONDS,
783+
)
784+
cls._pending_tool_list_notifications.discard(session)
785+
return True
786+
except (anyio.BrokenResourceError, anyio.ClosedResourceError):
787+
_untrack_mcp_session(session)
788+
cls._pending_tool_list_notifications.discard(session)
789+
logger.debug(
790+
"Failed to notify MCP session of tool list change; removed stale session",
791+
exc_info=True,
792+
)
793+
except asyncio.TimeoutError:
794+
cls._pending_tool_list_notifications.add(session)
795+
logger.debug(
796+
"Timed out notifying MCP session of tool list change; "
797+
"keeping session for the next publication",
798+
exc_info=True,
799+
)
647800
except Exception:
801+
cls._pending_tool_list_notifications.add(session)
648802
logger.debug(
649-
"Failed to notify MCP session of tool list change",
803+
"Failed to notify MCP session of tool list change; "
804+
"keeping session because closure was not confirmed",
650805
exc_info=True,
651806
)
807+
return False
808+
809+
notified = sum(await asyncio.gather(*(notify_session(session) for session in sessions)))
652810
logger.info(
653-
"Sent tools/list_changed notification to %d MCP session(s)",
654-
len(sessions),
811+
"Sent tools/list_changed notification to %d MCP session(s); %d active",
812+
notified,
813+
len(_active_mcp_sessions),
655814
)
656815

657816
async def _handle_command_result(self, payload: CommandResultMessage) -> None:

0 commit comments

Comments
 (0)