feat(api): add MCP user-identity forwarding (#36839)

Co-authored-by: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
This commit is contained in:
Charles Yao
2026-06-08 04:32:11 +00:00
committed by GitHub
co-authored by Claude Opus 4.8 autofix-ci[bot]
parent db1aa683bc
commit 37e1d452b8
21 changed files with 673 additions and 41 deletions
@@ -380,3 +380,53 @@ def test_tool_labels_list(app: Flask, controller_module, monkeypatch: pytest.Mon
resp = controller_module.ToolLabelsApi().get()
assert resp == ["a", "b"]
# --- _resolve_identity_mode: gating + None-resolution (PR #36839 review) ---
def test_resolve_identity_mode_none_keeps_current_when_enterprise(controller_module, monkeypatch: pytest.MonkeyPatch):
"""None means 'leave unchanged' — fall back to the stored mode (update path)."""
identity_mode = importlib.import_module("core.entities.mcp_provider").IdentityMode
monkeypatch.setattr(controller_module.dify_config, "ENTERPRISE_ENABLED", True)
resolved = controller_module._resolve_identity_mode(None, current=identity_mode.IDP_TOKEN)
assert resolved == identity_mode.IDP_TOKEN
def test_resolve_identity_mode_explicit_value_overrides_current(controller_module, monkeypatch: pytest.MonkeyPatch):
"""An explicit value wins over the stored mode."""
identity_mode = importlib.import_module("core.entities.mcp_provider").IdentityMode
monkeypatch.setattr(controller_module.dify_config, "ENTERPRISE_ENABLED", True)
resolved = controller_module._resolve_identity_mode(identity_mode.OFF, current=identity_mode.IDP_TOKEN)
assert resolved == identity_mode.OFF
def test_resolve_identity_mode_coerces_non_off_to_off_when_not_enterprise(
controller_module, monkeypatch: pytest.MonkeyPatch
):
"""Gate: a non-EE deployment must never persist a non-OFF mode — the
runtime won't forward, so the stored row must not imply it does."""
identity_mode = importlib.import_module("core.entities.mcp_provider").IdentityMode
monkeypatch.setattr(controller_module.dify_config, "ENTERPRISE_ENABLED", False)
# Both an explicit idp_token request AND an inherited non-OFF current
# must collapse to OFF.
assert (
controller_module._resolve_identity_mode(identity_mode.IDP_TOKEN, current=identity_mode.OFF)
== identity_mode.OFF
)
assert controller_module._resolve_identity_mode(None, current=identity_mode.IDP_TOKEN) == identity_mode.OFF
def test_resolve_identity_mode_off_is_passthrough_when_not_enterprise(
controller_module, monkeypatch: pytest.MonkeyPatch
):
"""OFF is always fine — the gate only neutralizes non-OFF values."""
identity_mode = importlib.import_module("core.entities.mcp_provider").IdentityMode
monkeypatch.setattr(controller_module.dify_config, "ENTERPRISE_ENABLED", False)
assert controller_module._resolve_identity_mode(None, current=identity_mode.OFF) == identity_mode.OFF
@@ -53,6 +53,7 @@ def test_from_db_model_maps_fields() -> None:
icon=None,
created_at=now,
updated_at=now,
identity_mode="off",
)
# Act
@@ -0,0 +1,41 @@
from __future__ import annotations
import pytest
from core.mcp.auth_client import MCPClientWithAuthRetry
from core.mcp.error import MCPAuthError
class TestForwardIdentityShortCircuit:
def test_forward_identity_active_reraises_without_retry(self):
client = MCPClientWithAuthRetry(
server_url="https://mcp.example.com",
headers={"Authorization": "Bearer user-jwt"},
forward_identity_active=True,
)
with pytest.raises(MCPAuthError):
client._handle_auth_error(MCPAuthError("unauthorized"))
assert client.headers["Authorization"] == "Bearer user-jwt"
assert client._has_retried is False
def test_forward_identity_active_takes_precedence_over_provider_entity(self):
sentinel_entity = object()
client = MCPClientWithAuthRetry(
server_url="https://mcp.example.com",
provider_entity=sentinel_entity, # type: ignore[arg-type]
forward_identity_active=True,
)
with pytest.raises(MCPAuthError, match="forwarded-id-401"):
client._handle_auth_error(MCPAuthError("forwarded-id-401"))
def test_default_path_unchanged_without_provider_entity(self):
client = MCPClientWithAuthRetry(server_url="https://mcp.example.com")
with pytest.raises(MCPAuthError, match="no-provider"):
client._handle_auth_error(MCPAuthError("no-provider"))
def test_default_constructor_defaults_forward_identity_to_false(self):
client = MCPClientWithAuthRetry(server_url="https://mcp.example.com")
assert client.forward_identity_active is False
@@ -148,3 +148,112 @@ def test_mcp_tool_handle_none_parameter_filters_empty_values():
tool = _build_mcp_tool()
cleaned = tool._handle_none_parameter({"a": 1, "b": None, "c": "", "d": " ", "e": "ok"})
assert cleaned == {"a": 1, "e": "ok"}
# ----- M2/M3 user-identity forwarding ---------------------------------------
def _build_forwarding_tool(*, mode: str = "idp_token") -> MCPTool:
"""Helper that builds an MCPTool with the identity_mode set."""
entity = ToolEntity(
identity=ToolIdentity(
author="author",
name="remote-tool",
label=I18nObject(en_US="remote-tool"),
provider="provider-id",
),
parameters=[],
output_schema={},
)
return MCPTool(
entity=entity,
runtime=ToolRuntime(tenant_id="tenant-1", invoke_from=InvokeFrom.DEBUGGER),
tenant_id="tenant-1",
icon="icon.svg",
server_url="https://mcp.example.com/mcp/",
provider_id="provider-id",
identity_mode=mode,
)
def test_inject_forwarded_identity_stamps_custom_header():
"""The minted SSO token must be placed in X-Dify-SSO-Access-Token; the
workspace-scoped Authorization header and any other custom headers must
pass through untouched so provider credentials keep working."""
from core.tools.mcp_tool.tool import FORWARDED_IDENTITY_HEADER
tool = _build_forwarding_tool()
headers: dict[str, str] = {"Authorization": "Bearer static-client-token", "X-Other": "keep"}
with patch(
"services.enterprise.enterprise_service.EnterpriseService.issue_mcp_token",
return_value=("forwarded.jwt.payload", 1900000000),
):
tool._inject_forwarded_identity(headers, user_id="alice", app_id=None, audience="https://mcp.example.com/mcp/")
assert headers[FORWARDED_IDENTITY_HEADER] == "forwarded.jwt.payload"
assert headers["Authorization"] == "Bearer static-client-token"
assert headers["X-Other"] == "keep"
def test_inject_forwarded_identity_translates_token_error_to_invoke_error():
"""EnterpriseService failures must surface as ToolInvokeError so the
workflow halts loudly instead of proceeding without identity."""
from core.tools.mcp_tool.tool import FORWARDED_IDENTITY_HEADER
from services.enterprise.base import MCPNoRefreshTokenError
tool = _build_forwarding_tool()
headers: dict[str, str] = {}
with patch(
"services.enterprise.enterprise_service.EnterpriseService.issue_mcp_token",
side_effect=MCPNoRefreshTokenError("please re-sso"),
):
with pytest.raises(ToolInvokeError, match="forwarded identity token"):
tool._inject_forwarded_identity(
headers, user_id="alice", app_id=None, audience="https://mcp.example.com/mcp/"
)
# Headers must NOT have been mutated when token-issuance failed.
assert FORWARDED_IDENTITY_HEADER not in headers
assert "Authorization" not in headers
def test_invoke_remote_mcp_tool_fails_closed_when_user_id_missing():
"""When forwarding is enabled AND the deployment is enterprise, missing
user_id must raise — never silently invoke as the static identity."""
tool = _build_forwarding_tool()
with patch("core.tools.mcp_tool.tool.dify_config") as cfg:
cfg.ENTERPRISE_ENABLED = True
with pytest.raises(ToolInvokeError, match="no end-user context"):
tool.invoke_remote_mcp_tool({}, user_id=None, app_id=None)
def test_invoke_skips_forwarding_when_enterprise_disabled():
"""Non-enterprise deployments treat the DB selector as a no-op: a stale
`identity_mode="idp_token"` row must NOT raise (fail-closed) AND must
NOT call the enterprise inner API. The runtime falls through to the
legacy provider-identity path."""
tool = _build_forwarding_tool()
with patch("core.tools.mcp_tool.tool.dify_config") as cfg:
cfg.ENTERPRISE_ENABLED = False
# The fail-closed branch must NOT fire (no enterprise → no forwarding).
# The function will still try the legacy DB-load path; we patch that
# to keep the test unit-scoped.
with patch("core.tools.mcp_tool.tool.MCPClientWithAuthRetry") as client_cls:
client_cls.return_value.__enter__.return_value.invoke_tool.return_value = CallToolResult(
content=[],
_meta=None,
)
with patch.object(tool, "_inject_forwarded_identity") as inject:
with patch("services.tools.mcp_tools_manage_service.MCPToolManageService"):
with patch("core.entities.mcp_provider.MCPProviderEntity.decrypt_server_url", return_value="u"):
with patch("core.entities.mcp_provider.MCPProviderEntity.decrypt_headers", return_value={}):
# Should not raise; should not call enterprise.
try:
tool.invoke_remote_mcp_tool({}, user_id=None, app_id=None)
except Exception:
pass
inject.assert_not_called()
@@ -596,7 +596,9 @@ def test_api_tool_create_records_id_mapping(monkeypatch):
def test_mcp_tool_import_restores_exported_tool_list(monkeypatch):
provider = type("Provider", (), {"id": "target-provider-id", "tools": "[]", "authed": False})()
provider = type(
"Provider", (), {"id": "target-provider-id", "tools": "[]", "authed": False, "identity_mode": "off"}
)()
report_items = []
class StubSession:
@@ -652,7 +654,9 @@ def test_mcp_tool_import_restores_exported_tool_list(monkeypatch):
@pytest.mark.parametrize("conflict_strategy", [ConflictStrategy.SKIP, ConflictStrategy.UPDATE])
def test_mcp_tool_existing_provider_records_id_mapping(monkeypatch, conflict_strategy):
provider = type("Provider", (), {"id": "target-mcp-provider-id", "tools": "[]", "authed": False})()
provider = type(
"Provider", (), {"id": "target-mcp-provider-id", "tools": "[]", "authed": False, "identity_mode": "off"}
)()
id_mapping = {}
id_mapping_details = []
@@ -712,7 +716,9 @@ def test_mcp_tool_existing_provider_records_id_mapping(monkeypatch, conflict_str
def test_mcp_tool_create_records_id_mapping(monkeypatch):
provider = type("Provider", (), {"id": "target-mcp-provider-id", "tools": "[]", "authed": False})()
provider = type(
"Provider", (), {"id": "target-mcp-provider-id", "tools": "[]", "authed": False, "identity_mode": "off"}
)()
id_mapping = {}
provider_created = False
@@ -466,3 +466,111 @@ class TestGetCachedLicenseStatus:
assert EnterpriseService.get_cached_license_status() is None
mock_redis.setex.assert_not_called()
class TestIssueMCPToken:
"""Coverage for EnterpriseService.issue_mcp_token (M2).
The function wraps `POST /inner/api/mcp/issue-token` and must map
EnterpriseServiceError subclasses to MCP-typed errors so the workflow
layer can halt with a precise message instead of leaking transport text.
"""
@staticmethod
def _call():
return EnterpriseService.issue_mcp_token(
user_id="user-uuid",
tenant_id="tenant-uuid",
app_id="app-uuid",
audience="https://mcp.example.com/mcp/",
)
def test_happy_path_returns_token_and_expiry(self):
with patch(f"{MODULE}.EnterpriseRequest") as req:
req.send_request.return_value = {"token": "abc.def.ghi", "expires_at": 1900000000}
token, exp = self._call()
assert token == "abc.def.ghi"
assert exp == 1900000000
req.send_request.assert_called_once_with(
"POST",
"/mcp/issue-token",
json={
"user_id": "user-uuid",
"tenant_id": "tenant-uuid",
"app_id": "app-uuid",
"audience": "https://mcp.example.com/mcp/",
},
)
def test_401_maps_to_identity_refresh_error(self):
from services.enterprise.base import MCPIdentityRefreshError
from services.errors.enterprise import EnterpriseAPIUnauthorizedError
with patch(f"{MODULE}.EnterpriseRequest") as req:
req.send_request.side_effect = EnterpriseAPIUnauthorizedError("refresh rejected by IdP")
with pytest.raises(MCPIdentityRefreshError, match="refresh rejected"):
self._call()
def test_428_maps_to_no_refresh_token_error(self):
from services.enterprise.base import MCPNoRefreshTokenError
from services.errors.enterprise import EnterpriseAPIError
with patch(f"{MODULE}.EnterpriseRequest") as req:
# 428 PreconditionRequired is what EE returns when there's no
# stored SSO refresh token for the user.
req.send_request.side_effect = EnterpriseAPIError("user has not completed SSO", status_code=428)
with pytest.raises(MCPNoRefreshTokenError, match="SSO"):
self._call()
def test_403_maps_to_identity_refresh_error_for_license(self):
from services.enterprise.base import MCPIdentityRefreshError
from services.errors.enterprise import EnterpriseAPIForbiddenError
with patch(f"{MODULE}.EnterpriseRequest") as req:
req.send_request.side_effect = EnterpriseAPIForbiddenError("not licensed for MCP forwarding")
with pytest.raises(MCPIdentityRefreshError, match="not licensed"):
self._call()
def test_other_status_maps_to_generic_token_error(self):
from services.enterprise.base import MCPTokenError
from services.errors.enterprise import EnterpriseAPIError
with patch(f"{MODULE}.EnterpriseRequest") as req:
req.send_request.side_effect = EnterpriseAPIError("upstream 502", status_code=502)
with pytest.raises(MCPTokenError, match="status=502"):
self._call()
def test_malformed_response_shape_raises_token_error(self):
from services.enterprise.base import MCPTokenError
with patch(f"{MODULE}.EnterpriseRequest") as req:
req.send_request.return_value = "not-a-dict"
with pytest.raises(MCPTokenError, match="invalid response shape"):
self._call()
def test_missing_token_field_raises_token_error(self):
from services.enterprise.base import MCPTokenError
with patch(f"{MODULE}.EnterpriseRequest") as req:
req.send_request.return_value = {"expires_at": 1700000000} # no token
with pytest.raises(MCPTokenError, match="missing or non-string token"):
self._call()
def test_float_expires_at_is_accepted(self):
"""expires_at may arrive as float (time.time()) — must be coerced."""
with patch(f"{MODULE}.EnterpriseRequest") as req:
req.send_request.return_value = {"token": "t", "expires_at": 1900000000.5}
token, exp = self._call()
assert token == "t"
assert exp == 1900000000
assert isinstance(exp, int)
def test_bool_expires_at_is_rejected(self):
"""bool is a subclass of int — must NOT be accepted as expires_at."""
from services.enterprise.base import MCPTokenError
with patch(f"{MODULE}.EnterpriseRequest") as req:
req.send_request.return_value = {"token": "t", "expires_at": True}
with pytest.raises(MCPTokenError, match="non-numeric expires_at"):
self._call()