diff --git a/rdmo/conditions/models.py b/rdmo/conditions/models.py index 04f1723ce3..935630650b 100644 --- a/rdmo/conditions/models.py +++ b/rdmo/conditions/models.py @@ -116,7 +116,7 @@ def is_locked(self): return self.locked def resolve(self, values, set_prefix=None, set_index=None): - source_values = filter(lambda value: value.attribute == self.source, values) + source_values = filter(lambda value: value.attribute_id == self.source_id, values) if set_prefix is not None: source_values = filter(lambda value: value.set_prefix == set_prefix, source_values) @@ -168,8 +168,8 @@ def _resolve_equal(self, values): results = [] for value in values: - if self.target_option: - results.append(value.option == self.target_option) + if self.target_option_id: + results.append(value.option_id == self.target_option_id) else: results.append(value.text == self.target_text) @@ -230,7 +230,7 @@ def _resolve_lesser_than_equal(self, values): def _resolve_not_empty(self, values): for value in values: - if bool(value.text) or bool(value.option): + if bool(value.text) or bool(value.option_id): return True return False diff --git a/rdmo/projects/answers.py b/rdmo/projects/answers.py index 722039d04a..b2d069a2e5 100644 --- a/rdmo/projects/answers.py +++ b/rdmo/projects/answers.py @@ -3,6 +3,7 @@ from rdmo.core.utils import markdown2html from .models.value import Value +from .utils import check_conditions, compute_attribute_values_map, compute_sets class AnswerTree: @@ -10,9 +11,14 @@ class AnswerTree: def __init__(self, catalog, values, verbose=None): self.catalog = catalog self.values = values + self.attribute_values_map = compute_attribute_values_map(self.values) + self.condition_results = {} self.verbose = tuple(verbose or ()) - self.sets = values.compute_sets() + self.sets = compute_sets( + (value.attribute_id, value.set_prefix, value.set_index) + for value in self.values + ) self.conditions = catalog.conditions.in_bulk() # buffer for the resolved conditions: self.resolved_conditions[element][parent_set] @@ -213,17 +219,11 @@ def compute_value_node(self, value=None): def resolve_conditions(self, element, parent_set): # cache each resolved condition in self.resolved_conditions if self.resolved_conditions.get(element, {}).get(parent_set) is None: - if parent_set: - set_prefix, set_index = parent_set - self.resolved_conditions[element][parent_set] = any( - self.conditions[condition.id].resolve(self.values, set_prefix, set_index) - for condition in element.conditions.all() - ) - else: - self.resolved_conditions[element][parent_set] = any( - self.conditions[condition.id].resolve(self.values) - for condition in element.conditions.all() - ) + conditions = [self.conditions[condition.id] for condition in element.conditions.all()] + set_prefix, set_index = parent_set if parent_set else (None, None) + self.resolved_conditions[element][parent_set] = bool(conditions) and check_conditions( + conditions, self.attribute_values_map, set_prefix, set_index, self.condition_results + ) return self.resolved_conditions[element][parent_set] diff --git a/rdmo/projects/managers.py b/rdmo/projects/managers.py index 39c3415367..d938924a2d 100644 --- a/rdmo/projects/managers.py +++ b/rdmo/projects/managers.py @@ -1,5 +1,3 @@ -from collections import defaultdict - from django.conf import settings from django.db import models from django.db.models import Q @@ -10,6 +8,8 @@ from rdmo.accounts.utils import is_site_manager from rdmo.core.managers import CurrentSiteManagerMixin +from .utils import compute_sets + class ProjectQuerySet(TreeQuerySet): @@ -212,10 +212,7 @@ def filter_set(self, set_value): ) def compute_sets(self): - sets = defaultdict(set) - for attribute, set_prefix, set_index in self.distinct_list(): - sets[attribute].add((set_prefix, set_index)) - return sets + return compute_sets(self.distinct_list()) class ProjectManager(CurrentSiteManagerMixin, TreeManager): diff --git a/rdmo/projects/tests/test_conditions.py b/rdmo/projects/tests/test_conditions.py index f81e5c549c..de0186b07f 100644 --- a/rdmo/projects/tests/test_conditions.py +++ b/rdmo/projects/tests/test_conditions.py @@ -3,12 +3,32 @@ from rdmo.conditions.models import Condition from ..models import Project, Value +from ..utils import check_conditions, compute_attribute_values_map project_id = 1 value_id = 86 set_indexes = (0, 1) +@pytest.mark.parametrize('condition_id', [1, 10]) +def test_check_conditions_matches_condition_resolve(db, condition_id): + condition = Condition.objects.get(id=condition_id) + values = Project.objects.get(id=project_id).values.filter(snapshot=None).order_by() + attribute_values_map = compute_attribute_values_map(values) + + assert check_conditions([condition], attribute_values_map) is True + assert condition.resolve(values) is True + + +def test_check_conditions_preserves_set_prefix_fallback(db): + condition = Condition.objects.get(uri='http://example.com/terms/conditions/text_contains_test') + values = Project.objects.get(id=project_id).values.filter(snapshot=None).order_by() + attribute_values_map = compute_attribute_values_map(values) + + assert check_conditions([condition], attribute_values_map, set_prefix='0', set_index=0) is True + assert condition.resolve(values, set_prefix='0', set_index=0) is True + + @pytest.mark.parametrize('set_index', set_indexes) def test_set_collection(db, set_index): # test the special case, when a condition of a question in a set is checked @@ -18,6 +38,7 @@ def test_set_collection(db, set_index): result = condition.resolve(values, set_index=set_index) assert result is True + assert check_conditions([condition], compute_attribute_values_map(values), set_index=set_index) is result @pytest.mark.parametrize('set_index', set_indexes) @@ -31,6 +52,7 @@ def test_set_collection_error_none(db, set_index): values = Project.objects.get(id=project_id).values.filter(snapshot=None) result = condition.resolve(values, set_index=set_index) assert result is (True if set_index == 0 else False) + assert check_conditions([condition], compute_attribute_values_map(values), set_index=set_index) is result @pytest.mark.parametrize('set_index', set_indexes) @@ -44,3 +66,4 @@ def test_set_collection_error_true(db, set_index): values = Project.objects.get(id=project_id).values.filter(snapshot=None) result = condition.resolve(values, set_index=set_index) assert result is (True if set_index == 0 else False) + assert check_conditions([condition], compute_attribute_values_map(values), set_index=set_index) is result diff --git a/rdmo/projects/tests/test_queries.py b/rdmo/projects/tests/test_queries.py index ab1427182c..a33b1fd5b1 100644 --- a/rdmo/projects/tests/test_queries.py +++ b/rdmo/projects/tests/test_queries.py @@ -44,12 +44,13 @@ def test_resolve_queries(db, client, django_assert_max_num_queries): 'element_id': 1, } - with django_assert_max_num_queries(19): + with django_assert_max_num_queries(16): response = client.post(url, [params, params, params], content_type='application/json') assert response.status_code == 200 assert len(response.json()) == 3 assert response.json()[0] == response.json()[1] == response.json()[2] + assert response.json()[0]['result'] is True @pytest.mark.performance @@ -57,7 +58,7 @@ def test_resolve_empty_queries(db, client, django_assert_max_num_queries): client.login(username='owner', password='owner') url = reverse(urlnames['resolve'], kwargs={'pk': 1}) - with django_assert_max_num_queries(17): + with django_assert_max_num_queries(14): response = client.post(url, [], content_type='application/json') assert response.status_code == 200 diff --git a/rdmo/projects/tests/test_viewset_project.py b/rdmo/projects/tests/test_viewset_project.py index fd784e13f6..e5565915a4 100644 --- a/rdmo/projects/tests/test_viewset_project.py +++ b/rdmo/projects/tests/test_viewset_project.py @@ -2,6 +2,8 @@ from django.contrib.auth.models import Group, User from django.contrib.sites.models import Site +from django.db import connection +from django.test.utils import CaptureQueriesContext from django.urls import reverse from rdmo.conditions.models import Condition @@ -742,25 +744,40 @@ def test_resolve_post(db, client, username, password, project_id, condition_id): assert response.status_code == 401 -def test_resolve_post_resolves_duplicate_condition_once(db, client, mocker): +def test_resolve_post_empty_payload_does_not_load_values(db, client): client.login(username='owner', password='owner') - resolve = mocker.spy(Condition, 'resolve') - params = { - 'set_prefix': '', - 'set_index': 0, - 'element_type': 'conditions', - 'element_id': 1 - } - response = client.post( - reverse(urlnames['resolve'], args=[1]), - [params, params, params], - content_type='application/json' - ) + with CaptureQueriesContext(connection) as queries: + response = client.post( + reverse(urlnames['resolve'], args=[1]), + [], + content_type='application/json' + ) + + assert response.status_code == 200 + assert response.json() == [] + assert not any(Value._meta.db_table in query['sql'] for query in queries) + + +def test_resolve_post_missing_condition_returns_false(db, client): + client.login(username='owner', password='owner') + latest_condition_id = Condition.objects.order_by('-id').values_list('id', flat=True).first() or 0 + + with CaptureQueriesContext(connection) as queries: + response = client.post( + reverse(urlnames['resolve'], args=[1]), + [{ + 'set_prefix': '', + 'set_index': 0, + 'element_type': 'conditions', + 'element_id': latest_condition_id + 1, + }], + content_type='application/json' + ) assert response.status_code == 200 - assert response.json()[0] == response.json()[1] == response.json()[2] - assert resolve.call_count == 1 + assert response.json()[0]['result'] is False + assert not any(Value._meta.db_table in query['sql'] for query in queries) @pytest.mark.parametrize('username,password', users) diff --git a/rdmo/projects/utils.py b/rdmo/projects/utils.py index 8289e0b97d..cc9eeb7f03 100644 --- a/rdmo/projects/utils.py +++ b/rdmo/projects/utils.py @@ -33,10 +33,34 @@ def is_last_owner(project, user): return False -def check_conditions(conditions, values, set_prefix=None, set_index=None): +def compute_attribute_values_map(values): + attribute_values_map = defaultdict(list) + for value in values: + attribute_values_map[value.attribute_id].append(value) + return attribute_values_map + + +def compute_sets(distinct_values): + sets = defaultdict(set) + for attribute_id, set_prefix, set_index in distinct_values: + sets[attribute_id].add((set_prefix, set_index)) + return sets + + +def resolve_condition(condition, attribute_values_map, resolved_conditions, set_prefix=None, set_index=None): + condition_index = (condition.pk, set_prefix, set_index) + if condition_index not in resolved_conditions: + values = attribute_values_map.get(condition.source_id, ()) + resolved_conditions[condition_index] = condition.resolve(values, set_prefix, set_index) + return resolved_conditions[condition_index] + + +def check_conditions(conditions, attribute_values_map, set_prefix=None, set_index=None, resolved_conditions=None): + if resolved_conditions is None: + resolved_conditions = {} if conditions: for condition in conditions: - if condition.resolve(values, set_prefix, set_index): + if resolve_condition(condition, attribute_values_map, resolved_conditions, set_prefix, set_index): return True return False else: diff --git a/rdmo/projects/viewsets.py b/rdmo/projects/viewsets.py index 8bdfd2b050..f3d29c3984 100644 --- a/rdmo/projects/viewsets.py +++ b/rdmo/projects/viewsets.py @@ -85,6 +85,7 @@ from .utils import ( check_conditions, check_options, + compute_attribute_values_map, compute_set_prefix_from_set_value, copy_project, get_contact_message, @@ -133,6 +134,8 @@ def get_queryset(self): if self.action == 'navigation': # navigation only needs the project catalog and visibility before computing the answer tree. return queryset.select_related('catalog', 'visibility') + elif self.action in ('resolve', 'resolve_post'): + return queryset queryset = queryset.prefetch_related( 'snapshots', @@ -206,14 +209,16 @@ def resolve(self, request, pk=None): set_prefix = request.GET.get('set_prefix') set_index = request.GET.get('set_index') - values = self.get_object().values.filter(snapshot_id=snapshot_id).select_related('attribute', 'option') + values = self.get_object().values.filter(snapshot_id=snapshot_id).order_by() + attribute_values_map = compute_attribute_values_map(values) + resolved_conditions = {} page_id = request.GET.get('page') if page_id: try: page = Page.objects.get(id=page_id) - conditions = page.conditions.select_related('source', 'target_option') - if check_conditions(conditions, values, set_prefix, set_index): + conditions = page.conditions.all() + if check_conditions(conditions, attribute_values_map, set_prefix, set_index, resolved_conditions): return Response({'result': True}) except Page.DoesNotExist: pass @@ -222,8 +227,8 @@ def resolve(self, request, pk=None): if questionset_id: try: questionset = QuestionSet.objects.get(id=questionset_id) - conditions = questionset.conditions.select_related('source', 'target_option') - if check_conditions(conditions, values, set_prefix, set_index): + conditions = questionset.conditions.all() + if check_conditions(conditions, attribute_values_map, set_prefix, set_index, resolved_conditions): return Response({'result': True}) except QuestionSet.DoesNotExist: pass @@ -232,8 +237,8 @@ def resolve(self, request, pk=None): if question_id: try: question = Question.objects.get(id=question_id) - conditions = question.conditions.select_related('source', 'target_option') - if check_conditions(conditions, values, set_prefix, set_index): + conditions = question.conditions.all() + if check_conditions(conditions, attribute_values_map, set_prefix, set_index, resolved_conditions): return Response({'result': True}) except Question.DoesNotExist: pass @@ -242,8 +247,8 @@ def resolve(self, request, pk=None): if optionset_id: try: optionset = OptionSet.objects.get(id=optionset_id) - conditions = optionset.conditions.select_related('source', 'target_option') - if check_conditions(conditions, values, set_prefix, set_index): + conditions = optionset.conditions.all() + if check_conditions(conditions, attribute_values_map, set_prefix, set_index, resolved_conditions): return Response({'result': True}) except OptionSet.DoesNotExist: pass @@ -251,8 +256,8 @@ def resolve(self, request, pk=None): condition_id = request.GET.get('condition') if condition_id: try: - condition = Condition.objects.select_related('source', 'target_option').get(id=condition_id) - if check_conditions([condition], values, set_prefix, set_index): + condition = Condition.objects.get(id=condition_id) + if check_conditions([condition], attribute_values_map, set_prefix, set_index, resolved_conditions): return Response({'result': True}) except Condition.DoesNotExist: pass @@ -310,20 +315,31 @@ def resolve_post(self, request, pk=None): for element_dict in elements.values() for conditions_set in element_dict.values() )) - conditions = Condition.objects.select_related('source', 'target_option').in_bulk(condition_ids) + conditions = Condition.objects.in_bulk(condition_ids) + missing_condition_ids = condition_ids.difference(conditions) - # get all values of the project - values = project.values.filter(snapshot=None).select_related('attribute', 'option') + if conditions: + values = project.values.filter(snapshot=None).order_by() + attribute_values_map = compute_attribute_values_map(values) + else: + attribute_values_map = {} # second pass: resolve conditions + resolved_conditions = {} for params in validated_data: set_prefix = params['set_prefix'] set_index = params['set_index'] element_type = params['element_type'] element_id = params['element_id'] - element_conditions = [conditions[condition_id] for condition_id in elements[element_type][element_id]] - params['result'] = check_conditions(element_conditions, values, set_prefix, set_index) + element_condition_ids = elements[element_type][element_id] + if element_condition_ids.isdisjoint(missing_condition_ids): + element_conditions = [conditions[condition_id] for condition_id in element_condition_ids] + params['result'] = check_conditions( + element_conditions, attribute_values_map, set_prefix, set_index, resolved_conditions + ) + else: + params['result'] = False return Response(validated_data)