From 3623ecd3ad76fe58ddba8493f8c6f4f993f5adaf Mon Sep 17 00:00:00 2001 From: "Guan-Ming (Wesley) Chiu" <105915352+guan404ming@users.noreply.github.com> Date: Tue, 4 Aug 2026 10:27:17 +0800 Subject: [PATCH 1/3] Render multi-branch review choices as a multi-select --- .../common/ai/docs/operators/llm_branch.rst | 4 +- .../providers/common/ai/mixins/approval.py | 22 ++++++-- .../common/ai/operators/llm_branch.py | 6 ++- .../unit/common/ai/mixins/test_approval.py | 51 +++++++++++++++++++ .../common/ai/operators/test_llm_branch.py | 34 +++++++++++-- 5 files changed, 107 insertions(+), 10 deletions(-) diff --git a/providers/common/ai/docs/operators/llm_branch.rst b/providers/common/ai/docs/operators/llm_branch.rst index 52affb9af8659..21558b2277f02 100644 --- a/providers/common/ai/docs/operators/llm_branch.rst +++ b/providers/common/ai/docs/operators/llm_branch.rst @@ -83,8 +83,8 @@ branch(es) and wait for a human reviewer to approve the choice before any downstream task is skipped. The review form shows the LLM's choice and the valid downstream task IDs. When ``allow_modifications=True``, the reviewer can also change the choice — rendered as a dropdown of the downstream task -IDs, or a free-text JSON list of task IDs (e.g. ``["task_a", "task_b"]``) -with ``allow_multiple_branches=True``. The reviewed branch(es) are validated +IDs, or a multi-select of them with ``allow_multiple_branches=True``. The +reviewed branch(es) are validated against the downstream task IDs before branching: .. exampleinclude:: /../../ai/src/airflow/providers/common/ai/example_dags/example_llm_branch.py diff --git a/providers/common/ai/src/airflow/providers/common/ai/mixins/approval.py b/providers/common/ai/src/airflow/providers/common/ai/mixins/approval.py index 21c2e4c9de4eb..aee737384165b 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/mixins/approval.py +++ b/providers/common/ai/src/airflow/providers/common/ai/mixins/approval.py @@ -17,6 +17,7 @@ from __future__ import annotations +import json import logging from datetime import timedelta from typing import TYPE_CHECKING, Any, Protocol @@ -109,7 +110,9 @@ def defer_for_approval( :param modification_schema: JSON schema for the editable ``output`` param when ``allow_modifications=True``. Defaults to ``{"type": "string"}``. Pass e.g. ``{"type": "string", "enum": [...]}`` to render a dropdown - of valid values in the review form. + of valid values in the review form, or ``{"type": "array", "enum": [...]}`` + to render a multi-select; a list submitted by the reviewer is returned + from ``execute_complete`` re-serialized as a compact JSON string. """ from airflow.providers.standard.triggers.hitl import HITLTrigger from airflow.sdk.execution_time.hitl import upsert_hitl_detail @@ -117,6 +120,7 @@ def defer_for_approval( self.validate_approval_prompt() + raw_output = output if isinstance(output, BaseModel): output = output.model_dump_json() elif not isinstance(output, str): @@ -133,9 +137,17 @@ def defer_for_approval( hitl_params: dict[str, dict[str, Any]] = {} if self.allow_modifications: + # The multi-select rendered for an array schema needs the list, not its JSON string + param_value: Any = output + if ( + modification_schema is not None + and modification_schema.get("type") == "array" + and isinstance(raw_output, list) + ): + param_value = raw_output hitl_params = { "output": { - "value": output, + "value": param_value, "description": "Edit the output before approving (optional).", "schema": modification_schema or {"type": "string"}, }, @@ -215,13 +227,17 @@ def execute_complete(self, context: Context, generated_output: str, event: dict[ # when allow_modifications=False, bypassing the read-only approval flow. if getattr(self, "allow_modifications", False) and params_input: modified = params_input.get("output") + if isinstance(modified, list) and all(isinstance(item, str) for item in modified): + # Compact so an unchanged selection compares equal to generated_output + modified = json.dumps(modified, separators=(",", ":")) if modified is not None and not isinstance(modified, str): # On the awaiting_input path nothing upstream schema-validates params_input # (HITLTrigger did on the legacy path), so enforce the string contract here # rather than returning a non-string as the task's output. raise HITLTriggerEventError( { - "error": f"Modified output must be a string, got {type(modified).__name__}.", + "error": f"Modified output must be a string or a list of strings, " + f"got {type(modified).__name__}.", "error_type": "validation", } ) diff --git a/providers/common/ai/src/airflow/providers/common/ai/operators/llm_branch.py b/providers/common/ai/src/airflow/providers/common/ai/operators/llm_branch.py index 620b36e6809fc..c97903477c931 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/operators/llm_branch.py +++ b/providers/common/ai/src/airflow/providers/common/ai/operators/llm_branch.py @@ -56,7 +56,7 @@ class LLMBranchOperator(LLMOperator, BranchMixIn): unselected downstream tasks once a reviewer approves. The review form lists the valid downstream task IDs; with ``allow_modifications=True`` the editable choice is rendered as a dropdown of those IDs (single-branch - mode) or a free-text JSON list (``allow_multiple_branches=True``), and + mode) or a multi-select of them (``allow_multiple_branches=True``), and the reviewed branch(es) are validated against the downstream task IDs before branching. """ @@ -121,7 +121,9 @@ def execute(self, context: Context) -> str | Iterable[str] | None: f"```\nPrompt: {self.prompt}\n\nChosen branch(es): {chosen}\n```" ) modification_schema = ( - None if self.allow_multiple_branches else {"type": "string", "enum": choices} + {"type": "array", "enum": choices} + if self.allow_multiple_branches + else {"type": "string", "enum": choices} ) self.defer_for_approval( # type: ignore[misc] context, branches, body=body, modification_schema=modification_schema diff --git a/providers/common/ai/tests/unit/common/ai/mixins/test_approval.py b/providers/common/ai/tests/unit/common/ai/mixins/test_approval.py index 23c97bd247073..1b56bff9447d4 100644 --- a/providers/common/ai/tests/unit/common/ai/mixins/test_approval.py +++ b/providers/common/ai/tests/unit/common/ai/mixins/test_approval.py @@ -158,6 +158,21 @@ def test_modification_schema_overrides_output_param_schema( param = mock_upsert.call_args[1]["params"]["output"] assert param["schema"] == schema + @patch(HITL_TRIGGER_PATH, autospec=True) + @patch(UPSERT_HITL_PATH) + def test_array_schema_passes_list_param_value( + self, mock_upsert, mock_trigger_cls, approval_op_with_modifications, context + ): + schema = {"type": "array", "enum": ["task_a", "task_b"]} + + approval_op_with_modifications.defer_for_approval(context, ["task_a"], modification_schema=schema) + + param = mock_upsert.call_args[1]["params"]["output"] + assert param["value"] == ["task_a"] + assert param["schema"] == schema + defer_kwargs = approval_op_with_modifications.defer.call_args[1] + assert defer_kwargs["kwargs"]["generated_output"] == '["task_a"]' + @patch(HITL_TRIGGER_PATH, autospec=True) @patch(UPSERT_HITL_PATH) def test_no_modifications_params_empty(self, mock_upsert, mock_trigger_cls, approval_op, context): @@ -292,6 +307,42 @@ def test_approved_with_non_string_modified_output_raises(self, approval_op_with_ {}, generated_output="original output", event=event ) + def test_approved_with_list_modified_output_is_serialized(self, approval_op_with_modifications): + event = { + "chosen_options": ["Approve"], + "responded_by_user": "editor", + "params_input": {"output": ["task_b", "task_c"]}, + } + + result = approval_op_with_modifications.execute_complete( + {}, generated_output='["task_a"]', event=event + ) + + assert result == '["task_b","task_c"]' + + def test_approved_with_unmodified_list_output_returns_original(self, approval_op_with_modifications): + event = { + "chosen_options": ["Approve"], + "responded_by_user": "editor", + "params_input": {"output": ["task_a"]}, + } + + result = approval_op_with_modifications.execute_complete( + {}, generated_output='["task_a"]', event=event + ) + + assert result == '["task_a"]' + + def test_approved_with_non_string_list_items_raises(self, approval_op_with_modifications): + event = { + "chosen_options": ["Approve"], + "responded_by_user": "editor", + "params_input": {"output": ["task_a", 2]}, + } + + with pytest.raises(HITLTriggerEventError, match="must be a string"): + approval_op_with_modifications.execute_complete({}, generated_output='["task_a"]', event=event) + def test_approved_with_unmodified_output(self, approval_op_with_modifications): event = { "chosen_options": ["Approve"], diff --git a/providers/common/ai/tests/unit/common/ai/operators/test_llm_branch.py b/providers/common/ai/tests/unit/common/ai/operators/test_llm_branch.py index 78107daf1386c..0b702f962d4f0 100644 --- a/providers/common/ai/tests/unit/common/ai/operators/test_llm_branch.py +++ b/providers/common/ai/tests/unit/common/ai/operators/test_llm_branch.py @@ -313,10 +313,10 @@ def test_review_form_lists_choices_and_renders_enum_dropdown( @patch("airflow.providers.standard.triggers.hitl.HITLTrigger", autospec=True) @patch("airflow.sdk.execution_time.hitl.upsert_hitl_detail") @patch("airflow.providers.common.ai.operators.llm.PydanticAIHook", autospec=True) - def test_review_form_multi_branch_keeps_string_schema( + def test_review_form_multi_branch_renders_multiselect( self, mock_hook_cls, mock_upsert, mock_trigger_cls, mock_do_branch ): - """With allow_multiple_branches the editable param stays free-text (JSON list).""" + """With allow_multiple_branches the editable param is an array enum (multi-select).""" downstream_enum = Enum("DownstreamTasks", {"task_a": "task_a", "task_b": "task_b"}) mock_agent = MagicMock(spec=["run_sync"]) @@ -338,7 +338,11 @@ def test_review_form_multi_branch_keeps_string_schema( call_kwargs = mock_upsert.call_args.kwargs assert "Valid branches: `task_a`, `task_b`" in call_kwargs["body"] - assert call_kwargs["params"]["output"]["schema"] == {"type": "string"} + assert call_kwargs["params"]["output"]["schema"] == { + "type": "array", + "enum": ["task_a", "task_b"], + } + assert call_kwargs["params"]["output"]["value"] == ["task_a"] @patch("airflow.providers.common.ai.operators.llm.PydanticAIHook", autospec=True) def test_execute_rejects_sequence_prompt_with_require_approval(self, mock_hook_cls): @@ -405,6 +409,30 @@ def test_execute_complete_with_modified_branch(self, mock_do_branch): assert result == "task_b" mock_do_branch.assert_called_once_with(ctx, "task_b") + @patch.object(LLMBranchOperator, "do_branch") + def test_execute_complete_with_multiselect_modified_branches(self, mock_do_branch): + """A list submitted by the multi-select review form branches into those tasks.""" + mock_do_branch.return_value = ["task_b", "task_c"] + op = LLMBranchOperator( + task_id="t", + prompt="p", + llm_conn_id="c", + allow_multiple_branches=True, + allow_modifications=True, + ) + op.downstream_task_ids = {"task_a", "task_b", "task_c"} + event = { + "chosen_options": ["Approve"], + "responded_by_user": "admin", + "params_input": {"output": ["task_b", "task_c"]}, + } + ctx = _make_context() + + result = op.execute_complete(ctx, generated_output='["task_a"]', event=event) + + assert result == ["task_b", "task_c"] + mock_do_branch.assert_called_once_with(ctx, ["task_b", "task_c"]) + @patch.object(LLMBranchOperator, "do_branch") def test_execute_complete_rejects_invalid_modified_branch(self, mock_do_branch): """A reviewer-modified branch outside downstream_task_ids fails validation.""" From 8e3a05a0e1c89800798e24e801f60df343519f80 Mon Sep 17 00:00:00 2001 From: "Guan-Ming (Wesley) Chiu" <105915352+guan404ming@users.noreply.github.com> Date: Tue, 4 Aug 2026 18:07:11 +0800 Subject: [PATCH 2/3] Fix multi-select review form validation --- .../providers/common/ai/mixins/approval.py | 26 ++++++++++++++++--- .../common/ai/operators/llm_branch.py | 2 +- .../unit/common/ai/mixins/test_approval.py | 15 +++++++++-- .../common/ai/operators/test_llm_branch.py | 10 ++++++- 4 files changed, 45 insertions(+), 8 deletions(-) diff --git a/providers/common/ai/src/airflow/providers/common/ai/mixins/approval.py b/providers/common/ai/src/airflow/providers/common/ai/mixins/approval.py index aee737384165b..cccff32bfd150 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/mixins/approval.py +++ b/providers/common/ai/src/airflow/providers/common/ai/mixins/approval.py @@ -110,9 +110,11 @@ def defer_for_approval( :param modification_schema: JSON schema for the editable ``output`` param when ``allow_modifications=True``. Defaults to ``{"type": "string"}``. Pass e.g. ``{"type": "string", "enum": [...]}`` to render a dropdown - of valid values in the review form, or ``{"type": "array", "enum": [...]}`` - to render a multi-select; a list submitted by the reviewer is returned - from ``execute_complete`` re-serialized as a compact JSON string. + of valid values in the review form, or ``{"type": "array", "items": + {"type": "string", "enum": [...]}, "examples": [...]}`` to render a + multi-select (JSON Schema forbids ``enum`` at the array level, so the + options come from ``examples``); a list submitted by the reviewer is + returned from ``execute_complete`` re-serialized as a compact JSON string. """ from airflow.providers.standard.triggers.hitl import HITLTrigger from airflow.sdk.execution_time.hitl import upsert_hitl_detail @@ -227,7 +229,23 @@ def execute_complete(self, context: Context, generated_output: str, event: dict[ # when allow_modifications=False, bypassing the read-only approval flow. if getattr(self, "allow_modifications", False) and params_input: modified = params_input.get("output") - if isinstance(modified, list) and all(isinstance(item, str) for item in modified): + if "output" in params_input and modified is None: + raise HITLTriggerEventError( + { + "error": "Modified output must not be empty; edit it or reject instead.", + "error_type": "validation", + } + ) + if isinstance(modified, list): + for item in modified: + if not isinstance(item, str): + raise HITLTriggerEventError( + { + "error": f"Modified output list items must be strings, " + f"got {type(item).__name__}.", + "error_type": "validation", + } + ) # Compact so an unchanged selection compares equal to generated_output modified = json.dumps(modified, separators=(",", ":")) if modified is not None and not isinstance(modified, str): diff --git a/providers/common/ai/src/airflow/providers/common/ai/operators/llm_branch.py b/providers/common/ai/src/airflow/providers/common/ai/operators/llm_branch.py index c97903477c931..88c299432594b 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/operators/llm_branch.py +++ b/providers/common/ai/src/airflow/providers/common/ai/operators/llm_branch.py @@ -121,7 +121,7 @@ def execute(self, context: Context) -> str | Iterable[str] | None: f"```\nPrompt: {self.prompt}\n\nChosen branch(es): {chosen}\n```" ) modification_schema = ( - {"type": "array", "enum": choices} + {"type": "array", "items": {"type": "string", "enum": choices}, "examples": choices} if self.allow_multiple_branches else {"type": "string", "enum": choices} ) diff --git a/providers/common/ai/tests/unit/common/ai/mixins/test_approval.py b/providers/common/ai/tests/unit/common/ai/mixins/test_approval.py index 1b56bff9447d4..ace3ce7d3d98e 100644 --- a/providers/common/ai/tests/unit/common/ai/mixins/test_approval.py +++ b/providers/common/ai/tests/unit/common/ai/mixins/test_approval.py @@ -163,7 +163,8 @@ def test_modification_schema_overrides_output_param_schema( def test_array_schema_passes_list_param_value( self, mock_upsert, mock_trigger_cls, approval_op_with_modifications, context ): - schema = {"type": "array", "enum": ["task_a", "task_b"]} + choices = ["task_a", "task_b"] + schema = {"type": "array", "items": {"type": "string", "enum": choices}, "examples": choices} approval_op_with_modifications.defer_for_approval(context, ["task_a"], modification_schema=schema) @@ -340,7 +341,17 @@ def test_approved_with_non_string_list_items_raises(self, approval_op_with_modif "params_input": {"output": ["task_a", 2]}, } - with pytest.raises(HITLTriggerEventError, match="must be a string"): + with pytest.raises(HITLTriggerEventError, match="items must be strings, got int"): + approval_op_with_modifications.execute_complete({}, generated_output='["task_a"]', event=event) + + def test_approved_with_cleared_output_raises(self, approval_op_with_modifications): + event = { + "chosen_options": ["Approve"], + "responded_by_user": "editor", + "params_input": {"output": None}, + } + + with pytest.raises(HITLTriggerEventError, match="must not be empty"): approval_op_with_modifications.execute_complete({}, generated_output='["task_a"]', event=event) def test_approved_with_unmodified_output(self, approval_op_with_modifications): diff --git a/providers/common/ai/tests/unit/common/ai/operators/test_llm_branch.py b/providers/common/ai/tests/unit/common/ai/operators/test_llm_branch.py index 0b702f962d4f0..7f5804eba15ef 100644 --- a/providers/common/ai/tests/unit/common/ai/operators/test_llm_branch.py +++ b/providers/common/ai/tests/unit/common/ai/operators/test_llm_branch.py @@ -26,6 +26,8 @@ from airflow.providers.common.ai.operators.llm import LLMOperator from airflow.providers.common.ai.operators.llm_branch import LLMBranchOperator from airflow.providers.common.compat.sdk import TaskDeferred +from airflow.sdk.definitions.param import Param +from airflow.sdk.exceptions import ParamValidationError from tests_common.test_utils.version_compat import AIRFLOW_V_3_1_PLUS, AIRFLOW_V_3_3_PLUS @@ -340,10 +342,16 @@ def test_review_form_multi_branch_renders_multiselect( assert "Valid branches: `task_a`, `task_b`" in call_kwargs["body"] assert call_kwargs["params"]["output"]["schema"] == { "type": "array", - "enum": ["task_a", "task_b"], + "items": {"type": "string", "enum": ["task_a", "task_b"]}, + "examples": ["task_a", "task_b"], } assert call_kwargs["params"]["output"]["value"] == ["task_a"] + schema = call_kwargs["params"]["output"]["schema"] + assert Param(schema=schema).resolve(["task_a"]) == ["task_a"] + with pytest.raises(ParamValidationError): + Param(schema=schema).resolve(["task_x"]) + @patch("airflow.providers.common.ai.operators.llm.PydanticAIHook", autospec=True) def test_execute_rejects_sequence_prompt_with_require_approval(self, mock_hook_cls): """Non-string prompt + require_approval=True fails before the agent runs.""" From 4f5742c4af263b47cd10a018cbb530554f982eb4 Mon Sep 17 00:00:00 2001 From: "Guan-Ming (Wesley) Chiu" <105915352+guan404ming@users.noreply.github.com> Date: Tue, 4 Aug 2026 23:11:33 +0800 Subject: [PATCH 3/3] Fix test imports for older Airflow compat --- .../ai/tests/unit/common/ai/operators/test_llm_branch.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/providers/common/ai/tests/unit/common/ai/operators/test_llm_branch.py b/providers/common/ai/tests/unit/common/ai/operators/test_llm_branch.py index 7f5804eba15ef..7f2d11bf26336 100644 --- a/providers/common/ai/tests/unit/common/ai/operators/test_llm_branch.py +++ b/providers/common/ai/tests/unit/common/ai/operators/test_llm_branch.py @@ -25,9 +25,7 @@ from airflow.providers.common.ai.mixins.approval import LLMApprovalMixin from airflow.providers.common.ai.operators.llm import LLMOperator from airflow.providers.common.ai.operators.llm_branch import LLMBranchOperator -from airflow.providers.common.compat.sdk import TaskDeferred -from airflow.sdk.definitions.param import Param -from airflow.sdk.exceptions import ParamValidationError +from airflow.providers.common.compat.sdk import Param, ParamValidationError, TaskDeferred from tests_common.test_utils.version_compat import AIRFLOW_V_3_1_PLUS, AIRFLOW_V_3_3_PLUS