This is an automated email from the ASF dual-hosted git repository.

kaxil pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/airflow.git


The following commit(s) were added to refs/heads/main by this push:
     new ed72629410a Add branch_descriptions to LLMBranchOperator so the model 
reads what each branch means (#73367)
ed72629410a is described below

commit ed72629410a96de1ac81df80fb1f3dc3048d1ee0
Author: Kaxil Naik <[email protected]>
AuthorDate: Sun Sep 20 14:14:23 2026 +0100

    Add branch_descriptions to LLMBranchOperator so the model reads what each 
branch means (#73367)
    
    The model chooses a branch from an enum of downstream task IDs and nothing
    else. When two branches could plausibly own the same input, the task ID 
gives
    it nothing to decide on, and the only place to explain the branches today is
    the system prompt, detached from the schema.
    
    ``branch_descriptions`` maps a task ID to a description of what choosing 
that
    branch means. The descriptions travel in the output schema as ``anyOf`` of
    ``{const, description}`` entries, the one JSON Schema shape that carries a
    description per value, so a text model's tool schema and pydantic-ai's
    TypeSafe adapter both read each option with its meaning. Validation and the
    output are unchanged: the model still answers with a task ID and the 
operator
    still receives an enum member. Without the parameter the schema is the bare
    enum it always was.
    
    A key that is not a downstream task ID fails before the model call, naming
    the key and the valid task IDs. Unlisted tasks are presented by ID alone.
---
 providers/common/ai/docs/operators/llm_branch.rst  |  52 +++++-
 .../common/ai/example_dags/example_llm_branch.py   |  43 +++++
 .../providers/common/ai/operators/llm_branch.py    |  72 ++++++++-
 .../unit/common/ai/operators/test_llm_branch.py    | 180 ++++++++++++++++++++-
 4 files changed, 337 insertions(+), 10 deletions(-)

diff --git a/providers/common/ai/docs/operators/llm_branch.rst 
b/providers/common/ai/docs/operators/llm_branch.rst
index 20ca0d748bf..0a05dad659b 100644
--- a/providers/common/ai/docs/operators/llm_branch.rst
+++ b/providers/common/ai/docs/operators/llm_branch.rst
@@ -42,6 +42,48 @@ execute based on the prompt:
     :start-after: [START howto_operator_llm_branch_basic]
     :end-before: [END howto_operator_llm_branch_basic]
 
+Describing the Branches
+-----------------------
+
+By default the model sees each branch as its task ID and nothing else. That is
+enough when the IDs speak for themselves and the prompt clearly fits one of
+them. It is not enough when two branches could plausibly own the same input:
+in the example above, a missing password-reset email is a sign-in problem to
+one team and an email problem to another, and nothing tells the model which
+team owns it.
+
+``branch_descriptions`` maps a downstream task ID to a short description of
+what choosing that branch means. The descriptions travel in the output schema
+next to the option they describe, so the model reads each option together
+with its meaning rather than matching prose in the system prompt back to a
+task ID by name:
+
+.. exampleinclude:: 
/../../ai/src/airflow/providers/common/ai/example_dags/example_llm_branch.py
+    :language: python
+    :start-after: [START howto_operator_llm_branch_descriptions]
+    :end-before: [END howto_operator_llm_branch_descriptions]
+
+Three fields, three roles. ``prompt`` is the thing being classified.
+``system_prompt`` is the decision to make and the rules that apply across all
+options, including how to break ties. Each ``branch_descriptions`` entry is
+what selecting that option means: its scope and its boundary cases. A rule
+that applies to one branch belongs in that branch's description; a rule that
+applies to the whole decision belongs in the system prompt. Say each thing
+once, in one place.
+
+A downstream task without an entry is presented by its ID alone, as before,
+so a partial mapping is fine. A key that is not a downstream task ID fails the
+task before the model is called, with the valid task IDs in the message; a
+misspelled key silently turning into an option with no description is
+exactly the problem this parameter exists to prevent. The mapping supports 
Jinja
+templating and works with ``allow_multiple_branches=True`` and with the
+``@task.llm_branch`` decorator.
+
+Descriptions explain the choices; they do not make the model more certain,
+and a text model's structured output carries no confidence to read. With a
+classifier model such as TypeSafe's, the descriptions become the criteria of
+its choice question, which is the text it weighs each option by.
+
 Multiple Branches
 -----------------
 
@@ -120,7 +162,11 @@ How It Works
 At execution time, the operator:
 
 1. Reads ``self.downstream_task_ids`` from the Dag topology.
-2. Creates a dynamic ``Enum`` with one member per downstream task ID.
+2. Creates a dynamic ``Enum`` with one member per downstream task ID, in sorted
+   order so every worker presents the options the same way. With
+   ``branch_descriptions``, the enum's JSON Schema is an ``anyOf`` of
+   ``{"const": <task_id>, "description": <text>}`` entries, which is the one
+   schema shape that carries a description per value.
 3. Passes that enum as ``output_type`` to ``pydantic-ai``, constraining the 
LLM to
    valid task IDs only.
 4. Converts the LLM's structured output to task ID string(s) and calls
@@ -134,6 +180,10 @@ Parameters
 - ``llm_conn_id``: Airflow connection ID for the LLM provider.
 - ``model_id``: Model identifier (e.g. ``"openai:gpt-5"``). Overrides the 
connection's extra field.
 - ``system_prompt``: System-level instructions for the agent. Supports Jinja 
templating.
+- ``branch_descriptions``: Optional mapping of downstream task ID to a 
description of
+  what choosing that branch means, sent to the model in the output schema next 
to the
+  option. Unlisted tasks are presented by ID alone; a key that is not a 
downstream task
+  ID fails the task before the model call. Supports Jinja templating. Default 
``None``.
 - ``allow_multiple_branches``: When ``False`` (default) the LLM returns a 
single
   task ID. When ``True`` the LLM may return one or more task IDs.
 - ``agent_params``: Additional keyword arguments passed to the pydantic-ai 
``Agent``
diff --git 
a/providers/common/ai/src/airflow/providers/common/ai/example_dags/example_llm_branch.py
 
b/providers/common/ai/src/airflow/providers/common/ai/example_dags/example_llm_branch.py
index c94e984a011..63fc4dd83e7 100644
--- 
a/providers/common/ai/src/airflow/providers/common/ai/example_dags/example_llm_branch.py
+++ 
b/providers/common/ai/src/airflow/providers/common/ai/example_dags/example_llm_branch.py
@@ -54,6 +54,49 @@ def example_llm_branch_operator():
 example_llm_branch_operator()
 
 
+# [START howto_operator_llm_branch_descriptions]
+@dag(tags=["example"])
+def example_llm_branch_descriptions():
+    route = LLMBranchOperator(
+        task_id="route_ticket",
+        prompt="User says: 'My password reset email never arrived.'",
+        llm_conn_id="pydanticai_default",
+        system_prompt=(
+            "Route the ticket to the team responsible for resolving it. "
+            "Use the reported problem rather than the team the user asks for."
+        ),
+        branch_descriptions={
+            "handle_auth": (
+                "Sign-in, passwords, 2FA and account lockouts. This team owns 
missing password-reset emails."
+            ),
+            "handle_billing": "Invoices, charges, refunds and plan changes.",
+            "handle_general": (
+                "General support triage: product questions, issues outside the 
other "
+                "teams' responsibilities, and tickets that need clarification."
+            ),
+        },
+    )
+
+    @task
+    def handle_billing():
+        return "Handling billing issue"
+
+    @task
+    def handle_auth():
+        return "Handling auth issue"
+
+    @task
+    def handle_general():
+        return "Handling general issue"
+
+    route >> [handle_billing(), handle_auth(), handle_general()]
+
+
+# [END howto_operator_llm_branch_descriptions]
+
+example_llm_branch_descriptions()
+
+
 # [START howto_operator_llm_branch_multi]
 @dag(tags=["example"])
 def example_llm_branch_multi():
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 3541efa85f9..e6c12eb369a 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
@@ -19,7 +19,7 @@
 from __future__ import annotations
 
 import json
-from collections.abc import Iterable, Sequence
+from collections.abc import Iterable, Mapping, Sequence
 from enum import Enum
 from typing import TYPE_CHECKING, Any
 
@@ -33,6 +33,50 @@ if TYPE_CHECKING:
     from airflow.sdk import Context
 
 
+def _downstream_tasks_enum(
+    task_id: str, downstream_task_ids: Iterable[str], descriptions: 
Mapping[str, str] | None
+) -> type[Enum]:
+    """
+    Build the enum of branch options the model chooses from.
+
+    Sorted so every worker sends the model the same option order: 
``downstream_task_ids``
+    is a set, and set order follows string hashing, which differs between 
processes.
+
+    With ``descriptions``, the enum renders as ``anyOf`` of ``{const, 
description}`` instead
+    of a bare ``enum`` list. That is the one JSON Schema shape that carries a 
description per
+    value, and it is what both a text model's tool schema and pydantic-ai's 
TypeSafe adapter
+    read an option's meaning from. Validation is unchanged: the model still 
has to answer
+    with one of the task IDs, and the output is still an enum member.
+    """
+    task_ids = sorted(downstream_task_ids)
+    if descriptions:
+        unknown = sorted(set(descriptions) - set(task_ids))
+        if unknown:
+            raise ValueError(
+                f"branch_descriptions for {task_id!r} names {unknown}, which 
are not downstream "
+                f"tasks. Downstream tasks: {task_ids}."
+            )
+    enum_cls: type[Enum] = Enum("DownstreamTasks", {name: name for name in 
task_ids})  # type: ignore[misc]
+    if not descriptions:
+        return enum_cls
+
+    described = {name: text for name, text in descriptions.items() if text}
+
+    def json_schema(cls: type[Enum], core_schema: Any, handler: Any) -> 
dict[str, Any]:
+        options: list[dict[str, Any]] = []
+        for member in cls:
+            option: dict[str, Any] = {"const": member.value, "type": "string"}
+            if text := described.get(member.value):
+                option["description"] = text
+            options.append(option)
+        return {"anyOf": options, "title": cls.__name__}
+
+    # pydantic looks this hook up on the type when it builds the schema, so 
attaching it to the
+    # functional-API enum is the same as defining it in a class body.
+    setattr(enum_cls, "__get_pydantic_json_schema__", classmethod(json_schema))
+    return enum_cls
+
+
 class LLMBranchOperator(LLMOperator, BranchMixIn):
     """
     Ask an LLM to choose which downstream task(s) to execute.
@@ -53,6 +97,13 @@ class LLMBranchOperator(LLMOperator, BranchMixIn):
         :class:`~airflow.providers.common.ai.hooks.pydantic_ai.PydanticAIHook`
         for how blank entries in the list are dropped.
     :param system_prompt: System-level instructions for the LLM agent.
+    :param branch_descriptions: Optional mapping of downstream task ID to a 
short
+        description of what choosing that branch means. Descriptions travel in 
the
+        output schema next to the option they describe, so the model reads 
"here is
+        an option, here is what it means" rather than guessing from the task 
ID. A
+        downstream task without an entry is presented by its ID alone, as 
today. A
+        key that is not a downstream task ID fails the task before the model is
+        called. Supports Jinja templating.
     :param allow_multiple_branches: When ``False`` (default) the LLM returns a
         single task ID. When ``True`` the LLM may return one or more task IDs.
     :param fail_on_reject: If ``True``, a rejected review fails the task
@@ -89,11 +140,12 @@ class LLMBranchOperator(LLMOperator, BranchMixIn):
 
     inherits_from_skipmixin = True
 
-    template_fields: Sequence[str] = LLMOperator.template_fields
+    template_fields: Sequence[str] = (*LLMOperator.template_fields, 
"branch_descriptions")
 
     def __init__(
         self,
         *,
+        branch_descriptions: Mapping[str, str] | None = None,
         allow_multiple_branches: bool = False,
         fail_on_reject: bool = False,
         ignore_downstream_trigger_rules: bool = False,
@@ -101,6 +153,7 @@ class LLMBranchOperator(LLMOperator, BranchMixIn):
     ) -> None:
         kwargs.pop("output_type", None)
         super().__init__(**kwargs)
+        self.branch_descriptions = branch_descriptions
         self.allow_multiple_branches = allow_multiple_branches
         self.fail_on_reject = fail_on_reject
         self.ignore_downstream_trigger_rules = ignore_downstream_trigger_rules
@@ -115,13 +168,16 @@ class LLMBranchOperator(LLMOperator, BranchMixIn):
                 "LLMBranchOperator requires at least one downstream task to 
branch into."
             )
 
-        # Sorted so every worker sends the model the same option order. 
downstream_task_ids
-        # is a set, and set order follows string hashing, which differs 
between processes.
-        downstream_tasks_enum = Enum(  # type: ignore[misc]
-            "DownstreamTasks",
-            {task_id: task_id for task_id in sorted(self.downstream_task_ids)},
+        downstream_tasks_enum = _downstream_tasks_enum(
+            self.task_id, self.downstream_task_ids, self.branch_descriptions
+        )
+        output_type: Any = (
+            list[downstream_tasks_enum] if self.allow_multiple_branches else 
downstream_tasks_enum  # type: ignore[valid-type]
         )
-        output_type = list[downstream_tasks_enum] if 
self.allow_multiple_branches else downstream_tasks_enum
+        if self.branch_descriptions:
+            undescribed = sorted(set(self.downstream_task_ids) - 
set(self.branch_descriptions))
+            if undescribed:
+                self.log.debug("Branches presented by task ID alone (no 
description): %s", undescribed)
 
         # Coerced first so a bad rendered value fails before the expensive 
setup below.
         usage_limits = coerce_usage_limits(self.usage_limits)
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 4fca5f34911..2f4b74ef406 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
@@ -22,6 +22,10 @@ from unittest.mock import MagicMock, patch
 from uuid import uuid4
 
 import pytest
+from pydantic import TypeAdapter
+from pydantic_ai import Agent
+from pydantic_ai.messages import ModelResponse, ToolCallPart
+from pydantic_ai.models.function import FunctionModel
 
 from airflow.providers.common.ai.mixins.approval import LLMApprovalMixin
 from airflow.providers.common.ai.operators.llm import LLMOperator
@@ -46,7 +50,7 @@ class TestLLMBranchOperator:
         assert LLMBranchOperator.inherits_from_skipmixin is True
 
     def test_template_fields(self):
-        assert set(LLMBranchOperator.template_fields) == 
set(LLMOperator.template_fields)
+        assert set(LLMBranchOperator.template_fields) == 
{*LLMOperator.template_fields, "branch_descriptions"}
 
     def test_output_type_ignored(self):
         """Passing output_type= doesn't break anything; it's silently 
dropped."""
@@ -237,6 +241,180 @@ class TestLLMBranchOperator:
         output_type = 
mock_hook_cls.get_hook.return_value.create_agent.call_args.kwargs["output_type"]
         assert [m.value for m in output_type] == ["task_a", "task_b", "task_c"]
 
+    @patch.object(LLMBranchOperator, "do_branch")
+    @patch("airflow.providers.common.ai.operators.llm.PydanticAIHook", 
autospec=True)
+    def test_branch_descriptions_land_in_the_output_schema(
+        self, mock_hook_cls, mock_do_branch, make_mock_run_result
+    ):
+        """Each described option carries its description in the schema; an 
undescribed one carries none.
+
+        This is the shape both a text model's tool schema and pydantic-ai's 
TypeSafe adapter read
+        a per-option description from: ``anyOf`` of ``{const, description}``, 
not a bare ``enum``.
+        """
+        mock_agent = MagicMock(spec=["run_sync"])
+        mock_hook_cls.get_hook.return_value.create_agent.return_value = 
mock_agent
+
+        op = LLMBranchOperator(
+            task_id="route",
+            prompt="Pick",
+            llm_conn_id="my_llm",
+            branch_descriptions={
+                "handle_auth": "Sign-in, passwords, 2FA. Owns missing reset 
emails.",
+                "handle_billing": "Invoices, charges, refunds.",
+            },
+        )
+        op.downstream_task_ids = {"handle_general", "handle_billing", 
"handle_auth"}
+        output_type = None
+
+        def capture(**kwargs):
+            nonlocal output_type
+            output_type = kwargs["output_type"]
+            mock_agent.run_sync.return_value = 
make_mock_run_result(output_type.handle_auth)
+            return mock_agent
+
+        mock_hook_cls.get_hook.return_value.create_agent.side_effect = capture
+
+        op.execute(MagicMock())
+
+        schema = TypeAdapter(output_type).json_schema()
+        assert "enum" not in schema
+        assert schema["anyOf"] == [
+            {
+                "const": "handle_auth",
+                "type": "string",
+                "description": "Sign-in, passwords, 2FA. Owns missing reset 
emails.",
+            },
+            {"const": "handle_billing", "type": "string", "description": 
"Invoices, charges, refunds."},
+            {"const": "handle_general", "type": "string"},
+        ]
+        # Validation is still the enum: the output handling downstream is 
unchanged.
+        assert [m.value for m in output_type] == ["handle_auth", 
"handle_billing", "handle_general"]
+        
mock_do_branch.assert_called_once_with(mock_do_branch.call_args.args[0], 
"handle_auth")
+
+    @patch.object(LLMBranchOperator, "do_branch")
+    @patch("airflow.providers.common.ai.operators.llm.PydanticAIHook", 
autospec=True)
+    def test_branch_descriptions_reach_the_model(self, mock_hook_cls, 
mock_do_branch):
+        """Through a real pydantic-ai Agent, the descriptions are in the 
request the model receives.
+
+        ``FunctionModel`` sits where every provider adapter sits and is handed 
the same
+        ``output_tools`` schema, so this is the request as a model sees it, 
not the operator's
+        view of it.
+        """
+        seen: dict = {}
+
+        def model_fn(messages, info):
+            tool = info.output_tools[0]
+            seen["schema"] = tool.parameters_json_schema
+            return ModelResponse(parts=[ToolCallPart(tool.name, {"response": 
"handle_billing"})])
+
+        def create_agent(*, output_type, instructions, **_):
+            return Agent(FunctionModel(model_fn), output_type=output_type, 
instructions=instructions)
+
+        mock_hook_cls.get_hook.return_value.create_agent.side_effect = 
create_agent
+
+        op = LLMBranchOperator(
+            task_id="route",
+            prompt="I was charged twice.",
+            llm_conn_id="my_llm",
+            system_prompt="Route the ticket.",
+            branch_descriptions={"handle_billing": "Invoices, charges, 
refunds."},
+        )
+        op.downstream_task_ids = {"handle_auth", "handle_billing"}
+
+        op.execute(MagicMock())
+
+        options = seen["schema"]["$defs"]["DownstreamTasks"]["anyOf"]
+        assert {o["const"]: o.get("description") for o in options} == {
+            "handle_auth": None,
+            "handle_billing": "Invoices, charges, refunds.",
+        }
+        
mock_do_branch.assert_called_once_with(mock_do_branch.call_args.args[0], 
"handle_billing")
+
+    @patch.object(LLMBranchOperator, "do_branch")
+    @patch("airflow.providers.common.ai.operators.llm.PydanticAIHook", 
autospec=True)
+    def test_branch_descriptions_with_multiple_branches(self, mock_hook_cls, 
mock_do_branch):
+        """With allow_multiple_branches the descriptions sit on the list's 
items."""
+        seen: dict = {}
+
+        def model_fn(messages, info):
+            tool = info.output_tools[0]
+            seen["schema"] = tool.parameters_json_schema
+            return ModelResponse(
+                parts=[ToolCallPart(tool.name, {"response": 
["handle_shipping", "handle_packaging"]})]
+            )
+
+        def create_agent(*, output_type, instructions, **_):
+            return Agent(FunctionModel(model_fn), output_type=output_type, 
instructions=instructions)
+
+        mock_hook_cls.get_hook.return_value.create_agent.side_effect = 
create_agent
+
+        op = LLMBranchOperator(
+            task_id="classify",
+            prompt="Shipping was slow and the box was damaged.",
+            llm_conn_id="my_llm",
+            allow_multiple_branches=True,
+            branch_descriptions={"handle_shipping": "Late or lost 
deliveries."},
+        )
+        op.downstream_task_ids = {"handle_shipping", "handle_packaging"}
+
+        op.execute(MagicMock())
+
+        items = seen["schema"]["properties"]["response"]["items"]
+        assert items == {"$ref": "#/$defs/DownstreamTasks"}
+        options = seen["schema"]["$defs"]["DownstreamTasks"]["anyOf"]
+        assert [o["const"] for o in options] == ["handle_packaging", 
"handle_shipping"]
+        assert options[1]["description"] == "Late or lost deliveries."
+        mock_do_branch.assert_called_once_with(
+            mock_do_branch.call_args.args[0], ["handle_shipping", 
"handle_packaging"]
+        )
+
+    @patch("airflow.providers.common.ai.operators.llm.PydanticAIHook", 
autospec=True)
+    def test_branch_descriptions_unknown_key_fails_before_the_model_call(self, 
mock_hook_cls):
+        """A key that is not a downstream task is a ValueError naming it and 
the valid choices."""
+        op = LLMBranchOperator(
+            task_id="route",
+            prompt="Pick",
+            llm_conn_id="my_llm",
+            branch_descriptions={"handle_genral": "typo", "handle_auth": "ok"},
+        )
+        op.downstream_task_ids = {"handle_auth", "handle_general"}
+
+        with pytest.raises(
+            ValueError, match=r"'route' names 
\['handle_genral'\].*\['handle_auth', 'handle_general'\]"
+        ):
+            op.execute(MagicMock())
+
+        mock_hook_cls.get_hook.return_value.create_agent.assert_not_called()
+
+    @patch.object(LLMBranchOperator, "do_branch")
+    @patch("airflow.providers.common.ai.operators.llm.PydanticAIHook", 
autospec=True)
+    def test_no_branch_descriptions_keeps_the_plain_enum_schema(
+        self, mock_hook_cls, mock_do_branch, make_mock_run_result
+    ):
+        """Without descriptions the schema is the bare enum it always was."""
+        mock_agent = MagicMock(spec=["run_sync"])
+        mock_hook_cls.get_hook.return_value.create_agent.return_value = 
mock_agent
+        op = LLMBranchOperator(task_id="route", prompt="Pick", 
llm_conn_id="my_llm")
+        op.downstream_task_ids = {"task_b", "task_a"}
+        output_type = None
+
+        def capture(**kwargs):
+            nonlocal output_type
+            output_type = kwargs["output_type"]
+            mock_agent.run_sync.return_value = 
make_mock_run_result(output_type.task_a)
+            return mock_agent
+
+        mock_hook_cls.get_hook.return_value.create_agent.side_effect = capture
+
+        op.execute(MagicMock())
+
+        schema = TypeAdapter(output_type).json_schema()
+        assert schema["enum"] == ["task_a", "task_b"]
+        assert "anyOf" not in schema
+
+    def test_branch_descriptions_is_a_template_field(self):
+        assert "branch_descriptions" in LLMBranchOperator.template_fields
+
     def test_execute_raises_on_no_downstream_tasks(self):
         """ValueError when the operator has no downstream tasks."""
         op = LLMBranchOperator(

Reply via email to