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()