Skip to content

Commit 166a021

Browse files
authored
fix(trigger): pass task_name when Trigger.get_details fetches details (#1349)
Signed-off-by: J.J. Montgomery <monty@union.ai>
1 parent bfcdc21 commit 166a021

2 files changed

Lines changed: 86 additions & 2 deletions

File tree

src/flyte/remote/_trigger.py

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -383,8 +383,7 @@ async def get_details(self) -> TriggerDetails:
383383
Get detailed information about this trigger.
384384
"""
385385
if not self.details:
386-
details = await TriggerDetails.get.aio(name=self.pb2.id.name.name)
387-
self.details = details
386+
self.details = await TriggerDetails.get.aio(name=self.name, task_name=self.task_name)
388387
return self.details
389388

390389
@property
Lines changed: 85 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,85 @@
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

Comments
 (0)