Skip to content

Commit 01cfa16

Browse files
committed
fix: address some review comments
1 parent 4d1c571 commit 01cfa16

4 files changed

Lines changed: 113 additions & 95 deletions

File tree

src/dve/core_engine/backends/metadata/contract.py

Lines changed: 2 additions & 51 deletions
Original file line numberDiff line numberDiff line change
@@ -1,61 +1,13 @@
11
"""Metadata classes for the data contract."""
22

3-
from typing import Any, Optional, Union
3+
from typing import Any
44

5-
from pydantic import BaseModel, Field, PrivateAttr, model_validator
5+
from pydantic import BaseModel, PrivateAttr, model_validator
66

77
from dve.core_engine.type_hints import EntityName, ReportingFields
88
from dve.core_engine.validation import RowValidator
9-
from dve.metadata_parser.exc import EntityNotFoundError
109
from dve.parser.type_hints import Extension
1110

12-
class HierarchyNode(BaseModel):
13-
entity_name: str
14-
children: Optional[list["HierarchyNode"]] = Field(default_factory=list)
15-
16-
def get_descendents(self) -> list[str]:
17-
"""Recursively list all descendents of the node"""
18-
descendents = []
19-
for node in self.children:
20-
descendents.append(node.entity_name)
21-
descendents.extend(node.get_descendents())
22-
return descendents
23-
24-
def get_node(self, entity_name:str) -> Union["HierarchyNode", None]:
25-
"""Recursively search for node and return if found"""
26-
node = None
27-
if self.entity_name == entity_name:
28-
return self
29-
else:
30-
for child in self.children:
31-
node = child.get_node(entity_name)
32-
if node:
33-
break
34-
return node
35-
36-
def add_child_node(self, parent_entity: str, child_info: "HierarchyNode") -> None:
37-
"""Add a child node if the parent exists in the hierarchy"""
38-
try:
39-
self.get_node(parent_entity).children.append(child_info)
40-
except AttributeError:
41-
raise EntityNotFoundError(f"Can't find parent node {parent_entity} in {self.entity_name}")
42-
43-
def as_dict(self):
44-
ret_dict = {}
45-
for node in self.children:
46-
ret_dict.update(node.as_dict())
47-
return {self.entity_name: {"children": ret_dict}}
48-
49-
50-
class ChildHierarchyNode(HierarchyNode):
51-
join_fields: list[str]
52-
53-
def as_dict(self):
54-
ret_value = {self.entity_name: {"join_fields": self.join_fields}}
55-
for node in self.children:
56-
ret_value[self.entity_name] |= {"children": node.as_dict()}
57-
return ret_value
58-
5911

6012
class ReaderConfig(BaseModel):
6113
"""Configuration options for a given reader."""
@@ -86,7 +38,6 @@ class DataContractMetadata(BaseModel, frozen=True, arbitrary_types_allowed=True)
8638
"""Whether to cache the original entities after loading."""
8739
_schemas: dict[EntityName, type[BaseModel]] = PrivateAttr(default_factory=dict)
8840
"""The pydantic models of the schmas."""
89-
linkage_hierarchy: dict[EntityName, HierarchyNode] = Field(default_factor=dict)
9041

9142
@property
9243
def schemas(self) -> dict[EntityName, type[BaseModel]]:

src/dve/core_engine/configuration/v1/__init__.py

Lines changed: 17 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -7,7 +7,7 @@
77
from typing_extensions import Literal
88

