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
82 changes: 81 additions & 1 deletion tests/experimental/trajectory/file_store_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,13 +5,22 @@
from absl.testing import absltest
from absl.testing import parameterized
from etils import epath
import pydantic
from tunix.experimental.trajectory import file_store
from tunix.experimental.trajectory import store
from tunix.experimental.trajectory import store_testing
from tunix.experimental.trajectory import trajectory as trajectory_lib
from tunix.experimental.trajectory import trajectory_testing


class _CustomMetadata(trajectory_lib.TrajectoryMetadata):
custom_tag: str = ""


class _CustomTrajectory(_CustomMetadata):
steps: list[trajectory_lib.Step] = pydantic.Field(default_factory=list)


class FileTrajectoryReaderTest(store_testing.TrajectoryReaderTestCase):
"""Contract tests for FileTrajectoryStore's TrajectoryReader implementation."""

Expand Down Expand Up @@ -199,7 +208,8 @@ def blocking_process_task(task):
with mock.patch.object(
self.file_s._writer, "_process_task", side_effect=blocking_process_task
):
# add_step should enqueue task and return immediately while worker loop is blocked.
# add_step should enqueue task and return immediately while worker loop
# is blocked.
self.file_s.add_step(
trajectory_testing.STEP_1_1, trajectory_testing.METADATA_1
)
Expand Down Expand Up @@ -471,6 +481,76 @@ def test_get_trajectories_metadata_nonexistent_root_dir_returns_empty(
store_instance = file_store.FileTrajectoryStore(root_dir=nonexistent_root)
self.assertEmpty(store_instance.get_trajectories_metadata())

def test_tunix_trajectory_with_step_zero(self) -> None:
"""Verifies storing and retrieving TunixTrajectoryMetadata and TunixTrajectory with step_id=0."""
tunix_store: file_store.FileTrajectoryStore[
trajectory_lib.TunixTrajectoryMetadata, trajectory_lib.TunixTrajectory
] = file_store.FileTrajectoryStore(
root_dir=self.tmp_dir,
run_id="tunix_run",
metadata_cls=trajectory_lib.TunixTrajectoryMetadata,
)
meta = trajectory_lib.TunixTrajectoryMetadata(
trajectory_id="tunix_1",
agent=trajectory_lib.Agent(name="a1", version="1.0"),
status="RUNNING",
)
step0 = trajectory_lib.TunixEnvStep(
step_id=0, source=trajectory_lib.Source.USER, message="prompt"
)
step1 = trajectory_lib.TunixAgentStep(
step_id=1, source=trajectory_lib.Source.AGENT, message="response"
)
tunix_store.add_step(step0, meta)
tunix_store.add_step(step1, meta)
tunix_store.flush()

metas = tunix_store.get_trajectories_metadata(["tunix_1"])
self.assertLen(metas, 1)
self.assertIsInstance(metas[0], trajectory_lib.TunixTrajectoryMetadata)
self.assertEqual(metas[0].status, "RUNNING")

trajs = tunix_store.get_trajectories(["tunix_1"])
self.assertLen(trajs, 1)
self.assertIsInstance(trajs[0], trajectory_lib.TunixTrajectory)
self.assertEqual(trajs[0].steps[0].step_id, 0)
self.assertEqual(trajs[0].steps[1].step_id, 1)
self.assertIsInstance(trajs[0].steps[0], trajectory_lib.TunixEnvStep)
self.assertIsInstance(trajs[0].steps[1], trajectory_lib.TunixAgentStep)

def test_custom_metadata_and_trajectory_subclass(self) -> None:
"""Verifies FileTrajectoryStore supports custom metadata and trajectory types."""
custom_store: file_store.FileTrajectoryStore[
_CustomMetadata, _CustomTrajectory
] = file_store.FileTrajectoryStore(
root_dir=self.tmp_dir,
run_id="custom_run",
metadata_cls=_CustomMetadata,
trajectory_cls=_CustomTrajectory,
)
meta = _CustomMetadata(
trajectory_id="custom_1",
agent=trajectory_lib.Agent(name="custom_agent", version="1.0"),
custom_tag="experiment_42",
)
step = trajectory_lib.Step(
step_id=1, source=trajectory_lib.Source.AGENT, message="custom step"
)
custom_store.add_step(step, meta)
custom_store.flush()

metas = custom_store.get_trajectories_metadata(["custom_1"])
self.assertLen(metas, 1)
self.assertIsInstance(metas[0], _CustomMetadata)
self.assertEqual(metas[0].custom_tag, "experiment_42")

trajs = custom_store.get_trajectories(["custom_1"])
self.assertLen(trajs, 1)
self.assertIsInstance(trajs[0], _CustomTrajectory)
self.assertEqual(trajs[0].custom_tag, "experiment_42")
self.assertLen(trajs[0].steps, 1)
self.assertEqual(trajs[0].steps[0].message, "custom step")


if __name__ == "__main__":
absltest.main()
49 changes: 46 additions & 3 deletions tests/experimental/trajectory/in_memory_store_test.py
Original file line number Diff line number Diff line change
@@ -1,10 +1,19 @@
from absl.testing import absltest
import pydantic
from tunix.experimental.trajectory import in_memory_store
from tunix.experimental.trajectory import store
from tunix.experimental.trajectory import store_testing
from tunix.experimental.trajectory import trajectory as trajectory_lib


class _CustomMetadata(trajectory_lib.TrajectoryMetadata):
custom_tag: str = ""


class _CustomTrajectory(_CustomMetadata):
steps: list[trajectory_lib.Step] = pydantic.Field(default_factory=list)


class InMemoryTrajectoryReaderTest(store_testing.TrajectoryReaderTestCase):
"""Contract tests for InMemoryTrajectoryStore's TrajectoryReader implementation."""

Expand Down Expand Up @@ -37,7 +46,8 @@ def _create_reader_and_writer(
mem_store = in_memory_store.InMemoryTrajectoryStore()
return mem_store, mem_store

def test_update_metadata(self):
def test_update_metadata(self) -> None:
"""Verifies that updating metadata in-memory updates the stored metadata."""
mem_store = in_memory_store.InMemoryTrajectoryStore()
meta = trajectory_lib.TrajectoryMetadata(
trajectory_id="t1",
Expand All @@ -53,7 +63,8 @@ def test_update_metadata(self):
read_meta = mem_store.get_trajectories_metadata()[0]
self.assertEqual(read_meta.extra["status"], "SUCCEEDED")

def test_tunix_trajectory_with_step_zero(self):
def test_tunix_trajectory_with_step_zero(self) -> None:
"""Verifies storing and retrieving TunixTrajectoryMetadata and TunixTrajectory with step_id=0."""
mem_store = in_memory_store.InMemoryTrajectoryStore()
meta = trajectory_lib.TunixTrajectoryMetadata(
trajectory_id="tunix_1",
Expand All @@ -73,8 +84,40 @@ def test_tunix_trajectory_with_step_zero(self):
self.assertIsInstance(trajs[0], trajectory_lib.TunixTrajectory)
self.assertEqual(trajs[0].steps[0].step_id, 0)
self.assertEqual(trajs[0].steps[1].step_id, 1)
self.assertIsInstance(trajs[0].steps[0], trajectory_lib.TunixEnvStep)
self.assertIsInstance(trajs[0].steps[1], trajectory_lib.TunixAgentStep)

def test_custom_metadata_and_trajectory_subclass(self) -> None:
"""Verifies InMemoryTrajectoryStore supports custom metadata and trajectory types."""
custom_store: in_memory_store.InMemoryTrajectoryStore[
_CustomMetadata, _CustomTrajectory
] = in_memory_store.InMemoryTrajectoryStore(
trajectory_cls=_CustomTrajectory
)
meta = _CustomMetadata(
trajectory_id="custom_1",
agent=trajectory_lib.Agent(name="custom_agent", version="1.0"),
custom_tag="experiment_42",
)
step = trajectory_lib.Step(
step_id=1, source=trajectory_lib.Source.AGENT, message="custom step"
)
custom_store.add_step(step, meta)

metas = custom_store.get_trajectories_metadata(["custom_1"])
self.assertLen(metas, 1)
self.assertIsInstance(metas[0], _CustomMetadata)
self.assertEqual(metas[0].custom_tag, "experiment_42")

trajs = custom_store.get_trajectories(["custom_1"])
self.assertLen(trajs, 1)
self.assertIsInstance(trajs[0], _CustomTrajectory)
self.assertEqual(trajs[0].custom_tag, "experiment_42")
self.assertLen(trajs[0].steps, 1)
self.assertEqual(trajs[0].steps[0].message, "custom step")

def test_metadata_mutation_isolation(self):
def test_metadata_mutation_isolation(self) -> None:
"""Verifies that mutating returned metadata does not alter internal store state."""
mem_store = in_memory_store.InMemoryTrajectoryStore()
meta = trajectory_lib.TrajectoryMetadata(
trajectory_id="iso_1",
Expand Down
62 changes: 62 additions & 0 deletions tests/experimental/trajectory/store_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@
from tunix.experimental.trajectory import file_store
from tunix.experimental.trajectory import in_memory_store
from tunix.experimental.trajectory import store as store_lib
from tunix.experimental.trajectory import trajectory as trajectory_lib
from tunix.experimental.trajectory import trajectory_testing


Expand Down Expand Up @@ -192,5 +193,66 @@ def test_file_store_without_run_id_raises_on_to_config(self):
store.close()


class GenericTypeParametersTest(absltest.TestCase):
"""Tests that TrajectoryStore and backends preserve generic type parameters without Generic."""

def test_classes_inherit_parameters_without_explicit_generic(self) -> None:
"""Verifies that Python typing automatically discovers (T, TrajT)."""
self.assertLen(store_lib.TrajectoryStore.__parameters__, 2)
self.assertLen(in_memory_store.InMemoryTrajectoryStore.__parameters__, 2)
self.assertLen(file_store.FileTrajectoryStore.__parameters__, 2)

self.assertEqual(
store_lib.TrajectoryStore.__parameters__,
(store_lib.T, store_lib.TrajT),
)

def test_classes_are_runtime_subscriptable(self) -> None:
"""Verifies that classes are subscriptable with concrete models at runtime."""
subscripted_store = store_lib.TrajectoryStore[
trajectory_lib.TunixTrajectoryMetadata, trajectory_lib.TunixTrajectory
]
subscripted_mem = in_memory_store.InMemoryTrajectoryStore[
trajectory_lib.TunixTrajectoryMetadata, trajectory_lib.TunixTrajectory
]
subscripted_file = file_store.FileTrajectoryStore[
trajectory_lib.TunixTrajectoryMetadata, trajectory_lib.TunixTrajectory
]
self.assertIsNotNone(subscripted_store)
self.assertIsNotNone(subscripted_mem)
self.assertIsNotNone(subscripted_file)

def test_instances_satisfy_protocols_and_abc(self) -> None:
"""Verifies protocol and ABC conformance on instantiated instances."""
mem = in_memory_store.InMemoryTrajectoryStore()
self.assertIsInstance(mem, store_lib.TrajectoryStore)
self.assertIsInstance(mem, store_lib.TrajectoryReader)
self.assertIsInstance(mem, store_lib.TrajectoryWriter)

tmp_dir = self.create_tempdir().full_path
f_store = file_store.FileTrajectoryStore(root_dir=tmp_dir)
self.assertIsInstance(f_store, store_lib.TrajectoryStore)
self.assertIsInstance(f_store, store_lib.TrajectoryReader)
self.assertIsInstance(f_store, store_lib.TrajectoryWriter)
f_store.close()

def test_subclassing_closes_type_parameters(self) -> None:
"""Verifies that concrete subclassing binds and closes type parameters."""

class ConcreteFileStore(
file_store.FileTrajectoryStore[
trajectory_lib.TunixTrajectoryMetadata,
trajectory_lib.TunixTrajectory,
]
):
pass

self.assertEmpty(ConcreteFileStore.__parameters__)
tmp_dir = self.create_tempdir().full_path
c_store = ConcreteFileStore(root_dir=tmp_dir)
self.assertIsInstance(c_store, store_lib.TrajectoryStore)
c_store.close()


if __name__ == "__main__":
absltest.main()
Loading
Loading