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"