99
from dve.core_engine.backends.base.reference_data import ReferenceConfig, ReferenceConfigUnion
10-
from dve.core_engine.backends.metadata.contract import ChildHierarchyNode, DataContractMetadata, HierarchyNode, ReaderConfig
10+
from dve.core_engine.backends.metadata.contract import DataContractMetadata, ReaderConfig
1111
from dve.core_engine.backends.metadata.rules import AbstractStep, Rule, RuleMetadata
1212
from dve.core_engine.configuration.base import BaseEngineConfig
1313
from dve.core_engine.configuration.v1.filters import (
@@ -20,6 +20,7 @@
2020
BusinessFilterSpecConfig,
2121
BusinessRuleSpecConfig,
2222
)
23+
from dve.core_engine.configuration.v1.hierarchy import HierarchyNode, ChildHierarchyNode
2324
from dve.core_engine.configuration.v1.steps import StepConfigUnion
2425
from dve.core_engine.message import DataContractErrorDetail
2526
from dve.core_engine.type_hints import EntityName, ErrorCategory, ErrorType, TemplateVariables
@@ -90,6 +91,8 @@ class _LinkageConfig(BaseModel):
9091
"""The name of the parent entity"""
9192
join_fields: JoinFields
9293
"""The fields that can be used to link back to the parent entity"""
94+
mandatory: Optional[bool] = False
95+
"""If the entity is a child, is it a mandatory field of the parent"""
9396

9497

9598
class _SchemaConfig(BaseModel):
@@ -123,7 +126,6 @@ class _ModelConfig(_SchemaConfig):
123126
"""Reader configuration options for the model."""
124127
aliases: dict[FieldName, FieldName] = Field(default_factory=dict)
125128
"""An alias field name mapping."""
126-
linkage_details: Optional[_LinkageConfig] = None
127129

128130

129131
class _RuleStoreConfig(BaseModel):
@@ -189,6 +191,8 @@ class V1EngineConfig(BaseEngineConfig):
189191
default_factory=dict
190192
)
191193
"""Rule store rules from the loaded rule stores."""
194+
entity_relationships: dict[EntityName, _LinkageConfig] = Field(default_factory=dict)
195+
"""The parent-child relationships linking the defined entities"""
192196

193197
@validate_call
194198
def _update_rule_store(self, rule_store: dict[RuleName, BusinessComponentSpecConfigUnion]):
@@ -341,8 +345,7 @@ def get_contract_metadata(self) -> DataContractMetadata:
341345
reader_metadata=reader_metadata,
342346
validators=validators,
343347
reporting_fields=reporting_fields,
344-
cache_originals=self.contract.cache_originals,
345-
linkage_hierarchy=self.determine_entity_hierarchy()
348+
cache_originals=self.contract.cache_originals
346349
)
347350

348351
def load_error_message_info(self, uri):
@@ -365,20 +368,22 @@ def get_rule_metadata(self) -> RuleMetadata:
365368
reference_data_config=self.get_reference_data_config(),
366369
)
367370

