From 39093c0749da9cdac7464a2d3097bc26ed2e37bf Mon Sep 17 00:00:00 2001 From: fatelei Date: Mon, 31 Aug 2026 17:18:15 +0800 Subject: [PATCH] feat: agent support rbac --- api/controllers/console/agent/roster.py | 56 +++++++ api/controllers/console/workspace/rbac.py | 40 +++++ api/services/enterprise/rbac_service.py | 155 +++++++++++++++++- ...initialize_created_app_rbac_access_task.py | 67 +++++++- .../console/agent/test_agent_controllers.py | 57 +++++++ .../console/workspace/test_rbac.py | 61 +++++++ .../services/enterprise/test_rbac_service.py | 137 +++++++++++++++- ...initialize_created_app_rbac_access_task.py | 47 ++++++ 8 files changed, 616 insertions(+), 4 deletions(-) diff --git a/api/controllers/console/agent/roster.py b/api/controllers/console/agent/roster.py index f95f9d8b00f..53c07e7e2c8 100644 --- a/api/controllers/console/agent/roster.py +++ b/api/controllers/console/agent/roster.py @@ -7,6 +7,7 @@ from pydantic import AliasChoices, BaseModel, Field, field_validator from sqlalchemy import func, or_, select from sqlalchemy.orm import Session +from configs import dify_config from controllers.common.schema import ( query_params_from_model, query_params_from_request, @@ -78,9 +79,11 @@ from services.agent.observability_service import ( ) from services.agent.roster_service import AgentRosterService from services.app_service import AgentAppPublicationCounts, AppListParams, AppService, CreateAppParams +from services.enterprise import rbac_service as enterprise_rbac_service from services.enterprise.enterprise_service import EnterpriseService from services.entities.agent_entities import ComposerSavePayload, RosterListQuery from services.feature_service import FeatureService +from tasks.initialize_created_app_rbac_access_task import initialize_created_app_rbac_access_task AgentPublicationStatus = Literal["published", "drafts"] @@ -429,6 +432,14 @@ def _serialize_agent_app_detail( payload["debug_conversation_message_count"] = message_count payload["role"] = agent.role or "" payload["access_ready"] = agent_has_workflow_callable_active_snapshot(session=session, agent=agent) + if dify_config.RBAC_ENABLED: + permission_keys_map = enterprise_rbac_service.RBACService.AgentPermissions.batch_get( + str(app_model.tenant_id), + current_user.id, + [str(agent.id)], + session=session, + ) + payload["permission_keys"] = permission_keys_map.get(str(agent.id), []) return payload @@ -454,6 +465,15 @@ def _serialize_agent_app_pagination( tenant_id=tenant_id, app_ids=app_ids, ) + agent_ids = [str(agent.id) for agent in agents_by_app_id.values()] + permission_keys_by_agent_id: dict[str, list[str]] = {} + if dify_config.RBAC_ENABLED: + permission_keys_by_agent_id = enterprise_rbac_service.RBACService.AgentPermissions.batch_get( + tenant_id, + current_user.id, + agent_ids, + session=session, + ) active_config_is_published_by_agent_id = roster_service.load_active_config_is_published_by_agent_id( tenant_id=tenant_id, agents=list(agents_by_app_id.values()), @@ -496,6 +516,7 @@ def _serialize_agent_app_pagination( item["id"] = agent.id item["debug_conversation_id"] = debug_conversation_ids_by_agent_id.get(agent.id) item["role"] = agent.role or "" + item["permission_keys"] = permission_keys_by_agent_id.get(str(agent.id), []) item["active_config_is_published"] = active_config_is_published_by_agent_id.get(agent.id, False) item["reference_count"] = reference_counts_by_agent_id.get(agent.id, 0) published_references = published_references_by_agent_id.get(agent.id, []) @@ -567,6 +588,29 @@ def _serialize_agent_api_access(session: Session, app_model: App) -> dict: return response.model_dump(mode="json") +def _initialize_created_agent_rbac_access( + session: Session, + *, + tenant_id: str, + account_id: str, + app_model: App, +) -> None: + if not dify_config.RBAC_ENABLED: + return + + agent = _agent_roster_service(session).get_app_backing_agent(tenant_id=tenant_id, app_id=str(app_model.id)) + if not agent: + raise AgentNotFoundError() + + enterprise_rbac_service.RBACService.AgentAccess.replace_whitelist( + tenant_id, + account_id, + str(agent.id), + enterprise_rbac_service.ReplaceMemberBindings(automatic_include_workspace_members=True), + ) + initialize_created_app_rbac_access_task.delay(tenant_id, account_id, agent_id=str(agent.id)) + + def _agent_observability_service(session: Session) -> AgentObservabilityService: return AgentObservabilityService(session) @@ -684,6 +728,12 @@ class AgentAppListApi(Resource): ) app = AppService().create_app(current_tenant_id, params, current_user, session=session) + _initialize_created_agent_rbac_access( + session, + tenant_id=current_tenant_id, + account_id=current_user.id, + app_model=app, + ) return _serialize_agent_app_detail(session, app, current_user=current_user), 201 @@ -957,6 +1007,12 @@ class AgentAppCopyApi(Resource): icon=req_data.icon, icon_background=req_data.icon_background, ) + _initialize_created_agent_rbac_access( + session, + tenant_id=tenant_id, + account_id=current_user.id, + app_model=copied_app, + ) return _serialize_agent_app_detail(session, copied_app, current_user=current_user), 201 diff --git a/api/controllers/console/workspace/rbac.py b/api/controllers/console/workspace/rbac.py index 96b882871bf..f4651a36677 100644 --- a/api/controllers/console/workspace/rbac.py +++ b/api/controllers/console/workspace/rbac.py @@ -70,16 +70,19 @@ _LEGACY_ROLE_PERMISSION_KEYS: dict[str, list[str]] = { *svc._LEGACY_WORKSPACE_OWNER_KEYS, *svc._LEGACY_APP_OWNER_KEYS, *svc._LEGACY_DATASET_OWNER_KEYS, + *svc._LEGACY_AGENT_OWNER_KEYS, ], "admin": [ *svc._LEGACY_WORKSPACE_ADMIN_KEYS, *svc._LEGACY_APP_ADMIN_KEYS, *svc._LEGACY_DATASET_ADMIN_KEYS, + *svc._LEGACY_AGENT_ADMIN_KEYS, ], "editor": [ *svc._LEGACY_WORKSPACE_EDITOR_KEYS, *svc._LEGACY_APP_EDITOR_KEYS, *svc._LEGACY_DATASET_EDITOR_KEYS, + *svc._LEGACY_AGENT_EDITOR_KEYS, ], "normal": [ *svc._LEGACY_WORKSPACE_NORMAL_KEYS, @@ -797,6 +800,43 @@ class RBACDatasetUserAccessPoliciesApi(Resource): return _dump(result) +# --------------------------------------------------------------------------- +# Per-agent access (Agent Access Config). +# --------------------------------------------------------------------------- + + +@console_ns.route("/workspaces/current/rbac/agents//whitelist_config") +class RBACAgentWhitelistConfigApi(Resource): + @login_required + @console_ns.response(200, "Success", console_ns.models[svc.ResourceWhitelistConfig.__name__]) + def get(self, agent_id): + tenant_id, account_id = _current_ids() + return _dump(svc.RBACService.AgentAccess.whitelist_config(tenant_id, account_id, str(agent_id))) + + +@console_ns.route("/workspaces/current/rbac/agents//whitelist") +class RBACAgentWhitelistApi(Resource): + @login_required + @console_ns.response(200, "Success", console_ns.models[svc.ResourceWhitelist.__name__]) + def get(self, agent_id): + tenant_id, account_id = _current_ids() + return _dump(svc.RBACService.AgentAccess.whitelist(tenant_id, account_id, str(agent_id))) + + @login_required + @console_ns.expect(console_ns.models[_ResourceAccessScopeRequest.__name__]) + @console_ns.response(200, "Success", console_ns.models[svc.ResourceWhitelist.__name__]) + def put(self, agent_id): + tenant_id, account_id = _current_ids() + request = _payload(_ResourceAccessScopeRequest) + result = svc.RBACService.AgentAccess.replace_whitelist( + tenant_id, + account_id, + str(agent_id), + svc.ReplaceMemberBindings(automatic_include_workspace_members=request.automatic_include_workspace_members), + ) + return _dump(result) + + @console_ns.route("/workspaces/current/rbac/datasets//users//access-policies") class RBACDatasetUserAccessPolicyAssignmentApi(Resource): @login_required diff --git a/api/services/enterprise/rbac_service.py b/api/services/enterprise/rbac_service.py index 5f206548cc5..629e481ded8 100644 --- a/api/services/enterprise/rbac_service.py +++ b/api/services/enterprise/rbac_service.py @@ -58,6 +58,7 @@ class RBACResourceType(StrEnum): APP = "app" DATASET = "dataset" + AGENT = "agent" class RBACRoleType(StrEnum): @@ -339,6 +340,12 @@ class AppendDatasetWhitelistMembersBatchItem(_RBACModel): policy_id: str +class AppendAgentWhitelistMembersBatchItem(_RBACModel): + agent_id: str + account_ids: list[str] = Field(default_factory=list) + policy_id: str + + class MemberRolesResponse(_RBACModel): account_id: str roles: list[RBACRole] = Field(default_factory=list) @@ -369,6 +376,7 @@ class MyPermissionsResponse(_RBACModel): workspace: WorkspacePermissionSnapshot = Field(default_factory=WorkspacePermissionSnapshot) app: ResourcePermissionSnapshot = Field(default_factory=ResourcePermissionSnapshot) dataset: ResourcePermissionSnapshot = Field(default_factory=ResourcePermissionSnapshot) + agent: ResourcePermissionSnapshot = Field(default_factory=ResourcePermissionSnapshot) # Fallback permission snapshots for legacy Dify tenant roles when external RBAC is disabled. @@ -405,6 +413,7 @@ _LEGACY_WORKSPACE_OWNER_KEYS: list[str] = [ "tool.manage", "mcp.manage", "agent.manage", + "agent.create", ] _LEGACY_WORKSPACE_ADMIN_KEYS: list[str] = [ @@ -437,6 +446,7 @@ _LEGACY_WORKSPACE_ADMIN_KEYS: list[str] = [ "tool.manage", "mcp.manage", "agent.manage", + "agent.create", ] _LEGACY_WORKSPACE_EDITOR_KEYS: list[str] = [ @@ -456,6 +466,7 @@ _LEGACY_WORKSPACE_EDITOR_KEYS: list[str] = [ "snippets.create_and_modify", "tool.manage", "agent.manage", + "agent.create", ] _LEGACY_WORKSPACE_NORMAL_KEYS: list[str] = [ @@ -576,21 +587,55 @@ _LEGACY_DATASET_DATASET_OPERATOR_KEYS: list[str] = [ "dataset.acl.pipeline_release", ] +_LEGACY_AGENT_OWNER_KEYS: list[str] = [ + "agent.acl.preview", + "agent.acl.edit", + "agent.acl.release_and_version", + "agent.acl.access_point_manage", + "agent.acl.log_manage", + "agent.acl.monitor", + "agent.acl.access_config", + "agent.acl.import_export_dsl", + "agent.acl.delete", +] + +_LEGACY_AGENT_ADMIN_KEYS: list[str] = [ + "agent.acl.preview", + "agent.acl.edit", + "agent.acl.release_and_version", + "agent.acl.access_point_manage", + "agent.acl.log_manage", + "agent.acl.monitor", + "agent.acl.access_config", + "agent.acl.import_export_dsl", + "agent.acl.delete", +] + +_LEGACY_AGENT_EDITOR_KEYS: list[str] = [ + "agent.acl.preview", + "agent.acl.edit", + "agent.acl.release_and_version", + "agent.acl.access_point_manage", +] + _LEGACY_MY_PERMISSIONS: dict[TenantAccountRole, dict[str, list[str]]] = { TenantAccountRole.OWNER: { "workspace": _LEGACY_WORKSPACE_OWNER_KEYS, "app": _LEGACY_APP_OWNER_KEYS, "dataset": _LEGACY_DATASET_OWNER_KEYS, + "agent": _LEGACY_AGENT_OWNER_KEYS, }, TenantAccountRole.ADMIN: { "workspace": _LEGACY_WORKSPACE_ADMIN_KEYS, "app": _LEGACY_APP_ADMIN_KEYS, "dataset": _LEGACY_DATASET_ADMIN_KEYS, + "agent": _LEGACY_AGENT_ADMIN_KEYS, }, TenantAccountRole.EDITOR: { "workspace": _LEGACY_WORKSPACE_EDITOR_KEYS, "app": _LEGACY_APP_EDITOR_KEYS, "dataset": _LEGACY_DATASET_EDITOR_KEYS, + "agent": _LEGACY_AGENT_EDITOR_KEYS, }, TenantAccountRole.NORMAL: { "workspace": _LEGACY_WORKSPACE_NORMAL_KEYS, @@ -611,6 +656,7 @@ def _legacy_role_permission_keys(role: TenantAccountRole) -> list[str]: *permissions.get("workspace", []), *permissions.get("app", []), *permissions.get("dataset", []), + *permissions.get("agent", []), ] ) ) @@ -667,6 +713,7 @@ def _legacy_my_permissions(tenant_id: str, account_id: str | None, *, session: S workspace=WorkspacePermissionSnapshot(permission_keys=list(permissions.get("workspace", []))), app=ResourcePermissionSnapshot(default_permission_keys=list(permissions.get("app", []))), dataset=ResourcePermissionSnapshot(default_permission_keys=list(permissions.get("dataset", []))), + agent=ResourcePermissionSnapshot(default_permission_keys=list(permissions.get("agent", []))), ) @@ -681,6 +728,8 @@ def _legacy_resource_permission_keys_batch( snapshot = _legacy_my_permissions(tenant_id, account_id, session=session) if resource_type == RBACResourceType.APP: permission_keys = snapshot.app.default_permission_keys + elif resource_type == RBACResourceType.AGENT: + permission_keys = snapshot.agent.default_permission_keys else: permission_keys = snapshot.dataset.default_permission_keys return {str(resource_id): list(permission_keys) for resource_id in resource_ids} @@ -1605,6 +1654,81 @@ class RBACService: ) return AccessMatrixItem.model_validate(data or {}) + # ------------------------------------------------------------------ + # Per-agent access. + # ------------------------------------------------------------------ + class AgentAccess: + @staticmethod + def replace_user_access_policies( + tenant_id: str, + account_id: str | None, + agent_id: str, + target_account_id: str | None, + payload: ReplaceUserAccessPolicies, + ) -> ReplaceUserAccessPoliciesResponse: + data = _inner_call( + "PUT", + f"{_INNER_PREFIX}/agents/user-access-policies", + tenant_id=tenant_id, + account_id=account_id, + params={"agent_id": agent_id, "account_id": target_account_id}, + json=payload.model_dump(mode="json", exclude_unset=True), + ) + return ReplaceUserAccessPoliciesResponse.model_validate(data or {}) + + @staticmethod + def whitelist(tenant_id: str, account_id: str | None, agent_id: str) -> ResourceWhitelist: + data = _inner_call( + "GET", + f"{_INNER_PREFIX}/agents/whitelist", + tenant_id=tenant_id, + account_id=account_id, + params={"agent_id": agent_id}, + ) + return ResourceWhitelist.model_validate(data or {}) + + @staticmethod + def whitelist_config(tenant_id: str, account_id: str | None, agent_id: str) -> ResourceWhitelistConfig: + data = _inner_call( + "GET", + f"{_INNER_PREFIX}/agents/whitelist", + tenant_id=tenant_id, + account_id=account_id, + params={"agent_id": agent_id}, + ) + return ResourceWhitelistConfig.model_validate(data or {}) + + @staticmethod + def replace_whitelist( + tenant_id: str, + account_id: str | None, + agent_id: str, + payload: ReplaceMemberBindings, + ) -> ResourceWhitelist: + data = _inner_call( + "PUT", + f"{_INNER_PREFIX}/agents/whitelist", + tenant_id=tenant_id, + account_id=account_id, + params={"agent_id": agent_id}, + json=payload.model_dump(mode="json"), + ) + return ResourceWhitelist.model_validate(data or {}) + + @staticmethod + def append_whitelist_members_batch( + tenant_id: str, + account_id: str | None, + data: Sequence[AppendAgentWhitelistMembersBatchItem], + ) -> None: + _inner_call( + "POST", + f"{_INNER_PREFIX}/agents/whitelist/members/batch", + tenant_id=tenant_id, + account_id=account_id, + json={"data": [item.model_dump(mode="json") for item in data]}, + ) + # ------------------------------------------------------------------ # Workspace-level access (screenshot 2: Settings > Access Rules). # ------------------------------------------------------------------ @@ -1957,6 +2081,30 @@ class RBACService: ) return _parse_resource_permission_keys_batch(data, resource_id_key="dataset_id") + class AgentPermissions: + @staticmethod + def batch_get( + tenant_id: str, + account_id: str | None, + agent_ids: list[str], + *, + session: Session, + ) -> dict[str, list[str]]: + if not agent_ids: + return {} + if not dify_config.RBAC_ENABLED: + return _legacy_resource_permission_keys_batch( + tenant_id, account_id, agent_ids, RBACResourceType.AGENT, session=session + ) + data = _inner_call( + "POST", + f"{_INNER_PREFIX}/agents/permission-keys/batch", + tenant_id=tenant_id, + account_id=account_id, + json={"agent_ids": agent_ids}, + ) + return _parse_resource_permission_keys_batch(data, resource_id_key="agent_id") + class MyPermissions: @staticmethod def get( @@ -2001,7 +2149,12 @@ def _parse_resource_permission_keys_batch(data: Any, *, resource_id_key: str) -> if items is None: items = data.get("items") if items is None: - items = data.get("apps") if resource_id_key == "app_id" else data.get("datasets") + collection_key_by_resource_id_key = { + "app_id": "apps", + "dataset_id": "datasets", + "agent_id": "agents", + } + items = data.get(collection_key_by_resource_id_key.get(resource_id_key, "")) if isinstance(items, dict): items = [{"resource_id": key, "permission_keys": value} for key, value in items.items()] elif isinstance(data, list): diff --git a/api/tasks/initialize_created_app_rbac_access_task.py b/api/tasks/initialize_created_app_rbac_access_task.py index ad69f2374ac..07ba230f008 100644 --- a/api/tasks/initialize_created_app_rbac_access_task.py +++ b/api/tasks/initialize_created_app_rbac_access_task.py @@ -9,6 +9,7 @@ from sqlalchemy import select from configs import dify_config from extensions.ext_database import db from models import App, Dataset, TenantAccountJoin, TenantAccountRole +from models.agent import Agent, AgentScope, AgentStatus from services.account_service import TenantService from services.enterprise import rbac_service as enterprise_rbac_service @@ -68,6 +69,32 @@ def _iter_resource_config_batches( ] last_dataset_id = dataset_ids[-1] + last_agent_id: str | None = None + while True: + stmt = ( + select(Agent.id) + .where( + Agent.tenant_id == tenant_id, + Agent.scope == AgentScope.ROSTER, + Agent.status == AgentStatus.ACTIVE, + ) + .order_by(Agent.id.asc()) + .limit(batch_size) + ) + if last_agent_id: + stmt = stmt.where(Agent.id > last_agent_id) + agent_ids = [str(agent_id) for agent_id in db.session().scalars(stmt).all()] + if not agent_ids: + break + yield [ + enterprise_rbac_service.ResourceWhitelistConfigResource( + resource_type=enterprise_rbac_service.RBACResourceType.AGENT, + resource_id=agent_id, + ) + for agent_id in agent_ids + ] + last_agent_id = agent_ids[-1] + def _chunks[T](items: list[T], chunk_size: int) -> Iterator[list[T]]: for index in range(0, len(items), chunk_size): @@ -76,7 +103,12 @@ def _chunks[T](items: list[T], chunk_size: int) -> Iterator[list[T]]: @shared_task(queue=APP_RBAC_QUEUE, bind=True, max_retries=3, default_retry_delay=60) def initialize_created_app_rbac_access_task( - self, tenant_id: str, account_id: str, app_id: str | None = None, dataset_id: str | None = None + self, + tenant_id: str, + account_id: str, + app_id: str | None = None, + dataset_id: str | None = None, + agent_id: str | None = None, ) -> None: """Grant the default app policy to current workspace members. @@ -115,11 +147,25 @@ def initialize_created_app_rbac_access_task( account_ids=account_ids, ), ) + elif agent_id is not None: + enterprise_rbac_service.RBACService.AgentAccess.replace_user_access_policies( + tenant_id=tenant_id, + account_id=account_id, + agent_id=agent_id, + target_account_id=None, + payload=enterprise_rbac_service.ReplaceUserAccessPolicies( + access_policy_ids=[APP_RBAC_DEFAULT_ACCESS_POLICY_ID], + account_ids=account_ids, + ), + ) except Exception as exc: logger.exception( - "Failed to initialize app RBAC access; retrying: tenant_id=%s app_id=%s attempt=%s", + "Failed to initialize app RBAC access; retrying: " + "tenant_id=%s app_id=%s dataset_id=%s agent_id=%s attempt=%s", tenant_id, app_id, + dataset_id, + agent_id, self.request.retries + 1, ) raise self.retry(exc=exc) @@ -148,6 +194,7 @@ def sync_joined_workspace_member_rbac_access_task( app_ids: list[str] = [] dataset_ids: list[str] = [] + agent_ids: list[str] = [] for resources in _iter_resource_config_batches(tenant_id, APP_RBAC_RESOURCE_CONFIG_BATCH_SIZE): configs = enterprise_rbac_service.RBACService.ResourceWhitelistConfigs.batch_get( tenant_id=tenant_id, @@ -161,6 +208,8 @@ def sync_joined_workspace_member_rbac_access_task( app_ids.append(config.resource_id) elif config.resource_type == enterprise_rbac_service.RBACResourceType.DATASET: dataset_ids.append(config.resource_id) + elif config.resource_type == enterprise_rbac_service.RBACResourceType.AGENT: + agent_ids.append(config.resource_id) for app_id_batch in _chunks(app_ids, APP_RBAC_MEMBER_APPEND_BATCH_SIZE): enterprise_rbac_service.RBACService.AppAccess.append_whitelist_members_batch( @@ -176,6 +225,20 @@ def sync_joined_workspace_member_rbac_access_task( ], ) + for agent_id_batch in _chunks(agent_ids, APP_RBAC_MEMBER_APPEND_BATCH_SIZE): + enterprise_rbac_service.RBACService.AgentAccess.append_whitelist_members_batch( + tenant_id=tenant_id, + account_id=actor_account_id, + data=[ + enterprise_rbac_service.AppendAgentWhitelistMembersBatchItem( + agent_id=agent_id, + account_ids=[member_account_id], + policy_id=APP_RBAC_DEFAULT_ACCESS_POLICY_ID, + ) + for agent_id in agent_id_batch + ], + ) + for dataset_id_batch in _chunks(dataset_ids, APP_RBAC_MEMBER_APPEND_BATCH_SIZE): enterprise_rbac_service.RBACService.DatasetAccess.append_whitelist_members_batch( tenant_id=tenant_id, diff --git a/api/tests/unit_tests/controllers/console/agent/test_agent_controllers.py b/api/tests/unit_tests/controllers/console/agent/test_agent_controllers.py index 112b3a0264f..0b4e4b4d97f 100644 --- a/api/tests/unit_tests/controllers/console/agent/test_agent_controllers.py +++ b/api/tests/unit_tests/controllers/console/agent/test_agent_controllers.py @@ -310,6 +310,7 @@ def test_agent_app_list_and_create_use_agent_route( app: Flask, monkeypatch: pytest.MonkeyPatch, account_id: str, sqlite_session: Session ) -> None: captured: dict[str, object] = {} + monkeypatch.setattr(roster_controller.dify_config, "RBAC_ENABLED", True) class FakeAppService: def get_app(self, app_obj: object, *, session: object) -> object: @@ -397,6 +398,25 @@ def test_agent_app_list_and_create_use_agent_route( monkeypatch.setattr( roster_controller.AgentRosterService, "count_agent_app_debug_conversation_messages", lambda _self, **kwargs: 0 ) + monkeypatch.setattr( + roster_controller.enterprise_rbac_service.RBACService.AgentPermissions, + "batch_get", + lambda _tenant_id, _account_id, agent_ids, **_kwargs: { + agent_id: [f"permission:{agent_id}"] for agent_id in agent_ids + }, + ) + replace_agent_whitelist = MagicMock() + initialize_agent_rbac_access = MagicMock() + monkeypatch.setattr( + roster_controller.enterprise_rbac_service.RBACService.AgentAccess, + "replace_whitelist", + replace_agent_whitelist, + ) + monkeypatch.setattr( + roster_controller.initialize_created_app_rbac_access_task, + "delay", + initialize_agent_rbac_access, + ) def get_or_create_debug_conversation(_self: object, **kwargs: object) -> str: captured["get_or_create_debug_conversation"] = kwargs @@ -427,6 +447,7 @@ def test_agent_app_list_and_create_use_agent_route( assert listed["data"][0]["app_id"] == "app-list" assert listed["data"][0]["debug_conversation_id"] == "debug-conversation-list" assert listed["data"][0]["role"] == "List role" + assert listed["data"][0]["permission_keys"] == ["permission:agent-list"] assert listed["data"][0]["active_config_is_published"] is False assert listed["data"][0]["reference_count"] == 2 assert listed["data"][0]["published_reference_count"] == 1 @@ -468,12 +489,17 @@ def test_agent_app_list_and_create_use_agent_route( assert created["app_id"] == "app-created" assert created["debug_conversation_id"] == "debug-conversation-created" assert created["role"] == "Created role" + assert created["permission_keys"] == ["permission:agent-created"] assert "active_config_is_published" not in created assert "bound_agent_id" not in created create_call = cast(dict[str, object], captured["create"]) create_params = cast(Any, create_call["params"]) assert create_params.mode == "agent" assert create_params.agent_role == "Coordinator" + replace_agent_whitelist.assert_called_once() + assert replace_agent_whitelist.call_args.args[:3] == ("tenant-1", account_id, "agent-created") + assert replace_agent_whitelist.call_args.args[3].automatic_include_workspace_members is True + initialize_agent_rbac_access.assert_called_once_with("tenant-1", account_id, agent_id="agent-created") assert captured["get_or_create_debug_conversation"] == { "tenant_id": "tenant-1", "agent_id": "agent-created", @@ -534,6 +560,7 @@ def test_agent_app_create_omits_optional_role_as_empty_string( def test_agent_app_detail_update_delete_resolve_app_from_agent_id( app: Flask, monkeypatch: pytest.MonkeyPatch, account_id: str, sqlite_session: Session ) -> None: + monkeypatch.setattr(roster_controller.dify_config, "RBAC_ENABLED", True) agent_id = "00000000-0000-0000-0000-000000000001" tenant_id = "00000000-0000-0000-0000-000000000002" app_id = "00000000-0000-0000-0000-000000000003" @@ -573,6 +600,13 @@ def test_agent_app_detail_update_delete_resolve_app_from_agent_id( "agent_has_workflow_callable_active_snapshot", lambda **_kwargs: False, ) + monkeypatch.setattr( + roster_controller.enterprise_rbac_service.RBACService.AgentPermissions, + "batch_get", + lambda _tenant_id, _account_id, agent_ids, **_kwargs: { + agent_id: [f"permission:{agent_id}"] for agent_id in agent_ids + }, + ) class FakeAppService: def get_app(self, app_obj: object, *, session: object) -> object: @@ -597,6 +631,7 @@ def test_agent_app_detail_update_delete_resolve_app_from_agent_id( assert detail["debug_conversation_message_count"] == 2 assert detail["role"] == "Resolved role" assert detail["access_ready"] is False + assert detail["permission_keys"] == [f"permission:{agent_id}"] assert "active_config_is_published" not in detail assert "bound_agent_id" not in detail assert captured["get_app"] == {"app": app_model, "session": session} @@ -642,12 +677,29 @@ def test_agent_app_copy_uses_agent_id_and_returns_agent_detail( captured.update(kwargs) return copied_app + def get_app_backing_agent(self, **kwargs: object) -> Agent: + captured["get_app_backing_agent"] = kwargs + return Agent(id="copied-agent", app_id="copied-app") + monkeypatch.setattr(roster_controller, "_agent_roster_service", lambda *_args: FakeRosterService()) monkeypatch.setattr( roster_controller, "_serialize_agent_app_detail", lambda _session, app_model, **_kwargs: {"id": "copied-agent", "app_id": app_model.id, "name": app_model.name}, ) + monkeypatch.setattr(roster_controller.dify_config, "RBAC_ENABLED", True) + replace_agent_whitelist = MagicMock() + initialize_agent_rbac_access = MagicMock() + monkeypatch.setattr( + roster_controller.enterprise_rbac_service.RBACService.AgentAccess, + "replace_whitelist", + replace_agent_whitelist, + ) + monkeypatch.setattr( + roster_controller.initialize_created_app_rbac_access_task, + "delay", + initialize_agent_rbac_access, + ) with app.test_request_context( "/console/api/agent/00000000-0000-0000-0000-000000000001/copy", json={ @@ -686,7 +738,12 @@ def test_agent_app_copy_uses_agent_id_and_returns_agent_detail( "icon_type": "emoji", "icon": "sparkles", "icon_background": "#fff", + "get_app_backing_agent": {"tenant_id": "tenant-1", "app_id": "copied-app"}, } + replace_agent_whitelist.assert_called_once() + assert replace_agent_whitelist.call_args.args[:3] == ("tenant-1", account_id, "copied-agent") + assert replace_agent_whitelist.call_args.args[3].automatic_include_workspace_members is True + initialize_agent_rbac_access.assert_called_once_with("tenant-1", account_id, agent_id="copied-agent") def test_agent_debug_conversation_refresh_resets_build_for_current_user( diff --git a/api/tests/unit_tests/controllers/console/workspace/test_rbac.py b/api/tests/unit_tests/controllers/console/workspace/test_rbac.py index 861cce48d24..63230db29dd 100644 --- a/api/tests/unit_tests/controllers/console/workspace/test_rbac.py +++ b/api/tests/unit_tests/controllers/console/workspace/test_rbac.py @@ -389,6 +389,49 @@ class TestResourceAccessScopeBindings: mock_sync_task.delay.assert_not_called() + def test_agent_whitelist_forwards_to_agent_access(self, app): + result = rbac_mod.svc.ResourceWhitelist(account_ids=["acct-2"], automatic_include_workspace_members=False) + with ( + app.test_request_context("/workspaces/current/rbac/agents/agent-1/whitelist"), + patch("controllers.console.workspace.rbac._current_ids", return_value=("tenant-1", "acct-1")), + patch( + "controllers.console.workspace.rbac.svc.RBACService.AgentAccess.whitelist", + return_value=result, + ) as mock_get, + ): + response = inspect.unwrap(rbac_mod.RBACAgentWhitelistApi.get)( + rbac_mod.RBACAgentWhitelistApi(), + "agent-1", + ) + + assert response == {"account_ids": ["acct-2"]} + mock_get.assert_called_once_with("tenant-1", "acct-1", "agent-1") + + def test_agent_whitelist_put_forwards_to_agent_access(self, app): + result = rbac_mod.svc.ResourceWhitelist(account_ids=["acct-1"], automatic_include_workspace_members=True) + with ( + app.test_request_context( + "/workspaces/current/rbac/agents/agent-1/whitelist", + method="PUT", + json={"automatic_include_workspace_members": True}, + ), + patch("controllers.console.workspace.rbac._current_ids", return_value=("tenant-1", "acct-actor")), + patch( + "controllers.console.workspace.rbac.svc.RBACService.AgentAccess.replace_whitelist", + return_value=result, + ) as mock_put, + ): + response = inspect.unwrap(rbac_mod.RBACAgentWhitelistApi.put)( + rbac_mod.RBACAgentWhitelistApi(), + "agent-1", + ) + + assert response == {"account_ids": ["acct-1"]} + mock_put.assert_called_once() + args = mock_put.call_args.args + assert args[:3] == ("tenant-1", "acct-actor", "agent-1") + assert args[3].automatic_include_workspace_members is True + def test_app_whitelist_config_returns_switch_state_only(self, app): result = rbac_mod.svc.ResourceWhitelistConfig(automatic_include_workspace_members=True) with ( @@ -425,6 +468,24 @@ class TestResourceAccessScopeBindings: assert response == {"automatic_include_workspace_members": False} mock_get.assert_called_once_with("tenant-1", "acct-1", "dataset-1") + def test_agent_whitelist_config_returns_switch_state_only(self, app): + result = rbac_mod.svc.ResourceWhitelistConfig(automatic_include_workspace_members=True) + with ( + app.test_request_context("/workspaces/current/rbac/agents/agent-1/whitelist_config"), + patch("controllers.console.workspace.rbac._current_ids", return_value=("tenant-1", "acct-1")), + patch( + "controllers.console.workspace.rbac.svc.RBACService.AgentAccess.whitelist_config", + return_value=result, + ) as mock_get, + ): + response = inspect.unwrap(rbac_mod.RBACAgentWhitelistConfigApi.get)( + rbac_mod.RBACAgentWhitelistConfigApi(), + "agent-1", + ) + + assert response == {"automatic_include_workspace_members": True} + mock_get.assert_called_once_with("tenant-1", "acct-1", "agent-1") + def test_app_user_access_policy_assignment_forwards_ids(self, app): with ( app.test_request_context( diff --git a/api/tests/unit_tests/services/enterprise/test_rbac_service.py b/api/tests/unit_tests/services/enterprise/test_rbac_service.py index 93e417dbec1..3584ae48ac2 100644 --- a/api/tests/unit_tests/services/enterprise/test_rbac_service.py +++ b/api/tests/unit_tests/services/enterprise/test_rbac_service.py @@ -436,6 +436,23 @@ class TestResourceAccess: assert call.json == {"access_policy_ids": ["policy-1"]} assert out.access_policies[0].id == "policy-1" + def test_agent_replace_user_access_policies(self, mock_send: MagicMock): + mock_send.return_value = { + "access_policies": [{"id": "policy-1", "resource_type": "agent", "name": "Can preview"}] + } + payload = svc.ReplaceUserAccessPolicies(access_policy_ids=["policy-1"], account_ids=["acct-2"]) + + out = svc.RBACService.AgentAccess.replace_user_access_policies( + "tenant-1", "acct-actor", "agent-1", None, payload + ) + + call = _call_args(mock_send) + assert call.method == "PUT" + assert call.endpoint == "/rbac/agents/user-access-policies" + assert call.params == {"agent_id": "agent-1", "account_id": None} + assert call.json == {"access_policy_ids": ["policy-1"], "account_ids": ["acct-2"]} + assert out.access_policies[0].id == "policy-1" + def test_app_append_whitelist_members_batch(self, mock_send: MagicMock): mock_send.return_value = None @@ -480,6 +497,28 @@ class TestResourceAccess: "data": [{"dataset_id": "dataset-1", "account_ids": ["acct-1", "acct-2"], "policy_id": "policy-1"}] } + def test_agent_append_whitelist_members_batch(self, mock_send: MagicMock): + mock_send.return_value = None + + svc.RBACService.AgentAccess.append_whitelist_members_batch( + "tenant-1", + "acct-actor", + [ + svc.AppendAgentWhitelistMembersBatchItem( + agent_id="agent-1", + account_ids=["acct-1", "acct-2"], + policy_id="policy-1", + ) + ], + ) + + call = _call_args(mock_send) + assert call.method == "POST" + assert call.endpoint == "/rbac/agents/whitelist/members/batch" + assert call.json == { + "data": [{"agent_id": "agent-1", "account_ids": ["acct-1", "acct-2"], "policy_id": "policy-1"}] + } + def test_dataset_whitelist(self, mock_send: MagicMock): mock_send.return_value = {"account_ids": ["acct-2"], "automatic_include_workspace_members": False} @@ -491,6 +530,17 @@ class TestResourceAccess: assert call.params == {"dataset_id": "dataset-1"} assert out.account_ids == ["acct-2"] + def test_agent_whitelist(self, mock_send: MagicMock): + mock_send.return_value = {"account_ids": ["acct-2"], "automatic_include_workspace_members": False} + + out = svc.RBACService.AgentAccess.whitelist("tenant-1", "acct-1", "agent-1") + + call = _call_args(mock_send) + assert call.method == "GET" + assert call.endpoint == "/rbac/agents/whitelist" + assert call.params == {"agent_id": "agent-1"} + assert out.account_ids == ["acct-2"] + def test_app_whitelist_config(self, mock_send: MagicMock): mock_send.return_value = { "account_ids": ["acct-1"], @@ -520,6 +570,38 @@ class TestResourceAccess: assert call.params == {"dataset_id": "dataset-1"} assert out.model_dump(mode="json") == {"automatic_include_workspace_members": False} + def test_agent_whitelist_config(self, mock_send: MagicMock): + mock_send.return_value = { + "account_ids": ["acct-1"], + "automatic_include_workspace_members": True, + "scope": "all", + } + + out = svc.RBACService.AgentAccess.whitelist_config("tenant-1", "acct-1", "agent-1") + + call = _call_args(mock_send) + assert call.method == "GET" + assert call.endpoint == "/rbac/agents/whitelist" + assert call.params == {"agent_id": "agent-1"} + assert out.model_dump(mode="json") == {"automatic_include_workspace_members": True} + + def test_agent_replace_whitelist(self, mock_send: MagicMock): + mock_send.return_value = {"account_ids": ["acct-1"], "automatic_include_workspace_members": True} + + out = svc.RBACService.AgentAccess.replace_whitelist( + "tenant-1", + "acct-1", + "agent-1", + svc.ReplaceMemberBindings(automatic_include_workspace_members=True), + ) + + call = _call_args(mock_send) + assert call.method == "PUT" + assert call.endpoint == "/rbac/agents/whitelist" + assert call.params == {"agent_id": "agent-1"} + assert call.json == {"automatic_include_workspace_members": True} + assert out.account_ids == ["acct-1"] + def test_dataset_legacy_whitelist_config_reads_old_scope_without_public_dump(self, mock_send: MagicMock): mock_send.return_value = { "account_ids": ["acct-1"], @@ -748,37 +830,42 @@ class TestMyPermissions: assert out.workspace.permission_keys == ["workspace.member.manage"] @pytest.mark.parametrize( - ("role", "workspace_keys", "app_keys", "dataset_keys"), + ("role", "workspace_keys", "app_keys", "dataset_keys", "agent_keys"), [ ( "owner", svc._LEGACY_WORKSPACE_OWNER_KEYS, svc._LEGACY_APP_OWNER_KEYS, svc._LEGACY_DATASET_OWNER_KEYS, + svc._LEGACY_AGENT_OWNER_KEYS, ), ( "admin", svc._LEGACY_WORKSPACE_ADMIN_KEYS, svc._LEGACY_APP_ADMIN_KEYS, svc._LEGACY_DATASET_ADMIN_KEYS, + svc._LEGACY_AGENT_ADMIN_KEYS, ), ( "editor", svc._LEGACY_WORKSPACE_EDITOR_KEYS, svc._LEGACY_APP_EDITOR_KEYS, svc._LEGACY_DATASET_EDITOR_KEYS, + svc._LEGACY_AGENT_EDITOR_KEYS, ), ( "normal", svc._LEGACY_WORKSPACE_NORMAL_KEYS, svc._LEGACY_APP_NORMAL_KEYS, [], + [], ), ( "dataset_operator", svc._LEGACY_WORKSPACE_DATASET_OPERATOR_KEYS, [], svc._LEGACY_DATASET_DATASET_OPERATOR_KEYS, + [], ), ], ) @@ -789,6 +876,7 @@ class TestMyPermissions: workspace_keys: list[str], app_keys: list[str], dataset_keys: list[str], + agent_keys: list[str], sqlite_session: Session, config_overrides, ): @@ -804,17 +892,21 @@ class TestMyPermissions: assert len(out.workspace.permission_keys) == len(set(out.workspace.permission_keys)) assert out.app.default_permission_keys == app_keys assert out.dataset.default_permission_keys == dataset_keys + assert out.agent.default_permission_keys == agent_keys assert out.app.overrides == [] assert out.dataset.overrides == [] + assert out.agent.overrides == [] if role == "owner": assert "snippets.management" in out.workspace.permission_keys assert "app.acl.preview" in out.workspace.permission_keys assert "dataset.acl.preview" in out.workspace.permission_keys assert "app.acl.preview" in out.app.default_permission_keys assert "dataset.acl.preview" in out.dataset.default_permission_keys + assert "agent.acl.preview" in out.agent.default_permission_keys assert not any(key.startswith("billing.") for key in out.workspace.permission_keys) if role == "editor": assert "app.acl.log_and_annotation" in out.app.default_permission_keys + assert "agent.acl.log_manage" not in out.agent.default_permission_keys assert "app.acl.deploy" not in out.app.default_permission_keys @pytest.mark.parametrize( @@ -924,12 +1016,16 @@ class TestMemberRoles: *svc._LEGACY_WORKSPACE_EDITOR_KEYS, *svc._LEGACY_APP_EDITOR_KEYS, *svc._LEGACY_DATASET_EDITOR_KEYS, + *svc._LEGACY_AGENT_EDITOR_KEYS, ] ) ) assert "snippets.create_and_modify" in out.roles[0].permission_keys assert "app.acl.preview" in out.roles[0].permission_keys assert "dataset.acl.preview" in out.roles[0].permission_keys + assert "agent.create" in out.roles[0].permission_keys + assert "agent.acl.preview" in out.roles[0].permission_keys + assert "agent.acl.log_manage" not in out.roles[0].permission_keys assert "app.acl.deploy" not in out.roles[0].permission_keys def test_replace(self, mock_send: MagicMock, sqlite_session: Session): @@ -1110,6 +1206,45 @@ class TestResourcePermissions: "ds-2": svc._LEGACY_DATASET_DATASET_OPERATOR_KEYS, } + def test_agent_permissions_batch_get(self, mock_send: MagicMock, sqlite_session: Session): + mock_send.return_value = { + "data": [ + {"resource_id": "agent-1", "permission_keys": ["agent.acl.view", "agent.acl.edit"]}, + {"resource_id": "agent-2", "permission_keys": []}, + ] + } + + out = svc.RBACService.AgentPermissions.batch_get( + "tenant-1", "acct-1", ["agent-1", "agent-2"], session=sqlite_session + ) + + call = _call_args(mock_send) + assert call.method == "POST" + assert call.endpoint == "/rbac/agents/permission-keys/batch" + assert call.json == {"agent_ids": ["agent-1", "agent-2"]} + assert out == { + "agent-1": ["agent.acl.view", "agent.acl.edit"], + "agent-2": [], + } + + def test_agent_permissions_batch_get_uses_legacy_agent_acl_permissions_when_rbac_disabled( + self, mock_send: MagicMock, sqlite_session: Session, config_overrides + ): + config_overrides(RBAC_ENABLED=False) + sqlite_session.add( + TenantAccountJoin(tenant_id="tenant-1", account_id="acct-1", role=svc.TenantAccountRole.EDITOR) + ) + sqlite_session.commit() + out = svc.RBACService.AgentPermissions.batch_get( + "tenant-1", "acct-1", ["agent-1", "agent-2"], session=sqlite_session + ) + + mock_send.assert_not_called() + assert out == { + "agent-1": svc._LEGACY_AGENT_EDITOR_KEYS, + "agent-2": svc._LEGACY_AGENT_EDITOR_KEYS, + } + class TestListOption: def test_empty_produces_empty_params(self): diff --git a/api/tests/unit_tests/tasks/test_initialize_created_app_rbac_access_task.py b/api/tests/unit_tests/tasks/test_initialize_created_app_rbac_access_task.py index 54eaf46f37b..fa9c3fc671c 100644 --- a/api/tests/unit_tests/tasks/test_initialize_created_app_rbac_access_task.py +++ b/api/tests/unit_tests/tasks/test_initialize_created_app_rbac_access_task.py @@ -50,6 +50,36 @@ def test_initialize_created_app_rbac_access_task_batches_workspace_members(monke assert call.kwargs["payload"].access_policy_ids == [task_module.APP_RBAC_DEFAULT_ACCESS_POLICY_ID] +def test_initialize_created_app_rbac_access_task_batches_agent_workspace_members(monkeypatch: pytest.MonkeyPatch): + import tasks.initialize_created_app_rbac_access_task as task_module + from tasks.initialize_created_app_rbac_access_task import initialize_created_app_rbac_access_task + + monkeypatch.setattr(task_module.dify_config, "RBAC_ENABLED", True) + monkeypatch.setattr( + task_module.TenantService, + "iter_member_account_id_batches", + lambda tenant_id, batch_size, session: iter([["acct-1", "acct-2"], ["acct-3"]]), + ) + replace_user_access_policies = MagicMock() + monkeypatch.setattr( + task_module.enterprise_rbac_service.RBACService.AgentAccess, + "replace_user_access_policies", + replace_user_access_policies, + ) + + initialize_created_app_rbac_access_task.run("tenant-1", "actor-1", agent_id="agent-1") + + assert replace_user_access_policies.call_count == 2 + assert replace_user_access_policies.call_args_list[0].kwargs["payload"].account_ids == ["acct-1", "acct-2"] + assert replace_user_access_policies.call_args_list[1].kwargs["payload"].account_ids == ["acct-3"] + for call in replace_user_access_policies.call_args_list: + assert call.kwargs["tenant_id"] == "tenant-1" + assert call.kwargs["account_id"] == "actor-1" + assert call.kwargs["agent_id"] == "agent-1" + assert call.kwargs["target_account_id"] is None + assert call.kwargs["payload"].access_policy_ids == [task_module.APP_RBAC_DEFAULT_ACCESS_POLICY_ID] + + def test_initialize_created_app_rbac_access_task_retries_on_failure(monkeypatch: pytest.MonkeyPatch): import tasks.initialize_created_app_rbac_access_task as task_module from tasks.initialize_created_app_rbac_access_task import initialize_created_app_rbac_access_task @@ -86,6 +116,7 @@ def test_sync_joined_workspace_member_rbac_access_task_appends_auto_included_res rbac.ResourceWhitelistConfigResource(resource_type=rbac.RBACResourceType.APP, resource_id="app-1"), rbac.ResourceWhitelistConfigResource(resource_type=rbac.RBACResourceType.DATASET, resource_id="dataset-1"), rbac.ResourceWhitelistConfigResource(resource_type=rbac.RBACResourceType.APP, resource_id="app-2"), + rbac.ResourceWhitelistConfigResource(resource_type=rbac.RBACResourceType.AGENT, resource_id="agent-1"), ] configs = rbac.ResourceWhitelistConfigsResponse( data=[ @@ -104,17 +135,24 @@ def test_sync_joined_workspace_member_rbac_access_task_appends_auto_included_res resource_id="app-2", automatic_include_workspace_members=False, ), + rbac.ResourceWhitelistConfigItem( + resource_type=rbac.RBACResourceType.AGENT, + resource_id="agent-1", + automatic_include_workspace_members=True, + ), ] ) batch_get = MagicMock(return_value=configs) app_append = MagicMock() dataset_append = MagicMock() + agent_append = MagicMock() monkeypatch.setattr(task_module.dify_config, "RBAC_ENABLED", True) monkeypatch.setattr(task_module, "_iter_resource_config_batches", lambda tenant_id, batch_size: iter([resources])) monkeypatch.setattr(rbac.RBACService.ResourceWhitelistConfigs, "batch_get", batch_get) monkeypatch.setattr(rbac.RBACService.AppAccess, "append_whitelist_members_batch", app_append) monkeypatch.setattr(rbac.RBACService.DatasetAccess, "append_whitelist_members_batch", dataset_append) + monkeypatch.setattr(rbac.RBACService.AgentAccess, "append_whitelist_members_batch", agent_append) sync_joined_workspace_member_rbac_access_task.run("tenant-1", "member-1", "actor-1") @@ -140,3 +178,12 @@ def test_sync_joined_workspace_member_rbac_access_task_appends_auto_included_res assert dataset_call["data"][0].dataset_id == "dataset-1" assert dataset_call["data"][0].account_ids == ["member-1"] assert dataset_call["data"][0].policy_id == task_module.APP_RBAC_DEFAULT_ACCESS_POLICY_ID + + agent_append.assert_called_once() + agent_call = agent_append.call_args.kwargs + assert agent_call["tenant_id"] == "tenant-1" + assert agent_call["account_id"] == "actor-1" + assert len(agent_call["data"]) == 1 + assert agent_call["data"][0].agent_id == "agent-1" + assert agent_call["data"][0].account_ids == ["member-1"] + assert agent_call["data"][0].policy_id == task_module.APP_RBAC_DEFAULT_ACCESS_POLICY_ID