diff --git a/src/vtlengine/API/_InternalApi.py b/src/vtlengine/API/_InternalApi.py index 3787e55e9..41f70c239 100644 --- a/src/vtlengine/API/_InternalApi.py +++ b/src/vtlengine/API/_InternalApi.py @@ -36,15 +36,16 @@ to_vtl_json, ) from vtlengine.Model import ( - Component as VTL_Component, -) -from vtlengine.Model import ( + CaseInsensitiveDict, Dataset, ExternalRoutine, Role, Scalar, ValueDomain, ) +from vtlengine.Model import ( + Component as VTL_Component, +) # Cache SCALAR_TYPES keys for performance _SCALAR_TYPE_KEYS = SCALAR_TYPES.keys() @@ -118,27 +119,34 @@ def _load_dataset_from_structure( """ _validate_json(structures, schema, kind="DataStructures") - datasets = { - dataset_json["name"]: Dataset( - name=dataset_json["name"], - components={ - c["name"]: _build_component(c) - for c in _resolve_components(dataset_json, structures) - }, - data=None, - ) - for dataset_json in structures.get("datasets", []) - } + # VTL regular names are case-insensitive: keep datasets/scalars in + # CaseInsensitiveDict so lookups match regardless of the written casing + # (component dicts are wrapped by Dataset.__post_init__). + datasets: CaseInsensitiveDict[Dataset] = CaseInsensitiveDict( + { + dataset_json["name"]: Dataset( + name=dataset_json["name"], + components={ + c["name"]: _build_component(c) + for c in _resolve_components(dataset_json, structures) + }, + data=None, + ) + for dataset_json in structures.get("datasets", []) + } + ) - scalars = { - scalar_json["name"]: Scalar( - name=scalar_json["name"], - data_type=_extract_data_type(scalar_json)[1], - value=None, - nullable=scalar_json.get("nullable", True), - ) - for scalar_json in structures.get("scalars", []) - } + scalars: CaseInsensitiveDict[Scalar] = CaseInsensitiveDict( + { + scalar_json["name"]: Scalar( + name=scalar_json["name"], + data_type=_extract_data_type(scalar_json)[1], + value=None, + nullable=scalar_json.get("nullable", True), + ) + for scalar_json in structures.get("scalars", []) + } + ) return datasets, scalars @@ -344,15 +352,15 @@ def _load_datastructure_single( if not data_structure.exists(): raise DataLoadError(code="0-3-1-1", file=data_structure) if data_structure.is_dir(): - datasets: Dict[str, Dataset] = {} - scalars: Dict[str, Scalar] = {} + dir_datasets: CaseInsensitiveDict[Dataset] = CaseInsensitiveDict() + dir_scalars: CaseInsensitiveDict[Scalar] = CaseInsensitiveDict() for f in data_structure.iterdir(): if f.suffix not in (".json", ".xml"): continue ds, sc = _load_datastructure_single(f, sdmx_mappings=sdmx_mappings) - datasets = {**datasets, **ds} - scalars = {**scalars, **sc} - return datasets, scalars + dir_datasets.update(ds) + dir_scalars.update(sc) + return dir_datasets, dir_scalars else: suffix = data_structure.suffix.lower() # Handle SDMX-ML structure files (.xml) - strict, must be SDMX diff --git a/src/vtlengine/API/__init__.py b/src/vtlengine/API/__init__.py index 1cc028695..cd542f62b 100644 --- a/src/vtlengine/API/__init__.py +++ b/src/vtlengine/API/__init__.py @@ -41,6 +41,13 @@ pd.options.mode.chained_assignment = None +def _convert_to_regular_dicts(result: Dict[str, Any]) -> None: + """Convert internal CaseInsensitiveDict components to a plain dict for external consumers.""" + for obj in result.values(): + if isinstance(obj, Dataset): + obj.components = dict(obj.components) + + def _extract_input_datasets(script: Union[str, TransformationScheme, Path]) -> List[str]: if isinstance(script, TransformationScheme): vtl_script = _check_script(script) @@ -247,6 +254,7 @@ def semantic_analysis( scalars=scalars, ) result = interpreter.visit(ast) + _convert_to_regular_dicts(result) return result @@ -479,6 +487,9 @@ def run( format_date_iso8601(obj) format_time_period_external_representation(obj, time_period_representation) + # Convert internal CaseInsensitiveDict to plain dict for external consumers + _convert_to_regular_dicts(results) + return results diff --git a/src/vtlengine/AST/DAG/__init__.py b/src/vtlengine/AST/DAG/__init__.py index 5e94566dd..f055cba6e 100644 --- a/src/vtlengine/AST/DAG/__init__.py +++ b/src/vtlengine/AST/DAG/__init__.py @@ -48,7 +48,7 @@ TO, ) from vtlengine.Exceptions import SemanticError -from vtlengine.Model import Component +from vtlengine.Model import Component, normalize_name @dataclass @@ -82,13 +82,15 @@ def _ds_usage_analysis(self) -> DatasetSchedule: deletion: Dict[int, List[str]] = defaultdict(list) insertion: Dict[int, List[str]] = defaultdict(list) all_outputs: Set[str] = set() + # Casefolded mirror of all_outputs, used only for case-insensitive matching. + all_outputs_norm: Set[str] = set() persistent_datasets: List[str] = [] - # Reverse index: dataset_name -> last statement that uses it as input + # Reverse index: dataset_name (casefolded) -> last statement that uses it as input last_consumer: Dict[str, int] = {} for key, statement in self.dependencies.items(): for input_name in statement.inputs: - last_consumer[input_name] = key + last_consumer[normalize_name(input_name)] = key # Schedule deletion for statement outputs at their last consumer for key, statement in self.dependencies.items(): @@ -100,17 +102,19 @@ def _ds_usage_analysis(self) -> DatasetSchedule: persistent_datasets.append(statement.persistent[0]) ds_name = reference[0] all_outputs.add(ds_name) - deletion[last_consumer.get(ds_name, key)].append(ds_name) + all_outputs_norm.add(normalize_name(ds_name)) + deletion[last_consumer.get(normalize_name(ds_name), key)].append(ds_name) # Schedule insertion (first use) and deletion (last use) for global inputs global_inputs: List[str] = [] global_set: Set[str] = set() for key, statement in self.dependencies.items(): for element in statement.inputs: - if element not in all_outputs and element not in global_set: - global_set.add(element) + norm = normalize_name(element) + if norm not in all_outputs_norm and norm not in global_set: + global_set.add(norm) global_inputs.append(element) - deletion[last_consumer.get(element, key)].append(element) + deletion[last_consumer.get(norm, key)].append(element) insertion[key].append(element) return DatasetSchedule( @@ -183,12 +187,12 @@ def load_edges(self) -> None: for key, statement in self.dependencies.items(): reference = statement.outputs + statement.persistent if reference: - ref_to_keys[reference[0]] = key + ref_to_keys[normalize_name(reference[0])] = key for sub_key, sub_statement in self.dependencies.items(): for input_val in sub_statement.inputs: - if input_val in ref_to_keys: - key = ref_to_keys[input_val] + if normalize_name(input_val) in ref_to_keys: + key = ref_to_keys[normalize_name(input_val)] self.edges[count_edges] = (key, sub_key) count_edges += 1 @@ -196,11 +200,12 @@ def sort_elements(self, statements: list) -> list: return [statements[x - 1] for x in self.sorting] # type: ignore[union-attr] def check_overwriting(self, statements: list) -> None: + # Regular names are case-insensitive: DS_r and DS_R are the same output. seen: Set[str] = set() for statement in statements: - if statement.left.value in seen: + if normalize_name(statement.left.value) in seen: raise SemanticError("1-2-2", varId_value=statement.left.value) - seen.add(statement.left.value) + seen.add(normalize_name(statement.left.value)) def sort_ast(self, ast: AST) -> None: statements_nodes = ast.children @@ -290,13 +295,11 @@ def visit_RegularAggregation(self, node: RegularAggregation) -> None: self.visit(node.dataset) if node.op in [KEEP, DROP, RENAME]: return - saved_is_dataset = self.is_dataset self.is_dataset = False for child in node.children: self.is_from_regular_aggregation = True self.visit(child) self.is_from_regular_aggregation = False - self.is_dataset = saved_is_dataset def visit_BinOp(self, node: BinOp) -> None: if node.op == MEMBERSHIP: @@ -306,20 +309,20 @@ def visit_BinOp(self, node: BinOp) -> None: self.visit(node.right) elif node.op == AS or node.op == TO: self.visit(node.left) - self.alias.add(node.right.value) + self.alias.add(normalize_name(node.right.value)) else: self.visit(node.left) self.visit(node.right) def visit_VarID(self, node: VarID) -> None: - if ( - not self.is_from_regular_aggregation or self.is_dataset - ) and node.value not in self.alias: + if (not self.is_from_regular_aggregation or self.is_dataset) and normalize_name( + node.value + ) not in self.alias: if node.value not in self.current_deps.inputs: self.current_deps.inputs.append(node.value) elif ( self.is_from_regular_aggregation - and node.value not in self.alias + and normalize_name(node.value) not in self.alias and not self.is_dataset and node.value not in self.current_deps.unknown_variables ): @@ -328,7 +331,7 @@ def visit_VarID(self, node: VarID) -> None: def visit_Identifier(self, node: Identifier) -> None: if ( node.kind == "DatasetID" - and node.value not in self.alias + and normalize_name(node.value) not in self.alias and node.value not in self.current_deps.inputs ): self.current_deps.inputs.append(node.value) diff --git a/src/vtlengine/Interpreter/__init__.py b/src/vtlengine/Interpreter/__init__.py index f798e609c..433626dd2 100644 --- a/src/vtlengine/Interpreter/__init__.py +++ b/src/vtlengine/Interpreter/__init__.py @@ -46,6 +46,7 @@ ) from vtlengine.Exceptions import SemanticError from vtlengine.Model import ( + CaseInsensitiveDict, Component, DataComponent, Dataset, @@ -54,6 +55,8 @@ Scalar, ScalarSet, ValueDomain, + names_equal, + normalize_name, ) from vtlengine.Operators.Aggregation import extract_grouping_identifiers from vtlengine.Operators.Assignment import Assignment @@ -131,8 +134,17 @@ class InterpreterAnalyzer(ASTTemplate): signature_values: Optional[Dict[str, Any]] = None def __post_init__(self) -> None: - self.datasets_inputs = set(self.datasets.keys()) - self.scalars_inputs = set(self.scalars.keys()) if self.scalars else set() + # VTL regular names are case-insensitive: keep all symbol tables in + # CaseInsensitiveDict so lookups match regardless of the written casing. + self.datasets = CaseInsensitiveDict(self.datasets) + if self.scalars is not None: + self.scalars = CaseInsensitiveDict(self.scalars) + if self.value_domains is not None: + self.value_domains = CaseInsensitiveDict(self.value_domains) + if self.external_routines is not None: + self.external_routines = CaseInsensitiveDict(self.external_routines) + self.datasets_inputs = {normalize_name(k) for k in self.datasets} + self.scalars_inputs = {normalize_name(k) for k in self.scalars} if self.scalars else set() # ********************************** # * * @@ -155,9 +167,9 @@ def visit_Start(self, node: AST.Start) -> Any: ) and not isinstance(child, (AST.Assignment, AST.PersistentAssignment)): raise SemanticError("1-2-5") result = self.visit(child) - if isinstance(result, Dataset) and result.name in self.datasets_inputs: + if isinstance(result, Dataset) and normalize_name(result.name) in self.datasets_inputs: invalid_dataset_outputs.append(result.name) - if isinstance(result, Scalar) and result.name in self.scalars_inputs: + if isinstance(result, Scalar) and normalize_name(result.name) in self.scalars_inputs: invalid_scalar_outputs.append(result.name) self.is_from_join = False @@ -171,7 +183,7 @@ def visit_Start(self, node: AST.Start) -> Any: results[result.name] = result if isinstance(result, Scalar): if self.scalars is None: - self.scalars = {} + self.scalars = CaseInsensitiveDict() self.scalars[result.name] = copy(result) if invalid_dataset_outputs: raise SemanticError("0-1-2-8", names=", ".join(invalid_dataset_outputs)) @@ -184,7 +196,7 @@ def visit_Start(self, node: AST.Start) -> Any: def visit_Operator(self, node: AST.Operator) -> None: if self.udos is None: - self.udos = {} + self.udos = CaseInsensitiveDict() elif node.op in self.udos: raise ValueError(f"User Defined Operator {node.op} already exists") @@ -234,7 +246,7 @@ def visit_DPRuleset(self, node: AST.DPRuleset) -> None: ) # Signature has the actual parameters names or aliases if provided - signature_actual_names = {} + signature_actual_names: CaseInsensitiveDict[str] = CaseInsensitiveDict() if not isinstance(node.params, AST.DefIdentifier): for param in node.params: if param.alias is not None: @@ -255,7 +267,7 @@ def visit_DPRuleset(self, node: AST.DPRuleset) -> None: # Adding the ruleset to the dprs dictionary if self.dprs is None: - self.dprs = {} + self.dprs = CaseInsensitiveDict() elif node.name in self.dprs: raise ValueError(f"Datapoint Ruleset {node.name} already exists") @@ -263,7 +275,7 @@ def visit_DPRuleset(self, node: AST.DPRuleset) -> None: def visit_HRuleset(self, node: AST.HRuleset) -> None: if self.hrs is None: - self.hrs = {} + self.hrs = CaseInsensitiveDict() if node.name in self.hrs: raise ValueError(f"Hierarchical Ruleset {node.name} already exists") @@ -720,16 +732,16 @@ def visit_VarID(self, node: AST.VarID) -> Any: # noqa: C901 return copy(self.scalars[node.value]) if ( self.is_from_join - and node.value not in self.regular_aggregation_dataset.get_components_names() + and node.value not in self.regular_aggregation_dataset.components ): is_partial_present = 0 found_comp = None for comp_name in self.regular_aggregation_dataset.get_components_names(): if ( "#" in comp_name - and comp_name.split("#")[1] == node.value + and names_equal(comp_name.split("#")[1], node.value) or "#" in node.value - and node.value.split("#")[1] == comp_name + and names_equal(node.value.split("#")[1], comp_name) ): is_partial_present += 1 found_comp = comp_name @@ -1163,7 +1175,9 @@ def visit_HROperation(self, node: AST.HROperation) -> Any: # noqa: C901 if len(cond_components) != len(hr_info["condition"]): raise SemanticError("1-1-10-2", op=node.op) - if hr_info["node"].signature_type == "variable" and hr_info["signature"] != component: + if hr_info["node"].signature_type == "variable" and not names_equal( + hr_info["signature"], component + ): raise SemanticError( "1-1-10-3", op=node.op, @@ -1223,7 +1237,9 @@ def visit_HROperation(self, node: AST.HROperation) -> Any: # noqa: C901 # Set up interpreter state for rule processing self.ruleset_dataset = dataset - self.ruleset_signature = {**{"RULE_COMPONENT": component}, **cond_info} + self.ruleset_signature = CaseInsensitiveDict( + {**{"RULE_COMPONENT": component}, **cond_info} + ) self.ruleset_mode = mode rule_output_values = {} @@ -1284,7 +1300,7 @@ def visit_DPValidation(self, node: AST.DPValidation) -> Any: and dpr_info["params"] ): for i, comp_name in enumerate(node.components): - if comp_name != dpr_info["params"][i]: + if not names_equal(comp_name, dpr_info["params"][i]): raise SemanticError( "1-1-10-3", op=CHECK_DATAPOINT, diff --git a/src/vtlengine/Model/__init__.py b/src/vtlengine/Model/__init__.py index 6bcc737d2..6412557d5 100644 --- a/src/vtlengine/Model/__init__.py +++ b/src/vtlengine/Model/__init__.py @@ -1,9 +1,10 @@ import inspect import json from collections import Counter +from copy import deepcopy from dataclasses import dataclass from enum import Enum -from typing import Any, Dict, List, Optional, Type, Union +from typing import Any, Dict, Iterator, List, Optional, Tuple, Type, TypeVar, Union import pandas as pd import sqlglot @@ -16,6 +17,160 @@ from vtlengine.DataTypes.TimeHandling import TimePeriodHandler from vtlengine.Exceptions import InputValidationException, SemanticError +V = TypeVar("V") + + +def normalize_name(name: str) -> str: + """Canonical form of a VTL regular name, used for case-insensitive matching.""" + return name.casefold() + + +def names_equal(a: Optional[str], b: Optional[str]) -> bool: + """Case-insensitive equality of two VTL regular names (None-safe).""" + if a is None or b is None: + return a is b + return normalize_name(a) == normalize_name(b) + + +class CaseInsensitiveDict(Dict[str, V]): + """A dict subclass that treats string keys as case-insensitive.""" + + def __init__(self, *args: Any, **kwargs: V) -> None: + self._key_map: Dict[str, str] = {} # lowercase -> original key + super().__init__() + if args: + arg = args[0] + if isinstance(arg, dict): + for k, v in arg.items(): + self[k] = v + elif hasattr(arg, "__iter__"): + for k, v in arg: + self[k] = v + for k, v in kwargs.items(): + self[k] = v + + def _normalize(self, key: str) -> str: + return normalize_name(key) + + def __setitem__(self, key: str, value: V) -> None: + norm = self._normalize(key) + if norm not in self._key_map: + self._key_map[norm] = key + original = self._key_map[norm] + super().__setitem__(original, value) + + def __getitem__(self, key: str) -> V: + norm = self._normalize(key) + if norm not in self._key_map: + raise KeyError(key) + return super().__getitem__(self._key_map[norm]) + + def __contains__(self, key: object) -> bool: + if not isinstance(key, str): + return False + return self._normalize(key) in self._key_map + + def __delitem__(self, key: str) -> None: + norm = self._normalize(key) + if norm not in self._key_map: + raise KeyError(key) + original = self._key_map.pop(norm) + super().__delitem__(original) + + def get(self, key: str, default: Optional[V] = None) -> Optional[V]: # type: ignore[override] + norm = self._normalize(key) + if norm not in self._key_map: + return default + return super().__getitem__(self._key_map[norm]) + + def pop(self, key: str, *args: V) -> V: # type: ignore[override] + norm = self._normalize(key) + if norm not in self._key_map: + if args: + return args[0] + raise KeyError(key) + original = self._key_map.pop(norm) + return super().pop(original) + + def setdefault(self, key: str, default: Optional[V] = None) -> V: + norm = self._normalize(key) + if norm not in self._key_map: + self[key] = default # type: ignore[assignment] + return self[key] + + def update(self, *args: Any, **kwargs: V) -> None: + if args: + other = args[0] + if isinstance(other, dict): + for k, v in other.items(): + self[k] = v + elif hasattr(other, "__iter__"): + for k, v in other: + self[k] = v + for k, v in kwargs.items(): + self[k] = v + + def canonical_key(self, key: str) -> str: + """Return the original-case key for a given (possibly different-case) key. + + Raises KeyError if the key doesn't exist. + """ + norm = self._normalize(key) + if norm not in self._key_map: + raise KeyError(key) + return self._key_map[norm] + + def __iter__(self) -> Iterator[str]: + return super().__iter__() + + def copy(self) -> "CaseInsensitiveDict[V]": + result: CaseInsensitiveDict[V] = CaseInsensitiveDict() + result._key_map = self._key_map.copy() + for key in dict.keys(self): + dict.__setitem__(result, key, dict.__getitem__(self, key)) + return result + + def __repr__(self) -> str: + return f"CaseInsensitiveDict({dict(self.items())})" + + def __eq__(self, other: object) -> bool: + if isinstance(other, CaseInsensitiveDict): + return dict.__eq__(self, other) + if isinstance(other, dict): + if len(self) != len(other): + return False + return all(k in self and self[k] == v for k, v in other.items()) + return NotImplemented + + def __deepcopy__(self, memo: Dict[int, Any]) -> "CaseInsensitiveDict[V]": + new: CaseInsensitiveDict[V] = CaseInsensitiveDict.__new__(CaseInsensitiveDict) + memo[id(self)] = new + dict.__init__(new) + new._key_map = deepcopy(self._key_map, memo) + for key in dict.keys(self): + dict.__setitem__(new, key, deepcopy(dict.__getitem__(self, key), memo)) + return new + + def __copy__(self) -> "CaseInsensitiveDict[V]": + new: CaseInsensitiveDict[V] = CaseInsensitiveDict.__new__(CaseInsensitiveDict) + dict.__init__(new) + new._key_map = self._key_map.copy() + for key in dict.keys(self): + dict.__setitem__(new, key, dict.__getitem__(self, key)) + return new + + @classmethod + def from_dict(cls, d: Dict[str, V]) -> "CaseInsensitiveDict[V]": + """Create a CaseInsensitiveDict from a regular dict.""" + return cls(d) + + def to_dict(self) -> Dict[str, V]: + """Convert back to a regular dict with original-cased keys.""" + return dict(self.items()) + + def __reduce__(self) -> Tuple[type, Tuple[Dict[str, V]]]: + return (CaseInsensitiveDict, (dict(self.items()),)) + @dataclass class Scalar: @@ -220,6 +375,8 @@ class Dataset: persistent: bool = False def __post_init__(self) -> None: + if not isinstance(self.components, CaseInsensitiveDict): + self.components = CaseInsensitiveDict(self.components) if self.data is not None: if len(self.components) != len(self.data.columns): raise ValueError( @@ -339,7 +496,8 @@ def add_component(self, component: Component) -> None: self.components[component.name] = component def delete_component(self, component_name: str) -> None: - self.components.pop(component_name, None) + if component_name in self.components: + del self.components[component_name] if self.data is not None: self.data.drop(columns=[component_name], inplace=True) diff --git a/src/vtlengine/Operators/Aggregation.py b/src/vtlengine/Operators/Aggregation.py index 5b5fa9382..6ac8d1904 100644 --- a/src/vtlengine/Operators/Aggregation.py +++ b/src/vtlengine/Operators/Aggregation.py @@ -21,7 +21,7 @@ unary_implicit_promotion, ) from vtlengine.Exceptions import SemanticError -from vtlengine.Model import Component, Dataset, Role +from vtlengine.Model import Component, Dataset, Role, normalize_name def extract_grouping_identifiers( @@ -30,7 +30,9 @@ def extract_grouping_identifiers( if group_op == "group by": return grouping_components elif group_op == "group except": - return [comp for comp in identifier_names if comp not in grouping_components] + # Regular names are case-insensitive. + excluded = {normalize_name(comp) for comp in grouping_components} + return [comp for comp in identifier_names if normalize_name(comp) not in excluded] elif group_op == "group all": return identifier_names if grouping_components else [] else: @@ -69,8 +71,9 @@ def validate( # type: ignore[override] identifiers_to_keep = extract_grouping_identifiers( operand.get_identifiers_names(), group_op, grouping_columns ) + keep_norm = {normalize_name(name) for name in identifiers_to_keep} for comp_name, comp in operand.components.items(): - if comp.role == Role.IDENTIFIER and comp_name not in identifiers_to_keep: + if comp.role == Role.IDENTIFIER and normalize_name(comp_name) not in keep_norm: del result_components[comp_name] else: for comp_name, comp in operand.components.items(): diff --git a/src/vtlengine/Operators/Clause.py b/src/vtlengine/Operators/Clause.py index ed051520b..02b68c96f 100644 --- a/src/vtlengine/Operators/Clause.py +++ b/src/vtlengine/Operators/Clause.py @@ -12,7 +12,15 @@ unary_implicit_promotion, ) from vtlengine.Exceptions import SemanticError -from vtlengine.Model import Component, DataComponent, Dataset, Role, Scalar +from vtlengine.Model import ( + CaseInsensitiveDict, + Component, + DataComponent, + Dataset, + Role, + Scalar, + normalize_name, +) from vtlengine.Operators import Operator from vtlengine.Utils.__Virtual_Assets import VirtualCounter @@ -152,17 +160,21 @@ class Rename(Operator): @classmethod def validate(cls, operands: List[RenameNode], dataset: Dataset) -> Dataset: dataset_name = VirtualCounter._new_ds_name() + # Regular names are case-insensitive: Me_1 and ME_1 are the same component, + # so duplicate detection must compare normalized names. from_names = [operand.old_name for operand in operands] - if len(from_names) != len(set(from_names)): - duplicates = set([name for name in from_names if from_names.count(name) > 1]) + norm_from = [normalize_name(name) for name in from_names] + if len(norm_from) != len(set(norm_from)): + duplicates = {name for name in from_names if norm_from.count(normalize_name(name)) > 1} raise SemanticError("1-1-6-9", op=cls.op, from_components=duplicates) to_names = [operand.new_name for operand in operands] - if len(to_names) != len(set(to_names)): # If duplicates - duplicates = set([name for name in to_names if to_names.count(name) > 1]) + norm_to = [normalize_name(name) for name in to_names] + if len(norm_to) != len(set(norm_to)): # If duplicates + duplicates = {name for name in to_names if norm_to.count(normalize_name(name)) > 1} raise SemanticError("1-2-1", alias=duplicates) - from_names_set = set(from_names) + from_names_set = {normalize_name(name) for name in from_names} for operand in operands: if operand.old_name not in dataset.components: raise SemanticError( @@ -171,7 +183,10 @@ def validate(cls, operands: List[RenameNode], dataset: Dataset) -> Dataset: comp_name=operand.old_name, dataset_name=dataset_name, ) - if operand.new_name in dataset.components and operand.new_name not in from_names_set: + if ( + operand.new_name in dataset.components + and normalize_name(operand.new_name) not in from_names_set + ): raise SemanticError( "1-1-6-8", op=cls.op, @@ -179,7 +194,9 @@ def validate(cls, operands: List[RenameNode], dataset: Dataset) -> Dataset: dataset_name=dataset_name, ) - rename_map = {op.old_name: op.new_name for op in operands} + rename_map: CaseInsensitiveDict[str] = CaseInsensitiveDict( + {op.old_name: op.new_name for op in operands} + ) result_components = {} for comp in dataset.components.values(): if comp.name in rename_map: diff --git a/src/vtlengine/Operators/General.py b/src/vtlengine/Operators/General.py index cd151db43..b25618209 100644 --- a/src/vtlengine/Operators/General.py +++ b/src/vtlengine/Operators/General.py @@ -5,7 +5,7 @@ from vtlengine.DataTypes import COMP_NAME_MAPPING from vtlengine.Exceptions import RunTimeError, SemanticError -from vtlengine.Model import Component, Dataset, ExternalRoutine, Role, Scalar +from vtlengine.Model import Component, Dataset, ExternalRoutine, Role, Scalar, names_equal from vtlengine.Operators import Binary, Unary from vtlengine.Utils.__Virtual_Assets import VirtualCounter @@ -41,7 +41,7 @@ def validate(cls, left_operand: Any, right_operand: Any) -> Union[Dataset, Scala for name, comp in left_operand.components.items() if comp.role == Role.IDENTIFIER or comp.role == Role.VIRAL_ATTRIBUTE - or (not promote_to_measure and comp.name == right_operand) + or (not promote_to_measure and names_equal(comp.name, right_operand)) } if promote_to_measure: measure_name = COMP_NAME_MAPPING[component.data_type] diff --git a/src/vtlengine/duckdb_transpiler/Transpiler/__init__.py b/src/vtlengine/duckdb_transpiler/Transpiler/__init__.py index efb9326fe..89da29be0 100644 --- a/src/vtlengine/duckdb_transpiler/Transpiler/__init__.py +++ b/src/vtlengine/duckdb_transpiler/Transpiler/__init__.py @@ -43,7 +43,15 @@ _try_normalize_time_period, ) from vtlengine.Exceptions import RunTimeError, SemanticError -from vtlengine.Model import Component, Dataset, ExternalRoutine, Role, Scalar, ValueDomain +from vtlengine.Model import ( + CaseInsensitiveDict, + Component, + Dataset, + ExternalRoutine, + Role, + Scalar, + ValueDomain, +) from vtlengine.Operators.Join import merged_viral_attribute_names from vtlengine.ViralPropagation import get_current_registry from vtlengine.ViralPropagation.sql import ( @@ -255,25 +263,35 @@ class SQLTranspiler(StructureVisitor, ASTTemplate): _consumed_join_aliases: Set[str] = field(default_factory=set, init=False) # UDO definitions - _udos: Dict[str, Dict[str, Any]] = field(default_factory=dict, init=False) + _udos: Dict[str, Dict[str, Any]] = field(default_factory=CaseInsensitiveDict, init=False) # UDO parameter stack _udo_params: Optional[List[Dict[str, Any]]] = field(default=None, init=False) # Datapoint rulesets - _dprs: Dict[str, Dict[str, Any]] = field(default_factory=dict, init=False) + _dprs: Dict[str, Dict[str, Any]] = field(default_factory=CaseInsensitiveDict, init=False) # Datapoint ruleset context _dp_signature: Optional[Dict[str, str]] = field(default=None, init=False) # Hierarchical rulesets - _hrs: Dict[str, Dict[str, Any]] = field(default_factory=dict, init=False) + _hrs: Dict[str, Dict[str, Any]] = field(default_factory=CaseInsensitiveDict, init=False) def __post_init__(self) -> None: - """Initialize available tables.""" - self.datasets = {**self.input_datasets, **self.output_datasets} - self.scalars = {**self.input_scalars, **self.output_scalars} - self.available_tables = dict(self.datasets) + """Initialize available tables. + + VTL regular names are case-insensitive: keep all name-keyed lookups in + CaseInsensitiveDict so references resolve regardless of the written casing. + """ + self.input_datasets = CaseInsensitiveDict(self.input_datasets) + self.output_datasets = CaseInsensitiveDict(self.output_datasets) + self.input_scalars = CaseInsensitiveDict(self.input_scalars) + self.output_scalars = CaseInsensitiveDict(self.output_scalars) + self.value_domains = CaseInsensitiveDict(self.value_domains) + self.external_routines = CaseInsensitiveDict(self.external_routines) + self.datasets = CaseInsensitiveDict({**self.input_datasets, **self.output_datasets}) + self.scalars = CaseInsensitiveDict({**self.input_scalars, **self.output_scalars}) + self.available_tables = CaseInsensitiveDict(self.datasets) # Helper methods @@ -1670,7 +1688,9 @@ def visit_RegularAggregation_calc(self, node: AST.RegularAggregation) -> str: resolved = self._resolve_clause_dataset(node) ds, table_src = resolved - calc_exprs: Dict[str, str] = {} + # calc_exprs is case-insensitive: a calc target that matches an existing + # component case-insensitively overrides it, and the written casing wins. + calc_exprs: CaseInsensitiveDict[str] = CaseInsensitiveDict() with self._clause_scope(ds): for child in node.children: assignment = self._unwrap_assignment(child) @@ -1690,7 +1710,9 @@ def visit_RegularAggregation_calc(self, node: AST.RegularAggregation) -> str: select_cols: List[str] = [] for name in ds.components: if name in calc_exprs: - select_cols.append(f"{calc_exprs[name]} AS {quote_name(name)}") + # Alias to the written casing (canonical key in calc_exprs). + written = calc_exprs.canonical_key(name) + select_cols.append(f"{calc_exprs[name]} AS {quote_name(written)}") else: select_cols.append(quote_name(name)) @@ -1748,7 +1770,7 @@ def visit_RegularAggregation_rename(self, node: AST.RegularAggregation) -> str: resolved = self._resolve_clause_dataset(node) ds, table_src = resolved - renames: Dict[str, str] = {} + renames: CaseInsensitiveDict[str] = CaseInsensitiveDict() for child in node.children: if isinstance(child, AST.RenameNode): old = self._resolve_membership_name(child.old_name) diff --git a/src/vtlengine/duckdb_transpiler/Transpiler/structure_visitor.py b/src/vtlengine/duckdb_transpiler/Transpiler/structure_visitor.py index c67e3c182..4e50795e3 100644 --- a/src/vtlengine/duckdb_transpiler/Transpiler/structure_visitor.py +++ b/src/vtlengine/duckdb_transpiler/Transpiler/structure_visitor.py @@ -19,7 +19,7 @@ from vtlengine.DataTypes import String as StringType from vtlengine.DataTypes.TimeHandling import TimePeriodHandler from vtlengine.duckdb_transpiler.Transpiler.sql_builder import quote_name -from vtlengine.Model import Component, Dataset, Role +from vtlengine.Model import CaseInsensitiveDict, Component, Dataset, Role from vtlengine.Operators.Join import merged_viral_attribute_names @@ -53,12 +53,13 @@ def __init__( output_datasets: Optional[Dict[str, Dataset]] = None, scalars: Optional[Dict[str, Any]] = None, ) -> None: - self.output_datasets: Dict[str, Dataset] = output_datasets or {} - self.available_tables: Dict[str, Dataset] = { - **(available_tables or {}), - **self.output_datasets, - } - self.scalars: Dict[str, Any] = scalars or {} + # VTL regular names are case-insensitive: keep name-keyed lookups in + # CaseInsensitiveDict so references resolve regardless of the written casing. + self.output_datasets: Dict[str, Dataset] = CaseInsensitiveDict(output_datasets or {}) + self.available_tables: Dict[str, Dataset] = CaseInsensitiveDict( + {**(available_tables or {}), **self.output_datasets} + ) + self.scalars: Dict[str, Any] = CaseInsensitiveDict(scalars or {}) self.current_assignment: str = "" self._in_clause: bool = False self._current_dataset: Optional[Dataset] = None diff --git a/src/vtlengine/duckdb_transpiler/io/_io.py b/src/vtlengine/duckdb_transpiler/io/_io.py index 7c51f2e4f..13c35d181 100644 --- a/src/vtlengine/duckdb_transpiler/io/_io.py +++ b/src/vtlengine/duckdb_transpiler/io/_io.py @@ -31,7 +31,7 @@ is_sdmx_datapoint_file, load_sdmx_datapoints, ) -from vtlengine.Model import Component, Dataset, Role, Scalar +from vtlengine.Model import CaseInsensitiveDict, Component, Dataset, Role, Scalar def _skip_load_validation() -> bool: @@ -468,8 +468,10 @@ def extract_datapoint_paths( if datapoints is None: return None, {} - path_dict: Dict[str, Path] = {} - df_dict: Dict[str, pd.DataFrame] = {} + # Regular names are case-insensitive: key by-name lookups must match the + # dataset name regardless of the casing the user used for the datapoint key. + path_dict: CaseInsensitiveDict[Path] = CaseInsensitiveDict() + df_dict: CaseInsensitiveDict[pd.DataFrame] = CaseInsensitiveDict() # Handle dictionary input if isinstance(datapoints, dict): diff --git a/tests/CaseInsensitive/__init__.py b/tests/CaseInsensitive/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/tests/CaseInsensitive/test_case_insensitive.py b/tests/CaseInsensitive/test_case_insensitive.py new file mode 100644 index 000000000..ff688c64a --- /dev/null +++ b/tests/CaseInsensitive/test_case_insensitive.py @@ -0,0 +1,476 @@ +"""Tests for case-insensitive regular name resolution (VTL 2.1 spec).""" + +import pandas as pd +import pytest + +from vtlengine import run, semantic_analysis +from vtlengine.Exceptions import SemanticError + +# --------------------------------------------------------------------------- +# Shared fixtures +# --------------------------------------------------------------------------- + +BASE_STRUCTURES = { + "datasets": [ + { + "name": "DS_1", + "DataStructure": [ + {"name": "Id_1", "type": "Integer", "role": "Identifier", "nullable": False}, + {"name": "Id_2", "type": "String", "role": "Identifier", "nullable": False}, + {"name": "Me_1", "type": "Number", "role": "Measure", "nullable": True}, + ], + } + ] +} + +BASE_DATAPOINTS = { + "DS_1": pd.DataFrame({"Id_1": [1, 1, 1], "Id_2": ["A", "B", "C"], "Me_1": [10.0, 20.0, 30.0]}) +} + +HR_DATAPOINTS = { + "DS_1": pd.DataFrame( + {"Id_1": [1, 1, 1, 1], "Id_2": ["A", "B", "C", "D"], "Me_1": [10.0, 20.0, 30.0, None]} + ) +} + +HR_RULE_BODY = """\ + E = A + B errorcode "e1" errorlevel 1 +end hierarchical ruleset""" + +DPR_RULE_BODY = """\ + when Id_2 = "A" then Me_1 >= 0 errorcode "err1" +end datapoint ruleset""" + + +def _run(script: str, datapoints: dict = BASE_DATAPOINTS) -> dict: + return run(script=script, data_structures=BASE_STRUCTURES, datapoints=datapoints) + + +# --------------------------------------------------------------------------- +# 1. Dataset name resolution +# --------------------------------------------------------------------------- + +dataset_name_params = [ + pytest.param("ds_1", id="all_lower"), + pytest.param("Ds_1", id="mixed_1"), + pytest.param("DS_1", id="original"), + pytest.param("dS_1", id="mixed_2"), +] + + +@pytest.mark.parametrize("alias", dataset_name_params) +def test_dataset_name_case_variants(alias): + result = _run(f"DS_r <- {alias};") + assert "DS_r" in result + assert list(result["DS_r"].data.columns) == ["Id_1", "Id_2", "Me_1"] + + +def test_dataset_preserves_original_name(): + result = _run("My_Result <- ds_1;") + assert "My_Result" in result + + +def test_dataset_chained_resolution(): + result = _run("DS_r <- DS_1; DS_r2 <- ds_r;") + assert "DS_r2" in result + pd.testing.assert_frame_equal(result["DS_r"].data, result["DS_r2"].data) + + +# --------------------------------------------------------------------------- +# 2. Duplicate assignment detection +# --------------------------------------------------------------------------- + +duplicate_params = [ + pytest.param("DS_r <- DS_1; DS_R <- DS_1;", id="different_case"), + pytest.param("DS_r <- DS_1; DS_r <- DS_1;", id="same_case"), + pytest.param("DS_r <- DS_1; ds_r <- DS_1;", id="all_lower"), +] + + +@pytest.mark.parametrize("script", duplicate_params) +def test_duplicate_assignment_raises(script): + with pytest.raises(SemanticError, match="1-2-2"): + _run(script) + + +# --------------------------------------------------------------------------- +# 3. Component name resolution (calc, filter, rename) +# --------------------------------------------------------------------------- + +component_calc_params = [ + pytest.param( + "DS_r <- ds_1[calc me_2 := me_1 * 2];", + ["me_2"], + id="calc_lowercase", + ), + pytest.param( + "DS_r <- ds_1[calc me_2 := Me_1, mE_3 := ME_1 + me_1];", + ["me_2", "mE_3"], + id="calc_mixed_case", + ), +] + + +@pytest.mark.parametrize("script, expected_comps", component_calc_params) +def test_calc_case_insensitive(script, expected_comps): + result = _run(script) + for comp in expected_comps: + assert comp in result["DS_r"].components + + +calc_override_params = [ + pytest.param( + "DS_r <- DS_1[calc ME_1 := Me_1 * 2];", + ["Id_1", "Id_2", "ME_1"], + id="calc_override_upper", + ), + pytest.param( + "DS_r <- DS_1[calc me_1 := Me_1 * 2];", + ["Id_1", "Id_2", "me_1"], + id="calc_override_lower", + ), +] + + +@pytest.mark.parametrize("script, expected_cols", calc_override_params) +def test_calc_output_columns_case_insensitive(script, expected_cols): + """Output DataFrame columns must match component names (no duplicates).""" + result = _run(script) + ds = result["DS_r"] + assert list(ds.components.keys()) == expected_cols + assert list(ds.data.columns) == expected_cols + + +filter_params = [ + pytest.param("DS_r <- ds_1[filter me_1 > 15];", 2, id="lowercase_measure"), + pytest.param("DS_r <- ds_1[filter ME_1 > 15];", 2, id="uppercase_measure"), + pytest.param("DS_r <- ds_1[filter Me_1 > 25];", 1, id="original_case"), +] + + +@pytest.mark.parametrize("script, expected_rows", filter_params) +def test_filter_case_insensitive(script, expected_rows): + result = _run(script) + assert len(result["DS_r"].data) == expected_rows + + +rename_params = [ + pytest.param("me_1", "Me_New", id="lowercase_old"), + pytest.param("ME_1", "Me_New", id="uppercase_old"), + pytest.param("Me_1", "Me_Renamed", id="original_case_old"), +] + + +@pytest.mark.parametrize("old_name, new_name", rename_params) +def test_rename_case_insensitive(old_name, new_name): + result = _run(f"DS_r <- ds_1[rename {old_name} to {new_name}];") + assert new_name in result["DS_r"].components + assert "Me_1" not in result["DS_r"].components + + +# --------------------------------------------------------------------------- +# 4. Hierarchical ruleset name resolution +# --------------------------------------------------------------------------- + +hr_params = [ + pytest.param("hr1", "HR1", "Id_2", "Id_2", id="name_upper"), + pytest.param("hr1", "hr1", "Id_2", "id_2", id="comp_lower"), + pytest.param("My_HR", "MY_HR", "Id_2", "ID_2", id="both_different"), + pytest.param("hr1", "Hr1", "id_2", "ID_2", id="all_mixed"), +] + + +@pytest.mark.parametrize("def_name, call_name, def_comp, call_comp", hr_params) +def test_hr_case_insensitive(def_name, call_name, def_comp, call_comp): + script = f""" + define hierarchical ruleset {def_name} (variable rule {def_comp}) is + {HR_RULE_BODY}; + DS_r <- hierarchy(DS_1, {call_name} rule {call_comp} computed); + """ + result = _run(script, datapoints=HR_DATAPOINTS) + assert "DS_r" in result + assert "Id_2" in result["DS_r"].components + + +# --------------------------------------------------------------------------- +# 5. Datapoint ruleset name resolution +# --------------------------------------------------------------------------- + +dpr_params = [ + pytest.param("dpr1", "DPR1", "Id_2, Me_1", "Id_2, Me_1", id="name_upper"), + pytest.param("dpr1", "dpr1", "ID_2, ME_1", "id_2, me_1", id="comps_swapped"), + pytest.param("My_DPR", "MY_DPR", "ID_2, ME_1", "id_2, me_1", id="both_different"), +] + + +@pytest.mark.parametrize("def_name, call_name, def_comps, call_comps", dpr_params) +def test_dpr_case_insensitive(def_name, call_name, def_comps, call_comps): + script = f""" + define datapoint ruleset {def_name} (variable {def_comps}) is + {DPR_RULE_BODY}; + DS_r := check_datapoint(DS_1, {call_name} components {call_comps} invalid); + """ + result = _run(script) + assert result is not None + + +# --------------------------------------------------------------------------- +# 6. UDO name resolution +# --------------------------------------------------------------------------- + +udo_params = [ + pytest.param( + """ + define operator my_op (ds dataset) returns dataset is ds end operator; + DS_r <- MY_OP(DS_1); + """, + "DS_r", + id="simple_upper", + ), + pytest.param( + """ + define operator my_op (ds dataset) returns dataset is ds end operator; + DS_r <- My_Op(ds_1); + """, + "DS_r", + id="simple_mixed", + ), + pytest.param( + """ + define operator suma (ds1 dataset, ds2 dataset) returns dataset is ds1 + ds2 end operator; + define operator drop_id (ds dataset, comp component) + returns dataset is max(ds group except comp) end operator; + DS_r <- DROP_ID(SUMA(ds_1, Ds_1), Id_2); + """, + "DS_r", + id="nested_mixed", + ), +] + + +@pytest.mark.parametrize("script, expected_ds", udo_params) +def test_udo_case_insensitive(script, expected_ds): + result = _run(script) + assert expected_ds in result + + +def test_udo_duplicate_definition_different_case(): + script = """ + define operator my_op (ds dataset) returns dataset is ds end operator; + define operator MY_OP (ds dataset) returns dataset is ds end operator; + DS_r <- my_op(DS_1); + """ + with pytest.raises((ValueError, SemanticError)): + _run(script) + + +# --------------------------------------------------------------------------- +# 7. Aggregation with case-insensitive component refs +# --------------------------------------------------------------------------- + +agg_params = [ + pytest.param("sum(ds_1 group by id_1)", ["Id_1", "Me_1"], id="group_by_lower"), + pytest.param("sum(DS_1 group by Id_1)", ["Id_1", "Me_1"], id="group_by_original"), + pytest.param("max(ds_1 group except id_2)", ["Id_1", "Me_1"], id="group_except_lower"), + pytest.param("max(ds_1 group except ID_2)", ["Id_1", "Me_1"], id="group_except_upper"), + pytest.param("sum(ds_1 group by ID_1)", ["Id_1", "Me_1"], id="group_by_upper"), +] + + +@pytest.mark.parametrize("expr, expected_comps", agg_params) +def test_aggregation_case_insensitive(expr, expected_comps): + result = _run(f"DS_r <- {expr};") + assert "DS_r" in result + for comp in expected_comps: + assert comp in result["DS_r"].components + + +aggr_override_params = [ + pytest.param( + "DS_r <- DS_1[aggr ME_1 := sum(Me_1) group by Id_1];", + ["Id_1", "ME_1"], + id="aggr_override_upper", + ), + pytest.param( + "DS_r <- DS_1[aggr me_1 := sum(Me_1) group by Id_1];", + ["Id_1", "me_1"], + id="aggr_override_lower", + ), +] + + +@pytest.mark.parametrize("script, expected_cols", aggr_override_params) +def test_aggr_output_columns_case_insensitive(script, expected_cols): + """Output DataFrame columns must match component names (no duplicates).""" + result = _run(script) + ds = result["DS_r"] + assert list(ds.components.keys()) == expected_cols + assert list(ds.data.columns) == expected_cols + + +# --------------------------------------------------------------------------- +# 8. Scalar name resolution +# --------------------------------------------------------------------------- + +scalar_params = [ + pytest.param("my_sc", "my_sc", 43, id="lower"), + pytest.param("My_Sc", "my_sc", 43, id="mixed"), + pytest.param("MY_SC", "my_sc", 43, id="upper"), +] + + +@pytest.mark.parametrize("def_name, ref_name, expected_value", scalar_params) +def test_scalar_case_insensitive(def_name, ref_name, expected_value): + script = f"{def_name} <- 42; DS_r <- {ref_name} + 1;" + result = run(script=script, data_structures={"datasets": []}, datapoints={}) + assert "DS_r" in result + assert result["DS_r"].value == expected_value + + +# --------------------------------------------------------------------------- +# 9. End-to-end mixed operations +# --------------------------------------------------------------------------- + +e2e_params = [ + pytest.param( + "DS_r <- DS_1; DS_r2 <- ds_r;", + ["DS_r", "DS_r2"], + id="chained_datasets", + ), + pytest.param( + "DS_r <- ds_1[calc me_2 := Me_1, mE_3 := ME_1 + me_1];", + ["DS_r"], + id="calc_mixed_refs", + ), + pytest.param( + "DS_r <- DS_1; DS_r2 <- ds_r; DS_r3 <- ds_1[calc me_2 := Me_1];", + ["DS_r", "DS_r2", "DS_r3"], + id="full_pipeline", + ), +] + + +@pytest.mark.parametrize("script, expected_datasets", e2e_params) +def test_end_to_end(script, expected_datasets): + result = _run(script) + for ds in expected_datasets: + assert ds in result + + +# --------------------------------------------------------------------------- +# 10. Input loading: datapoint dict key matched case-insensitively +# --------------------------------------------------------------------------- + + +def test_datapoints_dict_key_case_insensitive(): + # A datapoints dict keyed 'ds_1' must bind to structure 'DS_1'. + result = _run("DS_r <- DS_1;", datapoints={"ds_1": BASE_DATAPOINTS["DS_1"]}) + assert len(result["DS_r"].data) == 3 + + +def test_calc_mixed_case_ref_computes_values(): + # A mixed-case measure reference must actually feed the computation. + result = _run("DS_r <- ds_1[calc Me_2 := ME_1 * 2];") + assert sorted(result["DS_r"].data["Me_2"].tolist()) == [20.0, 40.0, 60.0] + + +# --------------------------------------------------------------------------- +# 11. semantic_analysis (structure-only public API) +# --------------------------------------------------------------------------- + +semantic_params = [ + pytest.param("DS_r <- ds_1;", id="dataset_alias"), + pytest.param("DS_r <- DS_1[filter ME_1 > 15];", id="component_ref"), + pytest.param("DS_r <- DS_1; DS_r2 <- ds_r;", id="chained"), +] + + +@pytest.mark.parametrize("script", semantic_params) +def test_semantic_analysis_case_insensitive(script): + result = semantic_analysis(script=script, data_structures=BASE_STRUCTURES) + assert any(name.startswith("DS_r") for name in result) + + +def test_returned_components_are_plain_dict_for_external_consumers(): + # The internal CaseInsensitiveDict must not leak to callers. + result = semantic_analysis(script="DS_r <- ds_1;", data_structures=BASE_STRUCTURES) + assert type(result["DS_r"].components) is dict + + +# --------------------------------------------------------------------------- +# 12. Rename duplicate detection (case-insensitive) +# --------------------------------------------------------------------------- + + +def test_rename_duplicate_source_different_case_raises(): + # Me_1 and ME_1 are the same component; renaming it twice must be rejected. + with pytest.raises(SemanticError, match="1-1-6-9"): + _run("DS_r <- DS_1[rename Me_1 to A, ME_1 to B];") + + +def test_rename_duplicate_target_different_case_raises(): + with pytest.raises(SemanticError, match="1-2-1"): + _run("DS_r <- DS_1[rename Id_2 to X, Me_1 to x];") + + +# --------------------------------------------------------------------------- +# 13. Operand case-insensitivity in a normal VTL run (data is computed) +# --------------------------------------------------------------------------- + + +def test_binary_op_dataset_operands_mixed_case(): + # Both operands of '+' reference DS_1 with different casing. + result = _run("DS_r <- ds_1 + DS_1;") + assert result["DS_r"].data["Me_1"].tolist() == [20.0, 40.0, 60.0] + + +def test_membership_component_operand_mixed_case(): + # The membership operand references Me_1 as ME_1. + result = _run("DS_r <- ds_1#ME_1;") + assert "Me_1" in result["DS_r"].components + assert result["DS_r"].data["Me_1"].tolist() == [10.0, 20.0, 30.0] + + +def test_comparison_operand_mixed_case(): + result = _run("DS_r <- ds_1[calc gt := ME_1 > 15];") + assert result["DS_r"].data["gt"].tolist() == [False, True, True] + + +def test_if_then_else_operands_mixed_case(): + result = _run("DS_r <- ds_1[calc m2 := if ME_1 > 15 then me_1 else 0];") + assert result["DS_r"].data["m2"].tolist() == [0.0, 20.0, 30.0] + + +_JOIN_STRUCT = { + "datasets": [ + { + "name": "DS_1", + "DataStructure": [ + {"name": "Id_1", "type": "Integer", "role": "Identifier", "nullable": False}, + {"name": "Me_1", "type": "Number", "role": "Measure", "nullable": True}, + ], + }, + { + "name": "DS_2", + "DataStructure": [ + {"name": "Id_1", "type": "Integer", "role": "Identifier", "nullable": False}, + {"name": "Me_2", "type": "Number", "role": "Measure", "nullable": True}, + ], + }, + ] +} +_JOIN_DATAPOINTS = { + "DS_1": pd.DataFrame({"Id_1": [1, 2, 3], "Me_1": [10.0, 20.0, 30.0]}), + "DS_2": pd.DataFrame({"Id_1": [1, 2, 3], "Me_2": [1.0, 2.0, 3.0]}), +} + + +def test_join_body_component_operands_mixed_case(): + # Component operands inside a join body must resolve case-insensitively + # against the join's (virtual) dataset. + result = run( + script="DS_r <- inner_join(ds_1 as a, DS_2 as b calc Me_3 := ME_1 + me_2);", + data_structures=_JOIN_STRUCT, + datapoints=_JOIN_DATAPOINTS, + ) + assert result["DS_r"].data["Me_3"].tolist() == [11.0, 22.0, 33.0] diff --git a/tests/DAG/data/references/2.json b/tests/DAG/data/references/2.json index 44dfb5825..7999a8470 100644 --- a/tests/DAG/data/references/2.json +++ b/tests/DAG/data/references/2.json @@ -1,7 +1,7 @@ { "insertion": { "2": [ - "A" + "Z" ] }, "deletion": { @@ -19,11 +19,11 @@ "mdl2.f" ], "2": [ - "A" + "Z" ] }, "global_inputs": [ - "A" + "Z" ], "persistent": [ "a", @@ -31,4 +31,4 @@ "d", "g" ] -} \ No newline at end of file +} diff --git a/tests/DAG/data/vtl/2.vtl b/tests/DAG/data/vtl/2.vtl index e4f8f5d3b..ff843f6ef 100644 --- a/tests/DAG/data/vtl/2.vtl +++ b/tests/DAG/data/vtl/2.vtl @@ -1,7 +1,7 @@ mdl2.e := 2; a <- - A; + Z; b <- 1; mdl1.c := diff --git a/tests/Model/test_case_insensitive_dict.py b/tests/Model/test_case_insensitive_dict.py new file mode 100644 index 000000000..6853a0482 --- /dev/null +++ b/tests/Model/test_case_insensitive_dict.py @@ -0,0 +1,75 @@ +import copy +import pickle + +import pytest + +import vtlengine.DataTypes as DataTypes +from vtlengine.Model import CaseInsensitiveDict, Component, Dataset, Role + + +def test_case_insensitive_lookup(): + d: CaseInsensitiveDict[int] = CaseInsensitiveDict() + d["Me_1"] = 10 + assert d["me_1"] == 10 + assert d["ME_1"] == 10 + assert "mE_1" in d + + +def test_preserves_first_seen_key(): + d: CaseInsensitiveDict[int] = CaseInsensitiveDict() + d["Me_1"] = 1 + d["ME_1"] = 2 # updates value, keeps original display key + assert list(d.keys()) == ["Me_1"] + assert d["me_1"] == 2 + + +def test_canonical_key(): + d: CaseInsensitiveDict[object] = CaseInsensitiveDict({"DS_1": object()}) + assert d.canonical_key("ds_1") == "DS_1" + with pytest.raises(KeyError): + d.canonical_key("missing") + + +def test_get_pop_setdefault(): + d: CaseInsensitiveDict[int] = CaseInsensitiveDict({"A": 1}) + assert d.get("a") == 1 + assert d.get("z", 99) == 99 + assert d.setdefault("a", 5) == 1 + popped = d.pop("A") + assert popped == 1 + assert "a" not in d + + +def test_to_dict_is_plain_with_original_keys(): + d: CaseInsensitiveDict[int] = CaseInsensitiveDict({"Me_1": 1}) + d["me_1"] = 2 + plain = d.to_dict() + assert type(plain) is dict + assert plain == {"Me_1": 2} + + +def test_deepcopy_and_pickle_roundtrip(): + d: CaseInsensitiveDict[list[int]] = CaseInsensitiveDict({"Me_1": [1, 2]}) + dc = copy.deepcopy(d) + dc["ME_1"].append(3) + assert d["me_1"] == [1, 2] # independent + assert dc["me_1"] == [1, 2, 3] + rt = pickle.loads(pickle.dumps(d)) # noqa: S301 + assert rt["me_1"] == [1, 2] + assert isinstance(rt, CaseInsensitiveDict) + + +def test_dataset_components_is_case_insensitive_dict(): + comps = { + "Id_1": Component( + name="Id_1", data_type=DataTypes.Integer, role=Role.IDENTIFIER, nullable=False + ), + "Me_1": Component( + name="Me_1", data_type=DataTypes.Number, role=Role.MEASURE, nullable=True + ), + } + ds = Dataset(name="DS_1", components=comps, data=None) + assert isinstance(ds.components, CaseInsensitiveDict) + assert "me_1" in ds.components + assert ds.components["ME_1"].name == "Me_1" + assert ds.get_component("me_1").name == "Me_1"