from unittest.mock import MagicMock, patch import numpy as np import pytest from aarecommon.models.beamline import MXBeamline from aarecommon.models.models import MLBoxModel, MLBoxType, MLOutputModel from aare.daq.mlbox import MlBox @pytest.fixture(autouse=True) def no_mlbox_sleep(monkeypatch): monkeypatch.setattr(MlBox, "RETRY_SLEEP_S", 0.0) @pytest.fixture def mock_wrapper(): return MagicMock() @pytest.fixture def mlbox(mock_wrapper): return MlBox(MXBeamline.X10SA, wrapper=mock_wrapper) def test_init_x10sa(mock_wrapper): ml = MlBox(MXBeamline.X10SA, wrapper=mock_wrapper) assert ml._MlBox__beamline == MXBeamline.X10SA def test_predict_uses_injected_wrapper(): fake = MagicMock() fake.get_latest_prediction_bundle.return_value = MagicMock() ml = MlBox(MXBeamline.X10SA, wrapper=fake) ml._fetch_prediction_bundle() fake.get_latest_prediction_bundle.assert_called_once() def test_init_unsupported(): with pytest.raises(NotImplementedError): MlBox(MXBeamline.SIMULATED) def test_prediction_score(): assert MlBox._prediction_score(None) == (0, 0.0) out = MLOutputModel() assert MlBox._prediction_score(out) == (0, 0.0) out.add_box(MLBoxType.CRYSTAL, (0, 0, 10, 10), conf=0.9) assert MlBox._prediction_score(out) == (1, 0.9) def test_check_overlap(): box_a = MLBoxModel.from_tuple(MLBoxType.CRYSTAL, (0, 0, 10, 10), conf=1.0) box_b = MLBoxModel.from_tuple(MLBoxType.CRYSTAL, (5, 5, 15, 15), conf=1.0) assert MlBox._check_overlap(box_a, box_b) > 0 box_c = MLBoxModel.from_tuple(MLBoxType.CRYSTAL, (20, 20, 30, 30), conf=1.0) assert MlBox._check_overlap(box_a, box_c) == 0 def test_filter_predictions(mlbox): out = MLOutputModel() out.add_box(MLBoxType.CRYSTAL, (0, 0, 10, 10), conf=0.4) out.add_box(MLBoxType.CRYSTAL, (20, 20, 30, 30), conf=0.8) mlbox._filter_predictions(out, confidence_min=0.5) assert len(out.boxes) == 1 assert list(out.boxes.values())[0].conf == 0.8 @patch("cv2.imdecode") def test_decode_bundle_image(mock_imdecode): mock_imdecode.return_value = np.zeros((100, 100, 3)) img = MlBox._decode_bundle_image(b"fake_jpeg") assert img is not None assert img.shape == (100, 100, 3) def test_extract_bundle_meta(): pred = MagicMock() pred.target_point = [100.0, 200.0] pred.focus_score = 0.5 meta = MlBox._extract_bundle_meta(pred) assert meta.target_point == (100.0, 200.0) assert meta.focus == 0.5 def test_preferred_class_prefers_loop_over_pin_when_margin_not_exceeded(): boxes = MLOutputModel() boxes.add_box(MLBoxType.PIN, (0, 0, 10, 10), 0.8) boxes.add_box(MLBoxType.LOOP_FACE, (1, 1, 9, 9), 0.75) best = MlBox.get_preferred_class_box_with_confidence_threshold( boxes, loop_preference_margin=0.1 ) assert best is not None assert best.cls == MLBoxType.LOOP_FACE def test_prediction_score_empty(): assert MlBox._prediction_score(None) == (0, 0.0) def test_prediction_score_uses_count_and_max_confidence(): model = MLOutputModel() model.add_box(MLBoxType.PIN, (0, 0, 10, 10), 0.4) model.add_box(MLBoxType.CRYSTAL, (1, 1, 9, 9), 0.8) assert MlBox._prediction_score(model) == (2, 0.8) def test_best_by_class_keeps_highest_confidence_per_class(): results = [ {"class": 1, "confidence": 0.4, "box": {"x1": 0, "y1": 0, "x2": 10, "y2": 10}}, {"class": 1, "confidence": 0.8, "box": {"x1": 1, "y1": 1, "x2": 11, "y2": 11}}, {"class": 2, "confidence": 0.6, "box": {"x1": 2, "y1": 2, "x2": 12, "y2": 12}}, ] out = MlBox._best_by_class(results) assert out is not None assert out.get_best_for_class(MLBoxType.PIN).conf == 0.8 assert out.get_best_for_class(MLBoxType.CRYSTAL).conf == 0.6