150 lines
5.6 KiB
Python
150 lines
5.6 KiB
Python
from __future__ import annotations
|
|
|
|
import unittest
|
|
from datetime import UTC, datetime
|
|
|
|
from sqlalchemy import create_engine, event
|
|
from sqlalchemy.orm import sessionmaker
|
|
|
|
from govoplan_access.backend.administration import SqlAccessAdministration
|
|
from govoplan_access.backend.api.v1.admin_common import (
|
|
_accounts_by_user_id,
|
|
_group_member_ids_by_group_id,
|
|
_groups_by_user_id,
|
|
_roles_by_group_id,
|
|
_roles_by_user_id,
|
|
_tenant_role_assignment_counts,
|
|
)
|
|
from govoplan_access.backend.db.models import Account, ApiKey, Group, GroupRoleAssignment, Role, User, UserGroupMembership, UserRoleAssignment
|
|
from govoplan_core.db.base import Base
|
|
|
|
|
|
class AdminBatchHelperTests(unittest.TestCase):
|
|
def setUp(self) -> None:
|
|
self.engine = create_engine("sqlite:///:memory:")
|
|
Base.metadata.create_all(bind=self.engine)
|
|
self.Session = sessionmaker(bind=self.engine)
|
|
self.session = self.Session()
|
|
|
|
def tearDown(self) -> None:
|
|
self.session.close()
|
|
Base.metadata.drop_all(bind=self.engine)
|
|
self.engine.dispose()
|
|
|
|
def test_access_admin_batch_maps_related_rows(self) -> None:
|
|
session = self.session
|
|
account = Account(id="account-1", email="ada@example.test", normalized_email="ada@example.test")
|
|
user = User(id="user-1", tenant_id="tenant-1", account_id=account.id, email=account.email)
|
|
group = Group(id="group-1", tenant_id="tenant-1", slug="clerks", name="Clerks")
|
|
role = Role(id="role-1", tenant_id="tenant-1", slug="reader", name="Reader", permissions=["admin:users:read"])
|
|
session.add_all([account, user, group, role])
|
|
session.flush()
|
|
session.add(UserGroupMembership(tenant_id="tenant-1", user_id=user.id, group_id=group.id))
|
|
session.add(UserRoleAssignment(tenant_id="tenant-1", user_id=user.id, role_id=role.id))
|
|
session.add(GroupRoleAssignment(tenant_id="tenant-1", group_id=group.id, role_id=role.id))
|
|
session.commit()
|
|
|
|
accounts_by_user = _accounts_by_user_id(session, [user.id])
|
|
group_member_ids = _group_member_ids_by_group_id(session, tenant_id="tenant-1", group_ids=[group.id])
|
|
groups_by_user = _groups_by_user_id(session, tenant_id="tenant-1", user_ids=[user.id])
|
|
roles_by_group = _roles_by_group_id(session, tenant_id="tenant-1", group_ids=[group.id])
|
|
roles_by_user = _roles_by_user_id(session, tenant_id="tenant-1", user_ids=[user.id])
|
|
role_counts = _tenant_role_assignment_counts(session, [role.id])
|
|
|
|
self.assertEqual(accounts_by_user[user.id].id, account.id)
|
|
self.assertEqual(group_member_ids, {group.id: [user.id]})
|
|
self.assertEqual([item.id for item in groups_by_user[user.id]], [group.id])
|
|
self.assertEqual([item.id for item in roles_by_group[group.id]], [role.id])
|
|
self.assertEqual([item.id for item in roles_by_user[user.id]], [role.id])
|
|
self.assertEqual(role_counts, {role.id: (1, 1)})
|
|
|
|
def test_tenant_counts_many_uses_three_grouped_queries(self) -> None:
|
|
accounts = [
|
|
Account(
|
|
id=f"account-{index}",
|
|
email=f"user-{index}@example.test",
|
|
normalized_email=f"user-{index}@example.test",
|
|
)
|
|
for index in range(3)
|
|
]
|
|
users = [
|
|
User(
|
|
id="user-1",
|
|
tenant_id="tenant-1",
|
|
account_id=accounts[0].id,
|
|
email=accounts[0].email,
|
|
),
|
|
User(
|
|
id="user-2",
|
|
tenant_id="tenant-1",
|
|
account_id=accounts[1].id,
|
|
email=accounts[1].email,
|
|
is_active=False,
|
|
),
|
|
User(
|
|
id="user-3",
|
|
tenant_id="tenant-2",
|
|
account_id=accounts[2].id,
|
|
email=accounts[2].email,
|
|
),
|
|
]
|
|
self.session.add_all(
|
|
[
|
|
*accounts,
|
|
*users,
|
|
Group(id="group-1", tenant_id="tenant-1", slug="one", name="One"),
|
|
Group(id="group-2", tenant_id="tenant-2", slug="two", name="Two"),
|
|
ApiKey(
|
|
id="key-1",
|
|
tenant_id="tenant-1",
|
|
user_id="user-1",
|
|
name="Active",
|
|
prefix="active",
|
|
key_hash="hash-1",
|
|
),
|
|
ApiKey(
|
|
id="key-2",
|
|
tenant_id="tenant-1",
|
|
user_id="user-2",
|
|
name="Revoked",
|
|
prefix="revoked",
|
|
key_hash="hash-2",
|
|
revoked_at=datetime.now(UTC),
|
|
),
|
|
]
|
|
)
|
|
self.session.commit()
|
|
|
|
query_count = 0
|
|
|
|
def count_query(*_args: object) -> None:
|
|
nonlocal query_count
|
|
query_count += 1
|
|
|
|
event.listen(self.engine, "before_cursor_execute", count_query)
|
|
try:
|
|
counts = SqlAccessAdministration().tenant_counts_many(
|
|
self.session,
|
|
["tenant-1", "tenant-2", "tenant-empty"],
|
|
)
|
|
finally:
|
|
event.remove(self.engine, "before_cursor_execute", count_query)
|
|
|
|
self.assertEqual(3, query_count)
|
|
self.assertEqual(
|
|
{
|
|
"users": 2,
|
|
"active_users": 1,
|
|
"groups": 1,
|
|
"api_keys": 2,
|
|
"active_api_keys": 1,
|
|
},
|
|
counts["tenant-1"],
|
|
)
|
|
self.assertEqual(1, counts["tenant-2"]["users"])
|
|
self.assertEqual(0, counts["tenant-empty"]["users"])
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|