Skip to content

Commit 6283117

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 6283117

28 files changed

Lines changed: 1450 additions & 90 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: b5d2e3f14a70
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] = "b5d2e3f14a70"
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: 58 additions & 2 deletions
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

@@ -53,9 +54,64 @@ class ModelResult(_BaseModel):
5354

5455

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

60116

61117
class JsonResult(BaseModel):

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: 152 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,152 @@
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+
10+
import acidwatch_api.database as db
11+
from acidwatch_api.authentication import OptionalCurrentUser
12+
from acidwatch_api.database import GetDB
13+
from acidwatch_api.models import InputError
14+
from acidwatch_api.models.datamodel import (
15+
Axis,
16+
CreateGridSimulation,
17+
GridSimulationResult,
18+
SimulationResult,
19+
)
20+
from acidwatch_api.routes.models import (
21+
AdapterSet,
22+
build_adapters,
23+
build_model_input_rows,
24+
build_simulation_result,
25+
get_adapters,
26+
run_adapters,
27+
)
28+
29+
router = APIRouter()
30+
31+
logger = logging.getLogger(__name__)
32+
33+
34+
def _cartesian_values(axes: list[Axis]) -> list[list[float]]:
35+
ranges = [axis.range.values() for axis in axes]
36+
return [list(point) for point in itertools.product(*ranges)]
37+
38+
39+
@router.post("/grid-simulations")
40+
async def run_grid_simulation(
41+
create: CreateGridSimulation,
42+
user: OptionalCurrentUser,
43+
request: Request,
44+
session: GetDB,
45+
background_tasks: BackgroundTasks,
46+
all_adapters: Annotated[AdapterSet, Depends(get_adapters)],
47+
) -> UUID:
48+
jwt_token = user.jwt_token if user else None
49+
50+
adapters = build_adapters(create.models, create.conditions, all_adapters, jwt_token)
51+
52+
for axis in create.axes:
53+
if axis.substance not in adapters[0].valid_substances:
54+
raise HTTPException(
55+
status_code=422,
56+
detail={
57+
"axes": [
58+
f"'{axis.substance}' is not supported by the selected model"
59+
]
60+
},
61+
)
62+
63+
test_concentrations = {
64+
**create.concentrations,
65+
**{axis.substance: 0 for axis in create.axes},
66+
}
67+
try:
68+
adapters[0].validate_concentrations(test_concentrations)
69+
except InputError as exc:
70+
raise HTTPException(status_code=422, detail=exc.detail)
71+
72+
grid_points = _cartesian_values(create.axes)
73+
74+
scheduled: list[tuple[dict[str, int | float], list, list[UUID]]] = []
75+
simulation_ids: list[str] = []
76+
77+
for coordinates in grid_points:
78+
point_concentrations = {
79+
**create.concentrations,
80+
**{axis.substance: value for axis, value in zip(create.axes, coordinates)},
81+
}
82+
model_input_rows = build_model_input_rows(create.models)
83+
simulation = db.Simulation(
84+
owner_id=UUID(user.id) if user else None,
85+
phases=[
86+
{
87+
"kind": "co2-rich",
88+
"fraction": 1.0,
89+
"concentrations": point_concentrations,
90+
}
91+
],
92+
conditions=create.conditions.model_dump(),
93+
model_inputs=model_input_rows,
94+
)
95+
session.add(simulation)
96+
session.flush()
97+
simulation_ids.append(str(simulation.id))
98+
99+
point_adapters = build_adapters(
100+
create.models, create.conditions, all_adapters, jwt_token
101+
)
102+
scheduled.append(
103+
(
104+
point_concentrations,
105+
point_adapters,
106+
[row.id for row in model_input_rows],
107+
)
108+
)
109+
110+
grid = db.GridSimulation(
111+
owner_id=UUID(user.id) if user else None,
112+
axes=[axis.model_dump() for axis in create.axes],
113+
simulation_ids=simulation_ids,
114+
)
115+
session.add(grid)
116+
session.commit()
117+
118+
for point_concentrations, point_adapters, model_input_ids in scheduled:
119+
background_tasks.add_task(
120+
run_adapters,
121+
request.state.session,
122+
point_concentrations,
123+
point_adapters,
124+
model_input_ids,
125+
)
126+
127+
return grid.id
128+
129+
130+
@router.get("/grid-simulations/{grid_id}/result")
131+
def get_grid_simulation_result(
132+
grid_id: UUID,
133+
session: GetDB,
134+
) -> GridSimulationResult:
135+
grid = session.get_one(db.GridSimulation, grid_id)
136+
137+
axes = [Axis(**a) for a in grid.axes]
138+
sim_uuids = [UUID(sid) for sid in grid.simulation_ids]
139+
140+
simulations: list[SimulationResult] = [
141+
build_simulation_result(session, sim_id) for sim_id in sim_uuids
142+
]
143+
144+
overall_status = "done"
145+
if any(s.status == "pending" for s in simulations):
146+
overall_status = "pending"
147+
148+
return GridSimulationResult(
149+
status=overall_status,
150+
axes=axes,
151+
simulations=simulations,
152+
)

backend/src/acidwatch_api/routes/models.py

Lines changed: 20 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -253,11 +253,7 @@ async def _run_adapter(
253253
session.add(result_obj)
254254

255255

256-
@router.get("/simulations/{simulation_id}/result")
257-
def get_result_for_simulation(
258-
simulation_id: UUID,
259-
session: GetDB,
260-
) -> SimulationResult:
256+
def build_simulation_result(session: Session, simulation_id: UUID) -> SimulationResult:
261257
db_simulation = session.get_one(db.Simulation, simulation_id)
262258

263259
model_inputs: list[ModelInput] = []
@@ -278,9 +274,17 @@ def get_result_for_simulation(
278274

279275
if result.error is not None:
280276
logger.error("Simulation %s failed: %s", simulation_id, result.error)
281-
raise HTTPException(
282-
status_code=500,
283-
detail=f"Simulation failed: {result.error}",
277+
return SimulationResult(
278+
status="error",
279+
input=Simulation(
280+
concentrations=_phases_to_concentrations(
281+
[Phase(**p) for p in db_simulation.phases]
282+
),
283+
conditions=Conditions(**(db_simulation.conditions or {})),
284+
models=model_inputs,
285+
),
286+
results=results,
287+
error=result.error,
284288
)
285289

286290
results.append(
@@ -320,6 +324,14 @@ def get_result_for_simulation(
320324
)
321325

322326

327+
@router.get("/simulations/{simulation_id}/result")
328+
def get_result_for_simulation(
329+
simulation_id: UUID,
330+
session: GetDB,
331+
) -> SimulationResult:
332+
return build_simulation_result(session, simulation_id)
333+
334+
323335
@router.post("/simulations")
324336
async def run_simulation(
325337
create_simulation: Simulation,

0 commit comments

Comments
 (0)