Skip to content

Commit ec7a0e2

Browse files
Facilitate matrix simulations
The backend allows for n x m number of values and components to vary. The frontend utilises this to make a 1D grid to vary single component over a linear range.
1 parent ade11c9 commit ec7a0e2

24 files changed

Lines changed: 1480 additions & 76 deletions
Lines changed: 35 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,35 @@
1+
"""add grid simulations
2+
3+
Revision ID: a7f3c9d21b84
4+
Revises: 7c2e1f4b8a90
5+
Create Date: 2026-06-12 00:00:00.000000
6+
7+
"""
8+
9+
from typing import Sequence, Union
10+
11+
import sqlalchemy as sa
12+
from alembic import op
13+
14+
15+
revision: str = "a7f3c9d21b84"
16+
down_revision: Union[str, Sequence[str], None] = "7c2e1f4b8a90"
17+
branch_labels: Union[str, Sequence[str], None] = None
18+
depends_on: Union[str, Sequence[str], None] = None
19+
20+
21+
def upgrade() -> None:
22+
op.create_table(
23+
"grid_simulations",
24+
sa.Column("owner_id", sa.Uuid(), nullable=True),
25+
sa.Column("axes", sa.JSON(), nullable=False),
26+
sa.Column("simulation_ids", sa.JSON(), nullable=False),
27+
sa.Column("id", sa.Uuid(), nullable=False),
28+
sa.Column("created_at", sa.DateTime(), nullable=False),
29+
sa.Column("updated_at", sa.DateTime(), nullable=False),
30+
sa.PrimaryKeyConstraint("id"),
31+
)
32+
33+
34+
def downgrade() -> None:
35+
op.drop_table("grid_simulations")

backend/src/acidwatch_api/database.py

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -49,6 +49,14 @@ class Simulation(Base):
4949
model_inputs: Mapped[list[ModelInput]] = relationship(back_populates="simulation")
5050

5151

52+
class GridSimulation(Base):
53+
__tablename__ = "grid_simulations"
54+
55+
owner_id: Mapped[UUID | None] = mapped_column(Uuid)
56+
axes: Mapped[list[dict]] = mapped_column(JSON)
57+
simulation_ids: Mapped[list[str]] = mapped_column(JSON)
58+
59+
5260
class ModelInput(Base):
5361
__tablename__ = "model_inputs"
5462

backend/src/acidwatch_api/models/datamodel.py

Lines changed: 56 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,8 +1,9 @@
11
from __future__ import annotations
22

33
from typing import Any, Literal, Optional, Dict, TypeAlias
4+
from uuid import UUID
45

5-
from pydantic import BaseModel, ConfigDict, Field
6+
from pydantic import BaseModel, ConfigDict, Field, model_validator
67
from pydantic.alias_generators import to_camel
78

89

@@ -58,6 +59,60 @@ class SimulationResult(_BaseModel):
5859
results: list[ModelResult]
5960

6061

62+
class AxisRange(_BaseModel):
63+
"""A linear, inclusive range that is sampled at ``steps`` points."""
64+
65+
min: float
66+
max: float
67+
steps: int = Field(default=10, ge=2, le=25)
68+
69+
@model_validator(mode="after")
70+
def _check_bounds(self) -> "AxisRange":
71+
if self.max <= self.min:
72+
raise ValueError("max must be greater than min")
73+
return self
74+
75+
def values(self) -> list[float]:
76+
step = (self.max - self.min) / (self.steps - 1)
77+
return [self.min + step * i for i in range(self.steps)]
78+
79+
80+
class Axis(_BaseModel):
81+
substance: str
82+
range: AxisRange
83+
84+
85+
class CreateGridSimulation(_BaseModel):
86+
"""Request body for starting a grid simulation.
87+
88+
Runs a model chain once for each point in the cartesian product of the
89+
axes, substituting each axis's substance in ``concentrations`` with the
90+
corresponding value.
91+
"""
92+
93+
axes: list[Axis] = Field(min_length=1, max_length=2)
94+
concentrations: dict[str, int | float]
95+
conditions: Conditions = Field(default_factory=Conditions)
96+
models: list[ModelInput] = Field(min_length=1)
97+
98+
99+
class GridPoint(_BaseModel):
100+
coordinates: list[float]
101+
simulation_id: UUID
102+
status: Literal["done", "pending", "error"]
103+
error: str | None = None
104+
concentrations: dict[str, int | float]
105+
106+
107+
class GridSimulationResult(_BaseModel):
108+
status: Literal["done", "pending"]
109+
axes: list[Axis]
110+
concentrations: dict[str, int | float]
111+
conditions: Conditions
112+
models: list[ModelInput]
113+
points: list[GridPoint]
114+
115+
61116
class JsonResult(BaseModel):
62117
type: Literal["json"] = "json"
63118
label: str | None = None