368-
def determine_entity_hierarchy(self) -> list[HierarchyNode]:
371+
def get_entity_hierarchy(self) -> dict[str, HierarchyNode]:
369372
"""Determine the linkage hierarchy using contact config"""
370-
linkage_hierarchy = {name: model_conf.linkage_details for name, model_conf in self.contract.datasets.items()}
371-
top_level_parents = {}
372-
for name, linkage_detail in linkage_hierarchy.items():
373-
if not linkage_detail:
374-
top_level_parents[name] = HierarchyNode(entity_name=name)
375-
continue
373+
top_level_parents = {
374+
entity_name: HierarchyNode(entity_name=entity_name)
375+
for entity_name in self.contract.datasets
376+
if not entity_name in self.entity_relationships
377+
}
378+
379+
for name, linkage_detail in self.entity_relationships.items():
376380
for main_entity, details in top_level_parents.items():
377381
if (linkage_detail.parent_entity == main_entity
378382
or linkage_detail.parent_entity in details.get_descendents()):
379383
top_level_parents[main_entity].add_child_node(linkage_detail.parent_entity,
380384
ChildHierarchyNode(entity_name=name,
381-
join_fields=linkage_detail.join_fields))
385+
join_fields=linkage_detail.join_fields,
386+
mandatory=linkage_detail.mandatory))
382387
break
383388
else:
384389
raise EntityNotFoundError(f"Can't find parent entity {linkage_detail.parent_entity} defined to establish hierarchy for {name} - please ensure it is defined above any child entities in the dischema.")
Lines changed: 56 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,56 @@
1+
"""Classes to help determine and store entity hierarchy information."""
2+
3+
from typing import Any, Optional, Union
4+
from pydantic import BaseModel, Field
5+
from dve.metadata_parser.exc import EntityNotFoundError
6+
7+
class HierarchyNode(BaseModel):
8+
"""Stores entity hierarchy information"""
9+
entity_name: str
10+
mandatory: Optional[bool] = False
11+
children: Optional[list["HierarchyNode"]] = Field(default_factory=list)
12+
13+
def get_descendents(self) -> list[str]:
14+
"""Recursively list all descendents of the node"""
15+
descendents = []
16+
for node in self.children:
17+
descendents.append(node.entity_name)
18+
descendents.extend(node.get_descendents())
19+
return descendents
20+
21+
def get_node(self, entity_name:str) -> Union["HierarchyNode", None]:
22+
"""Recursively search for node and return if found"""
23+
node = None
24+
if self.entity_name == entity_name:
25+
return self
26+
else:
27+
for child in self.children:
28+
node = child.get_node(entity_name)
29+
if node:
30+
break
31+
return node
32+
33+
def add_child_node(self, parent_entity: str, child_info: "HierarchyNode") -> None:
34+
"""Add a child node if the parent exists in the hierarchy"""
35+
try:
36+
self.get_node(parent_entity).children.append(child_info)
37+
except AttributeError:
38+
raise EntityNotFoundError(f"Can't find parent node {parent_entity} in {self.entity_name}")
39+
40+
def as_dict(self) -> dict[str, dict[str, Any]]:
41+
"""Get dictionary representation of entity hierarchy"""
42+
child_dict = {}
43+
for node in self.children:
44+
child_dict.update(node.as_dict())
45+
46+
ret_dict = {"children": child_dict,
47+
"mandatory": self.mandatory}
48+
if hasattr(self, "join_fields"):
49+
ret_dict |= {"join_fields": self.join_fields}
50+
51+
return {self.entity_name: ret_dict}
52+
53+
54+
class ChildHierarchyNode(HierarchyNode):
55+
"""Stores child entity hierarchy information"""
56+
join_fields: list[str]

tests/test_core_engine/test_config_load.py

