diff --git a/backend/src/acidwatch_api/routes/__init__.py b/backend/src/acidwatch_api/routes/__init__.py index 004c429e..de5347f7 100644 --- a/backend/src/acidwatch_api/routes/__init__.py +++ b/backend/src/acidwatch_api/routes/__init__.py @@ -2,14 +2,16 @@ from fastapi import APIRouter +from . import grid_simulations from . import models from . import oasis -from . import grid_simulations +from . import simulations router = APIRouter() router.include_router(models.router) -router.include_router(oasis.router) +router.include_router(simulations.router) router.include_router(grid_simulations.router) +router.include_router(oasis.router) __all__ = ["router"] diff --git a/backend/src/acidwatch_api/routes/_helpers.py b/backend/src/acidwatch_api/routes/_helpers.py new file mode 100644 index 00000000..e94cbf4d --- /dev/null +++ b/backend/src/acidwatch_api/routes/_helpers.py @@ -0,0 +1,267 @@ +from __future__ import annotations + +import logging +from collections import defaultdict +from datetime import datetime, timedelta +from typing import cast +from uuid import UUID, uuid4 + +from fastapi import HTTPException, Request +from pydantic import TypeAdapter, ValidationError +from sqlalchemy import select +from sqlalchemy.exc import IntegrityError +from sqlalchemy.orm import Session + +import acidwatch_api.database as db +from acidwatch_api.broker.heartbeat import HeartbeatRegistry +from acidwatch_api.settings import SETTINGS +from acidwatch_messaging import Transport +from acidwatch_models import ( + AdapterSet, + BaseAdapter, + InputError, +) +from acidwatch_models.datamodel import ( + AnyPanel, + Conditions, + ModelInput, + ModelResult, + Phase, + Simulation, + SimulationResult, +) + +logger = logging.getLogger(__name__) + + +def get_transport(request: Request) -> Transport: + return cast(Transport, request.state.transport) + + +def get_heartbeat_registry(request: Request) -> HeartbeatRegistry: + return cast(HeartbeatRegistry, request.state.heartbeat_registry) + + +def _now() -> datetime: + return datetime.now() + + +def build_adapters( + models: list[ModelInput], + conditions: Conditions, + all_adapters: AdapterSet, +) -> list[BaseAdapter]: + """Instantiate and validate the adapter chain for a set of model inputs. + + Raises: + HTTPException: 422 if a model is unknown or its parameters are invalid. + """ + adapters: list[BaseAdapter] = [] + for model in models: + adapter_class = all_adapters.get(model.model_id) + if adapter_class is None: + raise HTTPException( + status_code=422, + detail=f"Unknown model '{model.model_id}'", + ) + try: + adapter = adapter_class( + parameters=model.parameters, + conditions=conditions, + ) + adapters.append(adapter) + except InputError as exc: + raise HTTPException(status_code=422, detail=exc.detail) + except ValidationError as exc: + detail = defaultdict(list) + for err in exc.errors(): + for loc in err["loc"]: + detail[loc].append(err["msg"]) + + raise HTTPException(status_code=422, detail=dict(detail)) + except ValueError as exc: + raise HTTPException(status_code=422, detail=exc.args) + return adapters + + +def build_model_input_rows(models: list[ModelInput]) -> list[db.ModelInput]: + """Build the chained ``db.ModelInput`` rows for a simulation.""" + rows: list[db.ModelInput] = [] + previous_model_input_id: UUID | None = None + for model in models: + model_input_id = uuid4() + rows.append( + db.ModelInput( + id=model_input_id, + previous_model_input_id=previous_model_input_id, + model_id=model.model_id, + parameters=model.parameters, + ) + ) + previous_model_input_id = model_input_id + return rows + + +def order_chain( + rows: list[tuple[db.ModelInput, db.ModelResult | None]], +) -> list[tuple[db.ModelInput, db.ModelResult | None]]: + """Order ``(model_input, result)`` rows following the pipeline chain.""" + mapping: dict[UUID | None, UUID] = {} + rows_by_id: dict[UUID, tuple[db.ModelInput, db.ModelResult | None]] = {} + for model_input, result in rows: + mapping[model_input.previous_model_input_id] = model_input.id + rows_by_id[model_input.id] = (model_input, result) + + ordered: list[tuple[db.ModelInput, db.ModelResult | None]] = [] + current_id: UUID | None = mapping.get(None) + while current_id in rows_by_id: + assert current_id is not None + ordered.append(rows_by_id[current_id]) + current_id = mapping.get(current_id) + return ordered + + +def query_chain_rows( + session: Session, simulation_id: UUID +) -> list[tuple[db.ModelInput, db.ModelResult | None]]: + q = ( + select(db.ModelInput, db.ModelResult) + .where(db.ModelInput.simulation_id == simulation_id) + .outerjoin(db.ModelResult) + ) + return [(row[0], row[1]) for row in session.execute(q).fetchall()] + + +def _phases_to_concentrations(phases: list[Phase]) -> dict[str, int | float]: + merged: dict[str, int | float] = {} + for phase in phases: + if phase.kind == "co2-rich": + merged.update(phase.concentrations) + return merged + + +def build_simulation_result( + session: Session, + simulation_id: UUID, + registry: HeartbeatRegistry | None = None, +) -> SimulationResult: + db_simulation = session.get_one(db.Simulation, simulation_id) + + model_inputs: list[ModelInput] = [] + results: list[ModelResult] = [] + pending = False + processing = False + now = _now() + previous_result_created_at: datetime | None = None + + for model_input, result in order_chain(query_chain_rows(session, simulation_id)): + model_inputs.append( + ModelInput( + model_id=model_input.model_id, + parameters=model_input.parameters, + ) + ) + + if not result: + if pending: + continue + pending = True + if ( + registry is not None + and registry.job_status(str(model_input.id), now=now) == "processing" + ): + processing = True + continue + pending_since = previous_result_created_at or model_input.created_at + if now - pending_since >= timedelta( + minutes=SETTINGS.model_input_timeout_minutes + ): + result = db.ModelResult( + model_input_id=model_input.id, + phases=[], + panels=[], + error=f"Model {model_input.model_id} timed out", + ) + session.add(result) + try: + session.commit() + except IntegrityError: + session.rollback() + result = session.scalar( + select(db.ModelResult).where( + db.ModelResult.model_input_id == model_input.id + ) + ) + assert result is not None + logger.error( + "Simulation %s failed: %s", + simulation_id, + result.error, + ) + return SimulationResult( + status="error", + input=Simulation( + concentrations=_phases_to_concentrations( + [Phase(**p) for p in db_simulation.phases] + ), + conditions=Conditions(**(db_simulation.conditions or {})), + models=model_inputs, + ), + results=results, + error=result.error, + ) + continue + + previous_result_created_at = result.created_at + if result.error is not None: + logger.error("Simulation %s failed: %s", simulation_id, result.error) + return SimulationResult( + status="error", + input=Simulation( + concentrations=_phases_to_concentrations( + [Phase(**p) for p in db_simulation.phases] + ), + conditions=Conditions(**(db_simulation.conditions or {})), + models=model_inputs, + ), + results=results, + error=result.error, + ) + + results.append( + ModelResult( + phases=[Phase(**p) for p in result.phases], + panels=result.panels, + ) + ) + + simulation_input = Simulation( + concentrations=_phases_to_concentrations( + [Phase(**p) for p in db_simulation.phases] + ), + conditions=Conditions(**(db_simulation.conditions or {})), + models=model_inputs, + ) + + if pending: + return SimulationResult( + status="processing" if processing else "pending", + input=simulation_input, + results=results, + ) + + return SimulationResult( + status="done", + input=simulation_input, + results=[ + ModelResult( + phases=result.phases, + panels=[ + TypeAdapter(AnyPanel).validate_python(panel) + for panel in result.panels + ], + ) + for result in results + if result is not None + ], + ) diff --git a/backend/src/acidwatch_api/routes/grid_simulations.py b/backend/src/acidwatch_api/routes/grid_simulations.py index 70608f79..cddab623 100644 --- a/backend/src/acidwatch_api/routes/grid_simulations.py +++ b/backend/src/acidwatch_api/routes/grid_simulations.py @@ -19,7 +19,7 @@ GridSimulationResult, SimulationResult, ) -from acidwatch_api.routes.models import ( +from acidwatch_api.routes._helpers import ( build_adapters, build_model_input_rows, build_simulation_result, diff --git a/backend/src/acidwatch_api/routes/models.py b/backend/src/acidwatch_api/routes/models.py index 44f123e7..993fe1bb 100644 --- a/backend/src/acidwatch_api/routes/models.py +++ b/backend/src/acidwatch_api/routes/models.py @@ -1,144 +1,23 @@ from __future__ import annotations -import logging -from collections import defaultdict -from datetime import datetime, timedelta -from typing import Annotated, cast -from uuid import UUID, uuid4 +from typing import Annotated -from fastapi import APIRouter, Depends, HTTPException, Request -from pydantic import TypeAdapter, ValidationError -from sqlalchemy import select -from sqlalchemy.exc import IntegrityError -from sqlalchemy.orm import Session +from fastapi import APIRouter, Depends -import acidwatch_api.database as db -from acidwatch_api.authentication import OptionalCurrentUser from acidwatch_api.broker.heartbeat import HeartbeatRegistry -from acidwatch_api.database import GetDB -from acidwatch_api.settings import SETTINGS -from acidwatch_messaging import AdapterJob, Transport, job_queue_name from acidwatch_models import ( AdapterSet, - BaseAdapter, - InputError, get_adapters, get_parameters_schema, ) -from acidwatch_models.datamodel import ( - AnyPanel, - Conditions, - ModelInfo, - ModelInput, - ModelResult, - Phase, - Simulation, - SimulationResult, +from acidwatch_models.datamodel import ModelInfo +from acidwatch_api.routes import _helpers +from acidwatch_api.routes._helpers import ( + get_heartbeat_registry, ) - router = APIRouter() -logger = logging.getLogger(__name__) - - -def get_transport(request: Request) -> Transport: - return cast(Transport, request.state.transport) - - -def get_heartbeat_registry(request: Request) -> HeartbeatRegistry: - return cast(HeartbeatRegistry, request.state.heartbeat_registry) - - -def _now() -> datetime: - return datetime.now() - - -def build_adapters( - models: list[ModelInput], - conditions: Conditions, - all_adapters: AdapterSet, -) -> list[BaseAdapter]: - """Instantiate and validate the adapter chain for a set of model inputs. - - Raises: - HTTPException: 422 if a model is unknown or its parameters are invalid. - """ - adapters: list[BaseAdapter] = [] - for model in models: - adapter_class = all_adapters.get(model.model_id) - if adapter_class is None: - raise HTTPException( - status_code=422, - detail=f"Unknown model '{model.model_id}'", - ) - try: - adapter = adapter_class( - parameters=model.parameters, - conditions=conditions, - ) - adapters.append(adapter) - except InputError as exc: - raise HTTPException(status_code=422, detail=exc.detail) - except ValidationError as exc: - detail = defaultdict(list) - for err in exc.errors(): - for loc in err["loc"]: - detail[loc].append(err["msg"]) - - raise HTTPException(status_code=422, detail=dict(detail)) - except ValueError as exc: - raise HTTPException(status_code=422, detail=exc.args) - return adapters - - -def build_model_input_rows(models: list[ModelInput]) -> list[db.ModelInput]: - """Build the chained ``db.ModelInput`` rows for a simulation.""" - rows: list[db.ModelInput] = [] - previous_model_input_id: UUID | None = None - for model in models: - model_input_id = uuid4() - rows.append( - db.ModelInput( - id=model_input_id, - previous_model_input_id=previous_model_input_id, - model_id=model.model_id, - parameters=model.parameters, - ) - ) - previous_model_input_id = model_input_id - return rows - - -def order_chain( - rows: list[tuple[db.ModelInput, db.ModelResult | None]], -) -> list[tuple[db.ModelInput, db.ModelResult | None]]: - """Order ``(model_input, result)`` rows following the pipeline chain.""" - mapping: dict[UUID | None, UUID] = {} - rows_by_id: dict[UUID, tuple[db.ModelInput, db.ModelResult | None]] = {} - for model_input, result in rows: - mapping[model_input.previous_model_input_id] = model_input.id - rows_by_id[model_input.id] = (model_input, result) - - ordered: list[tuple[db.ModelInput, db.ModelResult | None]] = [] - current_id: UUID | None = mapping.get(None) - while current_id in rows_by_id: - assert current_id is not None - ordered.append(rows_by_id[current_id]) - current_id = mapping.get(current_id) - return ordered - - -def query_chain_rows( - session: Session, simulation_id: UUID -) -> list[tuple[db.ModelInput, db.ModelResult | None]]: - q = ( - select(db.ModelInput, db.ModelResult) - .where(db.ModelInput.simulation_id == simulation_id) - .outerjoin(db.ModelResult) - ) - return [(row[0], row[1]) for row in session.execute(q).fetchall()] - @router.get("/models") def get_models( @@ -166,198 +45,8 @@ def get_models_status( adapters: Annotated[AdapterSet, Depends(get_adapters)], registry: Annotated[HeartbeatRegistry, Depends(get_heartbeat_registry)], ) -> dict[str, dict[str, str]]: - now = _now() + now = _helpers._now() return { model_id: {"status": registry.status(model_id, now=now)} for model_id in adapters } - - -def _phases_to_concentrations(phases: list[Phase]) -> dict[str, int | float]: - merged: dict[str, int | float] = {} - for phase in phases: - if phase.kind == "co2-rich": - merged.update(phase.concentrations) - return merged - - -def build_simulation_result( - session: Session, - simulation_id: UUID, - registry: HeartbeatRegistry | None = None, -) -> SimulationResult: - db_simulation = session.get_one(db.Simulation, simulation_id) - - model_inputs: list[ModelInput] = [] - results: list[ModelResult] = [] - pending = False - processing = False - now = _now() - previous_result_created_at: datetime | None = None - - for model_input, result in order_chain(query_chain_rows(session, simulation_id)): - model_inputs.append( - ModelInput( - model_id=model_input.model_id, - parameters=model_input.parameters, - ) - ) - - if not result: - if pending: - continue - pending = True - if ( - registry is not None - and registry.job_status(str(model_input.id), now=now) == "processing" - ): - processing = True - continue - pending_since = previous_result_created_at or model_input.created_at - if now - pending_since >= timedelta( - minutes=SETTINGS.model_input_timeout_minutes - ): - result = db.ModelResult( - model_input_id=model_input.id, - phases=[], - panels=[], - error=f"Model {model_input.model_id} timed out", - ) - session.add(result) - try: - session.commit() - except IntegrityError: - session.rollback() - result = session.scalar( - select(db.ModelResult).where( - db.ModelResult.model_input_id == model_input.id - ) - ) - assert result is not None - logger.error( - "Simulation %s failed: %s", - simulation_id, - result.error, - ) - return SimulationResult( - status="error", - input=Simulation( - concentrations=_phases_to_concentrations( - [Phase(**p) for p in db_simulation.phases] - ), - conditions=Conditions(**(db_simulation.conditions or {})), - models=model_inputs, - ), - results=results, - error=result.error, - ) - continue - - previous_result_created_at = result.created_at - if result.error is not None: - logger.error("Simulation %s failed: %s", simulation_id, result.error) - return SimulationResult( - status="error", - input=Simulation( - concentrations=_phases_to_concentrations( - [Phase(**p) for p in db_simulation.phases] - ), - conditions=Conditions(**(db_simulation.conditions or {})), - models=model_inputs, - ), - results=results, - error=result.error, - ) - - results.append( - ModelResult( - phases=[Phase(**p) for p in result.phases], - panels=result.panels, - ) - ) - - simulation_input = Simulation( - concentrations=_phases_to_concentrations( - [Phase(**p) for p in db_simulation.phases] - ), - conditions=Conditions(**(db_simulation.conditions or {})), - models=model_inputs, - ) - - if pending: - return SimulationResult( - status="processing" if processing else "pending", - input=simulation_input, - results=results, - ) - - return SimulationResult( - status="done", - input=simulation_input, - results=[ - ModelResult( - phases=result.phases, - panels=[ - TypeAdapter(AnyPanel).validate_python(panel) - for panel in result.panels - ], - ) - for result in results - if result is not None - ], - ) - - -@router.get("/simulations/{simulation_id}/result") -def get_result_for_simulation( - simulation_id: UUID, - session: GetDB, - registry: Annotated[HeartbeatRegistry, Depends(get_heartbeat_registry)], -) -> SimulationResult: - return build_simulation_result(session, simulation_id, registry) - - -@router.post("/simulations") -async def run_simulation( - create_simulation: Simulation, - user: OptionalCurrentUser, - session: GetDB, - all_adapters: Annotated[AdapterSet, Depends(get_adapters)], - transport: Annotated[Transport, Depends(get_transport)], -) -> UUID: - adapters = build_adapters( - create_simulation.models, - create_simulation.conditions, - all_adapters, - ) - - concentrations = create_simulation.concentrations - try: - adapters[0].validate_concentrations(concentrations) - except InputError as exc: - raise HTTPException(status_code=422, detail=exc.detail) - - model_inputs = build_model_input_rows(create_simulation.models) - - simulation = db.Simulation( - owner_id=UUID(user.id) if user else None, - phases=[p.model_dump() for p in create_simulation.phases], - conditions=create_simulation.conditions.model_dump(), - model_inputs=model_inputs, - ) - session.add(simulation) - session.commit() - - first_input = simulation.model_inputs[0] - await transport.publish( - job_queue_name(first_input.model_id), - AdapterJob( - model_input_id=first_input.id, - model_id=first_input.model_id, - concentrations=concentrations, - parameters=first_input.parameters, - conditions=create_simulation.conditions, - ), - ) - - return simulation.id diff --git a/backend/src/acidwatch_api/routes/simulations.py b/backend/src/acidwatch_api/routes/simulations.py new file mode 100644 index 00000000..3af0a9f3 --- /dev/null +++ b/backend/src/acidwatch_api/routes/simulations.py @@ -0,0 +1,81 @@ +from __future__ import annotations + +from typing import Annotated +from uuid import UUID + +from fastapi import APIRouter, Depends, HTTPException + +import acidwatch_api.database as db +from acidwatch_api.authentication import OptionalCurrentUser +from acidwatch_api.broker.heartbeat import HeartbeatRegistry +from acidwatch_api.database import GetDB +from acidwatch_messaging import AdapterJob, Transport, job_queue_name +from acidwatch_models import AdapterSet, InputError, get_adapters +from acidwatch_models.datamodel import ( + Simulation, + SimulationResult, +) +from acidwatch_api.routes._helpers import ( + build_adapters, + build_model_input_rows, + build_simulation_result, + get_heartbeat_registry, + get_transport, +) + +router = APIRouter() + + +@router.get("/simulations/{simulation_id}/result") +def get_result_for_simulation( + simulation_id: UUID, + session: GetDB, + registry: Annotated[HeartbeatRegistry, Depends(get_heartbeat_registry)], +) -> SimulationResult: + return build_simulation_result(session, simulation_id, registry) + + +@router.post("/simulations") +async def run_simulation( + create_simulation: Simulation, + user: OptionalCurrentUser, + session: GetDB, + all_adapters: Annotated[AdapterSet, Depends(get_adapters)], + transport: Annotated[Transport, Depends(get_transport)], +) -> UUID: + adapters = build_adapters( + create_simulation.models, + create_simulation.conditions, + all_adapters, + ) + + concentrations = create_simulation.concentrations + try: + adapters[0].validate_concentrations(concentrations) + except InputError as exc: + raise HTTPException(status_code=422, detail=exc.detail) + + model_inputs = build_model_input_rows(create_simulation.models) + + simulation = db.Simulation( + owner_id=UUID(user.id) if user else None, + phases=[p.model_dump() for p in create_simulation.phases], + conditions=create_simulation.conditions.model_dump(), + model_inputs=model_inputs, + ) + session.add(simulation) + session.commit() + + first_input = simulation.model_inputs[0] + await transport.publish( + job_queue_name(first_input.model_id), + AdapterJob( + model_input_id=first_input.id, + model_id=first_input.model_id, + concentrations=concentrations, + parameters=first_input.parameters, + conditions=create_simulation.conditions, + ), + ) + + return simulation.id diff --git a/backend/tests/broker/integration/test_pipeline_execution.py b/backend/tests/broker/integration/test_pipeline_execution.py index 7706bf5a..4cd61061 100644 --- a/backend/tests/broker/integration/test_pipeline_execution.py +++ b/backend/tests/broker/integration/test_pipeline_execution.py @@ -24,7 +24,7 @@ wait_until, ) import acidwatch_api.database as db -from acidwatch_api.routes.models import build_simulation_result +from acidwatch_api.routes._helpers import build_simulation_result pytestmark = [requires_docker, pytest.mark.asyncio] diff --git a/backend/tests/broker/unit/test_grid_dispatch.py b/backend/tests/broker/unit/test_grid_dispatch.py index 0093d879..11be58b6 100644 --- a/backend/tests/broker/unit/test_grid_dispatch.py +++ b/backend/tests/broker/unit/test_grid_dispatch.py @@ -3,7 +3,7 @@ from sqlalchemy import func, select import acidwatch_api.database as db -from acidwatch_api.routes.models import get_transport +from acidwatch_api.routes._helpers import get_transport from acidwatch_messaging import AdapterJob, job_queue_name from acidwatch_models import BaseAdapter, get_adapters diff --git a/backend/tests/broker/unit/test_models_status_endpoint.py b/backend/tests/broker/unit/test_models_status_endpoint.py index e381361e..9476fcba 100644 --- a/backend/tests/broker/unit/test_models_status_endpoint.py +++ b/backend/tests/broker/unit/test_models_status_endpoint.py @@ -2,9 +2,9 @@ import pytest -import acidwatch_api.routes.models as models_route +import acidwatch_api.routes._helpers as helpers_route from acidwatch_api.broker.heartbeat import HeartbeatRegistry -from acidwatch_api.routes.models import get_heartbeat_registry +from acidwatch_api.routes._helpers import get_heartbeat_registry from acidwatch_models import get_adapters @@ -40,7 +40,7 @@ def test_models_status_reports_warm_and_cold_models(client, monkeypatch, frozen_ registry, {"model_a": object, "model_b": object}, ) - monkeypatch.setattr(models_route, "_now", lambda: frozen_now) + monkeypatch.setattr(helpers_route, "_now", lambda: frozen_now) response = client.get("/models/status") response.raise_for_status() @@ -65,7 +65,7 @@ def test_models_status_does_not_expose_active_job_id(client, monkeypatch, frozen registry, {"model_a": object}, ) - monkeypatch.setattr(models_route, "_now", lambda: frozen_now) + monkeypatch.setattr(helpers_route, "_now", lambda: frozen_now) response = client.get("/models/status") response.raise_for_status() @@ -80,7 +80,7 @@ def test_models_status_is_cold_before_first_heartbeat(client, monkeypatch, froze HeartbeatRegistry(timeout=timedelta(seconds=60)), {"model_a": object, "model_b": object}, ) - monkeypatch.setattr(models_route, "_now", lambda: frozen_now) + monkeypatch.setattr(helpers_route, "_now", lambda: frozen_now) response = client.get("/models/status") response.raise_for_status() diff --git a/backend/tests/broker/unit/test_processing_status.py b/backend/tests/broker/unit/test_processing_status.py index 36acd1a8..05b7811f 100644 --- a/backend/tests/broker/unit/test_processing_status.py +++ b/backend/tests/broker/unit/test_processing_status.py @@ -1,9 +1,9 @@ from datetime import datetime, timedelta import acidwatch_api.database as db -import acidwatch_api.routes.models as models_route +import acidwatch_api.routes._helpers as helpers_route from acidwatch_api.broker.heartbeat import HeartbeatRegistry -from acidwatch_api.routes.models import get_heartbeat_registry +from acidwatch_api.routes._helpers import get_heartbeat_registry def test_active_model_input_is_reported_as_processing(client, sql_session, monkeypatch): @@ -36,7 +36,7 @@ def test_active_model_input_is_reported_as_processing(client, sql_session, monke get_heartbeat_registry, lambda: registry, ) - monkeypatch.setattr(models_route, "_now", lambda: now) + monkeypatch.setattr(helpers_route, "_now", lambda: now) response = client.get(f"/simulations/{simulation.id}/result") response.raise_for_status() diff --git a/backend/tests/broker/unit/test_simulation_dispatch.py b/backend/tests/broker/unit/test_simulation_dispatch.py index a0ebe116..29f73470 100644 --- a/backend/tests/broker/unit/test_simulation_dispatch.py +++ b/backend/tests/broker/unit/test_simulation_dispatch.py @@ -2,7 +2,8 @@ from acidwatch_messaging import AdapterJob, job_queue_name import acidwatch_models.base as base -from acidwatch_api.routes.models import get_adapters, get_transport +from acidwatch_api.routes._helpers import get_transport +from acidwatch_models import get_adapters import acidwatch_api.database as db diff --git a/backend/tests/broker/unit/test_timeout_detection.py b/backend/tests/broker/unit/test_timeout_detection.py index 88b1dbdf..b972ddbd 100644 --- a/backend/tests/broker/unit/test_timeout_detection.py +++ b/backend/tests/broker/unit/test_timeout_detection.py @@ -6,9 +6,9 @@ import pytest import acidwatch_api.database as db -import acidwatch_api.routes.models as models_route +import acidwatch_api.routes._helpers as helpers_route from acidwatch_api.broker.heartbeat import HeartbeatRegistry -from acidwatch_api.routes.models import get_heartbeat_registry +from acidwatch_api.routes._helpers import get_heartbeat_registry from acidwatch_api.settings import SETTINGS @@ -70,7 +70,7 @@ def test_active_model_input_is_processing_and_does_not_time_out( get_heartbeat_registry, lambda: registry, ) - monkeypatch.setattr(models_route, "_now", lambda: now) + monkeypatch.setattr(helpers_route, "_now", lambda: now) response = client.get(f"/simulations/{simulation.id}/result") response.raise_for_status() @@ -124,7 +124,7 @@ def test_active_first_model_does_not_time_out_undispatched_second_model( get_heartbeat_registry, lambda: registry, ) - monkeypatch.setattr(models_route, "_now", lambda: now) + monkeypatch.setattr(helpers_route, "_now", lambda: now) response = client.get(f"/simulations/{simulation.id}/result") response.raise_for_status() diff --git a/backend/tests/models/test_all_adapters.py b/backend/tests/models/test_all_adapters.py index cf4d411d..1f9368ca 100644 --- a/backend/tests/models/test_all_adapters.py +++ b/backend/tests/models/test_all_adapters.py @@ -2,7 +2,7 @@ import pytest import re from acidwatch_models import BaseAdapter -from acidwatch_api.routes.models import get_adapters +from acidwatch_models import get_adapters ATOM_PATTERN = ( diff --git a/backend/tests/test_grid_simulations_endpoint.py b/backend/tests/test_grid_simulations_endpoint.py index 92afe144..03d351d4 100644 --- a/backend/tests/test_grid_simulations_endpoint.py +++ b/backend/tests/test_grid_simulations_endpoint.py @@ -8,7 +8,7 @@ from acidwatch_api.app import fastapi_app from acidwatch_api.authentication import authenticated_user_claims from acidwatch_api.broker.heartbeat import HeartbeatRegistry -from acidwatch_api.routes.models import get_heartbeat_registry +from acidwatch_api.routes._helpers import get_heartbeat_registry from acidwatch_messaging import AdapterJob, job_queue_name import acidwatch_models.base as base from acidwatch_models.datamodel import Phase diff --git a/backend/tests/test_models_endpoints.py b/backend/tests/test_models_endpoints.py index a3ca69d4..f23d4d4d 100644 --- a/backend/tests/test_models_endpoints.py +++ b/backend/tests/test_models_endpoints.py @@ -1,6 +1,6 @@ from enum import StrEnum -from acidwatch_api.routes.models import get_adapters +from acidwatch_models import get_adapters import pytest from fastapi.testclient import TestClient as _BaseTestClient from starlette.status import HTTP_422_UNPROCESSABLE_ENTITY