mirror of
https://github.com/tiennm99/DocsGPT.git
synced 2026-10-04 14:12:58 +00:00
Stop a removed MCP connection from saving its server as a custom tool
After removing a Linear connection, connecting again could skip signing in and save Linear as an unconnected custom tool: a client cached before the removal still held the old tokens, and a late token write re-created a connection for them. A token write for a named connection no longer creates one, removing or disconnecting a connection drops its cached clients, and a sign-in server is never saved without its connection.
This commit is contained in:
1 parent
159ac03904
commit
56ababae71
5 files changed
+124
-6
No files matched your search
@@ -37,6 +37,20 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
_mcp_clients_cache = {}
|
||||
|
||||
|
||||
def forget_cached_clients(*identities: str) -> None:
|
||||
"""Drop cached OAuth clients signed in as any of ``identities``.
|
||||
|
||||
OAuth cache keys name the connection or user whose tokens the client
|
||||
holds (see ``MCPTool._generate_cache_key``).
|
||||
|
||||
Args:
|
||||
identities: Connection ids and user ids.
|
||||
"""
|
||||
markers = tuple(f"#oauth:{identity}:" for identity in identities if identity)
|
||||
for key in [k for k in list(_mcp_clients_cache) if any(marker in k for marker in markers)]:
|
||||
_mcp_clients_cache.pop(key, None)
|
||||
|
||||
# A token expiry long past: a stored token of unknown age is renewed before use.
|
||||
_EXPIRED = 1.0
|
||||
|
||||
|
||||
@@ -399,6 +399,17 @@ class MCPServerSave(Resource):
|
||||
connection_id = _mcp_connection(
|
||||
user, storage_config, auth_type, auth_credentials, display_name,
|
||||
) or _previous_connection(existing_doc, storage_config, user)
|
||||
if auth_type == "oauth" and not connection_id:
|
||||
# A sign-in server's tokens live on its connection. Without one
|
||||
# (it was removed, and a client cached before that answered the
|
||||
# discovery) the tool would be saved unconnected.
|
||||
return make_response(
|
||||
jsonify({
|
||||
"success": False,
|
||||
"error": "Not signed in to this server. Sign in again to connect it.",
|
||||
}),
|
||||
400,
|
||||
)
|
||||
if connection_id and auth_type != "oauth":
|
||||
# The secret lives on the connection only.
|
||||
storage_config.pop("encrypted_credentials", None)
|
||||
|
||||
@@ -9,6 +9,7 @@ MCP server or a set of API credentials. Sources and tools point at it through
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import sys
|
||||
from typing import Any, Iterable, Optional
|
||||
|
||||
from docsgpt.connectors import catalog
|
||||
@@ -393,6 +394,7 @@ def disconnect(conn, row: dict) -> dict:
|
||||
except CredentialDecryptionError:
|
||||
secrets = {}
|
||||
revoke_at_provider(row, secrets)
|
||||
_forget_mcp_clients(row)
|
||||
kept = {"client_info": secrets["client_info"]} if secrets.get("client_info") else {}
|
||||
write_secrets(
|
||||
conn, row, kept, status=STATUS_DISCONNECTED, session_token=None, last_error=None,
|
||||
@@ -1135,16 +1137,24 @@ def update_mcp_secrets(
|
||||
"""Merge ``patch`` into an MCP connection's secrets (``None`` drops a key).
|
||||
|
||||
Creates the connection on first use, named after the matching preset or
|
||||
the server's host.
|
||||
the server's host. A write for a named connection never creates one: a
|
||||
client or sync still running when its connection was removed would
|
||||
otherwise bring it back, holding the tokens it just renewed.
|
||||
|
||||
Returns:
|
||||
The connection row after the update.
|
||||
|
||||
Raises:
|
||||
ConnectionUnavailable: ``connection_id`` names no connection of this
|
||||
user for this server (removed, or another server's).
|
||||
"""
|
||||
from docsgpt.security.encryption import CredentialDecryptionError
|
||||
|
||||
with db_session() as conn:
|
||||
repo = ConnectorSessionsRepository(conn)
|
||||
row = _mcp_row(conn, user_id, base_url, connection_id, lock=True)
|
||||
if row is None and connection_id:
|
||||
raise ConnectionUnavailable("Connection not found", connection_id=connection_id, status="missing")
|
||||
if row is None:
|
||||
row = repo.merge_session_data(user_id, mcp_provider(base_url), base_url, {})
|
||||
try:
|
||||
@@ -1389,6 +1399,22 @@ def ensure_connection_tools(
|
||||
return existing + created
|
||||
|
||||
|
||||
def _forget_mcp_clients(row: dict) -> None:
|
||||
"""Drop the MCP clients this process cached for a connection's tokens.
|
||||
|
||||
A cached client keeps the tokens it signed in with for a few minutes, so
|
||||
after the connection is removed or disconnected it would still answer as
|
||||
signed in (and save a tool with no connection behind it).
|
||||
"""
|
||||
if not str(row.get("provider") or "").startswith("mcp:") and row.get("auth_kind") != "mcp_oauth":
|
||||
return
|
||||
# Not imported here: a process that never loaded the MCP tool has no
|
||||
# clients cached, and importing it from this module would be circular.
|
||||
mcp_tool = sys.modules.get("docsgpt.agents.tools.mcp_tool")
|
||||
if mcp_tool is not None:
|
||||
mcp_tool.forget_cached_clients(str(row["id"]), str(row.get("user_id") or ""))
|
||||
|
||||
|
||||
def remove_connection(conn, row: dict, *, sources: str = "keep", tools: str = "delete") -> list[dict]:
|
||||
"""Delete a connection, choosing what happens to what it feeds.
|
||||
|
||||
@@ -1409,6 +1435,7 @@ def remove_connection(conn, row: dict, *, sources: str = "keep", tools: str = "d
|
||||
|
||||
repo = ConnectorSessionsRepository(conn)
|
||||
connection_id = str(row["id"])
|
||||
_forget_mcp_clients(row)
|
||||
linked_sources = repo.list_sources(connection_id)
|
||||
try:
|
||||
revoke_at_provider(row, read_secrets(row))
|
||||
|
||||
@@ -355,11 +355,19 @@ class TestMCPServerSave:
|
||||
"""Signed in before: the server answers with the saved tokens, no new handshake."""
|
||||
from docsgpt.api.user.tools.mcp import MCPServerSave
|
||||
|
||||
from docsgpt.storage.db.repositories.connector_sessions import ConnectorSessionsRepository
|
||||
|
||||
connection = ConnectorSessionsRepository(pg_conn).create(
|
||||
"u-signed-in", "mcp:https://mcp.linear.app", connector_key="mcp:linear", auth_kind="mcp_oauth",
|
||||
display_name="Linear", account_label="Linear", server_url="https://mcp.linear.app",
|
||||
)
|
||||
signed_in = MagicMock()
|
||||
signed_in.get_actions_metadata.return_value = [{"name": "search"}]
|
||||
with _patch_db(pg_conn), patch(
|
||||
"docsgpt.api.user.tools.mcp.MCPTool", return_value=signed_in,
|
||||
), patch("docsgpt.api.user.tools.mcp._mcp_connection", return_value=None), app.test_request_context(
|
||||
), patch(
|
||||
"docsgpt.api.user.tools.mcp._mcp_connection", return_value=str(connection["id"]),
|
||||
), app.test_request_context(
|
||||
"/api/mcp_server/save", method="POST",
|
||||
json={
|
||||
"displayName": "Linear",
|
||||
@@ -446,6 +454,33 @@ class TestMCPServerSave:
|
||||
assert team["filled_by_llm"] is False and team["value"] == "ENG"
|
||||
|
||||
|
||||
class TestSignInServerWithoutAConnection:
|
||||
def test_an_oauth_server_is_not_saved_without_its_connection(self, app, pg_conn):
|
||||
"""A client still cached from a removed connection can answer the
|
||||
discovery; the tool must not then be saved as an unconnected server."""
|
||||
from docsgpt.api.user.tools.mcp import MCPServerSave
|
||||
from docsgpt.storage.db.repositories.user_tools import UserToolsRepository
|
||||
|
||||
user = "u-mcp-no-connection"
|
||||
fake_tool = MagicMock()
|
||||
fake_tool.get_actions_metadata.return_value = [{"name": "list_issues", "parameters": {"properties": {}}}]
|
||||
with _patch_db(pg_conn), patch(
|
||||
"docsgpt.api.user.tools.mcp.MCPTool", return_value=fake_tool,
|
||||
), app.test_request_context(
|
||||
"/api/mcp_server/save", method="POST",
|
||||
json={
|
||||
"displayName": "Linear",
|
||||
"config": {"transport_type": "http", "server_url": "https://mcp.linear.app/mcp", "auth_type": "oauth"},
|
||||
"status": True,
|
||||
},
|
||||
):
|
||||
from flask import request
|
||||
request.decoded_token = {"sub": user}
|
||||
response = MCPServerSave().post()
|
||||
assert response.status_code == 400
|
||||
assert UserToolsRepository(pg_conn).list_for_user(user) == []
|
||||
|
||||
|
||||
class TestMCPOAuthCallback:
|
||||
def test_error_param_redirects_error(self, app):
|
||||
from docsgpt.api.user.tools.mcp import MCPOAuthCallback
|
||||
|
||||
@@ -521,13 +521,44 @@ class TestMcpConnectionScope:
|
||||
def test_writes_never_land_on_another_servers_connection(self, pg_conn):
|
||||
cid = self._mcp(pg_conn)
|
||||
with _patch_service_db(pg_conn):
|
||||
row = service.update_mcp_secrets(
|
||||
"alice", "https://attacker.example", {"tokens": {"access_token": "planted"}}, connection_id=cid,
|
||||
)
|
||||
assert str(row["id"]) != cid
|
||||
with pytest.raises(service.ConnectionUnavailable):
|
||||
service.update_mcp_secrets(
|
||||
"alice", "https://attacker.example", {"tokens": {"access_token": "planted"}}, connection_id=cid,
|
||||
)
|
||||
assert service.read_mcp_secrets("alice", "https://mcp.notion.com", cid)["tokens"]["access_token"] == (
|
||||
"alice-mcp-token"
|
||||
)
|
||||
assert service.read_mcp_secrets("alice", "https://attacker.example") == {}
|
||||
|
||||
def test_a_removed_connection_is_not_brought_back_by_a_late_token_write(self, pg_conn):
|
||||
"""A client or sync still running when its connection is removed may
|
||||
renew the tokens afterwards; that must not re-create the connection."""
|
||||
cid = self._mcp(pg_conn)
|
||||
with _patch_service_db(pg_conn):
|
||||
service.remove_connection(pg_conn, service.ConnectorSessionsRepository(pg_conn).get(cid))
|
||||
with pytest.raises(service.ConnectionUnavailable):
|
||||
service.update_mcp_secrets(
|
||||
"alice", "https://mcp.notion.com", {"tokens": {"access_token": "renewed"}}, connection_id=cid,
|
||||
)
|
||||
assert service.read_mcp_secrets("alice", "https://mcp.notion.com") == {}
|
||||
|
||||
def test_removing_a_connection_forgets_its_cached_mcp_clients(self, pg_conn):
|
||||
import docsgpt.api.user # noqa: F401 (loads mcp_tool without the circular import)
|
||||
from docsgpt.agents.tools import mcp_tool
|
||||
|
||||
cid = self._mcp(pg_conn)
|
||||
mcp_tool._mcp_clients_cache.update({
|
||||
f"https://mcp.notion.com/mcp#http#oauth:{cid}:DocsGPT:none:cb": {"client": object(), "created_at": 0},
|
||||
"https://mcp.notion.com/mcp#http#oauth:alice:DocsGPT:none:cb": {"client": object(), "created_at": 0},
|
||||
"https://mcp.notion.com/mcp#http#oauth:bob:DocsGPT:none:cb": {"client": object(), "created_at": 0},
|
||||
})
|
||||
try:
|
||||
with _patch_service_db(pg_conn):
|
||||
service.remove_connection(pg_conn, service.ConnectorSessionsRepository(pg_conn).get(cid))
|
||||
keys = [k for k in mcp_tool._mcp_clients_cache if "mcp.notion.com" in k]
|
||||
assert keys == ["https://mcp.notion.com/mcp#http#oauth:bob:DocsGPT:none:cb"]
|
||||
finally:
|
||||
mcp_tool._mcp_clients_cache.clear()
|
||||
|
||||
def test_mcp_routes_drop_client_supplied_connection_id(self):
|
||||
"""Only the tool executor may pick the connection whose tokens a tool uses."""
|
||||
|
||||
Reference in new issue
Block a user