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

pierrejeambrun 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 d7c60857bc0 Expose the resolved LLM model name on a dedicated XCom key 
(#74272)
d7c60857bc0 is described below

commit d7c60857bc01723f48be01c44cfffed8d98c65a3
Author: Pierre Jeambrun <[email protected]>
AuthorDate: Thu Oct 8 17:01:59 2026 +0200

    Expose the resolved LLM model name on a dedicated XCom key (#74272)
    
    * Expose the resolved LLM model name on a dedicated XCom key
    
    The model that actually answered a run was only reachable buried inside the
    decision record or the task logs. Downstream tasks, and the UI, need a 
stable
    first-class handle on which model ran, so publish it under a dedicated
    model_name XCom key from both the LLM and the agent operator.
    
    * Fold the resolved model name into the usage XCom instead of a separate key
    
    A separate model_name XCom key meant two lookups to see what answered a
    run and what it cost. Push model_name as a field inside the usage XCom
    instead, so both live under one key -- also giving LLMOperator its own
    usage/cost reporting for the first time, which it computed already but
    never exposed on XCom at all.
    
    * Revert the usage-XCom consolidation; namespace model_name instead
    
    Reverts "Fold the resolved model name into the usage XCom instead of a
    separate key" per @kaxil's review -- a plain "usage" field name is too
    easy for a user's own XCom to collide with. Keep model_name on its own
    key, but namespace it (__AIRFLOW__COMMON_AI_MODEL_NAME__) so it can't
    collide with a user's own "model_name" XCom either. Giving LLMOperator
    its own usage/cost reporting is deferred to a separate PR.
    
    * Drop a reviewer-attribution line from a source comment
    
    A code comment should carry the reasoning, not credit whoever raised it.
    
    * Push the resolved model name from LLMBranchOperator too, and document it
    
    LLMBranchOperator overrides execute() without calling super().execute(),
    so it never published the model_name XCom despite inheriting from
    LLMOperator -- @task.llm_branch silently never got the key the PR
    description promised. Moving the push into a helper both execute()s call
    keeps the two in sync going forward.
    
    Also documents the model_name XCom in observability.rst next to run_id
    and usage, since the "Scope" bullet there only mentioned AgentOperator
    and nothing described which operators actually publish this key.
    
    * Revert the LLMBranchOperator model_name fix for now
    
    Scoping back to AgentOperator and LLMOperator only for this PR;
    LLMBranchOperator (and the other LLMOperator subclasses that override
    execute() without calling super()) will get model_name support in a
    follow-up PR instead of a partial fix here.
    
    This reverts commit af02fc25b608f3649b26ac1d58b84d8e96b302e9.
    
    * Document the model_name XCom, scoped to AgentOperator and LLMOperator
    
    Addresses the other half of kaxil's review: the "Scope" bullet in
    observability.rst only covered run_id and usage, with nothing describing
    which operators publish model_name at all. Calls out that the four
    LLMOperator subclasses overriding execute() don't publish it yet.
    
    * Show the literal model name XCom key in the docs
    
    A Jinja template or REST call can't import MODEL_NAME_XCOM_KEY, so the
    docs need the actual key and an xcom_pull example like the run_id one.
---
 providers/common/ai/docs/observability.rst            | 18 ++++++++++++++++--
 .../airflow/providers/common/ai/operators/agent.py    |  5 ++++-
 .../src/airflow/providers/common/ai/operators/llm.py  | 18 ++++++++++++------
 .../src/airflow/providers/common/ai/utils/logging.py  |  6 ++++++
 .../ai/tests/unit/common/ai/operators/test_agent.py   | 12 +++++++-----
 .../ai/tests/unit/common/ai/operators/test_llm.py     | 19 ++++++++++++++++++-
 6 files changed, 63 insertions(+), 15 deletions(-)

diff --git a/providers/common/ai/docs/observability.rst 
b/providers/common/ai/docs/observability.rst
index 3a79c2c64b1..7e66f1f8197 100644
--- a/providers/common/ai/docs/observability.rst
+++ b/providers/common/ai/docs/observability.rst
@@ -84,11 +84,25 @@ How it works
   Airflow 2 has no task-instance id, so there the key is
   ``<dag_id>/<run_id>/<task_id>/<map_index>/<try_number>``, and spans carry 
the five
   identity keys without ``airflow.task_instance.id``.
+* **Model name.** The model that actually answered is exposed on XCom under the
+  namespaced key ``__AIRFLOW__COMMON_AI_MODEL_NAME__``, separate from 
``run_id`` /
+  ``usage`` above, so a downstream task, a Jinja template, or a REST client 
can read which
+  model responded without parsing the decision record or a span
+  (``ti.xcom_pull(task_ids="my_agent", 
key="__AIRFLOW__COMMON_AI_MODEL_NAME__")``).
+  Only ``AgentOperator`` /
+  ``@task.agent`` and ``LLMOperator`` / ``@task.llm`` push it today; like 
``run_id`` and
+  ``usage``, with ``enable_hitl_review`` it is the initial run's model, not a
+  human-feedback regeneration's, since ``regenerate_with_feedback`` doesn't 
re-push it.
+  ``LLMOperator``'s subclasses that override ``execute()`` instead of calling
+  ``super().execute()`` -- ``LLMBranchOperator``, ``LLMSQLQueryOperator``,
+  ``LLMSchemaCompareOperator``, ``LLMFileAnalysisOperator`` -- don't publish 
it yet;
+  this is upcoming work.
 * **Scope.** The ``run_id`` / ``usage`` XComs come only from ``AgentOperator`` 
and
   ``@task.agent``, and so do the ``airflow.*`` identity attributes, apart from 
a Strands or
-  ADK agent run inside ``agent_framework_tracing`` (see below). The other LLM
+  ADK agent run inside ``agent_framework_tracing`` (see below). 
``LLMOperator`` /
+  ``@task.llm`` additionally pushes the model name above. The other LLM
   operators still emit GenAI spans correlated to the task span by nesting, but
-  without the identity attributes or the run join key.
+  without the identity attributes, the run join key, or the model name.
 * **Content is off by default.** Only token counts, model id, latency, tool
   names, and finish reason are recorded. Prompt and completion text is never
   emitted unless you opt in (see below).
diff --git 
a/providers/common/ai/src/airflow/providers/common/ai/operators/agent.py 
b/providers/common/ai/src/airflow/providers/common/ai/operators/agent.py
index 26cc48d78e5..3c45bed198b 100644
--- a/providers/common/ai/src/airflow/providers/common/ai/operators/agent.py
+++ b/providers/common/ai/src/airflow/providers/common/ai/operators/agent.py
@@ -52,6 +52,7 @@ from airflow.providers.common.ai.observability import (
 )
 from airflow.providers.common.ai.toolsets.sandbox import SandboxToolset
 from airflow.providers.common.ai.utils.logging import (
+    MODEL_NAME_XCOM_KEY,
     format_usage_for_xcom,
     log_run_summary,
     log_run_usage,
@@ -1386,12 +1387,14 @@ class AgentOperator(CancellableAgentRunMixin, 
BaseOperator, HITLReviewMixin):
         context["task_instance"].xcom_push(key="message_history", 
value=transcript)
 
     def _emit_run_metadata(self, context: Context, result: Any, *, usage: 
RunUsage) -> None:
-        """Expose the pydantic-ai run id and token usage on XCom for 
downstream tasks."""
+        """Expose the pydantic-ai run id, resolved model name, and token usage 
on XCom."""
         if not self.do_xcom_push:
             return
         ti = context["task_instance"]
         ti.xcom_push(key="run_id", value=result.run_id)
         ti.xcom_push(key="usage", value=format_usage_for_xcom(usage))
+        if (model_name := getattr(result.response, "model_name", None)) is not 
None:
+            ti.xcom_push(key=MODEL_NAME_XCOM_KEY, value=model_name)
 
     def regenerate_with_feedback(self, *, feedback: str, message_history: Any) 
-> tuple[str, Any]:
         """
diff --git 
a/providers/common/ai/src/airflow/providers/common/ai/operators/llm.py 
b/providers/common/ai/src/airflow/providers/common/ai/operators/llm.py
index 069174ad132..85b8465295b 100644
--- a/providers/common/ai/src/airflow/providers/common/ai/operators/llm.py
+++ b/providers/common/ai/src/airflow/providers/common/ai/operators/llm.py
@@ -43,7 +43,7 @@ from airflow.providers.common.ai.utils.decision import (
     timed_out_record,
     validate_decision_policy,
 )
-from airflow.providers.common.ai.utils.logging import log_run_summary
+from airflow.providers.common.ai.utils.logging import MODEL_NAME_XCOM_KEY, 
log_run_summary
 from airflow.providers.common.ai.utils.output_type import 
rehydrate_pydantic_output
 from airflow.providers.common.ai.utils.usage import coerce_usage_limits
 from airflow.providers.common.compat.notifier import BaseNotifier
@@ -307,6 +307,8 @@ class LLMOperator(CancellableAgentRunMixin, BaseOperator, 
LLMApprovalMixin):
         output = result.output
 
         model_confidence = ModelConfidence.from_result(result)
+        if model_confidence.model is not None:
+            self._push_xcom(context, MODEL_NAME_XCOM_KEY, 
model_confidence.model)
         # Gated by the least confident of the fields that reported a 
confidence; a field whose type
         # reports none is not gated. A bare output type is the one field 
``response``.
         fields = list(model_confidence.confidence) or ["response"]
@@ -398,8 +400,8 @@ class LLMOperator(CancellableAgentRunMixin, BaseOperator, 
LLMApprovalMixin):
                 self._push_decision(context, timed_out_record(decision))
             raise
 
-    def _push_decision(self, context: Context, record: dict[str, Any]) -> None:
-        """Expose what the model proposed, its confidence, and what the gate 
decided, on XCom."""
+    def _push_xcom(self, context: Context, key: str, value: Any) -> None:
+        """Push ``value`` under ``key`` on XCom, honoring ``do_xcom_push`` and 
a missing task instance."""
         if not self.do_xcom_push:
             return
         try:
@@ -409,10 +411,14 @@ class LLMOperator(CancellableAgentRunMixin, BaseOperator, 
LLMApprovalMixin):
         push = getattr(ti, "xcom_push", None)
         if not callable(push):
             # A hand-built context (a dict, or no task instance at all) has 
nowhere to push to; the
-            # record is inspection output, so the run goes on without it.
-            self.log.warning("No task instance in the context; the decision 
record was not pushed to XCom.")
+            # value is inspection output, so the run goes on without it.
+            self.log.warning("No task instance in the context; %r was not 
pushed to XCom.", key)
             return
-        push(key=DECISION_XCOM_KEY, value=record)
+        push(key=key, value=value)
+
+    def _push_decision(self, context: Context, record: dict[str, Any]) -> None:
+        """Expose what the model proposed, its confidence, and what the gate 
decided, on XCom."""
+        self._push_xcom(context, DECISION_XCOM_KEY, record)
 
     def _finalize_decision(
         self, context: Context, event: dict[str, Any], decision: dict[str, 
Any] | None, *, action: Any
diff --git 
a/providers/common/ai/src/airflow/providers/common/ai/utils/logging.py 
b/providers/common/ai/src/airflow/providers/common/ai/utils/logging.py
index b442c870f13..57b1380705e 100644
--- a/providers/common/ai/src/airflow/providers/common/ai/utils/logging.py
+++ b/providers/common/ai/src/airflow/providers/common/ai/utils/logging.py
@@ -94,6 +94,12 @@ def _log_cache_and_cost(logger: Logger | logging.Logger, 
usage: RunUsage) -> Non
         logger.info("LLM run cost: $%s (USD, best-effort)", format(usage.cost, 
"f"))
 
 
+# XCom key the run's resolved model name is published under, so downstream 
tasks and the UI
+# can read which model actually answered without parsing the decision record. 
See the "Model
+# name" entry in docs/observability.rst for exactly which operators publish it.
+MODEL_NAME_XCOM_KEY = "__AIRFLOW__COMMON_AI_MODEL_NAME__"
+
+
 def format_usage_for_xcom(usage: RunUsage) -> dict[str, Any]:
     """Build the XCom ``usage`` payload -- shared by the success and failure 
paths."""
     return {
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 01273b99d8a..bbfcf8ad514 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
@@ -77,6 +77,7 @@ from airflow.providers.common.ai.toolsets.logging import 
LoggingToolset
 from airflow.providers.common.ai.toolsets.mcp import MCPToolset
 from airflow.providers.common.ai.toolsets.sandbox import SandboxToolset
 from airflow.providers.common.ai.toolsets.sql import SQLToolset
+from airflow.providers.common.ai.utils.logging import MODEL_NAME_XCOM_KEY
 from airflow.providers.common.ai.utils.prompt_cache import PromptCaching
 from airflow.providers.common.ai.utils.toolset_base import MaskingToolset
 from airflow.providers.common.ai.utils.toolsets import find_toolset
@@ -1826,10 +1827,10 @@ class TestAgentOperatorMessageHistory:
         op.execute(context=context)
 
         assert "message_history" not in mock_agent.run_sync.call_args.kwargs
-        # The transcript is not emitted without history, but run id + usage 
always are.
+        # The transcript is not emitted without history, but run id + usage + 
model always are.
         pushed_keys = {c.kwargs["key"] for c in 
context["task_instance"].xcom_push.call_args_list}
         assert "message_history" not in pushed_keys
-        assert pushed_keys == {"run_id", "usage"}
+        assert pushed_keys == {"run_id", "usage", MODEL_NAME_XCOM_KEY}
 
     @patch("airflow.providers.common.ai.operators.agent.PydanticAIHook", 
autospec=True)
     def test_transcript_emitted_to_xcom_when_history_set(self, mock_hook_cls, 
make_mock_run_result):
@@ -1844,7 +1845,7 @@ class TestAgentOperatorMessageHistory:
 
         ti = context["task_instance"]
         pushes = {c.kwargs["key"]: c.kwargs["value"] for c in 
ti.xcom_push.call_args_list}
-        assert set(pushes) == {"run_id", "usage", "message_history"}
+        assert set(pushes) == {"run_id", "usage", "message_history", 
MODEL_NAME_XCOM_KEY}
         restored = 
ModelMessagesTypeAdapter.validate_json(pushes["message_history"])
         assert len(restored) == 2
 
@@ -2232,8 +2233,8 @@ class TestAgentOperatorSandboxHandleTemplating:
 
 class TestAgentOperatorRunIdentity:
     @patch("airflow.providers.common.ai.operators.agent.PydanticAIHook", 
autospec=True)
-    def test_run_id_and_usage_pushed_to_xcom(self, mock_hook_cls, 
make_mock_run_result):
-        """The pydantic-ai run id and token usage are exposed on XCom for 
downstream tasks."""
+    def test_run_id_usage_and_model_pushed_to_xcom(self, mock_hook_cls, 
make_mock_run_result):
+        """The pydantic-ai run id, resolved model name, and token usage are 
exposed on XCom."""
         mock_agent = _make_mock_agent("ok", make_mock_run_result)
         mock_agent.run_sync.return_value.run_id = "the-run-id"
         mock_hook_cls.get_hook.return_value.create_agent.return_value = 
mock_agent
@@ -2246,6 +2247,7 @@ class TestAgentOperatorRunIdentity:
             c.kwargs["key"]: c.kwargs["value"] for c in 
context["task_instance"].xcom_push.call_args_list
         }
         assert pushes["run_id"] == "the-run-id"
+        assert pushes[MODEL_NAME_XCOM_KEY] == "test-model"
         assert pushes["usage"] == {
             "requests": 1,
             "input_tokens": 0,
diff --git a/providers/common/ai/tests/unit/common/ai/operators/test_llm.py 
b/providers/common/ai/tests/unit/common/ai/operators/test_llm.py
index b1cee57b5dd..51d79fb8baa 100644
--- a/providers/common/ai/tests/unit/common/ai/operators/test_llm.py
+++ b/providers/common/ai/tests/unit/common/ai/operators/test_llm.py
@@ -35,6 +35,7 @@ from airflow.providers.common.ai.mixins.approval import (
 )
 from airflow.providers.common.ai.operators import llm as llm_module
 from airflow.providers.common.ai.operators.llm import DecisionPolicy, 
LLMOperator
+from airflow.providers.common.ai.utils.logging import MODEL_NAME_XCOM_KEY
 
 from tests_common.test_utils.version_compat import AIRFLOW_V_3_1_PLUS, 
AIRFLOW_V_3_3_PLUS
 
@@ -403,7 +404,23 @@ class TestLLMOperatorConfidenceGate:
             output = op.execute(context)
 
         assert Summary.model_validate(output).text == "t"
-        assert "the decision record was not pushed to XCom" in caplog.text
+        assert "'decision' was not pushed to XCom" in caplog.text
+
+    @patch("airflow.providers.common.ai.operators.llm.PydanticAIHook", 
autospec=True)
+    def test_resolved_model_name_pushed_to_xcom(self, mock_hook_cls, 
make_mock_run_result):
+        """The model that actually answered is exposed on its own namespaced 
XCom key."""
+        mock_agent = MagicMock(spec=["run_sync"])
+        mock_agent.run_sync.return_value = self._result(make_mock_run_result, 
Summary(text="t"), None)
+        mock_hook_cls.get_hook.return_value.create_agent.return_value = 
mock_agent
+        op = LLMOperator(task_id="t", prompt="p", llm_conn_id="c", 
output_type=Summary)
+        context = MagicMock(spec=dict)
+
+        op.execute(context)
+
+        pushes = {
+            c.kwargs["key"]: c.kwargs["value"] for c in 
context["task_instance"].xcom_push.call_args_list
+        }
+        assert pushes[MODEL_NAME_XCOM_KEY] == "jev-1.13.0"
 
     @pytest.mark.skipif(
         not AIRFLOW_V_3_1_PLUS, reason="a reviewing decision_policy needs the 
HITL flow, Airflow >= 3.1"

Reply via email to