diff --git a/helm/daq-queuing-service/templates/deployment.yaml b/helm/daq-queuing-service/templates/deployment.yaml index f96d836..f4d59d3 100644 --- a/helm/daq-queuing-service/templates/deployment.yaml +++ b/helm/daq-queuing-service/templates/deployment.yaml @@ -60,6 +60,22 @@ spec: {{- with .Values.volumeMounts }} {{- toYaml . | nindent 12 }} {{- end }} + {{- if or .Values.udcSecret.enabled .Values.env }} + env: + {{- if .Values.udcSecret.enabled }} + - name: UDC_SECRET + valueFrom: + secretKeyRef: + name: {{ .Values.udcSecret.name }} + key: {{ .Values.udcSecret.key }} + + - name: UDC_CLIENT_ID + value: "{{ .Values.udcSecret.clientId }}" + {{- end }} + {{- with .Values.env }} + {{- toYaml . | nindent 12 }} + {{- end }} + {{- end }} volumes: {{- with .Values.volumes }} {{- toYaml . | nindent 8 }} diff --git a/helm/daq-queuing-service/values.yaml b/helm/daq-queuing-service/values.yaml index 79ca3c6..7f553c4 100644 --- a/helm/daq-queuing-service/values.yaml +++ b/helm/daq-queuing-service/values.yaml @@ -75,3 +75,11 @@ volumeMounts: [] volumes: [] nodeSelector: {} tolerations: [] + +env: [] + +udcSecret: + enabled: false + name: "" + key: udc-secret + clientId: "" diff --git a/pyproject.toml b/pyproject.toml index 233f095..57ace2c 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -13,7 +13,7 @@ classifiers = [ ] description = "A service to queue tasks and chain BlueAPI calls" dependencies = [ - "blueapi>=1.13.0", + "blueapi>=1.14.0", "fastapi>=0.136.0", "pydantic>=2.13.2", ] diff --git a/src/daq_queuing_service/api/api.py b/src/daq_queuing_service/api/api.py index 5899cf5..aaf02f8 100644 --- a/src/daq_queuing_service/api/api.py +++ b/src/daq_queuing_service/api/api.py @@ -1,6 +1,5 @@ import asyncio import json -import logging from collections.abc import AsyncGenerator, Callable from blueapi.client.rest import ( @@ -16,6 +15,7 @@ from daq_queuing_service.app._config import AppConfig, load_config from daq_queuing_service.blueapi_interaction.blueapi_call import BlueapiCallResponse from daq_queuing_service.broadcaster import Broadcaster +from daq_queuing_service.log import LOGGER from daq_queuing_service.task import ExperimentDefinition, Status, Task from daq_queuing_service.task_queue.queue import ( QUEUE_EVENTS, @@ -26,8 +26,6 @@ # pyright: reportUnusedFunction=false -LOGGER = logging.getLogger(__name__) - class InvalidExperimentDefinitionsError(Exception): def __init__(self, errors: dict[int, InvalidParametersError | UnknownPlanError]): diff --git a/src/daq_queuing_service/app/app.py b/src/daq_queuing_service/app/app.py index 12e7893..3cb3f67 100644 --- a/src/daq_queuing_service/app/app.py +++ b/src/daq_queuing_service/app/app.py @@ -3,14 +3,13 @@ from contextlib import asynccontextmanager from typing import NoReturn -from blueapi.client import BlueapiClient -from blueapi.client.rest import BlueapiRestClient from fastapi import FastAPI from fastapi.middleware.cors import CORSMiddleware from daq_queuing_service.api.api import create_api_router from daq_queuing_service.api.errors import register_exception_handlers from daq_queuing_service.blueapi_interaction.blueapi_adapter import BlueapiClientAdapter +from daq_queuing_service.blueapi_interaction.clients import get_blueapi_clients from daq_queuing_service.broadcaster import Broadcaster from daq_queuing_service.plugins.construct_task_request import ( construct_blueapi_task_request, @@ -68,8 +67,7 @@ def log_task_exception(task: asyncio.Task[NoReturn]): app.state.queue = TaskQueue(converter, broadcaster) - blueapi_rest_client = BlueapiRestClient(config=config.blueapi.api) - blueapi_client = BlueapiClient.from_config(config.blueapi) + blueapi_rest_client, blueapi_client = get_blueapi_clients(config.blueapi) blueapi_client_adapter = BlueapiClientAdapter(blueapi_client) app.state.worker = QueueWorker( diff --git a/src/daq_queuing_service/blueapi_interaction/blueapi_adapter.py b/src/daq_queuing_service/blueapi_interaction/blueapi_adapter.py index 87c4851..f312895 100644 --- a/src/daq_queuing_service/blueapi_interaction/blueapi_adapter.py +++ b/src/daq_queuing_service/blueapi_interaction/blueapi_adapter.py @@ -1,5 +1,4 @@ import asyncio -import logging from dataclasses import dataclass from typing import Generic, TypeVar @@ -14,7 +13,7 @@ from blueapi.service.model import TaskRequest from blueapi.worker import TaskStatus, WorkerState -LOGGER = logging.getLogger(__name__) +from daq_queuing_service.log import LOGGER T = TypeVar("T") E = TypeVar("E", bound=Exception) diff --git a/src/daq_queuing_service/blueapi_interaction/clients.py b/src/daq_queuing_service/blueapi_interaction/clients.py new file mode 100644 index 0000000..5a77357 --- /dev/null +++ b/src/daq_queuing_service/blueapi_interaction/clients.py @@ -0,0 +1,39 @@ +from unittest.mock import MagicMock + +from blueapi.client import BlueapiClient +from blueapi.client.event_bus import EventBusClient +from blueapi.client.rest import BlueapiRestClient +from blueapi.config import ApplicationConfig +from bluesky_stomp.messaging import Broker, StompClient + +from daq_queuing_service.blueapi_interaction.token_retriever import UDCTokenRetriever + + +def get_blueapi_clients( + blueapi_config: ApplicationConfig, +) -> tuple[BlueapiRestClient, BlueapiClient]: + if not blueapi_config.oidc: + blueapi_config.oidc = MagicMock() + + blueapi_rest_client = BlueapiRestClient( + config=blueapi_config.api, + # Waiting on https://github.com/DiamondLightSource/blueapi/pull/1553 + session_manager=UDCTokenRetriever(), # type: ignore + ) + + if blueapi_config.stomp.enabled: + assert blueapi_config.stomp.url.host is not None, "Stomp URL missing host" + assert blueapi_config.stomp.url.port is not None, "Stomp URL missing port" + stomp_client = StompClient.for_broker( + broker=Broker( + host=blueapi_config.stomp.url.host, + port=blueapi_config.stomp.url.port, + auth=blueapi_config.stomp.auth, + ) + ) + events = EventBusClient(stomp_client) + blueapi_client = BlueapiClient(blueapi_rest_client, events) + else: + blueapi_client = BlueapiClient(blueapi_rest_client) + + return blueapi_rest_client, blueapi_client diff --git a/src/daq_queuing_service/blueapi_interaction/token_retriever.py b/src/daq_queuing_service/blueapi_interaction/token_retriever.py new file mode 100644 index 0000000..8f82745 --- /dev/null +++ b/src/daq_queuing_service/blueapi_interaction/token_retriever.py @@ -0,0 +1,47 @@ +import os + +import requests + +from daq_queuing_service.log import LOGGER + + +class UDCTokenRetriever: + """Implements `get_valid_access_token` to get a token using a sealed secret.""" + + def __init__( + self, + secret_variable_name: str = "UDC_SECRET", + client_id_variable_name: str = "UDC_CLIENT_ID", + ): + self._secret_variable_name = secret_variable_name + self._client_id_variable_name = client_id_variable_name + + def get_valid_access_token(self) -> str: + token_url = ( + "https://identity.diamond.ac.uk/realms/dls/protocol/openid-connect/token" + ) + + client_id = os.environ.get(self._client_id_variable_name) + client_secret = os.environ.get(self._secret_variable_name) + + if not client_secret: + LOGGER.debug("No UDC secret found") + return "" + if not client_id: + LOGGER.debug("No UDC client ID found") + return "" + + LOGGER.debug("Found UDC secret") + + response = requests.post( + token_url, + data={ + "client_id": client_id, + "client_secret": client_secret, + "grant_type": "client_credentials", + }, + ) + response.raise_for_status() + token = response.json().get("access_token") + LOGGER.debug("Returning token") + return token diff --git a/src/daq_queuing_service/broadcaster.py b/src/daq_queuing_service/broadcaster.py index 23fca50..0faf5df 100644 --- a/src/daq_queuing_service/broadcaster.py +++ b/src/daq_queuing_service/broadcaster.py @@ -1,11 +1,10 @@ import asyncio -import logging from collections.abc import Iterable, Mapping from typing import Any, Generic, TypedDict, TypeVar from pydantic import BaseModel -LOGGER = logging.getLogger(__name__) +from daq_queuing_service.log import LOGGER T = TypeVar("T", bound=str) diff --git a/src/daq_queuing_service/log.py b/src/daq_queuing_service/log.py new file mode 100644 index 0000000..44629f7 --- /dev/null +++ b/src/daq_queuing_service/log.py @@ -0,0 +1,21 @@ +import logging + +import colorlog + +HANDLER = colorlog.StreamHandler() +HANDLER.setFormatter( + colorlog.ColoredFormatter( + "%(log_color)s%(asctime)s [%(name)s] %(levelname)s: %(message)s", + log_colors={ + "DEBUG": "cyan", + "INFO": "green", + "WARNING": "yellow", + "ERROR": "red", + "CRITICAL": "bold_red", + }, + ) +) +LOGGER = logging.getLogger("Queue") +LOGGER.addHandler(HANDLER) +LOGGER.setLevel(logging.DEBUG) +LOGGER.propagate = False diff --git a/src/daq_queuing_service/task_queue/queue.py b/src/daq_queuing_service/task_queue/queue.py index dbf3ab5..887a376 100644 --- a/src/daq_queuing_service/task_queue/queue.py +++ b/src/daq_queuing_service/task_queue/queue.py @@ -1,5 +1,4 @@ import asyncio -import logging from collections.abc import Callable, Sequence from types import TracebackType from typing import Any, Literal @@ -13,6 +12,7 @@ CallStatus, ) from daq_queuing_service.broadcaster import Broadcaster, Event +from daq_queuing_service.log import LOGGER from daq_queuing_service.plugins.converter_utils import Converter from daq_queuing_service.task import Status, Task, TaskWithPosition from daq_queuing_service.task_queue.queue_utils import ( @@ -24,8 +24,6 @@ TaskNotInQueueError, ) -LOGGER = logging.getLogger(__name__) - class TaskRegistry(dict[str, Task]): def __missing__(self, task_id: str) -> Task: diff --git a/src/daq_queuing_service/worker/worker.py b/src/daq_queuing_service/worker/worker.py index adc2802..40797a1 100644 --- a/src/daq_queuing_service/worker/worker.py +++ b/src/daq_queuing_service/worker/worker.py @@ -17,11 +17,14 @@ from daq_queuing_service.blueapi_interaction.blueapi_adapter import BlueapiClientAdapter from daq_queuing_service.blueapi_interaction.blueapi_call import BlueapiCall, CallStatus +from daq_queuing_service.log import HANDLER from daq_queuing_service.task import ExperimentDefinition from daq_queuing_service.task_queue.queue import TaskQueue -LOGGER = logging.getLogger(__name__) +LOGGER = logging.getLogger("Queue Worker") +LOGGER.addHandler(HANDLER) LOGGER.setLevel(logging.DEBUG) +LOGGER.propagate = False class QueueWorker: @@ -101,7 +104,8 @@ def _on_blueapi_event(event: AnyEvent, call: BlueapiCall): assert worker_event.task_status call.blueapi_id = worker_event.task_status.task_id LOGGER.info( - f"Call {call} is in progress, blueapi ID: {call.blueapi_id}" + f"Putting call in progress, blueapi ID: {call.blueapi_id}. " + + f"Call: ({call})" ) call.put_in_progress() case ProgressEvent(): diff --git a/tests/test_data/test_blueapi_config.yaml b/tests/test_data/test_blueapi_config.yaml index d4860ab..b681665 100644 --- a/tests/test_data/test_blueapi_config.yaml +++ b/tests/test_data/test_blueapi_config.yaml @@ -1,4 +1,4 @@ -api: +api: url: "http://localhost:8000" stomp: enabled: true # All other stomp settings will be ignored if this is false diff --git a/tests/unit_tests/conftest.py b/tests/unit_tests/conftest.py index f64360f..e09b10c 100644 --- a/tests/unit_tests/conftest.py +++ b/tests/unit_tests/conftest.py @@ -1,9 +1,11 @@ import pytest from blueapi.service.model import TaskRequest from blueapi.worker.event import TaskError, TaskResult +from pytest import MonkeyPatch from daq_queuing_service.blueapi_interaction.blueapi_call import BlueapiCall from daq_queuing_service.broadcaster import Broadcaster +from daq_queuing_service.log import LOGGER from daq_queuing_service.plugins.construct_task_request import ( construct_blueapi_call_list, ) @@ -11,6 +13,13 @@ from daq_queuing_service.task_queue.queue import TaskQueue +@pytest.fixture(autouse=True) +def propagate_logs(monkeypatch: MonkeyPatch): + # This is turned off in prod to avoid duplicate logs + # but needed in tests for caplog to receive logs + monkeypatch.setattr(LOGGER, "propagate", True) + + @pytest.fixture def tasks() -> list[Task]: return [ diff --git a/tests/unit_tests/test_get_blueapi_clients.py b/tests/unit_tests/test_get_blueapi_clients.py new file mode 100644 index 0000000..bf57b37 --- /dev/null +++ b/tests/unit_tests/test_get_blueapi_clients.py @@ -0,0 +1,46 @@ +from unittest.mock import MagicMock, patch + +from blueapi.config import ApplicationConfig, RestConfig, StompConfig +from pydantic import HttpUrl + +from daq_queuing_service.blueapi_interaction.clients import get_blueapi_clients + + +@patch("daq_queuing_service.blueapi_interaction.clients.UDCTokenRetriever") +@patch("daq_queuing_service.blueapi_interaction.clients.BlueapiClient") +@patch("daq_queuing_service.blueapi_interaction.clients.BlueapiRestClient") +def test_get_blueapi_clients_constructs_clients_with_expected_args_and_returns_clients( + mock_rest_client: MagicMock, + mock_blueapi_client: MagicMock, + mock_token_retriever: MagicMock, +): + rest_config = RestConfig(url=HttpUrl("http://test_url.com")) + rest_client, blueapi_client = get_blueapi_clients( + ApplicationConfig(api=rest_config) + ) + + mock_rest_client.assert_called_once_with( + config=rest_config, session_manager=mock_token_retriever.return_value + ) + mock_blueapi_client.assert_called_once_with(rest_client) + + assert rest_client is mock_rest_client.return_value + assert blueapi_client is mock_blueapi_client.return_value + + +@patch("daq_queuing_service.blueapi_interaction.clients.EventBusClient") +@patch("daq_queuing_service.blueapi_interaction.clients.BlueapiClient") +@patch("daq_queuing_service.blueapi_interaction.clients.BlueapiRestClient") +def test_get_blueapi_clients_constructs_blueapi_client_with_stomp_if_enabled_in_config( + mock_rest_client: MagicMock, + mock_blueapi_client: MagicMock, + mock_event_bus_client: MagicMock, +): + rest_config = RestConfig(url=HttpUrl("http://test_url.com")) + rest_client, _ = get_blueapi_clients( + ApplicationConfig(api=rest_config, stomp=StompConfig(enabled=True)) + ) + + mock_blueapi_client.assert_called_once_with( + rest_client, mock_event_bus_client.return_value + ) diff --git a/tests/unit_tests/test_udc_token_retriever.py b/tests/unit_tests/test_udc_token_retriever.py new file mode 100644 index 0000000..9dcef74 --- /dev/null +++ b/tests/unit_tests/test_udc_token_retriever.py @@ -0,0 +1,45 @@ +from unittest.mock import MagicMock, patch + +import pytest +from pytest import MonkeyPatch +from requests import Response + +from daq_queuing_service.blueapi_interaction.token_retriever import UDCTokenRetriever + + +@pytest.fixture(autouse=True) +def set_secret_env_vars(monkeypatch: MonkeyPatch): + monkeypatch.setenv("UDC_SECRET", "secret") + monkeypatch.setenv("UDC_CLIENT_ID", "ixxudc") + + +@patch("daq_queuing_service.blueapi_interaction.token_retriever.requests.post") +def test_get_valid_access_token_makes_expected_request_and_returns_result( + mock_post: MagicMock, +): + mock_post.return_value = Response() + mock_post.return_value.json = MagicMock( + return_value={"access_token": "valid_token"} + ) + mock_post.return_value.status_code = 200 + token_retriever = UDCTokenRetriever() + + assert token_retriever.get_valid_access_token() == "valid_token" + + mock_post.assert_called_once_with( + "https://identity.diamond.ac.uk/realms/dls/protocol/openid-connect/token", + data={ + "client_id": "ixxudc", + "client_secret": "secret", + "grant_type": "client_credentials", + }, + ) + + +@pytest.mark.parametrize("variable_name", ("UDC_SECRET", "UDC_CLIENT_ID")) +def test_get_valid_access_token_if_env_vars_not_found_then_returns_empty_string( + variable_name: str, monkeypatch: MonkeyPatch +): + monkeypatch.delenv(variable_name) + token_retriever = UDCTokenRetriever() + assert token_retriever.get_valid_access_token() == "" diff --git a/tests/unit_tests/test_worker.py b/tests/unit_tests/test_worker.py index 1b5804c..7e7d00f 100644 --- a/tests/unit_tests/test_worker.py +++ b/tests/unit_tests/test_worker.py @@ -14,7 +14,7 @@ from blueapi.core import DataEvent from blueapi.service.model import TaskRequest from blueapi.worker import ProgressEvent, TaskStatus, WorkerEvent, WorkerState -from pytest import LogCaptureFixture +from pytest import LogCaptureFixture, MonkeyPatch from daq_queuing_service.blueapi_interaction.blueapi_adapter import ( BlueapiClientAdapter, @@ -23,7 +23,14 @@ from daq_queuing_service.blueapi_interaction.blueapi_call import CallStatus from daq_queuing_service.task import ExperimentDefinition, Status from daq_queuing_service.task_queue.queue import TaskError, TaskQueue, TaskResult -from daq_queuing_service.worker.worker import QueueWorker +from daq_queuing_service.worker.worker import LOGGER, QueueWorker + + +@pytest.fixture(autouse=True) +def propagate_logs(monkeypatch: MonkeyPatch): + # This is turned off in prod to avoid duplicate logs + # but needed in tests for caplog to receive logs + monkeypatch.setattr(LOGGER, "propagate", True) def _get_mock_blueapi_client( @@ -228,7 +235,7 @@ async def test_when_parameter_error_then_call_failed_and_error_added_to_call( worker_with_parameter_error._client.run_task.assert_called_once() # type: ignore -async def test_when_plan_name_error_then_call_failed_and_error_added_to_task( +async def test_when_plan_name_error_then_call_failed_and_error_added_to_call( worker_with_unknown_plan_error: QueueWorker, only_loop_once: type[Exception] ): diff --git a/uv.lock b/uv.lock index dec3a65..2dff840 100644 --- a/uv.lock +++ b/uv.lock @@ -381,7 +381,7 @@ wheels = [ [[package]] name = "blueapi" -version = "1.13.0" +version = "1.14.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "aioca" }, @@ -409,9 +409,9 @@ dependencies = [ { name = "tomlkit" }, { name = "uvicorn" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/89/19/9ee222e1efef179c6975d070ef99b04752cecd70eac9479b670c6cb5ae3a/blueapi-1.13.0.tar.gz", hash = "sha256:3bf77f75e6bef65326786cc01be74fcf1eded37d463d9f24d8ad446cbbaccaa7", size = 1838838, upload-time = "2026-04-17T12:31:25.597Z" } +sdist = { url = "https://files.pythonhosted.org/packages/18/16/bb4e52cbbc59ba94d2ded1778206e6ce40664fca178f66aa3fe5b8bd7d5c/blueapi-1.14.0.tar.gz", hash = "sha256:d3cb6975fa8826fd9ec9adf47e75719427e9b992002956c0433c4abb887d86d4", size = 1840934, upload-time = "2026-04-24T10:10:35.367Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/d0/a8/bb3205aca03ddc4305130b92fda5dde7c69ef91b403f1a8aada324c50256/blueapi-1.13.0-py3-none-any.whl", hash = "sha256:99a17cd87d59a7dfebeac94b68fea701d0d9d1fc0f799cfd3b710a6259610a2f", size = 83960, upload-time = "2026-04-17T12:31:24.239Z" }, + { url = "https://files.pythonhosted.org/packages/f4/0a/9499d4ee173cab1f891c86f8332e8c2d4844ed5d76a72f84ced0f134c604/blueapi-1.14.0-py3-none-any.whl", hash = "sha256:e232529e795fad0f9dec36d1bd6dd9f5bcac21f15967a6a609b01b9306a80ee2", size = 84560, upload-time = "2026-04-24T10:10:33.59Z" }, ] [[package]] @@ -1030,7 +1030,7 @@ dev = [ [package.metadata] requires-dist = [ - { name = "blueapi", specifier = ">=1.13.0" }, + { name = "blueapi", specifier = ">=1.14.0" }, { name = "fastapi", specifier = ">=0.136.0" }, { name = "pydantic", specifier = ">=2.13.2" }, ]