Skip to content
Merged
67 changes: 44 additions & 23 deletions cdisc_rules_engine/models/actions.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
from typing import List, Optional, Set, Hashable
from typing import List, Optional, Set, Hashable, Iterable
from os import path
import pandas as pd
from business_rules.actions import BaseActions, rule_action
Expand Down Expand Up @@ -48,7 +48,8 @@ def generate_record_message(self, message, target=None):
def generate_dataset_error_objects(self, message: str, results: pd.Series):
# leave only those columns where errors have been found
rows_with_error = self.variable.dataset.get_error_rows(results)
target_names: Set[str] = RuleProcessor.extract_target_names_from_rule(
# get targets in the order they appear in rule.output_variables
target_names: List[str] = RuleProcessor.extract_target_names_from_rule(
self.rule,
self.dataset_metadata.domain_cleaned,
self.variable.dataset.columns.tolist(),
Expand All @@ -67,21 +68,29 @@ def generate_dataset_error_objects(self, message: str, results: pd.Series):
def generate_single_error(self, message):
self.output_container.append(message)

def _get_target_names_from_list_values(self, target_names, rows_with_error):
expanded_target_names = set(target_names)
expanded_target_names.update(
value
for target in target_names
if target in rows_with_error
for candidate_list in rows_with_error[target]
if isinstance(candidate_list, list)
for value in candidate_list
if value in self.variable.dataset.columns
)
return expanded_target_names
def _get_target_names_from_list_values(
self, target_names: List[str], rows_with_error: pd.DataFrame
) -> List[str]:
"""Expand target names with column names found inside list values, preserving order.

New targets are appended in the order they are discovered, duplicates are ignored.
"""
expanded: List[str] = list(target_names)
existing = set(expanded)
for target in target_names:
if target not in rows_with_error:
continue
for candidate_list in rows_with_error[target]:
if not isinstance(candidate_list, list):
continue
for value in candidate_list:
if value in self.variable.dataset.columns and value not in existing:
expanded.append(value)
existing.add(value)
return expanded

def generate_targeted_error_object( # noqa: C901
self, targets: Set[str], data: pd.DataFrame, message: str
self, targets: Iterable[str], data: pd.DataFrame, message: str
) -> ValidationErrorContainer:
"""
Generates a targeted error object.
Expand Down Expand Up @@ -116,17 +125,20 @@ def generate_targeted_error_object( # noqa: C901
"message": "AESTDY and DOMAIN are equal to test",
}
"""
# preserve incoming order for representation but use sets for membership tests
targets_list: List[str] = list(targets)
targets_set: Set[str] = set(targets_list)
df_columns: set = set(data)
targets_in_dataset = targets.intersection(df_columns)
targets_not_in_dataset = targets.difference(df_columns)
targets_in_dataset = targets_set.intersection(df_columns)
targets_not_in_dataset = targets_set.difference(df_columns)
all_targets_missing = (
len(targets_in_dataset) == 0 and len(targets_not_in_dataset) > 0
)
if targets_in_dataset:
errors_df = data[list(targets_in_dataset)]
else:
errors_df = data
if not targets:
if not targets_set:
errors_df = data

if self.rule.get("sensitivity") == Sensitivity.DATASET.value:
Expand Down Expand Up @@ -200,7 +212,7 @@ def generate_targeted_error_object( # noqa: C901
if self.dataset_metadata.is_supp
else (self.dataset_metadata.domain or self.dataset_metadata.name)
),
targets=sorted(targets),
targets=targets_list,
message="Invalid or undefined sensitivity in the rule",
errors=[error_entity],
)
Expand All @@ -217,7 +229,7 @@ def generate_targeted_error_object( # noqa: C901
dataset=", ".join(
sorted(set(error.dataset or "" for error in errors_list))
),
targets=sorted(targets),
targets=targets_list,
errors=errors_list,
message=message.replace("--", self.dataset_metadata.domain_cleaned or ""),
)
Expand Down Expand Up @@ -324,13 +336,14 @@ def _create_configuration_error(self, message, targets):
USUBJID="N/A",
SEQ=0,
)
targets_list = list(targets)
return ValidationErrorContainer(
domain=(
f"SUPP{self.dataset_metadata.rdomain}"
if self.dataset_metadata.is_supp
else (self.dataset_metadata.domain or self.dataset_metadata.name)
),
targets=sorted(targets),
targets=targets_list,
message=message,
errors=[error_entity],
)
Expand Down Expand Up @@ -485,8 +498,16 @@ def _create_error_object(
)
return error_object

