feat(enterprise): reject member invites that would exceed licensed seats (#38995)

This commit is contained in:
Xin Zhang 2026-07-15 15:59:16 +08:00 committed by GitHub
parent 10da5e8f9d
commit 120c38bad8
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
3 changed files with 238 additions and 23 deletions

View File

@ -23,7 +23,7 @@ from controllers.console.auth.error import (
NotOwnerError,
OwnerTransferLimitError,
)
from controllers.console.error import EmailSendIpLimitError, WorkspaceMembersLimitExceeded
from controllers.console.error import EmailSendIpLimitError, SeatsLimitExceeded, WorkspaceMembersLimitExceeded
from controllers.console.workspace.error import InvalidMemberRoleError
from controllers.console.wraps import (
account_initialization_required,
@ -40,7 +40,7 @@ from libs.login import current_account_with_tenant, login_required
from models.account import Account, TenantAccountJoin, TenantAccountRole
from services.account_service import AccountService, RegisterService, TenantService
from services.enterprise import rbac_service as enterprise_rbac_service
from services.errors.account import AccountAlreadyInTenantError, SeatsLimitExceededError
from services.errors.account import AccountAlreadyInTenantError
from services.feature_service import FeatureService
@ -160,12 +160,14 @@ def _normalize_enum_value(value: object) -> str:
return str(normalized) if normalized is not None else ""
def _count_new_member_invites(tenant_id: str, emails: list[str]) -> int:
def _count_new_member_invites(tenant_id: str, emails: list[str]) -> tuple[int, int]:
new_member_count = 0
new_account_count = 0
for email in emails:
account = AccountService.get_account_by_email_with_case_fallback(email, session=db.session())
if not account:
new_member_count += 1
new_account_count += 1
continue
exists = db.session.scalar(
@ -176,7 +178,7 @@ def _count_new_member_invites(tenant_id: str, emails: list[str]) -> int:
if not exists:
new_member_count += 1
return new_member_count
return new_member_count, new_account_count
def _count_current_members(tenant_id: str) -> int:
@ -185,7 +187,7 @@ def _count_current_members(tenant_id: str) -> int:
)
def _check_member_invite_limits(tenant_id: str, new_member_count: int) -> None:
def _check_member_invite_limits(tenant_id: str, new_member_count: int, new_account_count: int) -> None:
if new_member_count <= 0:
return
@ -195,6 +197,10 @@ def _check_member_invite_limits(tenant_id: str, new_member_count: int) -> None:
workspace_members = features.workspace_members
if workspace_members.enabled is True and not workspace_members.is_available(new_member_count):
raise WorkspaceMembersLimitExceeded()
if new_account_count > 0:
seats = FeatureService.get_system_features(is_authenticated=True).license.seats
if not seats.is_available(new_account_count):
raise SeatsLimitExceeded()
return
if dify_config.BILLING_ENABLED and features.billing.enabled is True:
@ -295,8 +301,8 @@ class MemberInviteEmailApi(Resource):
tenant_id = inviter.current_tenant.id
with redis_client.lock(f"workspace_member_invite:{tenant_id}", timeout=60):
if dify_config.ENTERPRISE_ENABLED is True or dify_config.BILLING_ENABLED is True:
new_member_count = _count_new_member_invites(tenant_id, invitee_emails)
_check_member_invite_limits(tenant_id, new_member_count)
new_member_count, new_account_count = _count_new_member_invites(tenant_id, invitee_emails)
_check_member_invite_limits(tenant_id, new_member_count, new_account_count)
for invitee_email in invitee_emails:
try:
@ -326,14 +332,6 @@ class MemberInviteEmailApi(Resource):
message="Account already in workspace.",
)
)
except SeatsLimitExceededError:
invitation_results.append(
MemberInviteFailedResponse(
status="failed",
email=invitee_email,
message="Licensed seats limit exceeded.",
)
)
except Exception as e:
invitation_results.append(
MemberInviteFailedResponse(status="failed", email=invitee_email, message=str(e))

View File

@ -52,7 +52,7 @@ class TestMemberInviteEmailApi:
with (
patch("controllers.console.workspace.members.dify_config.RBAC_ENABLED", False),
patch("controllers.console.workspace.members.dify_config.CONSOLE_WEB_URL", "https://console.example.com"),
patch("controllers.console.workspace.members._count_new_member_invites", return_value=1),
patch("controllers.console.workspace.members._count_new_member_invites", return_value=(1, 1)),
patch("controllers.console.workspace.members.dify_config.ENTERPRISE_ENABLED", False),
patch("controllers.console.workspace.members.dify_config.BILLING_ENABLED", False),
):

View File

@ -15,7 +15,7 @@ from controllers.console.auth.error import (
NotOwnerError,
OwnerTransferLimitError,
)
from controllers.console.error import EmailSendIpLimitError, WorkspaceMembersLimitExceeded
from controllers.console.error import EmailSendIpLimitError, SeatsLimitExceeded, WorkspaceMembersLimitExceeded
from controllers.console.workspace.error import InvalidMemberRoleError
from controllers.console.workspace.members import (
DatasetOperatorMemberListApi,
@ -26,9 +26,10 @@ from controllers.console.workspace.members import (
OwnerTransfer,
OwnerTransferCheckApi,
SendOwnerTransferEmailApi,
_count_new_member_invites,
)
from libs.external_api import ExternalApi
from services.errors.account import AccountAlreadyInTenantError
from services.errors.account import AccountAlreadyInTenantError, SeatsLimitExceededError
class TestMemberListApi:
@ -142,7 +143,7 @@ class TestMemberInviteEmailApi:
with (
app.test_request_context("/", json=payload),
patch("controllers.console.workspace.members.FeatureService.get_features", return_value=features),
patch("controllers.console.workspace.members._count_new_member_invites", return_value=1) as mock_count,
patch("controllers.console.workspace.members._count_new_member_invites", return_value=(1, 1)) as mock_count,
patch(
"controllers.console.workspace.members.RegisterService.invite_new_member", return_value="token"
) as mock_invite,
@ -178,7 +179,7 @@ class TestMemberInviteEmailApi:
with (
app.test_request_context("/", json=payload),
patch("controllers.console.workspace.members.FeatureService.get_features", return_value=features),
patch("controllers.console.workspace.members._count_new_member_invites", return_value=1),
patch("controllers.console.workspace.members._count_new_member_invites", return_value=(1, 1)),
patch("controllers.console.workspace.members.dify_config.ENTERPRISE_ENABLED", True),
patch("controllers.console.workspace.members.dify_config.BILLING_ENABLED", False),
):
@ -205,7 +206,7 @@ class TestMemberInviteEmailApi:
with (
app.test_request_context("/", json=payload),
patch("controllers.console.workspace.members.FeatureService.get_features", return_value=features),
patch("controllers.console.workspace.members._count_new_member_invites", return_value=2),
patch("controllers.console.workspace.members._count_new_member_invites", return_value=(2, 2)),
patch("controllers.console.workspace.members._count_current_members", return_value=9),
patch("controllers.console.workspace.members.dify_config.ENTERPRISE_ENABLED", False),
patch("controllers.console.workspace.members.dify_config.BILLING_ENABLED", True),
@ -232,7 +233,7 @@ class TestMemberInviteEmailApi:
with (
app.test_request_context("/", json=payload),
patch("controllers.console.workspace.members.FeatureService.get_features", return_value=features),
patch("controllers.console.workspace.members._count_new_member_invites", return_value=0),
patch("controllers.console.workspace.members._count_new_member_invites", return_value=(0, 0)),
patch(
"controllers.console.workspace.members.RegisterService.invite_new_member",
side_effect=AccountAlreadyInTenantError(),
@ -300,7 +301,7 @@ class TestMemberInviteEmailApi:
with (
app.test_request_context("/", json=payload),
patch("controllers.console.workspace.members.FeatureService.get_features", return_value=features),
patch("controllers.console.workspace.members._count_new_member_invites", return_value=1),
patch("controllers.console.workspace.members._count_new_member_invites", return_value=(1, 1)),
patch(
"controllers.console.workspace.members.RegisterService.invite_new_member",
side_effect=Exception("boom"),
@ -313,6 +314,222 @@ class TestMemberInviteEmailApi:
assert result["invitation_results"][0]["status"] == "failed"
def test_invite_seats_limit_exceeded(self, app: Flask):
api = MemberInviteEmailApi()
method = unwrap(api.post)
tenant = MagicMock(id="t1")
user = MagicMock(current_tenant=tenant)
features = MagicMock()
features.billing.enabled = False
features.workspace_members.enabled = False
system_features = MagicMock()
system_features.license.seats.is_available.return_value = False
payload = {
"emails": ["a@test.com", "b@test.com"],
"role": "normal",
}
with (
app.test_request_context("/", json=payload),
patch("controllers.console.workspace.members.FeatureService.get_features", return_value=features),
patch("controllers.console.workspace.members._count_new_member_invites", return_value=(2, 2)),
patch(
"controllers.console.workspace.members.FeatureService.get_system_features",
return_value=system_features,
) as mock_get_system_features,
patch("controllers.console.workspace.members.RegisterService.invite_new_member") as mock_invite,
patch("controllers.console.workspace.members.dify_config.ENTERPRISE_ENABLED", True),
patch("controllers.console.workspace.members.dify_config.BILLING_ENABLED", False),
):
with pytest.raises(SeatsLimitExceeded):
method(api, user)
mock_get_system_features.assert_called_once_with(is_authenticated=True)
system_features.license.seats.is_available.assert_called_once_with(2)
mock_invite.assert_not_called()
def test_invite_existing_accounts_do_not_consume_seats(self, app: Flask):
api = MemberInviteEmailApi()
method = unwrap(api.post)
tenant = MagicMock(id="t1")
user = MagicMock(current_tenant=tenant)
features = MagicMock()
features.billing.enabled = False
features.workspace_members.enabled = False
system_features = MagicMock()
system_features.license.seats.is_available.return_value = False
payload = {
"emails": ["a@test.com", "b@test.com"],
"role": "normal",
}
with (
app.test_request_context("/", json=payload),
patch("controllers.console.workspace.members.FeatureService.get_features", return_value=features),
patch("controllers.console.workspace.members._count_new_member_invites", return_value=(2, 0)),
patch(
"controllers.console.workspace.members.FeatureService.get_system_features",
return_value=system_features,
) as mock_get_system_features,
patch(
"controllers.console.workspace.members.RegisterService.invite_new_member", return_value="token"
) as mock_invite,
patch("controllers.console.workspace.members.dify_config.CONSOLE_WEB_URL", "http://x"),
patch("controllers.console.workspace.members.dify_config.ENTERPRISE_ENABLED", True),
patch("controllers.console.workspace.members.dify_config.BILLING_ENABLED", False),
):
result, status = method(api, user)
assert status == 201
assert len(result["invitation_results"]) == 2
mock_get_system_features.assert_not_called()
system_features.license.seats.is_available.assert_not_called()
assert mock_invite.call_count == 2
def test_invite_mixed_accounts_with_available_seats(self, app: Flask):
api = MemberInviteEmailApi()
method = unwrap(api.post)
tenant = MagicMock(id="t1")
user = MagicMock(current_tenant=tenant)
features = MagicMock()
features.billing.enabled = False
features.workspace_members.enabled = False
system_features = MagicMock()
system_features.license.seats.is_available.return_value = True
payload = {
"emails": ["a@test.com", "b@test.com"],
"role": "normal",
}
with (
app.test_request_context("/", json=payload),
patch("controllers.console.workspace.members.FeatureService.get_features", return_value=features),
patch("controllers.console.workspace.members._count_new_member_invites", return_value=(2, 1)),
patch(
"controllers.console.workspace.members.FeatureService.get_system_features",
return_value=system_features,
) as mock_get_system_features,
patch(
"controllers.console.workspace.members.RegisterService.invite_new_member", return_value="token"
) as mock_invite,
patch("controllers.console.workspace.members.dify_config.CONSOLE_WEB_URL", "http://x"),
patch("controllers.console.workspace.members.dify_config.ENTERPRISE_ENABLED", True),
patch("controllers.console.workspace.members.dify_config.BILLING_ENABLED", False),
):
result, status = method(api, user)
assert status == 201
assert len(result["invitation_results"]) == 2
mock_get_system_features.assert_called_once_with(is_authenticated=True)
system_features.license.seats.is_available.assert_called_once_with(1)
assert mock_invite.call_count == 2
def test_invite_skips_seats_limit_when_enterprise_disabled(self, app: Flask):
api = MemberInviteEmailApi()
method = unwrap(api.post)
tenant = MagicMock(id="t1")
user = MagicMock(current_tenant=tenant)
features = MagicMock()
features.billing.enabled = False
features.workspace_members.enabled = False
system_features = MagicMock()
system_features.license.seats.is_available.return_value = False
payload = {
"emails": ["a@test.com"],
"role": "normal",
}
with (
app.test_request_context("/", json=payload),
patch("controllers.console.workspace.members.FeatureService.get_features", return_value=features),
patch("controllers.console.workspace.members._count_new_member_invites", return_value=(1, 1)),
patch(
"controllers.console.workspace.members.FeatureService.get_system_features",
return_value=system_features,
) as mock_get_system_features,
patch("controllers.console.workspace.members.RegisterService.invite_new_member", return_value="token"),
patch("controllers.console.workspace.members.dify_config.CONSOLE_WEB_URL", "http://x"),
patch("controllers.console.workspace.members.dify_config.ENTERPRISE_ENABLED", False),
patch("controllers.console.workspace.members.dify_config.BILLING_ENABLED", False),
):
result, status = method(api, user)
assert status == 201
assert result["invitation_results"][0]["status"] == "success"
mock_get_system_features.assert_not_called()
system_features.license.seats.is_available.assert_not_called()
def test_invite_seats_error_is_reported_as_failed_result(self, app: Flask):
api = MemberInviteEmailApi()
method = unwrap(api.post)
tenant = MagicMock(id="t1")
user = MagicMock(current_tenant=tenant)
features = MagicMock()
features.billing.enabled = False
features.workspace_members.enabled = False
system_features = MagicMock()
system_features.license.seats.is_available.return_value = True
payload = {
"emails": ["a@test.com"],
"role": "normal",
}
with (
app.test_request_context("/", json=payload),
patch("controllers.console.workspace.members.FeatureService.get_features", return_value=features),
patch("controllers.console.workspace.members._count_new_member_invites", return_value=(1, 1)),
patch(
"controllers.console.workspace.members.FeatureService.get_system_features",
return_value=system_features,
),
patch(
"controllers.console.workspace.members.RegisterService.invite_new_member",
side_effect=SeatsLimitExceededError("licensed seats limit exceeded"),
),
patch("controllers.console.workspace.members.dify_config.CONSOLE_WEB_URL", "http://x"),
patch("controllers.console.workspace.members.dify_config.ENTERPRISE_ENABLED", True),
patch("controllers.console.workspace.members.dify_config.BILLING_ENABLED", False),
):
result, status = method(api, user)
assert status == 201
assert result["invitation_results"][0]["status"] == "failed"
assert result["invitation_results"][0]["message"] == "licensed seats limit exceeded"
class TestCountNewMemberInvites:
def test_count_new_member_invites(self):
new_account = None
existing_account_not_in_tenant = SimpleNamespace(id="account-2")
existing_account_in_tenant = SimpleNamespace(id="account-3")
with (
patch(
"controllers.console.workspace.members.AccountService.get_account_by_email_with_case_fallback",
side_effect=[new_account, existing_account_not_in_tenant, existing_account_in_tenant],
) as mock_get_account,
patch("controllers.console.workspace.members.db.session") as mock_session,
):
mock_session.scalar.side_effect = [None, "join-id"]
result = _count_new_member_invites(
"tenant-1",
["new@test.com", "existing@test.com", "member@test.com"],
)
assert result == (2, 1)
assert mock_get_account.call_count == 3
assert mock_session.scalar.call_count == 2
class TestMemberUpdateRoleApi:
def test_update_invalid_role(self, app: Flask):