Lee-W commented on code in PR #72786:
URL: https://github.com/apache/airflow/pull/72786#discussion_r4034929772
##########
providers/common/ai/tests/unit/common/ai/toolsets/test_managed_agent.py:
##########
@@ -435,3 +456,84 @@ async def test_model_retry_is_not_a_failover(self,
mock_stats):
with pytest.raises(ModelRetry):
await group.invoke("q")
mock_stats.incr.assert_not_called()
+
+
+class TestSafeAgentRef:
+ """agent_ref only labels a call; resolving it must never block the call
itself."""
+
+ async def _call(self, toolset, prompt="what is the number?"):
+ tools = await toolset.get_tools(ctx=None)
+ tool = tools[toolset._tool_name]
+ return await toolset.call_tool(toolset._tool_name, {"prompt": prompt},
None, tool)
+
+ @staticmethod
+ def _group(*members, **kwargs):
+ kwargs.setdefault("tool_name", "ask_resilient")
+ kwargs.setdefault("description", "Answers questions, on whichever
cloud is up.")
+ return FailoverManagedAgentToolset(members=list(members), **kwargs)
+
+ @pytest.mark.asyncio
+ async def
test_call_tool_survives_a_broken_agent_ref_on_a_lone_toolset(self):
+ toolset = BrokenAgentRefManagedAgentToolset(result="ok")
+ assert await self._call(toolset) == "ok"
+
+ @pytest.mark.asyncio
+ async def test_call_tool_survives_a_broken_standby_agent_ref(self):
+ # Kaxil's original repro: a group whose agent_ref join fails must still
+ # try its healthy primary through call_tool.
+ primary = FakeManagedAgentToolset(result="from primary")
+ standby = BrokenAgentRefManagedAgentToolset(result="from standby")
+ group = self._group(primary, standby)
+ assert await self._call(group) == "from primary"
+
+ @pytest.mark.asyncio
+ async def test_failover_survives_a_broken_primary_agent_ref(self):
+ primary = BrokenAgentRefManagedAgentToolset(result="from primary")
+ standby = FakeManagedAgentToolset(result="from standby")
+ group = self._group(primary, standby)
+ assert await group.invoke("q") == "from primary"
+
+ @pytest.mark.asyncio
+ async def test_failover_survives_a_broken_standby_agent_ref(self):
+ primary =
FakeManagedAgentToolset(raises=ManagedAgentInvocationError("down"))
+ standby = BrokenAgentRefManagedAgentToolset(result="from standby")
+ group = self._group(primary, standby)
+ assert await group.invoke("q") == "from standby"
+
+
+class TestInvokedMetric:
+ @pytest.mark.asyncio
+ @mock.patch("airflow.providers.common.ai.toolsets.managed_agent.Stats")
+ async def test_call_tool_emits_invoked_on_success(self, mock_stats):
+ toolset = FakeManagedAgentToolset()
+ tools = await toolset.get_tools(ctx=None)
+ await toolset.call_tool("ask_specialist", {"prompt": "q"}, None,
tools["ask_specialist"])
+ mock_stats.incr.assert_called_once_with(
+ "managed_agent.invoked",
+ tags={"tool": "ask_specialist", "platform": "fake.cloud"},
+ )
+
+ @pytest.mark.asyncio
+ @mock.patch("airflow.providers.common.ai.toolsets.managed_agent.Stats")
+ async def test_call_tool_does_not_emit_invoked_on_failure(self,
mock_stats):
+ toolset = FakeManagedAgentToolset(raises=RuntimeError("503"))
+ tools = await toolset.get_tools(ctx=None)
+ with pytest.raises(RuntimeError):
+ await toolset.call_tool("ask_specialist", {"prompt": "q"}, None,
tools["ask_specialist"])
+ mock_stats.incr.assert_not_called()
+
+ @pytest.mark.asyncio
+ @mock.patch("airflow.providers.common.ai.toolsets.managed_agent.Stats")
+ async def
test_group_call_tool_emits_invoked_once_regardless_of_failover(self,
mock_stats):
+ primary =
FakeManagedAgentToolset(raises=ManagedAgentInvocationError("down"))
+ standby = FakeManagedAgentToolset(result="from standby")
+ group = FailoverManagedAgentToolset(
+ members=[primary, standby],
+ tool_name="ask_resilient",
+ description="Answers questions, on whichever cloud is up.",
+ )
+ tools = await group.get_tools(ctx=None)
+ await group.call_tool("ask_resilient", {"prompt": "q"}, None,
tools["ask_resilient"])
+
+ kinds = [c.args[0] for c in mock_stats.incr.call_args_list]
+ assert kinds == ["managed_agent.failover", "managed_agent.served",
"managed_agent.invoked"]
Review Comment:
Now it asserts the full `call_args_list`, so `platform="failover"` on
`managed_agent.invoked` is pinned alongside the `failover` and `served` tags.
Added the Metrics sentence too.
--
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.
To unsubscribe, e-mail: [email protected]
For queries about this service, please contact Infrastructure at:
[email protected]