|
| 1 | +"""Tests for Trigger.get_details, which fetches (and caches) a trigger's TriggerDetails.""" |
| 2 | + |
| 3 | +from contextlib import contextmanager |
| 4 | +from unittest.mock import AsyncMock, MagicMock, patch |
| 5 | + |
| 6 | +import pytest |
| 7 | +from flyteidl2.common import identifier_pb2 |
| 8 | +from flyteidl2.trigger import trigger_definition_pb2, trigger_service_pb2 |
| 9 | + |
| 10 | +from flyte.remote._trigger import Trigger, TriggerDetails |
| 11 | + |
| 12 | + |
| 13 | +def _trigger() -> Trigger: |
| 14 | + pb2 = trigger_definition_pb2.Trigger(active=True) |
| 15 | + pb2.id.name.name = "t" |
| 16 | + pb2.id.name.task_name = "my_task" |
| 17 | + return Trigger(pb2=pb2) |
| 18 | + |
| 19 | + |
| 20 | +def _details_response() -> trigger_service_pb2.GetTriggerDetailsResponse: |
| 21 | + return trigger_service_pb2.GetTriggerDetailsResponse( |
| 22 | + trigger=trigger_definition_pb2.TriggerDetails( |
| 23 | + id=identifier_pb2.TriggerIdentifier( |
| 24 | + name=identifier_pb2.TriggerName(name="t", task_name="my_task", org="o", project="p", domain="d") |
| 25 | + ) |
| 26 | + ) |
| 27 | + ) |
| 28 | + |
| 29 | + |
| 30 | +@contextmanager |
| 31 | +def _mocked_client(client: MagicMock): |
| 32 | + cfg = MagicMock(org="o", project="p", domain="d") |
| 33 | + with ( |
| 34 | + patch("flyte.remote._trigger.ensure_client"), |
| 35 | + patch("flyte.remote._trigger.get_init_config", return_value=cfg), |
| 36 | + patch("flyte.remote._trigger.get_client", return_value=client), |
| 37 | + ): |
| 38 | + yield |
| 39 | + |
| 40 | + |
| 41 | +@pytest.mark.asyncio |
| 42 | +async def test_get_details_requests_the_trigger_by_name_and_task_name(): |
| 43 | + client = MagicMock() |
| 44 | + client.trigger_service.get_trigger_details = AsyncMock(return_value=_details_response()) |
| 45 | + |
| 46 | + with _mocked_client(client): |
| 47 | + details = await _trigger().get_details() |
| 48 | + |
| 49 | + assert isinstance(details, TriggerDetails) |
| 50 | + assert details.name == "t" |
| 51 | + |
| 52 | + # Both halves of the trigger's identity have to be sent; the name alone does not identify it. |
| 53 | + req = client.trigger_service.get_trigger_details.await_args.kwargs["request"] |
| 54 | + assert req.name.name == "t" |
| 55 | + assert req.name.task_name == "my_task" |
| 56 | + assert (req.name.org, req.name.project, req.name.domain) == ("o", "p", "d") |
| 57 | + |
| 58 | + |
| 59 | +@pytest.mark.asyncio |
| 60 | +async def test_get_details_caches_the_fetched_details(): |
| 61 | + client = MagicMock() |
| 62 | + client.trigger_service.get_trigger_details = AsyncMock(return_value=_details_response()) |
| 63 | + |
| 64 | + trigger = _trigger() |
| 65 | + with _mocked_client(client): |
| 66 | + first = await trigger.get_details() |
| 67 | + second = await trigger.get_details() |
| 68 | + |
| 69 | + assert first is second is trigger.details |
| 70 | + client.trigger_service.get_trigger_details.assert_awaited_once() |
| 71 | + |
| 72 | + |
| 73 | +@pytest.mark.asyncio |
| 74 | +async def test_get_details_returns_preloaded_details_without_a_request(): |
| 75 | + client = MagicMock() |
| 76 | + client.trigger_service.get_trigger_details = AsyncMock(return_value=_details_response()) |
| 77 | + |
| 78 | + preloaded = TriggerDetails(pb2=_details_response().trigger) |
| 79 | + trigger = _trigger() |
| 80 | + trigger.details = preloaded |
| 81 | + |
| 82 | + with _mocked_client(client): |
| 83 | + assert await trigger.get_details() is preloaded |
| 84 | + |
| 85 | + client.trigger_service.get_trigger_details.assert_not_awaited() |
0 commit comments