mirror of
https://github.com/langgenius/dify.git
synced 2026-09-24 23:22:26 +08:00
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:
co-authored by
Claude Opus 4.8
autofix-ci[bot]
parent
db1aa683bc
commit
37e1d452b8
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user