import asyncio import io from typing import Tuple, AsyncGenerator import cv2 import urllib3 import uvicorn from aaredaqlib.coordinate import SmargonCoordinate, Coordinate from aaredaqlib.models import SampleShortInfo, DAQStatusModel, BeamlineStateEnum, BeamlineSettingsModel, \ SampleShortInfoList, SessionStatus, SampleCameraSettings, AutofocusSettings, TokenData, \ CryojetSettingsModel from aaredaqlib.raster_grid import RasterGridRequest, CompletedRasterGrid from aaredaqlib.rotation_scan import RotationScanRequest, CompletedRotationScan from aaredaqlib.sample_geometry import SampleGeometryModel from fastapi import FastAPI, Depends from fastapi import HTTPException from fastapi import status as api_status from fastapi.security import OAuth2PasswordBearer, OAuth2PasswordRequestForm from starlette.responses import StreamingResponse from urllib3.exceptions import InsecureRequestWarning from aaredaq import auth from aaredaqlib.beamline import mx_beamline from aaredaq.config import BeamlineConfig from aaredaq.daq import AareDAQ from aaredaq.camera_stat_thread import start_image_stats_receiver, stop_image_stats_receiver app = FastAPI() # OAuth2 setup oauth2_scheme = OAuth2PasswordBearer(tokenUrl="token") bl = mx_beamline() cfg = BeamlineConfig(bl) daq = AareDAQ(cfg, bl) @app.post("/token") async def login(form_data: OAuth2PasswordRequestForm = Depends()): data = auth.authenticate_user(cfg, form_data) return {"access_token": data, "token_type": "bearer"} @app.get("/status") async def status(token: str = Depends(oauth2_scheme)) -> DAQStatusModel: data = auth.parse_token(token) auth.check_jwt_ro(cfg, data) try: ret = daq.status ret.session = SessionStatus(current_pgroup=cfg.pgroup, session=cfg.session_state(data.session), staff=data.staff) return ret except Exception as e: raise HTTPException( status_code=api_status.HTTP_500_INTERNAL_SERVER_ERROR, detail=f"Error getting status: {e}" ) @app.get("/beamline/geometry") async def sample_geometry( token: str = Depends(oauth2_scheme), ) -> SampleGeometryModel: auth.check_jwt_ro(cfg, auth.parse_token(token)) return daq.sample_geometry @app.put("/beamline/omega") async def omega(val: float, token: str = Depends(oauth2_scheme)): auth.check_jwt_rw(cfg, auth.parse_token(token)) daq.omega = val return "OK" @app.put("/beamline/light") async def light(val: float, token: str = Depends(oauth2_scheme)): auth.check_jwt_rw(cfg, auth.parse_token(token)) daq.light = val return "OK" @app.put("/beamline/zoom") async def zoom(val: float, token: str = Depends(oauth2_scheme)): auth.check_jwt_rw(cfg, auth.parse_token(token)) daq.zoom = val return "OK" @app.put("/beamline/smargon") async def smargon(val: SmargonCoordinate, token: str = Depends(oauth2_scheme)): auth.check_jwt_rw(cfg, auth.parse_token(token)) daq.smargon = val return "OK" @app.post("/beamline/tweak_abr_meas_pos") async def tweak_abr_meas_pos(val: Coordinate, token: str = Depends(oauth2_scheme)): auth.check_jwt_staff(auth.parse_token(token)) daq.tweak_abr_meas_pos(val) return "OK" @app.post("/beamline/save_abr_meas_pos") async def save_abr_meas_pos(token: str = Depends(oauth2_scheme)): auth.check_jwt_staff(auth.parse_token(token)) daq.save_abr_meas_pos() return "OK" @app.post("/beamline/anneal") async def anneal(time_s: float, token: str = Depends(oauth2_scheme)): auth.check_jwt_rw(cfg, auth.parse_token(token)) daq.anneal(time_s) return "OK" @app.post("/beamline/goto_abr_meas_pos") async def goto_abr_meas_pos(token: str = Depends(oauth2_scheme)): auth.check_jwt_rw(cfg, auth.parse_token(token)) daq.goto_abr_meas_pos() return "OK" @app.post("/beam_mark/add") async def mark_beam(x: float, y: float, token: str = Depends(oauth2_scheme)): auth.check_jwt_staff(auth.parse_token(token)) daq.mark_beam(x, y) return "OK" @app.post("/beam_mark/clear") async def clear_beam_mark(token: str = Depends(oauth2_scheme)): auth.check_jwt_staff(auth.parse_token(token)) daq.clear_mark_beam() return "OK" @app.post("/beamline/beam_center") async def beam_center(x: float, y: float, token: str = Depends(oauth2_scheme)): auth.check_jwt_staff(auth.parse_token(token)) daq.beam_center = (x, y) return "OK" @app.post("/beamline/beam_size_mm") async def beam_size_mm(x: float, y: float, token: str = Depends(oauth2_scheme)): auth.check_jwt_staff(auth.parse_token(token)) daq.beam_size_mm = Coordinate(x=x, y=y) return "OK" @app.put("/beamline/samcam") async def samcam_settings(s: SampleCameraSettings, token: str = Depends(oauth2_scheme)): auth.check_jwt_rw(cfg, auth.parse_token(token)) daq.samcam_settings = s return "OK" @app.post("/samcam/autofocus") async def samcam_autofocus(s: AutofocusSettings, token: str = Depends(oauth2_scheme)): auth.check_jwt_rw(cfg, auth.parse_token(token)) daq.autofocus(s) return "OK" @app.post("/beamline/shutter") async def shutter(val: bool, token: str = Depends(oauth2_scheme)): auth.check_jwt_rw(cfg, auth.parse_token(token)) daq.shutter = val return "OK" async def get_image(token: str = Depends(oauth2_scheme)): auth.check_jwt_ro(cfg, auth.parse_token(token)) _, encoded_image = cv2.imencode( ".jpg", daq.camera_image ) # Encodes the image in JPEG format image_bytes = io.BytesIO( encoded_image.tobytes() ) # Convert OpenCV byte format to a file-like object return StreamingResponse(image_bytes, media_type="image/jpeg") # TELL procedures @app.get("/sample/curr_sample") async def sample(token: str = Depends(oauth2_scheme)) -> SampleShortInfo: token_data = auth.parse_token(token) auth.check_jwt_ro(cfg, auth.parse_token(token)) s = daq.sample if token_data.staff or token_data.group == s.user: return daq.sample else: return SampleShortInfo( sample_name="Other user sample", puck_name="", dewar_name="", db_id=-1, pin=s.pin, location=s.location ) @app.post("/sample/mount") async def mount(dbid: int, token: str = Depends(oauth2_scheme)): token_data = auth.parse_token(token) auth.check_jwt_rw(cfg, auth.parse_token(token)) st = daq.sample_spreadsheet index = -1 for i in range(len(st.s)): if st.s[i].db_id == dbid: index = i if index == -1: raise RuntimeError("Sample not found") if token_data.staff or token_data.group == st.s[index].user: daq.sample = st.s[index] return "OK" else: raise HTTPException( status_code=api_status.HTTP_401_UNAUTHORIZED, detail="Sample belongs to a different user.", headers={"WWW-Authenticate": "Bearer"}, ) @app.post("/sample/unmount") async def unmount(token: str = Depends(oauth2_scheme)): auth.check_jwt_rw(cfg, auth.parse_token(token)) daq.sample = None return "OK" @app.post("/sample/manual") async def manual(s: SampleShortInfo, token: str = Depends(oauth2_scheme)): auth.check_jwt_rw(cfg, auth.parse_token(token)) print(f"DB ID prior creating {s.db_id}") daq.create_sample(s) print(f"DB ID after creating {s.db_id}") def get_spreadsheet(data: TokenData) -> SampleShortInfoList: if data.staff: return cfg.spreadsheet else: return cfg.spreadsheet_pgroup(data.group) async def spreadsheet_event_stream(data: TokenData) -> AsyncGenerator[str, None]: try: while True: yield get_spreadsheet(data).model_dump_json() await asyncio.sleep(10) except asyncio.CancelledError: return @app.get("/sse/spreadsheet") async def spreadsheet_sse(token: str = Depends(oauth2_scheme)): data = auth.parse_token(token) return StreamingResponse( spreadsheet_event_stream(data), media_type="text/event-stream", headers={ "Cache-Control": "no-cache", "Connection": "keep-alive", "Access-Control-Allow-Origin": "*", "Access-Control-Allow-Headers": "Cache-Control" } ) @app.get("/sample/spreadsheet") async def spreadsheet(token: str = Depends(oauth2_scheme)) -> SampleShortInfoList: return get_spreadsheet(auth.parse_token(token)) # State transitions @app.post("/state/dewar_exchange") async def dewar_exchange(token: str = Depends(oauth2_scheme)): auth.check_jwt_rw(cfg, auth.parse_token(token)) daq.state = BeamlineStateEnum.DewarTransfer return "OK" @app.post("/state/sample_exchange") async def dewar_exchange(token: str = Depends(oauth2_scheme)): auth.check_jwt_rw(cfg, auth.parse_token(token)) daq.state = BeamlineStateEnum.SampleExchange return "OK" @app.post("/state/sample_alignment") async def sample_alignment(token: str = Depends(oauth2_scheme)): auth.check_jwt_rw(cfg, auth.parse_token(token)) daq.state = BeamlineStateEnum.SampleAlignment @app.post("/state/beam_location") async def beam_location(token: str = Depends(oauth2_scheme)): auth.check_jwt_staff(auth.parse_token(token)) daq.state = BeamlineStateEnum.BeamLocation # Scans @app.post("/scan/raster") async def raster(val: RasterGridRequest, auto: bool = False, token: str = Depends(oauth2_scheme)) -> CompletedRasterGrid: auth.check_jwt_rw(cfg, auth.parse_token(token)) return daq.measure_raster(val, auto) @app.post("/scan/rotation") async def rotation(val: RotationScanRequest, token: str = Depends(oauth2_scheme)) -> CompletedRotationScan: auth.check_jwt_rw(cfg, auth.parse_token(token)) return daq.measure_rotation(val) @app.post("/scan/auto") async def auto(s: SampleShortInfo, token: str = Depends(oauth2_scheme)): auth.check_jwt_rw(cfg, auth.parse_token(token)) runtime = daq.measure(s) return f"{runtime:0.3f}" @app.post("/scan/cancel") async def cancel(token: str = Depends(oauth2_scheme)): auth.check_jwt_rw(cfg, auth.parse_token(token)) daq.cancel() # ALC routines @app.post("/alc/background") async def alc_background(token: str = Depends(oauth2_scheme)) -> str: auth.check_jwt_rw(cfg, auth.parse_token(token)) daq.get_background() return "OK" @app.post("/alc/center_loop") async def alc_center_loop(token: str = Depends(oauth2_scheme)) -> str: auth.check_jwt_rw(cfg, auth.parse_token(token)) daq.auto_loop_center() return "OK" @app.post("/alc/ml_bounding_box") async def alc_ml_bounding_box(token: str = Depends(oauth2_scheme)) -> RasterGridRequest | None: auth.check_jwt_rw(cfg, auth.parse_token(token)) return daq.ml_bounding_box() # Access management @app.get("/access/pgroup") async def pgroup(token: str = Depends(oauth2_scheme)) -> str: auth.check_jwt_ro(cfg, auth.parse_token(token)) return cfg.pgroup @app.put("/access/pgroup") async def set_pgroup(val: str, token: str = Depends(oauth2_scheme)) -> str: auth.check_jwt_staff(auth.parse_token(token)) cfg.pgroup = val return "OK" @app.delete("/access/pgroup") async def del_pgroup(token: str = Depends(oauth2_scheme)) -> str: auth.check_jwt_staff(auth.parse_token(token)) cfg.pgroup = None return "OK" @app.post("/access/end_session") async def end_session(token: str = Depends(oauth2_scheme)) -> str: # End active session will only delete session, if it is equal to token value # so no need to check R/W permissions token_data = auth.parse_token(token) cfg.end_active_session(token_data.session) return "OK" @app.post("/access/force_current_session") async def force_current_session(token: str = Depends(oauth2_scheme)) -> str: data = auth.parse_token(token) # Counterintuitive, this operation requires only R/O permission # as this is actually acquiring R/W permissions auth.check_jwt_ro(cfg, data) auth.force_current_sesion(cfg, data) return "OK" @app.get("/beamline/settings") async def get_settings(token: str = Depends(oauth2_scheme)) -> BeamlineSettingsModel: auth.check_jwt_staff(auth.parse_token(token)) return cfg.settings @app.put("/beamline/settings") async def put_settings(s: BeamlineSettingsModel, token: str = Depends(oauth2_scheme)): auth.check_jwt_staff(auth.parse_token(token)) cfg.settings = s @app.get("/beamline/cryo_settings") async def get_cryo_settings(token: str = Depends(oauth2_scheme)) -> CryojetSettingsModel: auth.check_jwt_staff(auth.parse_token(token)) return cfg.cryojet_settings @app.put("/beamline/cryo_settings") async def put_cryo_settings(s: CryojetSettingsModel, token: str = Depends(oauth2_scheme)): auth.check_jwt_staff(auth.parse_token(token)) cfg.cryojet_settings = s LOGGING_CONFIG = { "version": 1, "disable_existing_loggers": False, "formatters": { "default": { "()": "uvicorn.logging.DefaultFormatter", "fmt": "%(levelprefix)s %(message)s", "use_colors": True, }, }, "handlers": { "default": { "formatter": "default", "class": "logging.StreamHandler", "stream": "ext://sys.stdout", }, }, "loggers": { "uvicorn": { "handlers": ["default"], "level": "WARNING", }, "uvicorn.access": { "handlers": ["default"], "level": "WARNING", }, }, } urllib3.disable_warnings() def main(): # Remove in production! urllib3.disable_warnings() # Run the application using uvicorn uvicorn.run("aaredaq.server:app", host="0.0.0.0", port=5210, workers=4, log_config=LOGGING_CONFIG) if __name__ == "__main__": #start_image_stats_receiver(zmq_url="tcp://129.129.110.12:9089") main() #stop_image_stats_receiver()