Skip to content
Draft
23 changes: 12 additions & 11 deletions src/vtlengine/API/_InternalApi.py
Original file line number Diff line number Diff line change
Expand Up @@ -36,16 +36,17 @@
to_vtl_json,
)
from vtlengine.Model import (
Component as VTL_Component,
)
from vtlengine.Model import (
CaseInsensitiveDict,
Dataset,
ExternalRoutine,
Role,
Role_keys,
Scalar,
ValueDomain,
)
from vtlengine.Model import (
Component as VTL_Component,
)

# Cache SCALAR_TYPES keys for performance
_SCALAR_TYPE_KEYS = SCALAR_TYPES.keys()
Expand Down Expand Up @@ -87,13 +88,13 @@ def _load_dataset_from_structure(
"""
Loads a dataset with the structure given.
"""
datasets = {}
scalars = {}
datasets: CaseInsensitiveDict[Dataset] = CaseInsensitiveDict()
scalars: CaseInsensitiveDict[Scalar] = CaseInsensitiveDict()

if "datasets" in structures:
for dataset_json in structures["datasets"]:
dataset_name = dataset_json["name"]
components = {}
components: CaseInsensitiveDict[VTL_Component] = CaseInsensitiveDict()

if "structure" in dataset_json:
structure_name = dataset_json["structure"]
Expand Down Expand Up @@ -365,15 +366,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
Expand Down
11 changes: 11 additions & 0 deletions src/vtlengine/API/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -247,6 +254,7 @@ def semantic_analysis(
scalars=scalars,
)
result = interpreter.visit(ast)
_convert_to_regular_dicts(result)
return result


Expand Down Expand Up @@ -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


Expand Down
43 changes: 23 additions & 20 deletions src/vtlengine/AST/DAG/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -48,7 +48,7 @@
TO,
)
from vtlengine.Exceptions import SemanticError
from vtlengine.Model import Component
from vtlengine.Model import Component, normalize_name


@dataclass
Expand Down Expand Up @@ -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():
Expand All @@ -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(
Expand Down Expand Up @@ -183,24 +187,25 @@ 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

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
Expand Down Expand Up @@ -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:
Expand All @@ -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
):
Expand All @@ -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)
Expand Down
40 changes: 28 additions & 12 deletions src/vtlengine/Interpreter/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -46,6 +46,7 @@
)
from vtlengine.Exceptions import SemanticError
from vtlengine.Model import (
CaseInsensitiveDict,
Component,
DataComponent,
Dataset,
Expand All @@ -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
Expand Down Expand Up @@ -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()

# **********************************
# * *
Expand All @@ -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
Expand All @@ -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))
Expand All @@ -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")

Expand Down Expand Up @@ -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:
Expand All @@ -255,15 +267,15 @@ 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")

self.dprs[node.name] = ruleset_data

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")
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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 = {}

Expand Down Expand Up @@ -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,
Expand Down
Loading
Loading