diff --git a/src/murfey/server/api/session_info.py b/src/murfey/server/api/session_info.py index f647d874e..8ef6c8b7d 100644 --- a/src/murfey/server/api/session_info.py +++ b/src/murfey/server/api/session_info.py @@ -1,7 +1,6 @@ from datetime import datetime from logging import getLogger from pathlib import Path -from typing import Dict, List, Optional import requests from fastapi import APIRouter, Depends, Request @@ -12,6 +11,7 @@ import murfey.server import murfey.server.api.websocket as ws import murfey.server.prometheus as prom +import murfey.util.db as MurfeyDB from murfey.server.api import templates from murfey.server.api.auth import ( MurfeyInstrumentNameFrontend as MurfeyInstrumentName, @@ -34,22 +34,6 @@ from murfey.server.murfey_db import murfey_db from murfey.util import sanitise from murfey.util.config import get_machine_config -from murfey.util.db import ( - ClassificationFeedbackParameters, - ClientEnvironment, - DataCollection, - DataCollectionGroup, - FoilHole, - GridSquare, - Movie, - ProcessingJob, - RsyncInstance, - Session as MurfeySession, - SessionProcessingParameters, - SPARelionParameters, - Tilt, - TiltSeries, -) from murfey.util.models import UpstreamFileRequestInfo, Visit logger = getLogger("murfey.server.api.session_info") @@ -89,7 +73,7 @@ def check_smartem_availability(instrument_name: str): return {"available": bool(machine_config.smartem_api_url)} -@router.get("/instruments/{instrument_name}/visits_raw", response_model=List[Visit]) +@router.get("/instruments/{instrument_name}/visits_raw", response_model=list[Visit]) def get_current_visits(instrument_name: MurfeyInstrumentName, db=ispyb_db): logger.debug( f"Received request to look up ongoing visits for {sanitise(instrument_name)}" @@ -130,36 +114,44 @@ def all_visit_info( ) -@router.get("/sessions/{session_id}/rsyncers", response_model=List[RsyncInstance]) +@router.get( + "/sessions/{session_id}/rsyncers", response_model=list[MurfeyDB.RsyncInstance] +) def get_rsyncers_for_client( session_id: MurfeySessionID, db: SQLModelSession = murfey_db ): rsync_instances = db.exec( - select(RsyncInstance).where(RsyncInstance.session_id == session_id) + select(MurfeyDB.RsyncInstance).where( + MurfeyDB.RsyncInstance.session_id == session_id + ) ) return rsync_instances.all() class SessionClients(BaseModel): - session: MurfeySession - clients: List[ClientEnvironment] + session: MurfeyDB.Session + clients: list[MurfeyDB.ClientEnvironment] @router.get("/sessions/{session_id}") async def get_session( session_id: MurfeySessionID, db: SQLModelSession = murfey_db ) -> SessionClients: - session = db.exec(select(MurfeySession).where(MurfeySession.id == session_id)).one() + session = db.exec( + select(MurfeyDB.Session).where(MurfeyDB.Session.id == session_id) + ).one() clients = db.exec( - select(ClientEnvironment).where(ClientEnvironment.session_id == session_id) + select(MurfeyDB.ClientEnvironment).where( + MurfeyDB.ClientEnvironment.session_id == session_id + ) ).all() return SessionClients(session=session, clients=clients) @router.get("/sessions") async def get_sessions(db: SQLModelSession = murfey_db): - sessions = db.exec(select(MurfeySession)).all() - clients = db.exec(select(ClientEnvironment)).all() + sessions = db.exec(select(MurfeyDB.Session)).all() + clients = db.exec(select(MurfeyDB.ClientEnvironment)).all() res = [] for sess in sessions: r = {"session": sess, "clients": []} @@ -170,32 +162,34 @@ async def get_sessions(db: SQLModelSession = murfey_db): return res -class VisitEndTime(BaseModel): - end_time: Optional[datetime] = None +class NewSessionInfo(BaseModel): + visit: str + name: str + end_time: datetime | None = None -@router.post("/instruments/{instrument_name}/visits/{visit}/sessions/{name}") +@router.post("/instruments/{instrument_name}/sessions/new") def create_session( instrument_name: MurfeyInstrumentName, - visit: str, - name: str, - visit_end_time: VisitEndTime, + session_info: NewSessionInfo, db: SQLModelSession = murfey_db, ) -> int: - s = MurfeySession( - name=name, - visit=visit, + session = MurfeyDB.Session( + name=sanitise(session_info.name), + visit=sanitise(session_info.visit), instrument_name=instrument_name, - visit_end_time=visit_end_time.end_time, + visit_end_time=session_info.end_time, ) - db.add(s) + db.add(session) db.commit() - sid = s.id + session_id = session.id - if visit_end_time.end_time: - prom.alert_end_time.labels(visit=visit).set(visit_end_time.end_time.timestamp()) + if session_info.end_time: + prom.alert_end_time.labels(visit=session_info.visit).set( + session_info.end_time.timestamp() + ) - return sid + return session_id @router.post("/sessions/{session_id}") @@ -205,7 +199,9 @@ def update_session( smartem_acquisition_uuid: str | None = None, db: SQLModelSession = murfey_db, ) -> None: - session = db.exec(select(MurfeySession).where(MurfeySession.id == session_id)).one() + session = db.exec( + select(MurfeyDB.Session).where(MurfeyDB.Session.id == session_id) + ).one() session.process = process session.smartem_acquisition_uuid = smartem_acquisition_uuid db.add(session) @@ -223,11 +219,11 @@ def get_sessions_with_visit( instrument_name: MurfeyInstrumentName, visit_name: str, db: SQLModelSession = murfey_db, -) -> List[MurfeySession]: +) -> list[MurfeyDB.Session]: sessions = db.exec( - select(MurfeySession) - .where(MurfeySession.instrument_name == instrument_name) - .where(MurfeySession.visit == visit_name) + select(MurfeyDB.Session) + .where(MurfeyDB.Session.instrument_name == instrument_name) + .where(MurfeyDB.Session.visit == visit_name) ).all() return sessions @@ -235,9 +231,11 @@ def get_sessions_with_visit( @router.get("/instruments/{instrument_name}/sessions") async def get_sessions_by_instrument_name( instrument_name: MurfeyInstrumentName, db: SQLModelSession = murfey_db -) -> List[MurfeySession]: +) -> list[MurfeyDB.Session]: sessions = db.exec( - select(MurfeySession).where(MurfeySession.instrument_name == instrument_name) + select(MurfeyDB.Session).where( + MurfeyDB.Session.instrument_name == instrument_name + ) ).all() return sessions @@ -245,9 +243,11 @@ async def get_sessions_by_instrument_name( @router.get("/sessions/{session_id}/data_collection_groups") def get_dc_groups( session_id: MurfeySessionID, db: SQLModelSession = murfey_db -) -> Dict[str, DataCollectionGroup]: +) -> dict[str, MurfeyDB.DataCollectionGroup]: data_collection_groups = db.exec( - select(DataCollectionGroup).where(DataCollectionGroup.session_id == session_id) + select(MurfeyDB.DataCollectionGroup).where( + MurfeyDB.DataCollectionGroup.session_id == session_id + ) ).all() return {dcg.tag: dcg for dcg in data_collection_groups} @@ -255,16 +255,16 @@ def get_dc_groups( @router.get("/sessions/{session_id}/data_collection_groups/{dcgid}/data_collections") def get_data_collections( session_id: MurfeySessionID, dcgid: int, db: SQLModelSession = murfey_db -) -> List[DataCollection]: +) -> list[MurfeyDB.DataCollection]: data_collections = db.exec( - select(DataCollection).where(DataCollection.dcg_id == dcgid) + select(MurfeyDB.DataCollection).where(MurfeyDB.DataCollection.dcg_id == dcgid) ).all() return data_collections @router.get("/clients") async def get_clients(db: SQLModelSession = murfey_db): - clients = db.exec(select(ClientEnvironment)).all() + clients = db.exec(select(MurfeyDB.ClientEnvironment)).all() return clients @@ -278,13 +278,15 @@ def update_current_gain_ref( new_gain_ref: CurrentGainRef, db: SQLModelSession = murfey_db, ): - session = db.exec(select(MurfeySession).where(MurfeySession.id == session_id)).one() + session = db.exec( + select(MurfeyDB.Session).where(MurfeyDB.Session.id == session_id) + ).one() session.current_gain_ref = new_gain_ref.path db.add(session) session_processing_parameters = db.exec( - select(SessionProcessingParameters).where( - SessionProcessingParameters.session_id == session_id + select(MurfeyDB.SessionProcessingParameters).where( + MurfeyDB.SessionProcessingParameters.session_id == session_id ) ).all() if session_processing_parameters: @@ -390,30 +392,32 @@ def delete_silences(instrument_name: MurfeyInstrumentName): class ProcessingDetails(BaseModel): - data_collection_group: DataCollectionGroup - data_collections: List[DataCollection] - processing_jobs: List[ProcessingJob] - relion_params: SPARelionParameters - feedback_params: ClassificationFeedbackParameters + data_collection_group: MurfeyDB.DataCollectionGroup + data_collections: list[MurfeyDB.DataCollection] + processing_jobs: list[MurfeyDB.ProcessingJob] + relion_params: MurfeyDB.SPARelionParameters + feedback_params: MurfeyDB.ClassificationFeedbackParameters @spa_router.get("/sessions/{session_id}/spa_processing_parameters") def get_spa_proc_param_details( session_id: MurfeySessionID, db: SQLModelSession = murfey_db -) -> Optional[List[ProcessingDetails]]: +) -> list[ProcessingDetails] | None: params = db.exec( select( - DataCollectionGroup, - DataCollection, - ProcessingJob, - SPARelionParameters, - ClassificationFeedbackParameters, + MurfeyDB.DataCollectionGroup, + MurfeyDB.DataCollection, + MurfeyDB.ProcessingJob, + MurfeyDB.SPARelionParameters, + MurfeyDB.ClassificationFeedbackParameters, + ) + .where(MurfeyDB.DataCollectionGroup.session_id == session_id) + .where(MurfeyDB.DataCollectionGroup.id == MurfeyDB.DataCollection.dcg_id) + .where(MurfeyDB.DataCollection.id == MurfeyDB.ProcessingJob.dc_id) + .where(MurfeyDB.SPARelionParameters.pj_id == MurfeyDB.ProcessingJob.id) + .where( + MurfeyDB.ClassificationFeedbackParameters.pj_id == MurfeyDB.ProcessingJob.id ) - .where(DataCollectionGroup.session_id == session_id) - .where(DataCollectionGroup.id == DataCollection.dcg_id) - .where(DataCollection.id == ProcessingJob.dc_id) - .where(SPARelionParameters.pj_id == ProcessingJob.id) - .where(ClassificationFeedbackParameters.pj_id == ProcessingJob.id) ).all() if not params: return None @@ -451,14 +455,19 @@ def get_number_of_movies_from_foil_hole( session_id: int, dcgid: int, gsid: int, fhid: int, db: SQLModelSession = murfey_db ) -> int: movies = db.exec( - select(Movie, FoilHole, GridSquare, DataCollectionGroup) - .where(Movie.foil_hole_id == FoilHole.id) - .where(FoilHole.name == fhid) - .where(FoilHole.grid_square_id == GridSquare.id) - .where(GridSquare.name == gsid) - .where(GridSquare.session_id == session_id) - .where(GridSquare.tag == DataCollectionGroup.tag) - .where(DataCollectionGroup.id == dcgid) + select( + MurfeyDB.Movie, + MurfeyDB.FoilHole, + MurfeyDB.GridSquare, + MurfeyDB.DataCollectionGroup, + ) + .where(MurfeyDB.Movie.foil_hole_id == MurfeyDB.FoilHole.id) + .where(MurfeyDB.FoilHole.name == fhid) + .where(MurfeyDB.FoilHole.grid_square_id == MurfeyDB.GridSquare.id) + .where(MurfeyDB.GridSquare.name == gsid) + .where(MurfeyDB.GridSquare.session_id == session_id) + .where(MurfeyDB.GridSquare.tag == MurfeyDB.DataCollectionGroup.tag) + .where(MurfeyDB.DataCollectionGroup.id == dcgid) ).all() return len(movies) @@ -471,7 +480,7 @@ def get_grid_squares(session_id: MurfeySessionID, db: SQLModelSession = murfey_d @spa_router.get("/sessions/{session_id}/data_collection_groups/{dcgid}/grid_squares") def get_grid_squares_from_dcg( session_id: MurfeySessionID, dcgid: int, db: SQLModelSession = murfey_db -) -> List[GridSquare]: +) -> list[MurfeyDB.GridSquare]: return _get_grid_squares_from_dcg(session_id, dcgid, db) @@ -480,14 +489,14 @@ def get_grid_squares_from_dcg( ) def get_foil_holes_from_grid_square( session_id: MurfeySessionID, dcgid: int, gsid: int, db: SQLModelSession = murfey_db -) -> List[FoilHole]: +) -> list[MurfeyDB.FoilHole]: return _get_foil_holes_from_grid_square(session_id, dcgid, gsid, db) @spa_router.get("/sessions/{session_id}/foil_hole/{fh_name}") def get_foil_hole( session_id: MurfeySessionID, fh_name: int, db: SQLModelSession = murfey_db -) -> Dict[str, int]: +) -> dict[str, int]: return _get_foil_hole(session_id, fh_name, db) @@ -501,14 +510,14 @@ def get_foil_hole( @tomo_router.get("/sessions/{session_id}/tilt_series/{tilt_series_tag}/tilts") def get_tilts( session_id: MurfeySessionID, tilt_series_tag: str, db: SQLModelSession = murfey_db -) -> Dict[str, List[str]]: +) -> dict[str, list[str]]: res = db.exec( - select(TiltSeries, Tilt) - .where(TiltSeries.tag == tilt_series_tag) - .where(TiltSeries.session_id == session_id) - .where(Tilt.tilt_series_id == TiltSeries.id) + select(MurfeyDB.TiltSeries, MurfeyDB.Tilt) + .where(MurfeyDB.TiltSeries.tag == tilt_series_tag) + .where(MurfeyDB.TiltSeries.session_id == session_id) + .where(MurfeyDB.Tilt.tilt_series_id == MurfeyDB.TiltSeries.id) ).all() - tilts: Dict[str, List[str]] = {} + tilts: dict[str, list[str]] = {} for el in res: if tilts.get(el[1].rsync_source): tilts[el[1].rsync_source].append(el[2].movie_path) diff --git a/src/murfey/util/route_manifest.yaml b/src/murfey/util/route_manifest.yaml index 558600dec..6a39e71c4 100644 --- a/src/murfey/util/route_manifest.yaml +++ b/src/murfey/util/route_manifest.yaml @@ -1116,13 +1116,9 @@ murfey.server.api.session_info.router: path_params: [] methods: - GET - - path: /session_info/instruments/{instrument_name}/visits/{visit}/sessions/{name} + - path: /session_info/instruments/{instrument_name}/sessions/new function: create_session path_params: - - name: visit - type: str - - name: name - type: str - name: instrument_name type: str methods: diff --git a/tests/conftest.py b/tests/conftest.py index 69cbd8363..181f6885f 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -336,13 +336,15 @@ def murfey_db_session_factory(murfey_db_engine): @pytest.fixture(scope="session") def seed_murfey_db(murfey_db_session_factory): # Populate Murfey database with initial values + visit_name = f"{ExampleVisit.proposal_code}{ExampleVisit.proposal_number}-{ExampleVisit.visit_number}" session: SQLModelSession = murfey_db_session_factory() _ = get_or_create_db_entry( session=session, table=MurfeySession, lookup_kwargs={ - "id": ExampleVisit.murfey_session_id, - "name": f"{ExampleVisit.proposal_code}{ExampleVisit.proposal_number}-{ExampleVisit.visit_number}", + "name": visit_name, + "visit": visit_name, + "instrument_name": ExampleVisit.instrument_name, }, ) session.close() diff --git a/tests/server/api/test_session_info.py b/tests/server/api/test_session_info.py index 1e3cf8e51..0ac42bb14 100644 --- a/tests/server/api/test_session_info.py +++ b/tests/server/api/test_session_info.py @@ -1,12 +1,110 @@ +from datetime import datetime from pathlib import Path -from typing import Any +from typing import Any, Callable from unittest.mock import MagicMock import pytest +from fastapi import APIRouter, FastAPI +from fastapi.testclient import TestClient from pytest_mock import MockerFixture +from sqlmodel import Session as SQLModelSession, select -from murfey.server.api.session_info import gather_upstream_files +import murfey.util.db as MurfeyDB +from murfey.server.api.auth import ( + validate_frontend_session_access, + validate_token, + validate_user_instrument_access, +) +from murfey.server.api.session_info import gather_upstream_files, router +from murfey.server.murfey_db import murfey_db_session +from murfey.util.api import url_path_for from murfey.util.models import UpstreamFileRequestInfo +from tests.conftest import ExampleVisit + +instrument_name = ExampleVisit.instrument_name + + +def set_up_test_backend_client( + router: APIRouter, + session_id: int | None = None, + instrument_name: str | None = None, + mock_db_session: Callable | None = None, +): + """ + Helper function to set up a test backend server whose response can be inspected + to check that the endpoint function works as expected + """ + # Set up the backend server + backend_app = FastAPI() + + # Override validation and database dependencies as needed + backend_app.dependency_overrides[validate_token] = lambda: None + if instrument_name: + backend_app.dependency_overrides[validate_user_instrument_access] = ( + lambda: instrument_name + ) + if session_id: + backend_app.dependency_overrides[validate_frontend_session_access] = ( + lambda: session_id + ) + if mock_db_session: + backend_app.dependency_overrides[murfey_db_session] = mock_db_session + + # Attach router, initiate object, and return it + backend_app.include_router(router) + return TestClient(backend_app) + + +@pytest.mark.parametrize( + "test_params", + ( # Visit name | Session name | End time + ("cm23456-7", "Some string", "2026-10-01T11:13:00"), + ("cm34567-8", "New\r\nvisit", None), + ("cm45678-9", "New\nvisit", None), + ("cm56789-10", "", None), + ), +) +def test_create_session_with_db( + test_params: tuple[str, str, str | None], murfey_db_session: SQLModelSession +): + visit_name, session_name, visit_end_time = test_params + + # Set up a mock Murfey database session + def mock_get_db_session(): + yield murfey_db_session + + # Set up the backend server + backend_server = set_up_test_backend_client( + router=router, + instrument_name=instrument_name, + mock_db_session=mock_get_db_session, + ) + # Construct the URL path to poke + backend_url_path = url_path_for( + "api.session_info.router", + "create_session", + instrument_name=instrument_name, + ) + # Poke the backend + response = backend_server.post( + backend_url_path, + json={ + "visit": visit_name, + "name": session_name, + "end_time": visit_end_time, + }, + ) + assert response.status_code == 200 + + # Check that the database insert happened correctly + murfey_session = murfey_db_session.exec( + select(MurfeyDB.Session).where(MurfeyDB.Session.visit == visit_name) + ).one() + assert murfey_session.name == session_name.replace("\r\n", "").replace("\n", "") + if visit_end_time is not None: + assert murfey_session.visit_end_time == datetime.fromisoformat(visit_end_time) + else: + assert murfey_session.visit_end_time is None @pytest.mark.parametrize( diff --git a/tests/server/api/test_workflow.py b/tests/server/api/test_workflow.py index 74adbc2d6..b72a7f396 100644 --- a/tests/server/api/test_workflow.py +++ b/tests/server/api/test_workflow.py @@ -48,7 +48,7 @@ def test_register_dc_group_new_dcg(mock_transport, murfey_db_session: Session): "atlas_y_stage_position": None, "atlas_width": None, "atlas_height": None, - "microscope": "", + "microscope": ExampleVisit.instrument_name, "proposal_code": ExampleVisit.proposal_code, "proposal_number": str(ExampleVisit.proposal_number), "visit_number": str(ExampleVisit.visit_number), @@ -275,7 +275,7 @@ def test_register_dc_group_new_dcg_old_atlas( "atlas_y_stage_position": None, "atlas_width": None, "atlas_height": None, - "microscope": "", + "microscope": ExampleVisit.instrument_name, "proposal_code": ExampleVisit.proposal_code, "proposal_number": str(ExampleVisit.proposal_number), "visit_number": str(ExampleVisit.visit_number), @@ -353,7 +353,7 @@ def test_register_dc_group_new_atlas_with_searchmaps( """ mock_transport.feedback_queue = "mock_feedback_queue" mock_machine_config.return_value = { - "": MachineConfig(acquisition_software=["tomo"]) + ExampleVisit.instrument_name: MachineConfig(acquisition_software=["tomo"]) } # Make sure dcg is present with an atlas id @@ -464,7 +464,9 @@ def test_register_dc_group_new_atlas_with_sxt_roi( by adding an atlas, using the same tag, and also update sxt rois """ mock_transport.feedback_queue = "mock_feedback_queue" - mock_machine_config.return_value = {"": MachineConfig(acquisition_software=["sxt"])} + mock_machine_config.return_value = { + ExampleVisit.instrument_name: MachineConfig(acquisition_software=["sxt"]) + } # Make sure dcg is present with an atlas id dcg = DataCollectionGroup( @@ -596,7 +598,7 @@ def test_register_dc_group_roi_update( """ mock_transport.feedback_queue = "mock_feedback_queue" mock_machine_config.return_value = { - "": MachineConfig(acquisition_software=["tomo"]) + ExampleVisit.instrument_name: MachineConfig(acquisition_software=["tomo"]) } # Make sure dcg is present with an atlas id