diff --git a/src/blueapi/service/interface.py b/src/blueapi/service/interface.py index 6c0c5befbd..02b6447980 100644 --- a/src/blueapi/service/interface.py +++ b/src/blueapi/service/interface.py @@ -1,3 +1,4 @@ +import json import logging from collections.abc import Mapping from dataclasses import dataclass @@ -8,7 +9,9 @@ from bluesky.callbacks.tiled_writer import TiledWriter from bluesky_stomp.messaging import StompClient from bluesky_stomp.models import Broker, DestinationBase, MessageTopic +from fastapi import status from tiled.client import from_uri +from tiled.client.utils import ClientError from blueapi.cli.scratch import get_python_environment from blueapi.config import ApplicationConfig, OIDCConfig, ServiceAccount, StompConfig @@ -25,6 +28,7 @@ TaskRequest, WorkerTask, ) +from blueapi.utils import TILED_PROPOSAL_RE from blueapi.utils.serialization import access_blob from blueapi.worker.event import ProgressEvent, TaskStatusEnum, WorkerEvent, WorkerState from blueapi.worker.task import Task @@ -205,7 +209,43 @@ def begin_task( api_key=tiled_config.authentication, headers=pass_through_headers, ) - + if task.task_id is not None: + task_ = get_task_by_id(task_id=task.task_id) + if task_ is not None: + task_metadata = task_.task.metadata + instrument = active_context.run_engine.md["instrument"] + instrument_session = task_metadata["instrument_session"] + if not (match := TILED_PROPOSAL_RE.match(instrument_session)): + raise ValueError("Invalid instrument session") + proposal = match["proposal"] + # Each level's access blob is the prefix of the full + # (beamline, proposal, visit) one that access_blob() builds, + # matching the beamline/proposal/session tiers the tiled + # access policy expects a container to be tagged with. + session_blob = json.loads(access_blob(instrument_session, instrument)) + level_access_tags = [ + [json.dumps({"beamline": instrument})], + [json.dumps({"beamline": instrument, "proposal": proposal})], + [json.dumps(session_blob)], + ] + for key, access_tags in zip( + (instrument, proposal, instrument_session), + level_access_tags, + strict=True, + ): + if key not in tiled_client: + try: + tiled_client.create_container( + key=key, access_tags=access_tags + ) + except ClientError as e: + if ( + e.response.status_code == status.HTTP_409_CONFLICT + ): # already exists + ... + else: + raise + tiled_client = tiled_client[key] tiled_writer_token = active_context.run_engine.subscribe( TiledWriter(tiled_client, batch_size=1) ) @@ -230,7 +270,7 @@ def remove_callback_when_task_finished( if task.task_id is not None: try: active_worker.begin_task(task.task_id) - except: + except Exception: for channel, token in subscribers: channel.unsubscribe(token) raise diff --git a/src/blueapi/utils/__init__.py b/src/blueapi/utils/__init__.py index f722c5b42d..cba5f03f70 100644 --- a/src/blueapi/utils/__init__.py +++ b/src/blueapi/utils/__init__.py @@ -32,6 +32,8 @@ Return = TypeVar("Return") INSTRUMENT_SESSION_RE = re.compile(r"^[a-z]{2}(?P\d+)-(?P\d+)$") +# Full proposal code (e.g. "cm12345" from "cm12345-1") for building tiled node paths. +TILED_PROPOSAL_RE = re.compile(r"^(?P[a-z]{2}\d+)-\d+$") def report_successful_devices( diff --git a/src/blueapi/utils/serialization.py b/src/blueapi/utils/serialization.py index 8918cf8823..e79f6ace59 100644 --- a/src/blueapi/utils/serialization.py +++ b/src/blueapi/utils/serialization.py @@ -30,15 +30,18 @@ def serialize(obj: Any) -> Any: def access_blob(instrument_session: str, beamline: str) -> str: - m = utils.INSTRUMENT_SESSION_RE.match(instrument_session) - if m is None: + session_match = utils.INSTRUMENT_SESSION_RE.match(instrument_session) + proposal_match = utils.TILED_PROPOSAL_RE.match(instrument_session) + if session_match is None or proposal_match is None: raise ValueError( "Unable to extract proposal and visit from " f"instrument session {instrument_session}" ) blob = { - "proposal": int(m["proposal"]), - "visit": int(m["visit"]), + # The full proposal code (e.g. "cm12345"), not just its number - the + # tiled access policy strips the letters itself where it needs them. + "proposal": proposal_match["proposal"], + "visit": int(session_match["visit"]), "beamline": beamline, } return json.dumps(blob) diff --git a/tests/system_tests/services/opa_config/config.yaml b/tests/system_tests/services/opa_config/config.yaml index 914d18247e..dfb2759a2b 100644 --- a/tests/system_tests/services/opa_config/config.yaml +++ b/tests/system_tests/services/opa_config/config.yaml @@ -5,7 +5,7 @@ services: bundles: diamond-policies: service: ghcr - resource: ghcr.io/diamondlightsource/authz-policy:0.0.24 + resource: ghcr.io/zohebshaikh/authz-policy:0.0.25-alpha polling: min_delay_seconds: 30 max_delay_seconds: 120 diff --git a/tests/system_tests/services/tiled_config/dls.py b/tests/system_tests/services/tiled_config/dls.py index 875b559d5a..5696506b0f 100644 --- a/tests/system_tests/services/tiled_config/dls.py +++ b/tests/system_tests/services/tiled_config/dls.py @@ -2,7 +2,7 @@ import logging from fastapi import HTTPException -from pydantic import BaseModel, HttpUrl, TypeAdapter +from pydantic import BaseModel, HttpUrl, TypeAdapter, ValidationError from starlette.status import ( HTTP_401_UNAUTHORIZED, ) @@ -21,9 +21,35 @@ class DiamondAccessBlob(BaseModel): - proposal: int - visit: int beamline: str + # The full proposal code, e.g. "cm12345" - tiled.rego strips the leading + # letters itself where it needs the bare number. + proposal: str | None = None + visit: int | None = None + + +# Maps a composite access tag's "key" (as produced by tiled.rego's +# beamline_tag/proposal_tag/session_tag, e.g. +# "beamline:i22,proposal:cm111,session:cm111-1") to the corresponding OPA +# input field. All of these are strings on an existing node's tag - "session" +# here is the full "cm111-1" instrument session, not the internal numeric +# session id used by modify_session, so it's kept under its own field name +# rather than aliased to "visit". +_TAG_KEY_TO_INPUT_FIELD = { + "beamline": "beamline", + "proposal": "proposal", + "session": "session", +} + + +def _parse_composite_tag(tag: str) -> dict[str, str]: + fields: dict[str, str] = {} + for part in tag.split(","): + key, _, value = part.partition(":") + field = _TAG_KEY_TO_INPUT_FIELD.get(key) + if field is not None: + fields[field] = value + return fields def _check_principal(principal: Principal | None): @@ -131,9 +157,17 @@ def build_input( and "tags" in access_blob and len(access_blob["tags"]) > 0 ): - blob = self._type_adapter.validate_json(access_blob["tags"][0]) + tag = access_blob["tags"][0] + try: + blob = self._type_adapter.validate_json(tag) + except ValidationError: + # Not a create-request JSON blob - it's a composite tag + # already assigned to an existing node (e.g. tiled checking + # scopes on a parent before creating a child). + blob = None + _input.update(_parse_composite_tag(tag)) if isinstance(blob, DiamondAccessBlob): - _input.update(blob.model_dump()) + _input.update(blob.model_dump(exclude_none=True)) elif isinstance(blob, int): _input["session"] = str(blob) diff --git a/tests/system_tests/test_blueapi_system.py b/tests/system_tests/test_blueapi_system.py index 8c4dc19cf9..1ac76d4083 100644 --- a/tests/system_tests/test_blueapi_system.py +++ b/tests/system_tests/test_blueapi_system.py @@ -36,6 +36,7 @@ TaskResponse, WorkerTask, ) +from blueapi.utils import TILED_PROPOSAL_RE from blueapi.worker.event import ( TaskResult, TaskStatus, @@ -354,7 +355,7 @@ def test_task_metadata_propagated( "user": User.alice, "instrument_session": VALID_INSTRUMENT_SESSION[User.alice], "tiled_access_tags": [ - '{"proposal": 12345, "visit": 1, "beamline": "adsim"}', + '{"proposal": "cm12345", "visit": 1, "beamline": "adsim"}', ], "blueapi_task_id": response.task_id, } @@ -612,7 +613,12 @@ def on_event(event: AnyEvent) -> None: assert stream_resource["run_start"] == start_doc["uid"] assert stream_resource["uri"] == f"file://localhost/tmp/adsim-{scan_id}-det.h5" - tiled_url = f"http://localhost:8407/api/v1/metadata/{start_doc['uid']}" + proposal = TILED_PROPOSAL_RE.match(start_doc["instrument_session"])["proposal"] # type: ignore + tiled_url = ( + "http://localhost:8407/api/v1/metadata/" + f"{start_doc['instrument']}/{proposal}/{start_doc['instrument_session']}/" + f"{start_doc['uid']}" + ) response = requests.get( tiled_url, headers={"authorization": "Bearer " + get_access_token(user)} ) diff --git a/tests/unit_tests/service/test_interface.py b/tests/unit_tests/service/test_interface.py index 892c5ad2ea..1bc77c337e 100644 --- a/tests/unit_tests/service/test_interface.py +++ b/tests/unit_tests/service/test_interface.py @@ -10,10 +10,12 @@ from bluesky.protocols import Stoppable from bluesky.utils import MsgGenerator from bluesky_stomp.messaging import StompClient +from fastapi import status from ophyd_async.epics.motor import Motor -from pydantic import HttpUrl +from pydantic import HttpUrl, SecretStr from pytest_httpx import HTTPXMock from stomp.connect import StompConnection11 as Connection +from tiled.client.utils import ClientError from blueapi.config import ( ApplicationConfig, @@ -22,11 +24,13 @@ NumtrackerConfig, OIDCConfig, ScratchConfig, + ServiceAccount, StompConfig, TiledConfig, ) from blueapi.core.context import BlueskyContext from blueapi.service import interface +from blueapi.service.authentication import TiledAuth from blueapi.service.model import ( DeviceModel, PackageInfo, @@ -255,8 +259,8 @@ def test_subscribers_removed_when_task_not_found( # regression test for #1480 worker = worker_mock() ctx = context_mock() + worker.get_task_by_id.return_value = None worker.begin_task.side_effect = KeyError() - with pytest.raises(KeyError): interface.begin_task(WorkerTask(task_id="missing")) @@ -374,7 +378,7 @@ def test_get_task_by_id( if tiled_enabled: expected_access_tag = { - "proposal": 12345, + "proposal": "cm12345", "visit": 1, "beamline": "ixx", } @@ -436,6 +440,7 @@ def test_remove_tiled_subscriber(worker, context, from_uri, writer): context().tiled_conf = TiledConfig() context().run_engine.subscribe.return_value = 17 worker().worker_events.subscribe.return_value = 42 + worker().get_task_by_id.return_value = None interface.begin_task(task) @@ -476,6 +481,135 @@ def test_remove_tiled_subscriber(worker, context, from_uri, writer): worker().worker_events.unsubscribe.assert_called_once_with(42) +def _existing_task(instrument_session: str = FAKE_INSTRUMENT_SESSION) -> TrackableTask: + return TrackableTask( + task_id="foo_bar", + task=Task( + name="my_plan", + params={}, + metadata={"instrument_session": instrument_session}, + ), + ) + + +@patch("blueapi.service.interface.TiledWriter") +@patch("blueapi.service.interface.from_uri") +@patch("blueapi.service.interface.context") +@patch("blueapi.service.interface.worker") +def test_begin_task_creates_tiled_containers(worker, context, from_uri, writer): + context().numtracker = None + context().tiled_conf = TiledConfig() + context().run_engine.md = {"instrument": "p46"} + worker().get_task_by_id.return_value = _existing_task() + root_client = from_uri() + proposal_client = root_client.__getitem__.return_value + session_client = proposal_client.__getitem__.return_value + final_client = session_client.__getitem__.return_value + + interface.begin_task(WorkerTask(task_id="foo_bar")) + + root_client.create_container.assert_called_once_with( + key="p46", access_tags=[json.dumps({"beamline": "p46"})] + ) + proposal_client.create_container.assert_called_once_with( + key="cm12345", + access_tags=[json.dumps({"beamline": "p46", "proposal": "cm12345"})], + ) + session_client.create_container.assert_called_once_with( + key="cm12345-1", + access_tags=[ + json.dumps({"proposal": "cm12345", "visit": 1, "beamline": "p46"}) + ], + ) + writer.assert_called_once_with(final_client, batch_size=1) + + +@patch("blueapi.service.interface.TiledWriter") +@patch("blueapi.service.interface.from_uri") +@patch("blueapi.service.interface.context") +@patch("blueapi.service.interface.worker") +def test_begin_task_skips_existing_tiled_container(worker, context, from_uri, writer): + context().numtracker = None + context().tiled_conf = TiledConfig() + context().run_engine.md = {"instrument": "p46"} + worker().get_task_by_id.return_value = _existing_task() + root_client = from_uri() + proposal_client = root_client.__getitem__.return_value + session_client = proposal_client.__getitem__.return_value + for client in (root_client, proposal_client, session_client): + client.__contains__ = MagicMock(return_value=True) + + interface.begin_task(WorkerTask(task_id="foo_bar")) + + root_client.create_container.assert_not_called() + proposal_client.create_container.assert_not_called() + session_client.create_container.assert_not_called() + + +@patch("blueapi.service.interface.from_uri") +@patch("blueapi.service.interface.context") +@patch("blueapi.service.interface.worker") +def test_begin_task_raises_for_invalid_instrument_session(worker, context, from_uri): + context().numtracker = None + context().tiled_conf = TiledConfig() + context().run_engine.md = {"instrument": "p46"} + worker().get_task_by_id.return_value = _existing_task( + instrument_session="not-valid" + ) + + with pytest.raises(ValueError, match="Invalid instrument session"): + interface.begin_task(WorkerTask(task_id="foo_bar")) + + +@pytest.mark.parametrize( + "status_code,expect_raises", + [(status.HTTP_409_CONFLICT, False), (status.HTTP_500_INTERNAL_SERVER_ERROR, True)], +) +@patch("blueapi.service.interface.TiledWriter") +@patch("blueapi.service.interface.from_uri") +@patch("blueapi.service.interface.context") +@patch("blueapi.service.interface.worker") +def test_begin_task_handles_tiled_container_create_errors( + worker, context, from_uri, writer, status_code, expect_raises +): + context().numtracker = None + context().tiled_conf = TiledConfig() + context().run_engine.md = {"instrument": "p46"} + worker().get_task_by_id.return_value = _existing_task() + tiled_client = from_uri() + tiled_client.create_container.side_effect = ClientError( + "error", request=MagicMock(), response=MagicMock(status_code=status_code) + ) + + if expect_raises: + with pytest.raises(ClientError): + interface.begin_task(WorkerTask(task_id="foo_bar")) + else: + interface.begin_task(WorkerTask(task_id="foo_bar")) + + +@patch("blueapi.service.interface.TiledWriter") +@patch("blueapi.service.interface.from_uri") +@patch("blueapi.service.interface.context") +@patch("blueapi.service.interface.worker") +def test_begin_task_uses_service_account_auth_for_tiled( + worker, context, from_uri, writer +): + context().numtracker = None + context().tiled_conf = TiledConfig( + authentication=ServiceAccount( + client_id="tiled_writer", + client_secret=SecretStr("secret"), + token_url="https://example.com/token", + ) + ) + worker().get_task_by_id.return_value = None + + interface.begin_task(WorkerTask(task_id="foo_bar")) + + assert isinstance(from_uri.call_args.kwargs["auth"], TiledAuth) + + def test_get_oidc_config(oidc_config: OIDCConfig): interface.set_config(ApplicationConfig(oidc=oidc_config)) assert interface.get_oidc_config() == oidc_config diff --git a/tests/unit_tests/utils/test_access_blob.py b/tests/unit_tests/utils/test_access_blob.py index bb8c17d7c5..02ed00c962 100644 --- a/tests/unit_tests/utils/test_access_blob.py +++ b/tests/unit_tests/utils/test_access_blob.py @@ -8,27 +8,27 @@ [ ( "cm12345-1", - '{"proposal": 12345, "visit": 1, "beamline": "ixx"}', + '{"proposal": "cm12345", "visit": 1, "beamline": "ixx"}', ), ( "cm12345-111", - '{"proposal": 12345, "visit": 111, "beamline": "ixx"}', + '{"proposal": "cm12345", "visit": 111, "beamline": "ixx"}', ), ( "cv12345-1", - '{"proposal": 12345, "visit": 1, "beamline": "ixx"}', + '{"proposal": "cv12345", "visit": 1, "beamline": "ixx"}', ), ( "cm12345678-1", - '{"proposal": 12345678, "visit": 1, "beamline": "ixx"}', + '{"proposal": "cm12345678", "visit": 1, "beamline": "ixx"}', ), ( "cm12345678-111", - '{"proposal": 12345678, "visit": 111, "beamline": "ixx"}', + '{"proposal": "cm12345678", "visit": 111, "beamline": "ixx"}', ), ( "cv12345678-111", - '{"proposal": 12345678, "visit": 111, "beamline": "ixx"}', + '{"proposal": "cv12345678", "visit": 111, "beamline": "ixx"}', ), ], )