From 803cbd30a2915bec7c8f5801d6caf6abbae5bcd2 Mon Sep 17 00:00:00 2001 From: PoAn Yang Date: Sat, 10 Oct 2026 21:51:16 +0800 Subject: [PATCH] Fix agent HITL review for list, marker and function output types Signed-off-by: PoAn Yang --- .../providers/common/ai/utils/output_type.py | 17 ++++++++++- .../unit/common/ai/operators/test_agent.py | 25 +++++++++++++++++ .../unit/common/ai/utils/test_output_type.py | 28 +++++++++++++++++++ 3 files changed, 69 insertions(+), 1 deletion(-) diff --git a/providers/common/ai/src/airflow/providers/common/ai/utils/output_type.py b/providers/common/ai/src/airflow/providers/common/ai/utils/output_type.py index 44ca8505b5404..861f630eb465e 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/utils/output_type.py +++ b/providers/common/ai/src/airflow/providers/common/ai/utils/output_type.py @@ -18,7 +18,8 @@ from __future__ import annotations -from typing import Any +import json +from typing import Any, get_origin from pydantic import BaseModel, TypeAdapter, ValidationError @@ -39,12 +40,26 @@ def rehydrate_pydantic_output( pydantic ``TypeAdapter``. When validation fails (reviewer edited the string into something the type rejects), returns ``raw`` unchanged. + pydantic-ai also accepts an ``output_type`` that is neither a class nor a + generic alias: a list of output types, an output marker such as + ``ToolOutput``, or an output function. For these output types, ``raw`` is + parsed with ``json.loads`` instead, and ``raw`` is returned unchanged when + it is not JSON. + When ``serialize_output`` is ``True``, returns the model dumped to a ``dict`` -- matches the operator's ``serialize_output=True`` opt-in for consumers that want the dict shape. """ if output_type is str: return raw + if not isinstance(output_type, type) and get_origin(output_type) is None: + # TypeAdapter raises for a list of output types or for a marker. For an output function, TypeAdapter + # builds a schema for the function's arguments, so validating the reviewed text would call the + # function again with that text as its arguments. + try: + return json.loads(raw) + except (ValueError, TypeError): + return raw try: rehydrated = TypeAdapter(output_type).validate_json(raw) except (ValidationError, ValueError, TypeError): diff --git a/providers/common/ai/tests/unit/common/ai/operators/test_agent.py b/providers/common/ai/tests/unit/common/ai/operators/test_agent.py index d878ed546b279..43886e0d994c1 100644 --- a/providers/common/ai/tests/unit/common/ai/operators/test_agent.py +++ b/providers/common/ai/tests/unit/common/ai/operators/test_agent.py @@ -1133,6 +1133,31 @@ def test_execute_with_hitl_returns_the_approved_output_as_output_type( assert result == expected assert type(result) is type(expected) + @pytest.mark.skipif( + not AIRFLOW_V_3_1_PLUS, reason="Human in the loop is only compatible with Airflow >= 3.1.0" + ) + @patch("airflow.providers.common.ai.operators.agent.AgentOperator.run_hitl_review", autospec=True) + @patch("airflow.providers.common.ai.operators.agent.PydanticAIHook", autospec=True) + def test_execute_with_hitl_parses_the_approved_output_for_a_list_of_output_types( + self, mock_hook_cls, mock_run_hitl, make_mock_run_result + ): + mock_agent = MagicMock(spec=["run_sync", "instrument"]) + mock_agent.run_sync.return_value = make_mock_run_result(Summary(text="ok", score=1.0)) + mock_hook_cls.get_hook.return_value.create_agent.return_value = mock_agent + mock_run_hitl.return_value = '{"text":"ok","score":1.0}' + op = AgentOperator( + task_id="test", + prompt="Summarize", + llm_conn_id="my_llm", + output_type=[Summary, int], + enable_hitl_review=True, + hitl_timeout=timedelta(minutes=5), + ) + + result = op.execute(context=MagicMock()) + + assert result == {"text": "ok", "score": 1.0} + @pytest.mark.skipif( not AIRFLOW_V_3_1_PLUS, reason="Human in the loop is only compatible with Airflow >= 3.1.0" ) diff --git a/providers/common/ai/tests/unit/common/ai/utils/test_output_type.py b/providers/common/ai/tests/unit/common/ai/utils/test_output_type.py index f7e1a4799e269..40739163fb887 100644 --- a/providers/common/ai/tests/unit/common/ai/utils/test_output_type.py +++ b/providers/common/ai/tests/unit/common/ai/utils/test_output_type.py @@ -18,6 +18,7 @@ import pytest from pydantic import BaseModel +from pydantic_ai import ToolOutput from airflow.providers.common.ai.utils.output_type import rehydrate_pydantic_output @@ -26,6 +27,10 @@ class A(BaseModel): x: int +class B(BaseModel): + y: str + + class TestRehydratePydanticOutput: def test_returns_model_instance(self): result = rehydrate_pydantic_output(A, '{"x": 7}', serialize_output=False) @@ -60,3 +65,26 @@ def test_returns_raw_on_schema_mismatch(self): # ``A`` requires ``x: int`` -- this payload should fail validation result = rehydrate_pydantic_output(A, '{"y": "no-x-field"}', serialize_output=False) assert result == '{"y": "no-x-field"}' + + @pytest.mark.parametrize( + ("output_type", "raw", "expected"), + [ + pytest.param([A, B], '{"x": 7}', {"x": 7}, id="list-of-output-types"), + pytest.param(ToolOutput(A), '{"x": 7}', {"x": 7}, id="output-marker"), + pytest.param([A, B], "not-json", "not-json", id="not-json"), + ], + ) + def test_parses_output_type_that_is_not_a_type_as_json(self, output_type, raw, expected): + assert rehydrate_pydantic_output(output_type, raw, serialize_output=False) == expected + + def test_does_not_call_an_output_function_again(self): + calls = [] + + def make_a(x: int) -> A: + calls.append(x) + return A(x=x) + + result = rehydrate_pydantic_output(make_a, '{"x": 7}', serialize_output=False) + + assert result == {"x": 7} + assert calls == []