diff --git a/tests/experimental/trajectory/file_store_test.py b/tests/experimental/trajectory/file_store_test.py index cc6bd0189..552ebbec1 100644 --- a/tests/experimental/trajectory/file_store_test.py +++ b/tests/experimental/trajectory/file_store_test.py @@ -463,6 +463,14 @@ def test_context_manager_closes_store_on_exit(self) -> None: self.assertTrue(step_path.exists()) + def test_get_trajectories_metadata_nonexistent_root_dir_returns_empty( + self, + ) -> None: + """Verifies get_trajectories_metadata on a non-existent root_dir gracefully returns empty list.""" + nonexistent_root = self.tmp_dir / "does_not_exist" + store_instance = file_store.FileTrajectoryStore(root_dir=nonexistent_root) + self.assertEmpty(store_instance.get_trajectories_metadata()) + if __name__ == "__main__": absltest.main() diff --git a/tunix/experimental/trajectory/file_store.py b/tunix/experimental/trajectory/file_store.py index c72929991..1964b208e 100644 --- a/tunix/experimental/trajectory/file_store.py +++ b/tunix/experimental/trajectory/file_store.py @@ -114,24 +114,38 @@ def get_step_path(self, trajectory_id: str, step_id: int) -> epath.Path: return self.get_trajectory_dir(trajectory_id) / step_filename def get_trajectories_metadata( - self, + self, trajectory_ids: list[str] | None = None ) -> list[trajectory_lib.TrajectoryMetadata]: - """Retrieves metadata for each trajectory in the run.""" - metas: list[trajectory_lib.TrajectoryMetadata] = [] - if not self.root_dir.exists(): - return metas + """Retrieves metadata for trajectories in the run. - for entry in self.root_dir.iterdir(): - if not entry.is_dir(): - continue - if not (match := _TRAJECTORY_DIR_REGEX.match(entry.name)): - continue + Args: + trajectory_ids: Optional list of unique trajectory identifiers. If + specified, only metadata for these IDs is returned. If None, metadata + for all trajectories in the run is returned. - traj_id = match.group("trajectory_id") + Returns: + A list of TrajectoryMetadata objects for the requested trajectories. + + Raises: + store.TrajectoryMetadataNotFoundError: If any requested trajectory ID does + not exist. + """ + metas: list[trajectory_lib.TrajectoryMetadata] = [] + if trajectory_ids is None: + if not self.root_dir.exists(): + return metas + trajectory_ids = [] + for entry in self.root_dir.iterdir(): + if not entry.is_dir(): + continue + if not (match := _TRAJECTORY_DIR_REGEX.match(entry.name)): + continue + trajectory_ids.append(match.group("trajectory_id")) + + for traj_id in trajectory_ids: meta_path = self.get_trajectory_metadata_path(traj_id) if not meta_path.exists(): - raise store.TrajectoryMetadataNotFoundError(entry.name) - + raise store.TrajectoryMetadataNotFoundError(traj_id) meta = trajectory_lib.TrajectoryMetadata.model_validate_json( meta_path.read_text() ) diff --git a/tunix/experimental/trajectory/in_memory_store.py b/tunix/experimental/trajectory/in_memory_store.py index dd1c11341..5f0e1ad90 100644 --- a/tunix/experimental/trajectory/in_memory_store.py +++ b/tunix/experimental/trajectory/in_memory_store.py @@ -36,13 +36,31 @@ def __init__(self) -> None: ) def get_trajectories_metadata( - self, + self, trajectory_ids: list[str] | None = None ) -> list[trajectory_lib.TrajectoryMetadata]: - """Retrieves metadata for each trajectory in the run.""" - return [ - meta.model_copy(deep=True) - for meta in self._metadata_by_trajectory_id.values() - ] + """Retrieves metadata for trajectories in the run. + + Args: + trajectory_ids: Optional list of unique trajectory identifiers. If + specified, only metadata for these IDs is returned. If None, metadata + for all trajectories in the run is returned. + + Returns: + A list of TrajectoryMetadata objects for the requested trajectories. + + Raises: + store.TrajectoryMetadataNotFoundError: If any requested trajectory ID does + not exist. + """ + if trajectory_ids is None: + trajectory_ids = list(self._metadata_by_trajectory_id.keys()) + metas: list[trajectory_lib.TrajectoryMetadata] = [] + for traj_id in trajectory_ids: + if traj_id not in self._metadata_by_trajectory_id: + raise store.TrajectoryMetadataNotFoundError(traj_id) + meta = self._metadata_by_trajectory_id[traj_id].model_copy(deep=True) + metas.append(meta) + return metas def get_trajectories( self, trajectory_ids: list[str] diff --git a/tunix/experimental/trajectory/store.py b/tunix/experimental/trajectory/store.py index ca42645d8..3cd2e954c 100644 --- a/tunix/experimental/trajectory/store.py +++ b/tunix/experimental/trajectory/store.py @@ -1,7 +1,6 @@ """Protocols defining Trajectory Store interfaces.""" -import typing -from typing import Protocol +from typing import Protocol, runtime_checkable from tunix.experimental.trajectory import trajectory as trajectory_lib @@ -31,17 +30,26 @@ def __init__(self, trajectory_id: str) -> None: # ============================================================================== -@typing.runtime_checkable +@runtime_checkable class TrajectoryReader(Protocol): """Structural protocol defining read-only Trajectory Store operations.""" def get_trajectories_metadata( - self, + self, trajectory_ids: list[str] | None = None ) -> list[trajectory_lib.TrajectoryMetadata]: - """Retrieves metadata for each trajectory in the run. + """Retrieves metadata for trajectories in the run. + + Args: + trajectory_ids: Optional list of unique trajectory identifiers. If + specified, only metadata for these IDs is returned. If None, metadata + for all trajectories in the run is returned. Returns: - A list of TrajectoryMetadata objects for all trajectories in this run. + A list of TrajectoryMetadata objects for the requested trajectories. + + Raises: + TrajectoryMetadataNotFoundError: If any requested trajectory ID does not + exist. """ ... @@ -62,7 +70,7 @@ def get_trajectories( ... -@typing.runtime_checkable +@runtime_checkable class TrajectoryWriter(Protocol): """Structural protocol defining write Trajectory Store operations.""" diff --git a/tunix/experimental/trajectory/store_testing.py b/tunix/experimental/trajectory/store_testing.py index 16b675590..7f684ab62 100644 --- a/tunix/experimental/trajectory/store_testing.py +++ b/tunix/experimental/trajectory/store_testing.py @@ -53,18 +53,79 @@ def setUp(self) -> None: ], ) - def test_get_trajectories_metadata(self) -> None: - """Tests that metadata for all stored trajectories is retrieved.""" - metas = self.reader.get_trajectories_metadata() - self.assertCountEqual( - metas, - [trajectory_testing.METADATA_1, trajectory_testing.METADATA_2], - ) + @parameterized.named_parameters( + ( + "all_trajectory_ids", + None, + [trajectory_testing.METADATA_1, trajectory_testing.METADATA_2], + ), + ("empty_list", [], []), + ( + "single_trajectory", + [trajectory_testing.TRAJECTORY_ID_1], + [trajectory_testing.METADATA_1], + ), + ( + "multiple_trajectories", + [ + trajectory_testing.TRAJECTORY_ID_1, + trajectory_testing.TRAJECTORY_ID_2, + ], + [trajectory_testing.METADATA_1, trajectory_testing.METADATA_2], + ), + ) + def test_get_trajectories_metadata( + self, + trajectory_ids: list[str] | None, + expected_metas: list[trajectory_lib.TrajectoryMetadata], + ) -> None: + """Tests that metadata for trajectories is retrieved.""" + metas = self.reader.get_trajectories_metadata(trajectory_ids) + self.assertCountEqual(metas, expected_metas) - def test_get_trajectories_metadata_empty(self) -> None: + @parameterized.named_parameters( + ("all_trajectory_ids", None), + ("empty_list", []), + ) + def test_get_trajectories_metadata_empty_store( + self, trajectory_ids: list[str] | None + ) -> None: """Tests that metadata retrieval on an empty store returns an empty list.""" empty_reader = self._create_reader(initial_data=None) - self.assertEmpty(empty_reader.get_trajectories_metadata()) + self.assertEmpty(empty_reader.get_trajectories_metadata(trajectory_ids)) + + @parameterized.named_parameters( + ( + "single_trajectory", + [trajectory_testing.TRAJECTORY_ID_1], + ), + ( + "multiple_trajectories", + [ + trajectory_testing.TRAJECTORY_ID_1, + trajectory_testing.TRAJECTORY_ID_2, + ], + ), + ) + def test_get_trajectories_metadata_empty_store_with_ids_raises( + self, + trajectory_ids: list[str], + ) -> None: + """Tests that querying explicit IDs on an empty store raises TrajectoryMetadataNotFoundError.""" + empty_reader = self._create_reader(initial_data=None) + with self.assertRaisesRegex( + store.TrajectoryMetadataNotFoundError, + f"Trajectory metadata for ID '{trajectory_ids[0]}' not found.", + ): + empty_reader.get_trajectories_metadata(trajectory_ids) + + def test_get_trajectories_metadata_not_found(self) -> None: + """Tests that passing a non-existent trajectory ID raises TrajectoryMetadataNotFoundError.""" + with self.assertRaisesRegex( + store.TrajectoryMetadataNotFoundError, + "Trajectory metadata for ID 'non_existent_id' not found.", + ): + self.reader.get_trajectories_metadata(["non_existent_id"]) @parameterized.named_parameters( ("empty_list", [], []), @@ -93,7 +154,10 @@ def test_get_trajectories( def test_get_trajectories_not_found(self) -> None: """Tests that loading a non-existent trajectory ID raises TrajectoryNotFoundError.""" - with self.assertRaises(store.TrajectoryNotFoundError): + with self.assertRaisesRegex( + store.TrajectoryNotFoundError, + "Trajectory with ID 'non_existent_id' not found.", + ): self.reader.get_trajectories(["non_existent_id"]) @@ -358,5 +422,3 @@ def test_update_metadata_invalid_trajectory_id( ) with self.assertRaises(ValueError): self.writer.update_metadata(meta) - -