Skip to content

Commit 92f88ac

Browse files
committed
fix get_trajectory helper func + test
1 parent e09fd48 commit 92f88ac

2 files changed

Lines changed: 10 additions & 9 deletions

File tree

mp_api/client/routes/materials/tasks.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -24,7 +24,7 @@ class TaskRester(BaseRester):
2424
primary_key: str = "task_id"
2525
delta_backed = True
2626

27-
def get_trajectory(self, task_id: MPID | AlphaID | str) -> list[dict[str, Any]]:
27+
def get_trajectory(self, task_id: MPID | AlphaID | str) -> dict[str, Any]:
2828
"""Returns a Trajectory object containing the geometry of the
2929
material throughout a calculation. This is most useful for
3030
observing how a material relaxes during a geometry optimization.
@@ -33,7 +33,7 @@ def get_trajectory(self, task_id: MPID | AlphaID | str) -> list[dict[str, Any]]:
3333
task_id (str, MPID, AlphaID): Task ID
3434
3535
Returns:
36-
list of dict representing emmet.core.trajectory.Trajectory
36+
dict representing emmet.core.trajectory.RelaxTrajectory
3737
"""
3838
as_alpha = str(AlphaID(task_id, padlen=8)).split("-")[-1]
3939
traj_tbl = DeltaTable(
@@ -57,7 +57,7 @@ def get_trajectory(self, task_id: MPID | AlphaID | str) -> list[dict[str, Any]]:
5757
if not traj_data:
5858
raise MPRestError(f"No trajectory data for {task_id} found")
5959

60-
return RelaxTrajectory(**traj_data[0]).to_pmg().as_dict()
60+
return RelaxTrajectory(**traj_data[0]).model_dump()
6161

6262
def search(
6363
self,

tests/client/materials/test_tasks.py

Lines changed: 7 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -1,13 +1,14 @@
11
import os
2-
from ..conftest import client_search_testing, requires_api_key
3-
import pytest
42

3+
import pytest
54
from emmet.core.mpid import MPID, AlphaID
6-
from emmet.core.trajectory import Trajectory
5+
from emmet.core.trajectory import RelaxTrajectory
76
from emmet.core.utils import utcnow
87

98
from mp_api.client.routes.materials.tasks import TaskRester
109

10+
from ..conftest import client_search_testing, requires_api_key
11+
1112

1213
@pytest.fixture
1314
def rester():
@@ -57,11 +58,11 @@ def test_client(rester):
5758

5859
@pytest.mark.parametrize("mpid", ["mp-149", MPID("mp-149"), AlphaID("mp-149")])
5960
def test_get_trajectories(rester, mpid):
60-
trajectories = [traj for traj in rester.get_trajectory(mpid)]
61+
trajectory = rester.get_trajectory(mpid)
6162

6263
expected_model_fields = {
6364
field_name
64-
for field_name, field in Trajectory.model_fields.items()
65+
for field_name, field in RelaxTrajectory.model_fields.items()
6566
if not field.exclude
6667
}
67-
assert all(set(traj) == expected_model_fields for traj in trajectories)
68+
assert set(trajectory) == expected_model_fields

0 commit comments

Comments
 (0)