From e2afa211191ce204cac5530d1198bf4d28d32c25 Mon Sep 17 00:00:00 2001 From: seanmcculloch Date: Thu, 27 Aug 2026 14:33:28 -0700 Subject: [PATCH] fix: de-duplicate identical pipelines when combining Processing objects --- src/aind_data_schema/core/processing.py | 8 +++++-- src/aind_data_schema/utils/merge.py | 12 +++++++--- tests/test_composability_merge.py | 30 +++++++++++++++++++++++++ tests/test_utils_merge.py | 8 +++++++ 4 files changed, 53 insertions(+), 5 deletions(-) diff --git a/src/aind_data_schema/core/processing.py b/src/aind_data_schema/core/processing.py index c77d2a110..c51a94358 100644 --- a/src/aind_data_schema/core/processing.py +++ b/src/aind_data_schema/core/processing.py @@ -12,7 +12,7 @@ from aind_data_schema.base import AwareDatetimeWithDefault, DataCoreModel, DataModel, GenericModel from aind_data_schema.components.identifiers import Code, DataAsset # noqa: F401 from aind_data_schema.components.wrappers import AssetPath -from aind_data_schema.utils.merge import merge_notes, merge_optional_list, merge_process_graph +from aind_data_schema.utils.merge import merge_notes, merge_optional_list, merge_process_graph, remove_duplicates from aind_data_schema.utils.validators import TimeValidation @@ -260,8 +260,12 @@ def __add__(self, other: "Processing") -> "Processing": if merged_graph and len(self.data_processes) > 0 and len(other.data_processes) > 0: merged_graph[other.data_processes[0].name] = [self.data_processes[-1].name] + merged_pipelines = merge_optional_list(self.pipelines, other.pipelines) + if merged_pipelines: + merged_pipelines = remove_duplicates(merged_pipelines) + return Processing( - pipelines=merge_optional_list(self.pipelines, other.pipelines), + pipelines=merged_pipelines, data_processes=self.data_processes + other.data_processes, dependency_graph=merged_graph, notes=merge_notes(self.notes, other.notes), diff --git a/src/aind_data_schema/utils/merge.py b/src/aind_data_schema/utils/merge.py index 92e0f213f..a35025f14 100644 --- a/src/aind_data_schema/utils/merge.py +++ b/src/aind_data_schema/utils/merge.py @@ -71,9 +71,15 @@ def merge_str_tuple_lists( def remove_duplicates(lst: List[Any]) -> List[Any]: """Remove duplicates from a list while preserving order""" - seen = set() - - output_list = [x for x in lst if not (x in seen or seen.add(x))] + try: + seen = set() + output_list = [x for x in lst if not (x in seen or seen.add(x))] + except TypeError: + # Unhashable elements (e.g. pydantic models): fall back to equality + output_list = [] + for x in lst: + if x not in output_list: + output_list.append(x) if len(output_list) != len(lst): logger.info(f"Removed {len(lst) - len(output_list)} duplicates from list") diff --git a/tests/test_composability_merge.py b/tests/test_composability_merge.py index 0ee509dff..9c6859142 100644 --- a/tests/test_composability_merge.py +++ b/tests/test_composability_merge.py @@ -334,6 +334,36 @@ def test_add_processing_objects(self): self.assertEqual(combined.data_processes[2].name, "Denoising_2") self.assertEqual(combined.data_processes[3].name, "Denoising_3") + def test_add_deduplicates_pipelines(self): + """Test that __add__ collapses identical pipelines, keeps distinct ones""" + + t = datetime(2022, 11, 22, 8, 43, 00, tzinfo=timezone.utc) + + def _processing(pipeline): + """Build a one-process Processing carrying the given pipeline""" + return Processing.create_with_sequential_process_graph( + data_processes=[ + DataProcess( + experimenters=["Dr. Dan"], + process_type=ProcessName.DENOISING, + stage=ProcessStage.PROCESSING, + output_path="path/to/outputs", + start_date_time=t, + end_date_time=t, + code=Code(url="https://url/for/analysis", version="0.1.1"), + ), + ], + pipelines=[pipeline], + ) + + pipeline = Code(name="Pipeline", url="https://example.com/pipeline", version="1.0") + combined = _processing(pipeline) + _processing(pipeline) + self.assertEqual(len(combined.pipelines), 1) + + other = Code(name="Pipeline", url="https://example.com/pipeline", version="2.0") + combined = _processing(pipeline) + _processing(other) + self.assertEqual(len(combined.pipelines), 2) + def test_merge_dependency_graph(self): """Test merging dependency graphs""" diff --git a/tests/test_utils_merge.py b/tests/test_utils_merge.py index 0742449cf..1aed8c078 100644 --- a/tests/test_utils_merge.py +++ b/tests/test_utils_merge.py @@ -164,6 +164,14 @@ def test_all_duplicates(self): self.assertEqual(remove_duplicates([1, 1, 1, 1]), [1]) + def test_unhashable_elements(self): + """Test dedup of unhashable elements via the equality fallback""" + from aind_data_schema.utils.merge import remove_duplicates + + with self.assertLogs(level="INFO") as log: + self.assertEqual(remove_duplicates([[1], [1], [2]]), [[1], [2]]) + self.assertIn("Removed 1 duplicates from list", log.output[0]) + class MergeOptionalListTests(unittest.TestCase): """Tests for merge_optional_list"""