Skip to content

Commit bfe14bd

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 83bc534 commit bfe14bd

29 files changed

Lines changed: 1499 additions & 92 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: 61 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,7 @@
22

33
from typing import Any, Literal, Optional, Dict, TypeAlias
44

5-
from pydantic import BaseModel, ConfigDict, Field
5+
from pydantic import BaseModel, ConfigDict, Field, model_validator
66
from pydantic.alias_generators import to_camel
77

88

@@ -53,9 +53,68 @@ class ModelResult(_BaseModel):
5353

5454

5555
class SimulationResult(_BaseModel):
56-
status: Literal["done", "pending"]
56+
status: Literal["done", "pending", "error"]
5757
input: Simulation
5858
results: list[ModelResult]
59+
error: str | None = None
60+
61+
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+
@model_validator(mode="after")
99+
def _check_axes(self) -> "CreateGridSimulation":
100+
substances = [axis.substance for axis in self.axes]
101+
if len(substances) != len(set(substances)):
102+
raise ValueError("Each axis must use a unique substance")
103+
104+
total = 1
105+
for axis in self.axes:
106+
total *= axis.range.steps
107+
if total > 100:
108+
raise ValueError(
109+
f"Grid too large: {total} points exceeds the maximum of 100"
110+
)
111+
return self
112+
113+
114+
class GridSimulationResult(_BaseModel):
115+
status: Literal["done", "pending"]
116+
axes: list[Axis]
117+
simulations: list[SimulationResult]
59118

60119

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

backend/src/acidwatch_api/routes/models.py

Lines changed: 21 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -160,7 +160,7 @@ def query_chain_rows(
160160
.where(db.ModelInput.simulation_id == simulation_id)
161161
.outerjoin(db.ModelResult)
162162
)
163-
return list(session.execute(q).fetchall())
163+
return [(row[0], row[1]) for row in session.execute(q).fetchall()]
164164

165165

166166
@router.get("/models")
@@ -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)