Skip to content

Commit b49635f

Browse files
fix after data structure change
1 parent 5c9b6a3 commit b49635f

2 files changed

Lines changed: 32 additions & 6 deletions

File tree

backend/src/acidwatch_api/routes/sweeps.py

Lines changed: 10 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -16,11 +16,13 @@
1616
Conditions,
1717
CreateSweep,
1818
ModelInput,
19+
Phase,
1920
SweepPoint,
2021
SweepResult,
2122
)
2223
from acidwatch_api.routes.models import (
2324
AdapterSet,
25+
_phases_to_concentrations,
2426
build_adapters,
2527
build_model_input_rows,
2628
get_adapters,
@@ -52,7 +54,10 @@ def _summarize_point(
5254

5355
final_result = ordered[-1][1]
5456
assert final_result is not None
55-
return "done", final_result.concentrations, None
57+
concentrations = _phases_to_concentrations(
58+
[Phase(**p) for p in final_result.phases]
59+
)
60+
return "done", concentrations, None
5661

5762

5863
@router.post("/sweeps")
@@ -109,7 +114,7 @@ async def run_sweep(
109114
session.add(
110115
db.Simulation(
111116
owner_id=UUID(user.id) if user else None,
112-
concentrations=point_concentrations,
117+
phases=[{"kind": "co2-rich", "fraction": 1.0, "concentrations": point_concentrations}],
113118
conditions=create_sweep.conditions.model_dump(),
114119
sweep=sweep,
115120
sweep_value_index=index,
@@ -201,7 +206,9 @@ def get_sweep_result(
201206
ModelInput(model_id=mi.model_id, parameters=mi.parameters)
202207
for mi, _ in order_chain(rows_by_simulation.get(first_simulation.id, []))
203208
]
204-
base_concentrations = dict(first_simulation.concentrations or {})
209+
base_concentrations = _phases_to_concentrations(
210+
[Phase(**p) for p in (first_simulation.phases or [])]
211+
)
205212
conditions = Conditions(**(first_simulation.conditions or {}))
206213
else:
207214
models = []

backend/tests/test_sweeps_endpoint.py

Lines changed: 22 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,7 @@
55
from acidwatch_api.app import fastapi_app
66
from acidwatch_api.authentication import authenticated_user_claims
77
from acidwatch_api.models import base
8+
from acidwatch_api.models.datamodel import Phase
89
from acidwatch_api.routes.models import get_adapters
910

1011

@@ -38,7 +39,13 @@ class HalvingAdapter(base.BaseAdapter):
3839
valid_substances = ["H2O"]
3940

4041
async def run(self):
41-
return {key: value / 2 for key, value in self.concentrations.items()}
42+
return [
43+
Phase(
44+
kind="co2-rich",
45+
fraction=1.0,
46+
concentrations={key: value / 2 for key, value in self.concentrations.items()},
47+
)
48+
]
4249

4350

4451
class QuadruplingAdapter(base.BaseAdapter):
@@ -49,7 +56,13 @@ class QuadruplingAdapter(base.BaseAdapter):
4956
valid_substances = ["H2O"]
5057

5158
async def run(self):
52-
return {key: value * 4 for key, value in self.concentrations.items()}
59+
return [
60+
Phase(
61+
kind="co2-rich",
62+
fraction=1.0,
63+
concentrations={key: value * 4 for key, value in self.concentrations.items()},
64+
)
65+
]
5366

5467

5568
class FailingAdapter(base.BaseAdapter):
@@ -117,7 +130,13 @@ def test_sweep_points_are_individually_retrievable_simulations(client):
117130

118131
assert simulation["status"] == "done"
119132
assert simulation["input"]["concentrations"] == {"H2O": point["value"]}
120-
assert simulation["results"][-1]["concentrations"] == point["concentrations"]
133+
last_result_concentrations = {
134+
k: v
135+
for phase in simulation["results"][-1]["phases"]
136+
if phase["kind"] == "co2-rich"
137+
for k, v in phase["concentrations"].items()
138+
}
139+
assert last_result_concentrations == point["concentrations"]
121140

122141

123142
@pytest.mark.usefixtures("dummy_adapters")

0 commit comments

Comments
 (0)