From db913571f596717788f2e1c7cf825e2622efaa80 Mon Sep 17 00:00:00 2001 From: appleb_m Date: Mon, 9 Mar 2026 17:10:56 +0100 Subject: [PATCH] Sever and Daq worker: added new end poiints for the new panels and a screenshot button.Also added SSE for Face detection panel so it can update duering automation for comissioining purposes. --- src/aare/daq/server.py | 207 ++++++++++++++++++++++++++++- src/aare/gui/threads/daq_worker.py | 89 ++++++++++++- 2 files changed, 292 insertions(+), 4 deletions(-) diff --git a/src/aare/daq/server.py b/src/aare/daq/server.py index 84cc8591..e0a097ef 100644 --- a/src/aare/daq/server.py +++ b/src/aare/daq/server.py @@ -1,4 +1,5 @@ import asyncio +import hmac import io import os, time from typing import AsyncGenerator @@ -12,7 +13,7 @@ from aare.common.logger_config import setup_logger from aare.common.models import SampleShortInfo, DAQStatusModel, BeamlineStateEnum, BeamlineSettingsModel, \ SampleShortInfoList, SessionStatus, SampleCameraSettings, AutofocusSettings, TokenData, \ CryojetSettingsModel, SimpleScanParameters, CrystalSize, FluorescenceSpectrumParameterModel, \ - FluorescenceSpectrumOutputModel + FluorescenceSpectrumOutputModel, RecoveryActionRequest from aare.common.raster_grid import RasterGridRequest, CompletedRasterGrid from aare.common.rotation_scan import RotationScanRequest, CompletedRotationScan from aare.common.sample_geometry import SampleGeometryModel @@ -48,6 +49,67 @@ _all_pgroups_cache: dict[str, tuple[list[str], float]] = {} _ALL_PGROUPS_TTL_S = 60.0 # adjust TTL as needed logger = setup_logger("aareDAQ") +_face_detection_state: dict = { + "seq": 0, + "running": False, + "samples": [], + "height_fit": {}, + "area_fit": {}, +} +_face_detection_state_lock = asyncio.Lock() + +def _required_recovery_code() -> str: + code = os.getenv("AARE_RECOVERY_CODE", "").strip() + if not code: + logger.error("AARE_RECOVERY_CODE is not configured.") + raise HTTPException( + status_code=api_status.HTTP_503_SERVICE_UNAVAILABLE, + detail="Recovery confirmation code is not configured on the server.", + ) + return code + +def _validate_recovery_code(confirmation_code: str) -> None: + expected = _required_recovery_code() + provided = str(confirmation_code or "").strip() + if not hmac.compare_digest(provided, expected): + logger.warning("Invalid recovery confirmation code.") + raise HTTPException( + status_code=api_status.HTTP_403_FORBIDDEN, + detail="Invalid confirmation code.", + ) + +def _sample_is_mounted() -> bool: + try: + return daq.sample is not None + except Exception: + return False + +def _push_face_detection_progress(payload: dict) -> None: + global _face_detection_state + try: + next_seq = int(_face_detection_state.get("seq", 0)) + 1 + _face_detection_state = { + "seq": next_seq, + **payload, + } + except Exception as e: + logger.warning(f"Failed to update face detection progress: {e}") + + +async def face_detection_event_stream() -> AsyncGenerator[str, None]: + last_seq = -1 + try: + while True: + state = dict(_face_detection_state) + seq = int(state.get("seq", 0)) + if seq != last_seq: + last_seq = seq + yield f"data: {json.dumps(state, separators=(',', ':'))}\n\n" + await asyncio.sleep(0.15) + except asyncio.CancelledError: + return + +daq.set_face_detection_progress_callback(_push_face_detection_progress) @app.post("/token") async def login(form_data: OAuth2PasswordRequestForm = Depends()): @@ -382,6 +444,116 @@ async def beam_location(token: str = Depends(oauth2_scheme)): auth.check_jwt_staff(cfg, auth.parse_token(token)) daq.state = BeamlineStateEnum.BeamLocation +@app.post("/state/maintenance") +async def maintenance(token: str = Depends(oauth2_scheme)) -> str: + auth.check_jwt_staff(cfg, auth.parse_token(token)) + cfg.state = BeamlineStateEnum.Maintenance + logger.warning("Beamline state set to Maintenance via protected endpoint.") + return "OK" + +@app.post("/access/take_over_beamline") +async def take_over_beamline(payload: RecoveryActionRequest, token: str = Depends(oauth2_scheme)) -> str: + data = auth.parse_token(token) + auth.check_jwt_staff(cfg, data) + _validate_recovery_code(payload.confirmation_code) + auth.force_current_sesion(cfg, data) + logger.warning( + "Beamline session forcefully taken over.", + extra={"session": getattr(data, "session", None)}, + ) + return "OK" + +@app.post("/state/free_beamline") +async def free_beamline(payload: RecoveryActionRequest, token: str = Depends(oauth2_scheme)) -> str: + data = auth.parse_token(token) + auth.check_jwt_staff(cfg, data) + _validate_recovery_code(payload.confirmation_code) + cfg.state_busy = False + logger.warning( + "Beamline busy flag cleared via protected endpoint.", + extra={"session": getattr(data, "session", None)}, + ) + return "OK" + +@app.post("/recovery/recover_beamline") +async def recover_beamline(payload: RecoveryActionRequest, token: str = Depends(oauth2_scheme)) -> dict: + data = auth.parse_token(token) + auth.check_jwt_staff(cfg, data) + _validate_recovery_code(payload.confirmation_code) + + sample_mounted = _sample_is_mounted() + prev_state = cfg.state + prev_busy = cfg.state_busy + + auth.force_current_sesion(cfg, data) + cfg.state_busy = False + cfg.state = BeamlineStateEnum.Maintenance + + logger.warning( + "Beamline recovery action executed.", + extra={ + "session": getattr(data, "session", None), + "previous_state": getattr(prev_state, "name", str(prev_state)), + "previous_busy": prev_busy, + "sample_mounted": sample_mounted, + }, + ) + + return { + "ok": True, + "sample_mounted": sample_mounted, + "previous_state": getattr(prev_state, "name", str(prev_state)), + "previous_busy": prev_busy, + "new_state": BeamlineStateEnum.Maintenance.name, + } +@app.post("/recovery/unmount_sample") +async def recovery_unmount_sample(payload: RecoveryActionRequest, token: str = Depends(oauth2_scheme)) -> dict: + data = auth.parse_token(token) + auth.check_jwt_staff(cfg, data) + _validate_recovery_code(payload.confirmation_code) + + auth.force_current_sesion(cfg, data) + + if cfg.state_busy: + raise HTTPException( + status_code=api_status.HTTP_409_CONFLICT, + detail="Beamline is busy. Clear or recover the beamline before attempting recovery unmount.", + ) + + status = daq.status + if not getattr(status, "tell_connected", False): + raise HTTPException( + status_code=api_status.HTTP_409_CONFLICT, + detail=f"TELL is not connected: {getattr(status, 'tell_error', 'unknown error')}", + ) + + sample_mounted = _sample_is_mounted() + if not sample_mounted: + return { + "ok": True, + "sample_mounted": False, + "message": "No sample appears to be mounted.", + } + + prev_state = cfg.state + daq.recovery_unmount_sample() + + logger.warning( + "Recovery sample unmount executed.", + extra={ + "session": getattr(data, "session", None), + "previous_state": getattr(prev_state, "name", str(prev_state)), + }, + ) + + return { + "ok": True, + "sample_mounted": True, + "previous_state": getattr(prev_state, "name", str(prev_state)), + "new_state": getattr(cfg.state, "name", str(cfg.state)), + "message": "Recovery unmount completed.", + } + # Scans @app.post("/scan/raster") async def raster(val: RasterGridRequest, auto: bool = False, token: str = Depends(oauth2_scheme)) -> CompletedRasterGrid: @@ -436,14 +608,34 @@ async def alc_ml_bounding_box(token: str = Depends(oauth2_scheme)) -> RasterGrid auth.check_jwt_rw(cfg, auth.parse_token(token)) return daq.ml_bounding_box() - @app.post("/face_detection/run") async def face_detection_run(steps: int, step_size: int, token: str = Depends(oauth2_scheme)) -> dict: logger.debug(f"Face detection run: {steps} steps, {step_size} step size") auth.check_jwt_rw(cfg, auth.parse_token(token)) + _push_face_detection_progress({ + "running": True, + "status": "starting", + "samples": [], + "height_fit": {}, + "area_fit": {}, + }) result = daq.face_detection(steps=steps, step_size=step_size) return result +@app.get("/sse/face_detection") +async def sse_face_detection(token: str = Depends(oauth2_scheme)): + auth.check_jwt_ro(cfg, auth.parse_token(token)) + return StreamingResponse( + face_detection_event_stream(), + media_type="text/event-stream", + headers={ + "Cache-Control": "no-cache", + "Connection": "keep-alive", + "Access-Control-Allow-Origin": "*", + "Access-Control-Allow-Headers": "Cache-Control" + } + ) + # Access management @app.get("/access/pgroup") async def pgroup(token: str = Depends(oauth2_scheme)) -> str: @@ -627,6 +819,17 @@ async def sse_fluorimeter(token: str = Depends(oauth2_scheme)): } ) +@app.post("/samcam/send_screenshot_db") +async def send_screenshot_db( + filename: str | None = None, + message: str | None = None, + token: str = Depends(oauth2_scheme), +) -> str: + data = auth.parse_token(token) + auth.check_jwt_rw(cfg, data) + daq.send_screenshot_db(filename=filename, message=message) + return "OK" + LOGGING_CONFIG = { "version": 1, diff --git a/src/aare/gui/threads/daq_worker.py b/src/aare/gui/threads/daq_worker.py index 1cbccd4f..e688da5c 100644 --- a/src/aare/gui/threads/daq_worker.py +++ b/src/aare/gui/threads/daq_worker.py @@ -78,6 +78,11 @@ class DAQWorker(QObject): self._last_tell_connected: bool | None = None self._last_smargon_connected: bool | None = None + self._face_detection_stream_reply: QNetworkReply | None = None + + if self.__base_url is not None: + self.start_face_detection_stream() + def get_last_error_payload(self) -> dict: return dict(self._last_error_payload or {}) @@ -383,6 +388,34 @@ class DAQWorker(QObject): def beam_location(self): self.generic_post("state/beam_location") + @Slot(str) + def free_beamline(self, confirmation_code: str): + self.generic_post( + "state/free_beamline", + json.dumps({"confirmation_code": confirmation_code}), + ) + + @Slot(str) + def take_over_beamline(self, confirmation_code: str): + self.generic_post( + "access/take_over_beamline", + json.dumps({"confirmation_code": confirmation_code}), + ) + + @Slot(str) + def recover_beamline(self, confirmation_code: str): + self.generic_post( + "recovery/recover_beamline", + json.dumps({"confirmation_code": confirmation_code}), + ) + + @Slot(str) + def recovery_unmount_sample(self, confirmation_code: str): + self.generic_post( + "recovery/unmount_sample", + json.dumps({"confirmation_code": confirmation_code}), + ) + @Slot(str) def set_pgroup(self, val: str): if val == "": @@ -685,11 +718,43 @@ class DAQWorker(QObject): finally: reply.deleteLater() + def _read_face_detection_stream(self, reply: QNetworkReply): + try: + chunk = reply.readAll().data().decode("utf-8") + for line in chunk.splitlines(): + if line.startswith("data:"): + payload = line[5:].strip() + if payload: + data = json.loads(payload) + self.face_detection_result.emit(data) + except Exception as e: + logger.error(f"Face detection stream parse error: {e}") + + def _restart_face_detection_stream(self): + self._face_detection_stream_reply = None + if self.__base_url is not None: + QTimer.singleShot(1000, self.start_face_detection_stream) + + def start_face_detection_stream(self): + if self.__base_url is None: + return + + if self._face_detection_stream_reply is not None: + return + + request = QNetworkRequest(QUrl(f"{self.__base_url}/sse/face_detection")) + request.setRawHeader(b"Authorization", f"Bearer {self.__token}".encode("utf-8")) + reply = self.__net_manager.get(request) + reply.readyRead.connect(lambda: self._read_face_detection_stream(reply)) + reply.finished.connect(self._restart_face_detection_stream) + self._face_detection_stream_reply = reply + @Slot() - def face_detection(self, steps:int, step_size:int): + def face_detection(self, steps: int, step_size: int): if self.__base_url is None: logger.info(f"POST /face_detection/run?steps={steps}&step_size={step_size}") return + request = QNetworkRequest(QUrl(f"{self.__base_url}/face_detection/run?steps={steps}&step_size={step_size}")) request.setRawHeader(b"Authorization", f"Bearer {self.__token}".encode("utf-8")) request.setRawHeader(b"Content-Type", b"application/json") @@ -867,4 +932,24 @@ class DAQWorker(QObject): self.error_codes_loaded.emit(out) except Exception as e: logger.error(f"Failed to load error codes: {e}") - self.http_error.emit(str(e)) \ No newline at end of file + self.http_error.emit(str(e)) + + @Slot(str, str) + def send_screenshot_db(self, filename: str = "", message: str = ""): + if self.__base_url is None: + logger.info(f"POST /samcam/send_screenshot_db?filename={filename}&message={message}") + return + + from urllib.parse import quote + + query = [] + filename = filename.strip() + message = message.strip() + + if filename: + query.append(f"filename={quote(filename)}") + if message: + query.append(f"message={quote(message)}") + + suffix = f"?{'&'.join(query)}" if query else "" + self.generic_post(f"samcam/send_screenshot_db{suffix}") \ No newline at end of file