From 89bbb03628f6a5337d7c681c6c4c842af3c1ee55 Mon Sep 17 00:00:00 2001 From: Eu Pin Tien Date: Thu, 1 Oct 2026 10:18:50 +0100 Subject: [PATCH 01/13] Updated 'create_session' endpoint URL and pass visit and session names to it through JSON data payload instead --- src/murfey/server/api/session_info.py | 30 ++++++++++++++------------- src/murfey/util/route_manifest.yaml | 6 +----- 2 files changed, 17 insertions(+), 19 deletions(-) diff --git a/src/murfey/server/api/session_info.py b/src/murfey/server/api/session_info.py index f647d874e..b7b920f6b 100644 --- a/src/murfey/server/api/session_info.py +++ b/src/murfey/server/api/session_info.py @@ -170,32 +170,34 @@ async def get_sessions(db: SQLModelSession = murfey_db): return res -class VisitEndTime(BaseModel): +class NewSessionInfo(BaseModel): + visit: str + name: str end_time: Optional[datetime] = 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 = MurfeySession( + name=session_info.name, + visit=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}") 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: From bc8f824be7a8ea276a59de4185c482a4bed3bc60 Mon Sep 17 00:00:00 2001 From: Eu Pin Tien Date: Thu, 1 Oct 2026 10:23:55 +0100 Subject: [PATCH 02/13] Updated how Murfey's database tables are called --- src/murfey/server/api/session_info.py | 146 ++++++++++++++------------ 1 file changed, 77 insertions(+), 69 deletions(-) diff --git a/src/murfey/server/api/session_info.py b/src/murfey/server/api/session_info.py index b7b920f6b..00a6de5a0 100644 --- a/src/murfey/server/api/session_info.py +++ b/src/murfey/server/api/session_info.py @@ -12,6 +12,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 +35,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") @@ -130,36 +115,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": []} @@ -182,7 +175,7 @@ def create_session( session_info: NewSessionInfo, db: SQLModelSession = murfey_db, ) -> int: - session = MurfeySession( + session = MurfeyDB.Session( name=session_info.name, visit=session_info.visit, instrument_name=instrument_name, @@ -207,7 +200,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) @@ -225,11 +220,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 @@ -237,9 +232,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 @@ -247,9 +244,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} @@ -257,16 +256,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 @@ -280,13 +279,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: @@ -392,11 +393,11 @@ 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") @@ -405,17 +406,19 @@ def get_spa_proc_param_details( ) -> Optional[List[ProcessingDetails]]: 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 @@ -453,14 +456,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) @@ -473,7 +481,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) @@ -482,7 +490,7 @@ 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) @@ -505,10 +513,10 @@ def get_tilts( session_id: MurfeySessionID, tilt_series_tag: str, db: SQLModelSession = murfey_db ) -> 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]] = {} for el in res: From b00999d20712c6de95f77b806a461a46058b4f23 Mon Sep 17 00:00:00 2001 From: Eu Pin Tien Date: Thu, 1 Oct 2026 11:18:24 +0100 Subject: [PATCH 03/13] Added test for the 'create_session' function --- tests/server/api/test_session_info.py | 85 ++++++++++++++++++++++++++- 1 file changed, 83 insertions(+), 2 deletions(-) diff --git a/tests/server/api/test_session_info.py b/tests/server/api/test_session_info.py index 1e3cf8e51..efb0f12fc 100644 --- a/tests/server/api/test_session_info.py +++ b/tests/server/api/test_session_info.py @@ -1,12 +1,93 @@ +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 +visit_name = f"{ExampleVisit.proposal_code}{ExampleVisit.proposal_number}-{ExampleVisit.visit_number}" + + +def set_up_test_backend_client( + router: APIRouter, session_id: int, instrument_name: str, mock_db_session: Callable +): + """ + 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 + backend_app.dependency_overrides[validate_token] = lambda: None + backend_app.dependency_overrides[validate_user_instrument_access] = ( + lambda: instrument_name + ) + backend_app.dependency_overrides[validate_frontend_session_access] = ( + lambda: session_id + ) + backend_app.dependency_overrides[murfey_db_session] = mock_db_session + backend_app.include_router(router) + return TestClient(backend_app) + + +def test_create_session_with_db(murfey_db_session: SQLModelSession): + session_id = 10 + visit_end_time = "2026-10-01T11:13:00" + + # 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, + session_id=session_id, + 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, + data={ + "visit": visit_name, + "name": "Some string", + "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.id == session_id) + ).one_or_none() + assert murfey_session is not None + assert murfey_session.id == session_id + assert murfey_session.name == "Some string" + assert murfey_session.visit_end_time == datetime.fromisoformat(visit_end_time) @pytest.mark.parametrize( From 82f246f01848d5ab8dd88dc252a539cb01b9d48a Mon Sep 17 00:00:00 2001 From: Eu Pin Tien Date: Thu, 1 Oct 2026 11:23:49 +0100 Subject: [PATCH 04/13] Used wrong kwarg for passing data --- tests/server/api/test_session_info.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/server/api/test_session_info.py b/tests/server/api/test_session_info.py index efb0f12fc..e30f029d6 100644 --- a/tests/server/api/test_session_info.py +++ b/tests/server/api/test_session_info.py @@ -72,7 +72,7 @@ def mock_get_db_session(): # Poke the backend response = backend_server.post( backend_url_path, - data={ + json={ "visit": visit_name, "name": "Some string", "end_time": visit_end_time, From 2b4b5a6e9e4a4b88f92d60a0402efd67dc3d8f44 Mon Sep 17 00:00:00 2001 From: Eu Pin Tien Date: Thu, 1 Oct 2026 11:35:34 +0100 Subject: [PATCH 05/13] Updated type hints --- src/murfey/server/api/session_info.py | 33 +++++++++++++-------------- 1 file changed, 16 insertions(+), 17 deletions(-) diff --git a/src/murfey/server/api/session_info.py b/src/murfey/server/api/session_info.py index 00a6de5a0..499fd047a 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 @@ -74,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)}" @@ -116,7 +115,7 @@ def all_visit_info( @router.get( - "/sessions/{session_id}/rsyncers", response_model=List[MurfeyDB.RsyncInstance] + "/sessions/{session_id}/rsyncers", response_model=list[MurfeyDB.RsyncInstance] ) def get_rsyncers_for_client( session_id: MurfeySessionID, db: SQLModelSession = murfey_db @@ -131,7 +130,7 @@ def get_rsyncers_for_client( class SessionClients(BaseModel): session: MurfeyDB.Session - clients: List[MurfeyDB.ClientEnvironment] + clients: list[MurfeyDB.ClientEnvironment] @router.get("/sessions/{session_id}") @@ -166,7 +165,7 @@ async def get_sessions(db: SQLModelSession = murfey_db): class NewSessionInfo(BaseModel): visit: str name: str - end_time: Optional[datetime] = None + end_time: datetime | None = None @router.post("/instruments/{instrument_name}/sessions/new") @@ -220,7 +219,7 @@ def get_sessions_with_visit( instrument_name: MurfeyInstrumentName, visit_name: str, db: SQLModelSession = murfey_db, -) -> List[MurfeyDB.Session]: +) -> list[MurfeyDB.Session]: sessions = db.exec( select(MurfeyDB.Session) .where(MurfeyDB.Session.instrument_name == instrument_name) @@ -232,7 +231,7 @@ 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[MurfeyDB.Session]: +) -> list[MurfeyDB.Session]: sessions = db.exec( select(MurfeyDB.Session).where( MurfeyDB.Session.instrument_name == instrument_name @@ -244,7 +243,7 @@ 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, MurfeyDB.DataCollectionGroup]: +) -> dict[str, MurfeyDB.DataCollectionGroup]: data_collection_groups = db.exec( select(MurfeyDB.DataCollectionGroup).where( MurfeyDB.DataCollectionGroup.session_id == session_id @@ -256,7 +255,7 @@ 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[MurfeyDB.DataCollection]: +) -> list[MurfeyDB.DataCollection]: data_collections = db.exec( select(MurfeyDB.DataCollection).where(MurfeyDB.DataCollection.dcg_id == dcgid) ).all() @@ -394,8 +393,8 @@ def delete_silences(instrument_name: MurfeyInstrumentName): class ProcessingDetails(BaseModel): data_collection_group: MurfeyDB.DataCollectionGroup - data_collections: List[MurfeyDB.DataCollection] - processing_jobs: List[MurfeyDB.ProcessingJob] + data_collections: list[MurfeyDB.DataCollection] + processing_jobs: list[MurfeyDB.ProcessingJob] relion_params: MurfeyDB.SPARelionParameters feedback_params: MurfeyDB.ClassificationFeedbackParameters @@ -403,7 +402,7 @@ class ProcessingDetails(BaseModel): @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( MurfeyDB.DataCollectionGroup, @@ -481,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[MurfeyDB.GridSquare]: +) -> list[MurfeyDB.GridSquare]: return _get_grid_squares_from_dcg(session_id, dcgid, db) @@ -490,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[MurfeyDB.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) @@ -511,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(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) From 739536031c418422121833f037bc55bcfd7b3829 Mon Sep 17 00:00:00 2001 From: Eu Pin Tien Date: Thu, 1 Oct 2026 11:36:46 +0100 Subject: [PATCH 06/13] Updated 'set_up_test_backend_client' function so that dependency overrides are only implemented when needed by the router; used a different visit name for the test --- tests/server/api/test_session_info.py | 36 +++++++++++++++------------ 1 file changed, 20 insertions(+), 16 deletions(-) diff --git a/tests/server/api/test_session_info.py b/tests/server/api/test_session_info.py index e30f029d6..cb62ecdbe 100644 --- a/tests/server/api/test_session_info.py +++ b/tests/server/api/test_session_info.py @@ -22,11 +22,13 @@ from tests.conftest import ExampleVisit instrument_name = ExampleVisit.instrument_name -visit_name = f"{ExampleVisit.proposal_code}{ExampleVisit.proposal_number}-{ExampleVisit.visit_number}" def set_up_test_backend_client( - router: APIRouter, session_id: int, instrument_name: str, mock_db_session: Callable + 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 @@ -35,21 +37,26 @@ def set_up_test_backend_client( # Set up the backend server backend_app = FastAPI() - # Override validation and database dependencies + # Override validation and database dependencies as needed backend_app.dependency_overrides[validate_token] = lambda: None - backend_app.dependency_overrides[validate_user_instrument_access] = ( - lambda: instrument_name - ) - backend_app.dependency_overrides[validate_frontend_session_access] = ( - lambda: session_id - ) - backend_app.dependency_overrides[murfey_db_session] = mock_db_session + 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) def test_create_session_with_db(murfey_db_session: SQLModelSession): - session_id = 10 + visit_name = "cm23456-7" visit_end_time = "2026-10-01T11:13:00" # Set up a mock Murfey database session @@ -59,8 +66,6 @@ def mock_get_db_session(): # Set up the backend server backend_server = set_up_test_backend_client( router=router, - session_id=session_id, - instrument_name=instrument_name, mock_db_session=mock_get_db_session, ) # Construct the URL path to poke @@ -82,10 +87,9 @@ def mock_get_db_session(): # Check that the database insert happened correctly murfey_session = murfey_db_session.exec( - select(MurfeyDB.Session).where(MurfeyDB.Session.id == session_id) - ).one_or_none() + select(MurfeyDB.Session).where(MurfeyDB.Session.visit == visit_name) + ).one() assert murfey_session is not None - assert murfey_session.id == session_id assert murfey_session.name == "Some string" assert murfey_session.visit_end_time == datetime.fromisoformat(visit_end_time) From 29a96138e3338b6689486f0c06be2c7c1b5a7ff0 Mon Sep 17 00:00:00 2001 From: Eu Pin Tien Date: Thu, 1 Oct 2026 12:02:00 +0100 Subject: [PATCH 07/13] Instrument name validation required --- tests/server/api/test_session_info.py | 1 + 1 file changed, 1 insertion(+) diff --git a/tests/server/api/test_session_info.py b/tests/server/api/test_session_info.py index cb62ecdbe..7328ef4ca 100644 --- a/tests/server/api/test_session_info.py +++ b/tests/server/api/test_session_info.py @@ -66,6 +66,7 @@ def mock_get_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 From 1c902070922c88c6ede72ab4d94ecda8bb953c92 Mon Sep 17 00:00:00 2001 From: Eu Pin Tien Date: Thu, 1 Oct 2026 13:02:45 +0100 Subject: [PATCH 08/13] Let 'Session' table's 'id' key auto-increment --- src/murfey/util/db.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/murfey/util/db.py b/src/murfey/util/db.py index f49fed38b..9e4755df1 100644 --- a/src/murfey/util/db.py +++ b/src/murfey/util/db.py @@ -45,7 +45,7 @@ class RsyncInstance(SQLModel, table=True): # type: ignore class Session(SQLModel, table=True): # type: ignore - id: int = Field(primary_key=True) + id: int = Field(default=None, primary_key=True) name: str visit: str = Field(default="") started: bool = Field(default=False) From bcc1891440f444bd0825ff68c85a47108b62509d Mon Sep 17 00:00:00 2001 From: Eu Pin Tien Date: Thu, 1 Oct 2026 13:06:36 +0100 Subject: [PATCH 09/13] Revert "Let 'Session' table's 'id' key auto-increment" This reverts commit 1c902070922c88c6ede72ab4d94ecda8bb953c92. --- src/murfey/util/db.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/murfey/util/db.py b/src/murfey/util/db.py index 9e4755df1..f49fed38b 100644 --- a/src/murfey/util/db.py +++ b/src/murfey/util/db.py @@ -45,7 +45,7 @@ class RsyncInstance(SQLModel, table=True): # type: ignore class Session(SQLModel, table=True): # type: ignore - id: int = Field(default=None, primary_key=True) + id: int = Field(primary_key=True) name: str visit: str = Field(default="") started: bool = Field(default=False) From eef230470eddba4c3bf37fd22b5150be6c408da7 Mon Sep 17 00:00:00 2001 From: Eu Pin Tien Date: Thu, 1 Oct 2026 13:26:46 +0100 Subject: [PATCH 10/13] Updated initial seeding of Murfey database's 'Session' table --- tests/conftest.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) 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() From 8177a1135498702e1eb32f82bcee532a43dbc167 Mon Sep 17 00:00:00 2001 From: Eu Pin Tien Date: Thu, 1 Oct 2026 15:28:26 +0100 Subject: [PATCH 11/13] Fixed broken tests --- tests/server/api/test_workflow.py | 12 +++++++----- 1 file changed, 7 insertions(+), 5 deletions(-) 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 From c3c1242d357e559c718543e8b447b0b6c12fc11c Mon Sep 17 00:00:00 2001 From: Eu Pin Tien Date: Thu, 1 Oct 2026 15:50:13 +0100 Subject: [PATCH 12/13] Parametrise the tests --- tests/server/api/test_session_info.py | 24 +++++++++++++++++------- 1 file changed, 17 insertions(+), 7 deletions(-) diff --git a/tests/server/api/test_session_info.py b/tests/server/api/test_session_info.py index 7328ef4ca..d6c235f56 100644 --- a/tests/server/api/test_session_info.py +++ b/tests/server/api/test_session_info.py @@ -55,9 +55,17 @@ def set_up_test_backend_client( return TestClient(backend_app) -def test_create_session_with_db(murfey_db_session: SQLModelSession): - visit_name = "cm23456-7" - visit_end_time = "2026-10-01T11:13:00" +@pytest.mark.parametrize( + "test_params", + ( # Visit name | Session name | End time + ("cm23456-7", "Some string", "2026-10-01T11:13:00"), + ("cm34567-8", "", 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(): @@ -80,7 +88,7 @@ def mock_get_db_session(): backend_url_path, json={ "visit": visit_name, - "name": "Some string", + "name": session_name, "end_time": visit_end_time, }, ) @@ -90,9 +98,11 @@ def mock_get_db_session(): murfey_session = murfey_db_session.exec( select(MurfeyDB.Session).where(MurfeyDB.Session.visit == visit_name) ).one() - assert murfey_session is not None - assert murfey_session.name == "Some string" - assert murfey_session.visit_end_time == datetime.fromisoformat(visit_end_time) + assert murfey_session.name == session_name + 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( From 1f23e02fbe6b79f1672dfb2c4b00e4739355d89e Mon Sep 17 00:00:00 2001 From: Eu Pin Tien Date: Thu, 1 Oct 2026 17:27:29 +0100 Subject: [PATCH 13/13] Sanitise visit and session names when creating session --- src/murfey/server/api/session_info.py | 4 ++-- tests/server/api/test_session_info.py | 6 ++++-- 2 files changed, 6 insertions(+), 4 deletions(-) diff --git a/src/murfey/server/api/session_info.py b/src/murfey/server/api/session_info.py index 499fd047a..8ef6c8b7d 100644 --- a/src/murfey/server/api/session_info.py +++ b/src/murfey/server/api/session_info.py @@ -175,8 +175,8 @@ def create_session( db: SQLModelSession = murfey_db, ) -> int: session = MurfeyDB.Session( - name=session_info.name, - visit=session_info.visit, + name=sanitise(session_info.name), + visit=sanitise(session_info.visit), instrument_name=instrument_name, visit_end_time=session_info.end_time, ) diff --git a/tests/server/api/test_session_info.py b/tests/server/api/test_session_info.py index d6c235f56..0ac42bb14 100644 --- a/tests/server/api/test_session_info.py +++ b/tests/server/api/test_session_info.py @@ -59,7 +59,9 @@ def set_up_test_backend_client( "test_params", ( # Visit name | Session name | End time ("cm23456-7", "Some string", "2026-10-01T11:13:00"), - ("cm34567-8", "", None), + ("cm34567-8", "New\r\nvisit", None), + ("cm45678-9", "New\nvisit", None), + ("cm56789-10", "", None), ), ) def test_create_session_with_db( @@ -98,7 +100,7 @@ def mock_get_db_session(): murfey_session = murfey_db_session.exec( select(MurfeyDB.Session).where(MurfeyDB.Session.visit == visit_name) ).one() - assert murfey_session.name == session_name + 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: