diff --git a/dash/orgs/models.py b/dash/orgs/models.py index 6e7ad29..c7ec263 100644 --- a/dash/orgs/models.py +++ b/dash/orgs/models.py @@ -144,16 +144,19 @@ def get_org_users(self): return org_users.distinct() def get_user_org_group(self, user): - if user in self.get_org_admins(): + if hasattr(user, "_org_group"): + return user._org_group + + if self.administrators.filter(id=user.id).exists(): user._org_group = Group.objects.get(name="Administrators") - elif user in self.get_org_editors(): + elif self.editors.filter(id=user.id).exists(): user._org_group = Group.objects.get(name="Editors") - elif user in self.get_org_viewers(): + elif self.viewers.filter(id=user.id).exists(): user._org_group = Group.objects.get(name="Viewers") else: user._org_group = None - return getattr(user, "_org_group", None) + return user._org_group def get_user(self): user = self.administrators.filter(is_active=True).first() diff --git a/test_runner/tests.py b/test_runner/tests.py index 0dcfdf0..c541c16 100644 --- a/test_runner/tests.py +++ b/test_runner/tests.py @@ -1,17 +1,18 @@ import zoneinfo -from dash.tags.models import Tag -from unittest.mock import Mock, patch, call +from unittest.mock import Mock, call, patch import valkey from smartmin.tests import SmartminTest from temba_client.v2 import TembaClient from django.conf import settings -from django.contrib.auth.models import Group, User +from django.contrib.auth.models import AnonymousUser, Group, User from django.core import mail from django.core.exceptions import DisallowedHost +from django.db import connection from django.db.utils import IntegrityError from django.http import HttpRequest, HttpResponse +from django.test.utils import CaptureQueriesContext from django.urls import ResolverMatch, reverse from django.utils.encoding import force_str @@ -25,8 +26,9 @@ from dash.orgs.tasks import org_task from dash.orgs.templatetags.dashorgs import display_time, national_phone from dash.stories.models import Story, StoryImage -from dash.utils import random_string +from dash.tags.models import Tag from dash.test import MockResponse +from dash.utils import random_string class UserTest(SmartminTest): @@ -412,6 +414,41 @@ def setUp(self): self.org = self.create_org("uganda", self.admin) + def test_get_user_org_group_queries(self): + viewer = self.create_user("Viewer") + editor = self.create_user("Editor") + non_member = self.create_user("NonMember") + self.org.viewers.add(viewer) + self.org.editors.add(editor) + + self.assertIsNone(self.org.get_user_org_group(AnonymousUser())) + + def assert_membership_queries(user, group_name): + with CaptureQueriesContext(connection) as context: + group = self.org.get_user_org_group(user) + + self.assertEqual(group_name, group.name if group else None) + + # ignore the query that fetches the group itself + membership_queries = [q["sql"] for q in context.captured_queries if '"auth_group"' not in q["sql"]] + self.assertTrue(membership_queries) + + # membership is checked with existence queries rather than fetching entire member lists + for sql in membership_queries: + self.assertTrue(sql.startswith("SELECT 1 AS"), f"expected an existence query but got: {sql}") + self.assertNotIn('"auth_user"."password"', sql) + + assert_membership_queries(self.admin, "Administrators") + assert_membership_queries(editor, "Editors") + assert_membership_queries(viewer, "Viewers") + assert_membership_queries(non_member, None) + + # result is cached on the user object so repeated lookups don't hit the database + with self.assertNumQueries(0): + self.assertEqual(self.org.get_user_org_group(self.admin).name, "Administrators") + with self.assertNumQueries(0): + self.assertIsNone(self.org.get_user_org_group(non_member)) + def test_org_model(self): user = self.create_user("User")