def extract_target_names_from_value_level_metadata(self):
return set([item["define_variable_name"] for item in self.value_level_metadata])
def extract_target_names_from_value_level_metadata(self) -> List[str]:
"""Return target names from value-level metadata preserving input order."""
seen = set()
ordered: List[str] = []
for item in self.value_level_metadata:
name = item.get("define_variable_name")
if name and name not in seen:
seen.add(name)
ordered.append(name)
return ordered

@staticmethod
def _sequence_exists(sequence: pd.Series, row_name: Hashable) -> bool:
Expand Down
2 changes: 1 addition & 1 deletion cdisc_rules_engine/models/validation_error_container.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,7 +28,7 @@ def to_representation(self) -> dict:
"executionStatus": self.executionStatus,
"dataset": self.dataset,
"domain": self.domain,
"variables": sorted(self.targets),
"variables": self.targets,
"message": self.message,
"errors": [error.to_representation() for error in self.errors],
**({"entity": self.entity} if self.entity else {}),
Expand Down
79 changes: 47 additions & 32 deletions cdisc_rules_engine/utilities/rule_processor.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@
import copy
import os

from typing import Iterable, List, Optional, Set, Union, Tuple
from typing import Iterable, List, Optional, Union, Tuple
from cdisc_rules_engine.enums.rule_types import RuleTypes
from cdisc_rules_engine.interfaces.cache_service_interface import (
CacheServiceInterface,
Expand Down Expand Up @@ -736,10 +736,50 @@ def is_suitable_for_validation(
return False, reason
return self.log_suitable_for_validation(rule_id, dataset_name)

@staticmethod
def _extract_targets_from_output_variables(rule: dict, domain: str) -> List[str]:
output_variables: List[str] = rule.get("output_variables", [])
target_names: List[str] = []
seen: set[str] = set()
for var in output_variables:
name = var.replace("--", domain or "", 1)
if name not in seen:
seen.add(name)
target_names.append(name)
return target_names

@staticmethod
def _extract_targets_from_conditions(
rule: dict, domain: str, column_names: List[str]
) -> List[str]:
target_names: List[str] = []
seen: set[str] = set()
conditions: ConditionInterface = rule["conditions"]
for condition in conditions.values():
if condition.get("operator") == "not_exists":
continue
target: str = condition["value"].get("target")
if target is None:
continue
target = target.replace("--", domain or "")
op_related_pattern: str = RuleProcessor.get_operator_related_pattern(
condition.get("operator"), target
)
if op_related_pattern is not None:
for name in column_names:
if re.match(op_related_pattern, name) and name not in seen:
seen.add(name)
target_names.append(name)
else:
if target not in seen:
seen.add(target)
target_names.append(target)
return target_names

@staticmethod
def extract_target_names_from_rule(
rule: dict, domain: str, column_names: List[str]
) -> Set[str]:
) -> List[str]:
r"""
Extracts target from each item of condition list.

Expand All @@ -754,36 +794,11 @@ def extract_target_names_from_rule(
pattern: ^TSVAL\d+$ (starts with TSVAL and ends with number)
additional columns: TSVAL1, TSVAL2, TSVAL3 etc.
"""
output_variables: List[str] = rule.get("output_variables", [])
if output_variables:
target_names: List[str] = [
var.replace("--", domain or "", 1) for var in output_variables
]
else:
target_names: List[str] = []
conditions: ConditionInterface = rule["conditions"]
for condition in conditions.values():
if condition.get("operator") == "not_exists":
continue
target: str = condition["value"].get("target")
if target is None:
continue
target = target.replace("--", domain or "")
op_related_pattern: str = RuleProcessor.get_operator_related_pattern(
condition.get("operator"), target
)
if op_related_pattern is not None:
# if pattern exists -> return only matching column names
target_names.extend(
filter(
lambda name: re.match(op_related_pattern, name),
column_names,
)
)
else:
target_names.append(target)
target_names.sort()
return set(target_names)
if rule.get("output_variables"):
return RuleProcessor._extract_targets_from_output_variables(rule, domain)
return RuleProcessor._extract_targets_from_conditions(
rule, domain, column_names
)

@staticmethod
def extract_referenced_variables_from_rule(rule: dict):
Expand Down
14 changes: 7 additions & 7 deletions tests/unit/test_rules_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -952,7 +952,7 @@ def test_validate_variable_metadata_wrong_metadata(
{
"domain": "EC",
"dataset": "bundle",
"variables": ["variable_data_type", "variable_label", "variable_name"],
"variables": ["variable_name", "variable_label", "variable_data_type"],
"executionStatus": ExecutionStatus.SUCCESS.value,
"errors": [
{
Expand Down Expand Up @@ -1289,7 +1289,7 @@ def test_validate_single_dataset_not_equal_to(
"executionStatus": "success",
"dataset": "ae.xpt",
"domain": "AE",
"variables": ["dataset_label", "dataset_location", "dataset_name"],
"variables": ["dataset_label", "dataset_name", "dataset_location"],
"message": "Dataset metadata does not correspond to Define XML",
"errors": [
{
Expand Down Expand Up @@ -1887,7 +1887,7 @@ def test_validate_split_dataset_variables_metadata(
"domain": "EC",
"dataset": "ec_2.xpt",
"executionStatus": ExecutionStatus.SUCCESS.value,
"variables": ["variable_data_type", "variable_label", "variable_name"],
"variables": ["variable_name", "variable_label", "variable_data_type"],
"errors": [
{
"dataset": "ec_2.xpt",
Expand Down Expand Up @@ -2011,7 +2011,7 @@ def test_validate_record_in_parent_domain(
"executionStatus": "success",
"domain": "EC",
"dataset": "ec.xpt",
"variables": ["ECPRESP", "ECREASOC"],
"variables": ["ECREASOC", "ECPRESP"],
"message": "Dataset contents is wrong.",
"errors": [
{
Expand Down Expand Up @@ -2187,8 +2187,8 @@ def test_validate_dataset_contents_against_define_and_library_variable_metadata(
"dataset": "filename",
"domain": "AE",
"variables": [
"AESER",
"AESEV",
"AESER",
], # AELNKID must not be included since its core status is not "Perm"
"message": RuleProcessor.extract_message_from_rule(
rule_check_dataset_against_library_and_define
Expand Down Expand Up @@ -2686,10 +2686,10 @@ def mock_cached_method(*args, **kwargs):
"variables": [
"$column_order_from_dataset",
"$column_order_from_library",
"AESEQ",
"AETERM",
"DOMAIN",
"AESEQ",
"STUDYID",
"AETERM",
],
"message": "Order of variables is invalid",
"errors": [
Expand Down
2 changes: 1 addition & 1 deletion tests/unit/test_utilities/test_jsonata_processor.py
Original file line number Diff line number Diff line change
Expand Up @@ -53,7 +53,7 @@ class TestJSONataProcessor(TestCase):
"executionStatus": "success",
"dataset": None,
"domain": None,
"variables": ["A", "B", "row"],
"variables": ["row", "A", "B"],
"message": "A equals B",
"errors": [
{
Expand Down
14 changes: 7 additions & 7 deletions tests/unit/test_utilities/test_rule_processor.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
from typing import List, Set
from typing import List
from unittest.mock import MagicMock, patch

import pandas as pd
Expand Down Expand Up @@ -1101,7 +1101,7 @@ def test_extract_target_names_from_rule():
rule: dict = {
"conditions": ConditionCompositeFactory.get_condition_composite(conditions),
}
target_names: Set[str] = RuleProcessor.extract_target_names_from_rule(
target_names: List[str] = RuleProcessor.extract_target_names_from_rule(
rule,
"AE",
[
Expand All @@ -1110,11 +1110,11 @@ def test_extract_target_names_from_rule():
"TARGET",
],
)
assert target_names == {
assert target_names == [
"AESTDY",
"USUBJID",
"TARGET",
}
]


def test_extract_target_names_from_rule_output_variables():
Expand All @@ -1130,7 +1130,7 @@ def test_extract_target_names_from_rule_output_variables():
"output_variables": ["AESTDY", "USUBJID", "TARGET"],
"conditions": ConditionCompositeFactory.get_condition_composite(conditions),
}
target_names: Set[str] = RuleProcessor.extract_target_names_from_rule(
target_names: List[str] = RuleProcessor.extract_target_names_from_rule(
rule,
"AE",
[
Expand All @@ -1139,11 +1139,11 @@ def test_extract_target_names_from_rule_output_variables():
"TARGET",
],
)
assert target_names == {
assert target_names == [
"AESTDY",
"USUBJID",
"TARGET",
}
]


@pytest.mark.parametrize(
Expand Down
Loading