jeff3071 commented on code in PR #72202:
URL: https://github.com/apache/airflow/pull/72202#discussion_r4145446955


##########
providers/common/ai/tests/unit/common/ai/operators/test_agent.py:
##########
@@ -693,6 +695,75 @@ def test_enable_tool_logging_false_skips_wrapping(self, 
mock_hook_cls, make_mock
 
         create_call = 
mock_hook_cls.get_hook.return_value.create_agent.call_args
         assert create_call[1]["toolsets"] == [mock_toolset]
+        assert "capabilities" not in create_call[1]
+
+    @pytest.mark.parametrize("tool_source", ["tools", "capabilities"])
+    @patch("airflow.providers.common.ai.operators.agent.PydanticAIHook", 
autospec=True)
+    def test_tool_logging_wraps_agent_param_tools(self, mock_hook_cls, caplog, 
tool_source):
+        """Tool logging wraps the complete toolset assembled by pydantic-ai."""
+
+        def my_tool() -> str:
+            return "tool-result"
+
+        def model_fn(messages, info):
+            saw_return = any(isinstance(p, ToolReturnPart) for m in messages 
for p in getattr(m, "parts", []))
+            if saw_return:
+                return ModelResponse(parts=[TextPart(content="done")])
+            return ModelResponse(parts=[ToolCallPart(tool_name="my_tool", 
args={}, tool_call_id="c1")])
+
+        mock_hook_cls.get_hook.return_value.create_agent.side_effect = lambda 
**kw: Agent(
+            FunctionModel(model_fn), **kw
+        )
+        configured_tools = [my_tool] if tool_source == "tools" else 
[Toolset(FunctionToolset([my_tool]))]
+        op = AgentOperator(
+            task_id="test",
+            prompt="Do something",
+            llm_conn_id="my_llm",
+            agent_params={tool_source: configured_tools},
+        )
+
+        with caplog.at_level(logging.INFO):
+            result = op.execute(context=_make_context())
+
+        assert result == "done"
+        assert any(record.message == "::group::Tool call: my_tool" for record 
in caplog.records)
+
+    @patch("airflow.providers.common.ai.operators.agent.PydanticAIHook", 
autospec=True)
+    def test_tool_logging_handles_non_json_mapping_keys(self, mock_hook_cls, 
caplog):
+        def my_tool(values: dict[UUID, int]) -> int:
+            return sum(values.values())
+
+        def model_fn(messages, info):
+            saw_return = any(isinstance(p, ToolReturnPart) for m in messages 
for p in getattr(m, "parts", []))
+            if saw_return:
+                return ModelResponse(parts=[TextPart(content="done")])
+            return ModelResponse(
+                parts=[
+                    ToolCallPart(
+                        tool_name="my_tool",
+                        args={"values": 
{"12345678-1234-5678-1234-567812345678": 1}},
+                        tool_call_id="c1",
+                    )
+                ]
+            )
+
+        mock_hook_cls.get_hook.return_value.create_agent.side_effect = lambda 
**kw: Agent(
+            FunctionModel(model_fn), **kw
+        )
+        op = AgentOperator(
+            task_id="test",
+            prompt="Do something",
+            llm_conn_id="my_llm",
+            agent_params={"tools": [my_tool]},
+        )
+
+        with caplog.at_level(logging.DEBUG):

Review Comment:
   Updated logger level to fix test fail



-- 
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]

Reply via email to