backend/src/acidwatch_api/routes/__init__.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -4,10 +4,12 @@
44

55
from . import models
66
from . import oasis
7+
from . import grid_simulations
78

89
router = APIRouter()
910
router.include_router(models.router)
1011
router.include_router(oasis.router)
12+
router.include_router(grid_simulations.router)
1113

1214

1315
__all__ = ["router"]
Lines changed: 235 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,235 @@
1+
from __future__ import annotations
2+
3+
import itertools
4+
import logging
5+
from typing import Annotated
6+
from uuid import UUID
7+
8+
from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException, Request
9+
from sqlalchemy import select
10+
11+
import acidwatch_api.database as db
12+
from acidwatch_api.authentication import OptionalCurrentUser
13+
from acidwatch_api.database import GetDB
14+
from acidwatch_api.models import InputError
15+
from acidwatch_api.models.datamodel import (
16+
Axis,
17+
Conditions,
18+
CreateGridSimulation,
19+
GridPoint,
20+
GridSimulationResult,
21+
ModelInput,
22+
Phase,
23+
)
24+
from acidwatch_api.routes.models import (
25+
AdapterSet,
26+
_phases_to_concentrations,
27+
build_adapters,
28+
build_model_input_rows,
29+
get_adapters,
30+
order_chain,
31+
run_adapters,
32+
)
33+
34+
router = APIRouter()
35+
36+
logger = logging.getLogger(__name__)
37+
38+
39+
def _cartesian_values(axes: list[Axis]) -> list[list[float]]:
40+
ranges = [axis.range.values() for axis in axes]
41+
return [list(point) for point in itertools.product(*ranges)]
42+
43+
44+
def _summarize_point(
45+
ordered: list[tuple[db.ModelInput, db.ModelResult | None]],
46+
) -> tuple[str, dict[str, int | float], str | None]:
47+
if not ordered:
48+
return "pending", {}, None
49+
50+
for _, result in ordered:
51+
if result is not None and result.error is not None:
52+
return "error", {}, result.error
53+
54+
if any(result is None for _, result in ordered):
55+
return "pending", {}, None
56+
57+
final_result = ordered[-1][1]
58+
assert final_result is not None
59+
concentrations = _phases_to_concentrations(
60+
[Phase(**p) for p in final_result.phases]
61+
)
62+
return "done", concentrations, None
63+
64+
65+
@router.post("/grid-simulations")
66+
async def run_grid_simulation(
67+
create: CreateGridSimulation,
68+
user: OptionalCurrentUser,
69+
request: Request,
70+
session: GetDB,
71+
background_tasks: BackgroundTasks,
72+
all_adapters: Annotated[AdapterSet, Depends(get_adapters)],
73+
) -> UUID:
74+
jwt_token = user.jwt_token if user else None
75+
76+
adapters = build_adapters(
77+
create.models, create.conditions, all_adapters, jwt_token
78+
)
79+
80+
for axis in create.axes:
81+
if axis.substance not in adapters[0].valid_substances:
82+
raise HTTPException(
83+
status_code=422,
84+
detail={
85+
"axes": [
86+
f"'{axis.substance}' is not supported by the selected model"
87+
]
88+
},
89+
)
90+
91+
test_concentrations = {
92+
**create.concentrations,
93+
**{axis.substance: 0 for axis in create.axes},
94+
}
95+
try:
96+
adapters[0].validate_concentrations(test_concentrations)
97+
except InputError as exc:
98+
raise HTTPException(status_code=422, detail=exc.detail)
99+
100+
grid_points = _cartesian_values(create.axes)
101+
102+
scheduled: list[tuple[dict[str, int | float], list, list[UUID]]] = []
103+
simulation_ids: list[str] = []
104+
105+
for coordinates in grid_points:
106+
point_concentrations = {
107+
**create.concentrations,
108+
**{
109+
axis.substance: value
110+
for axis, value in zip(create.axes, coordinates)
111+
},
112+
}
113+
model_input_rows = build_model_input_rows(create.models)
114+
simulation = db.Simulation(
115+
owner_id=UUID(user.id) if user else None,
116+
phases=[{"kind": "co2-rich", "fraction": 1.0, "concentrations": point_concentrations}],
117+
conditions=create.conditions.model_dump(),
118+
model_inputs=model_input_rows,
119+
)
120+
session.add(simulation)
121+
session.flush()
122+
simulation_ids.append(str(simulation.id))
123+
124+
point_adapters = build_adapters(
125+
create.models, create.conditions, all_adapters, jwt_token
126+
)
127+
scheduled.append(
128+
(
129+
point_concentrations,
130+
point_adapters,
131+
[row.id for row in model_input_rows],
132+
)
133+
)
134+
135+
grid = db.GridSimulation(
136+
owner_id=UUID(user.id) if user else None,
137+
axes=[axis.model_dump() for axis in create.axes],
138+
simulation_ids=simulation_ids,
139+
)
140+
session.add(grid)
141+
session.commit()
142+
143+
for point_concentrations, point_adapters, model_input_ids in scheduled:
144+
background_tasks.add_task(
145+
run_adapters,
146+
request.state.session,
147+
point_concentrations,
148+
point_adapters,
149+
model_input_ids,
150+
)
151+
152+
return grid.id
153+
154+
155+
@router.get("/grid-simulations/{grid_id}/result")
156+
def get_grid_simulation_result(
157+
grid_id: UUID,
158+
session: GetDB,
159+
) -> GridSimulationResult:
160+
grid = session.get_one(db.GridSimulation, grid_id)
161+
162+
axes = [Axis(**a) for a in grid.axes]
163+
sim_uuids = [UUID(sid) for sid in grid.simulation_ids]
164+
165+
simulations_by_id: dict[UUID, db.Simulation] = {}
166+
if sim_uuids:
167+
rows = (
168+
session.execute(
169+
select(db.Simulation).where(db.Simulation.id.in_(sim_uuids))
170+
)
171+
.scalars()
172+
.all()
173+
)
174+
simulations_by_id = {sim.id: sim for sim in rows}
175+
176+
rows_by_simulation: dict[UUID, list[tuple[db.ModelInput, db.ModelResult | None]]] = {}
177+
if sim_uuids:
178+
q = (
179+
select(db.ModelInput, db.ModelResult)
180+
.where(db.ModelInput.simulation_id.in_(sim_uuids))
181+
.outerjoin(db.ModelResult)
182+
)
183+
for model_input, result in session.execute(q):
184+
rows_by_simulation.setdefault(model_input.simulation_id, []).append(
185+
(model_input, result)
186+
)
187+
188+
grid_points = _cartesian_values(axes)
189+
points: list[GridPoint] = []
190+
overall_pending = False
191+
192+
for index, coordinates in enumerate(grid_points):
193+
sim_id = sim_uuids[index] if index < len(sim_uuids) else None
194+
if sim_id is None or sim_id not in simulations_by_id:
195+
overall_pending = True
196+
continue
197+
198+
ordered = order_chain(rows_by_simulation.get(sim_id, []))
199+
status, concentrations, error = _summarize_point(ordered)
200+
if status == "pending":
201+
overall_pending = True
202+
203+
points.append(
204+
GridPoint(
205+
coordinates=coordinates,
206+
simulation_id=sim_id,
207+
status=status, # type: ignore[arg-type]
208+
error=error,
209+
concentrations=concentrations,
210+
)
211+
)
212+
213+
first_sim = simulations_by_id.get(sim_uuids[0]) if sim_uuids else None
214+
if first_sim is not None:
215+
models = [
216+
ModelInput(model_id=mi.model_id, parameters=mi.parameters)
217+
for mi, _ in order_chain(rows_by_simulation.get(first_sim.id, []))
218+
]
219+
base_concentrations = _phases_to_concentrations(
220+
[Phase(**p) for p in (first_sim.phases or [])]
221+
)
222+
conditions = Conditions(**(first_sim.conditions or {}))
223+
else:
224+
models = []
225+
base_concentrations = {}
226+
conditions = Conditions()
227+
228+
return GridSimulationResult(
229+
status="pending" if overall_pending else "done",
230+
axes=axes,
231+
concentrations=base_concentrations,
232+
conditions=conditions,
233+
models=models,
234+
points=points,
235+
)

0 commit comments

Comments
 (0)