Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 8 additions & 0 deletions tests/experimental/trajectory/file_store_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
40 changes: 27 additions & 13 deletions tunix/experimental/trajectory/file_store.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
)
Expand Down
30 changes: 24 additions & 6 deletions tunix/experimental/trajectory/in_memory_store.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Expand Down
22 changes: 15 additions & 7 deletions tunix/experimental/trajectory/store.py
Original file line number Diff line number Diff line change
@@ -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

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

Expand All @@ -62,7 +70,7 @@ def get_trajectories(
...


@typing.runtime_checkable
@runtime_checkable
class TrajectoryWriter(Protocol):
"""Structural protocol defining write Trajectory Store operations."""

Expand Down
86 changes: 74 additions & 12 deletions tunix/experimental/trajectory/store_testing.py
Original file line number Diff line number Diff line change
Expand Up @@ -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", [], []),
Expand Down Expand Up @@ -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"])


Expand Down Expand Up @@ -358,5 +422,3 @@ def test_update_metadata_invalid_trajectory_id(
)
with self.assertRaises(ValueError):
self.writer.update_metadata(meta)


Loading