Lines changed: 38 additions & 32 deletions
Original file line numberDiff line numberDiff line change
@@ -123,11 +123,7 @@
123123
"mandatory_fields": [
124124
"ds_003_id",
125125
"ds_001_id"
126-
],
127-
"linkage_details": {
128-
"parent_entity": "ds_001",
129-
"join_fields": ["ds_001_id"]
130-
}
126+
]
131127
},
132128
"ds_101": {
133129
"fields": {
@@ -147,11 +143,7 @@
147143
"mandatory_fields": [
148144
"referral_id",
149145
"ds_001_id"
150-
],
151-
"linkage_details": {
152-
"parent_entity": "ds_001",
153-
"join_fields": ["ds_001_id"]
154-
}
146+
]
155147
},
156148
"ds_201": {
157149
"fields": {
@@ -172,11 +164,7 @@
172164
"referral_id",
173165
"ds_201_id",
174166
"contact_date"
175-
],
176-
"linkage_details": {
177-
"parent_entity": "ds_101",
178-
"join_fields": ["referral_id"]
179-
}
167+
]
180168
},
181169
"ds_202": {
182170
"fields": {
@@ -196,11 +184,7 @@
196184
"mandatory_fields": [
197185
"ds_202_id",
198186
"ds_201_id"
199-
],
200-
"linkage_details": {
201-
"parent_entity": "ds_201",
202-
"join_fields": ["ds_201_id"]
203-
}
187+
]
204188
}
205189
}
206190
},
@@ -214,38 +198,60 @@
214198
"failure_message": "Record rejected - `{{ name }}` is not valid."
215199
}
216200
]
217-
}
201+
},
202+
"entity_relationships": {
203+
"ds_003": {
204+
"parent_entity": "ds_001",
205+
"join_fields": ["ds_001_id"],
206+
"mandatory": false
207+
},
208+
"ds_101": {
209+
"parent_entity": "ds_001",
210+
"join_fields": ["ds_001_id"],
211+
"mandatory_entity": true
212+
},
213+
"ds_201": {
214+
"parent_entity": "ds_101",
215+
"join_fields": ["referral_id"],
216+
"mandatory": false
217+
},
218+
"ds_202": {
219+
"parent_entity": "ds_201",
220+
"join_fields": ["ds_201_id"],
221+
"mandatory": true
222+
}
223+
}
218224
}"""
219225

220226
def test_no_linkage_config_load():
221227
config = V1EngineConfig(location="",
222228
**json.loads(CONFIG_WITHOUT_LINKAGE))
223229
assert len(config.contract.datasets) == 1
224-
dc_metadata = config.get_contract_metadata()
225-
assert len(dc_metadata.linkage_hierarchy) == 1
226-
assert not dc_metadata.linkage_hierarchy.get("animals").children
230+
hierarchy = config.get_entity_hierarchy()
231+
assert len(hierarchy) == 1
232+
assert not hierarchy.get("animals").children
227233

228234

229235
def test_linkage_config_load():
230236
config = V1EngineConfig(location="",
231237
**json.loads(CONFIG_WITH_LINKAGE))
232238
assert len(config.contract.datasets) == 6
233-
dc_metadata = config.get_contract_metadata()
234-
assert len(dc_metadata.linkage_hierarchy) == 2
235-
assert not dc_metadata.linkage_hierarchy.get("ds_002").children
236-
assert len(dc_metadata.linkage_hierarchy.get("ds_001").get_descendents()) == 4
237-
children_001 = sorted(dc_metadata.linkage_hierarchy.get("ds_001").children, key=lambda x: x.entity_name)
238-
dict_rep_001 = dc_metadata.linkage_hierarchy.get("ds_001").as_dict()
239+
hierarchy = config.get_entity_hierarchy()
240+
assert len(hierarchy) == 2
241+
assert not hierarchy.get("ds_002").children
242+
assert len(hierarchy.get("ds_001").get_descendents()) == 4
243+
children_001 = sorted(hierarchy.get("ds_001").children, key=lambda x: x.entity_name)
244+
dict_rep_001 = hierarchy.get("ds_001").as_dict()
239245
assert len(children_001) == 2
240246
assert children_001[0].entity_name == "ds_003"
241247
assert not children_001[0].children
242248
assert children_001[1].entity_name == "ds_101"
243-
assert dict_rep_001 == {'ds_001': {'children': {'ds_003': {'join_fields': ['ds_001_id']}, 'ds_101': {'join_fields': ['ds_001_id'], 'children': {'ds_201': {'join_fields': ['referral_id'], 'children': {'ds_202': {'join_fields': ['ds_201_id']}}}}}}}}
249+
assert dict_rep_001 == {'ds_001': {'children': {'ds_003': {'children': {}, 'mandatory': False, 'join_fields': ['ds_001_id']}, 'ds_101': {'children': {'ds_201': {'children': {'ds_202': {'children': {}, 'mandatory': True, 'join_fields': ['ds_201_id']}}, 'mandatory': False, 'join_fields': ['referral_id']}}, 'mandatory': False, 'join_fields': ['ds_001_id']}}, 'mandatory': False}}
244250
children_101 = children_001[1].children
245251
dict_rep_101 = dict_rep_001["ds_001"]["children"]["ds_101"]
246252
assert len(children_101) == 1
247253
assert children_101[0].entity_name == "ds_201"
248254
assert children_101[0].children[0].entity_name == "ds_202"
249255
assert not children_101[0].children[0].children
250-
assert dict_rep_101 == {'join_fields': ['ds_001_id'], 'children': {'ds_201': {'join_fields': ['referral_id'], 'children': {'ds_202': {'join_fields': ['ds_201_id']}}}}}
256+
assert dict_rep_101 == {'children': {'ds_201': {'children': {'ds_202': {'children': {}, 'mandatory': True, 'join_fields': ['ds_201_id']}}, 'mandatory': False, 'join_fields': ['referral_id']}}, 'mandatory': False, 'join_fields': ['ds_001_id']}
251257

0 commit comments

Comments
 (0)