192 lines
6.6 KiB
Python
192 lines
6.6 KiB
Python
from __future__ import annotations
|
|
|
|
import os
|
|
import sys
|
|
import unittest
|
|
from pathlib import Path
|
|
|
|
from fastapi import HTTPException
|
|
from sqlalchemy import select
|
|
from sqlalchemy.ext.asyncio import AsyncSession, create_async_engine
|
|
from sqlalchemy.orm import sessionmaker
|
|
|
|
|
|
BACKEND_DIR = Path(__file__).resolve().parents[1]
|
|
os.environ["KEFU_DB_TYPE"] = "sqlite"
|
|
os.environ["KEFU_DATABASE_URL"] = "sqlite+aiosqlite:///:memory:"
|
|
if str(BACKEND_DIR) not in sys.path:
|
|
sys.path.insert(0, str(BACKEND_DIR))
|
|
|
|
from models.models import Base, Role, User # noqa: E402
|
|
from auth.passwords import hash_password # noqa: E402
|
|
from auth.permissions import ( # noqa: E402
|
|
ACCOUNTS_WRITE,
|
|
ALL_PERMISSIONS,
|
|
MENU_USERS,
|
|
USERS_MANAGE,
|
|
expand_paired_permissions,
|
|
)
|
|
from auth.role_service import ( # noqa: E402
|
|
create_role,
|
|
delete_role,
|
|
guard_last_admin_change,
|
|
seed_builtin_roles,
|
|
update_role,
|
|
)
|
|
from auth.roles import has_permission, is_admin, permissions_for_role # noqa: E402
|
|
from auth import router as auth_router # noqa: E402
|
|
|
|
|
|
class RolesRbacTests(unittest.IsolatedAsyncioTestCase):
|
|
async def asyncSetUp(self):
|
|
self.engine = create_async_engine("sqlite+aiosqlite:///:memory:")
|
|
async with self.engine.begin() as conn:
|
|
await conn.run_sync(Base.metadata.create_all)
|
|
self.session_factory = sessionmaker(
|
|
self.engine, class_=AsyncSession, expire_on_commit=False
|
|
)
|
|
self.db = self.session_factory()
|
|
await seed_builtin_roles(self.db)
|
|
|
|
async def asyncTearDown(self):
|
|
await self.db.close()
|
|
await self.engine.dispose()
|
|
|
|
async def test_seed_creates_builtin_roles(self):
|
|
result = await self.db.execute(select(Role))
|
|
codes = {row.code for row in result.scalars().all()}
|
|
self.assertEqual(codes, {"admin", "operator", "viewer"})
|
|
self.assertTrue(is_admin("admin"))
|
|
self.assertFalse(is_admin("operator"))
|
|
self.assertIn(MENU_USERS, permissions_for_role("admin"))
|
|
self.assertNotIn(MENU_USERS, permissions_for_role("operator"))
|
|
self.assertFalse(has_permission("viewer", ACCOUNTS_WRITE))
|
|
|
|
async def test_custom_role_crud(self):
|
|
created = await create_role(
|
|
self.db,
|
|
code="ops_leader",
|
|
label="运营主管",
|
|
description="可管账号",
|
|
permissions=[MENU_USERS, ACCOUNTS_WRITE],
|
|
)
|
|
self.assertEqual(created.code, "ops_leader")
|
|
self.assertTrue(has_permission("ops_leader", ACCOUNTS_WRITE))
|
|
self.assertFalse(is_admin("ops_leader"))
|
|
|
|
updated = await update_role(
|
|
self.db,
|
|
"ops_leader",
|
|
label="主管",
|
|
permissions=[ACCOUNTS_WRITE],
|
|
)
|
|
self.assertEqual(updated.label, "主管")
|
|
self.assertFalse(has_permission("ops_leader", MENU_USERS))
|
|
|
|
await delete_role(self.db, "ops_leader")
|
|
self.assertFalse(has_permission("ops_leader", ACCOUNTS_WRITE))
|
|
|
|
async def test_cannot_delete_system_role(self):
|
|
with self.assertRaises(HTTPException) as caught:
|
|
await delete_role(self.db, "operator")
|
|
self.assertEqual(caught.exception.status_code, 400)
|
|
|
|
async def test_admin_permissions_always_full(self):
|
|
await update_role(
|
|
self.db,
|
|
"admin",
|
|
label="管理员",
|
|
permissions=[ACCOUNTS_WRITE],
|
|
)
|
|
self.assertEqual(permissions_for_role("admin"), list(ALL_PERMISSIONS))
|
|
|
|
async def test_me_payload_includes_permissions(self):
|
|
user = User(
|
|
username="u1",
|
|
password_hash=hash_password("password1"),
|
|
display_name="U1",
|
|
role="operator",
|
|
is_active=True,
|
|
email_verified=True,
|
|
)
|
|
self.db.add(user)
|
|
await self.db.commit()
|
|
await self.db.refresh(user)
|
|
payload = await auth_router._build_user_response(self.db, user)
|
|
self.assertEqual(payload.role_label, "运营")
|
|
self.assertFalse(payload.is_admin)
|
|
self.assertIn("accounts.create", payload.permissions)
|
|
self.assertNotIn(MENU_USERS, payload.permissions)
|
|
self.assertNotIn("data.scope_all", payload.permissions)
|
|
|
|
async def test_legacy_accounts_write_implies_granular(self):
|
|
created = await create_role(
|
|
self.db,
|
|
code="legacy_ops",
|
|
label="旧版运营",
|
|
description=None,
|
|
permissions=[ACCOUNTS_WRITE, "menu.accounts"],
|
|
)
|
|
self.assertTrue(has_permission("legacy_ops", "accounts.start"))
|
|
self.assertTrue(has_permission("legacy_ops", "accounts.cookie"))
|
|
self.assertIn(ACCOUNTS_WRITE, created.permissions)
|
|
|
|
async def test_data_scope_all_for_custom_role(self):
|
|
from auth.roles import has_global_scope
|
|
|
|
await create_role(
|
|
self.db,
|
|
code="auditor",
|
|
label="审计",
|
|
description=None,
|
|
permissions=["menu.accounts", "data.scope_all"],
|
|
)
|
|
self.assertTrue(has_global_scope("auditor"))
|
|
self.assertFalse(has_global_scope("operator"))
|
|
self.assertTrue(has_global_scope("admin"))
|
|
|
|
async def test_last_admin_cannot_be_demoted(self):
|
|
admin = User(
|
|
username="admin1",
|
|
password_hash=hash_password("password1"),
|
|
role="admin",
|
|
is_active=True,
|
|
email_verified=True,
|
|
)
|
|
self.db.add(admin)
|
|
await self.db.commit()
|
|
await self.db.refresh(admin)
|
|
|
|
with self.assertRaises(HTTPException) as caught:
|
|
await guard_last_admin_change(self.db, user=admin, new_role="operator")
|
|
self.assertEqual(caught.exception.status_code, 400)
|
|
|
|
async def test_menu_action_pairs_expand(self):
|
|
expanded = expand_paired_permissions([MENU_USERS])
|
|
self.assertIn(MENU_USERS, expanded)
|
|
self.assertIn(USERS_MANAGE, expanded)
|
|
created = await create_role(
|
|
self.db,
|
|
code="hr_desk",
|
|
label="人事台",
|
|
description=None,
|
|
permissions=[MENU_USERS],
|
|
)
|
|
self.assertIn(USERS_MANAGE, created.permissions)
|
|
self.assertNotIn("menu.roles", created.permissions)
|
|
self.assertNotIn("roles.manage", created.permissions)
|
|
|
|
roles_only = await create_role(
|
|
self.db,
|
|
code="role_editor",
|
|
label="角色编辑",
|
|
description=None,
|
|
permissions=["menu.roles"],
|
|
)
|
|
self.assertIn("roles.manage", roles_only.permissions)
|
|
self.assertNotIn(USERS_MANAGE, roles_only.permissions)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|