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 27cd88ca74c Cache repeated agent prompts by default in AgentOperator
(#73994)
27cd88ca74c is described below
commit 27cd88ca74c6d58c874870586b8a03ac6bdfa8b4
Author: Kaxil Naik <[email protected]>
AuthorDate: Thu Oct 1 12:43:57 2026 +0100
Cache repeated agent prompts by default in AgentOperator (#73994)
AgentOperator(cache_prompt=True) turns on each provider's prompt caching:
Anthropic, Bedrock Converse and OpenRouter get breakpoints on the tool
definitions, system prompt and latest message; OpenAI and Gemini already
cache on their own. A provider's own cache settings from the caller, the
connection's model or a spec file take that provider over entirely.
* Leave caching to a caller's CachePoint and document when caching costs
more
A CachePoint in the prompt or history now turns PromptCaching off, since
Anthropic needs a longer-lived entry before a shorter one. The docs cover
the final uncached tail write and map indexes that start together. The
failure-path log gets the cache line too.
---
docs/spelling_wordlist.txt | 2 +
providers/common/ai/docs/operators/agent.rst | 96 ++++++++++
.../providers/common/ai/durable/fingerprint.py | 8 +-
.../airflow/providers/common/ai/operators/agent.py | 18 ++
.../airflow/providers/common/ai/utils/logging.py | 18 +-
.../providers/common/ai/utils/prompt_cache.py | 140 ++++++++++++++
.../common/ai/tests/unit/common/ai/conftest.py | 11 +-
.../tests/unit/common/ai/decorators/test_agent.py | 3 +-
.../common/ai/decorators/test_llm_file_analysis.py | 5 +-
.../unit/common/ai/durable/test_fingerprint.py | 20 ++
.../tests/unit/common/ai/operators/test_agent.py | 77 +++++++-
.../common/ai/operators/test_llm_file_analysis.py | 5 +-
.../ai/tests/unit/common/ai/utils/test_logging.py | 55 +++++-
.../unit/common/ai/utils/test_prompt_cache.py | 209 +++++++++++++++++++++
14 files changed, 648 insertions(+), 19 deletions(-)
diff --git a/docs/spelling_wordlist.txt b/docs/spelling_wordlist.txt
index f4cc0e617dc..a52f6c650d7 100644
--- a/docs/spelling_wordlist.txt
+++ b/docs/spelling_wordlist.txt
@@ -1,5 +1,6 @@
aarch
abc
+AbstractCapability
AbstractFileSystem
AbstractToolset
accessor
@@ -1499,6 +1500,7 @@ rtype
ru
ruleset
runAsUser
+RunContext
runnable
RunQueryOperator
RunQuerySensor
diff --git a/providers/common/ai/docs/operators/agent.rst
b/providers/common/ai/docs/operators/agent.rst
index 2d6473e1966..3852d3de5f5 100644
--- a/providers/common/ai/docs/operators/agent.rst
+++ b/providers/common/ai/docs/operators/agent.rst
@@ -263,6 +263,99 @@ Durable execution
Moved to :doc:`../durable_execution`.
+.. _agent-prompt-caching:
+
+Prompt caching
+^^^^^^^^^^^^^^
+
+Every request an agent makes re-sends its tool definitions, its system prompt
and the
+conversation so far. An agent that calls three tools makes four requests, and
a mapped
+``@task.agent`` makes that many per map index, all starting with the same
system prompt.
+``cache_prompt`` (on by default) asks the provider to keep that repeated
prefix, so later
+requests read it back instead of paying the full input price for it again.
+
+What it turns on depends on the model the connection resolves to:
+
+.. list-table::
+ :header-rows: 1
+
+ * - Provider
+ - What ``cache_prompt=True`` does
+ * - Anthropic (``anthropic:``)
+ - Marks the end of the tool definitions, the end of the system prompt and
the latest
+ message as points to cache up to.
+ * - Bedrock (``bedrock:``) and OpenRouter (``openrouter:``)
+ - The same three marks, for models pydantic-ai knows support caching.
Nothing for the
+ rest.
+ * - OpenAI, Azure OpenAI, Gemini
+ - Nothing. These cache long prompts on their own.
+
+Because each model reads only its own provider's settings, the same flag
covers a
+:doc:`fallback chain <../provider_fallback>` that spans providers.
+
+On Anthropic, a 5-minute cache write costs 1.25x the normal input price and a
read costs
+0.1x (less on some newer models), so a prefix read back even once costs less
than sending it
+twice. A prompt shorter than the model's minimum length for caching (512 to
4,096 tokens,
+depending on the model) is not cached and costs nothing extra. See Anthropic's
+`prompt caching guide
<https://platform.claude.com/docs/en/build-with-claude/prompt-caching>`__
+for the per-model minimums and prices.
+
+Caching costs more than it saves in three cases:
+
+- **A single long request.** An agent with no tools that sends a long prompt
once and is not
+ run again within five minutes pays the write and never reads it back.
+- **A large final tool result.** The latest message is written to the cache on
every request,
+ including the last one, which nothing reads. When the last tool returns much
more than the
+ system prompt and tool definitions add up to, as a query returning thousands
of rows can,
+ that final write costs more than the earlier reads saved.
+- **Map indexes that start together.** A cache entry exists only once the
response that wrote
+ it has started, so map indexes that all start at the same moment each write
their own copy.
+ Indexes that start later read it back.
+
+Turn it off for such a task:
+
+.. code-block:: python
+
+ AgentOperator(
+ task_id="summarize_quarter",
+ prompt="Summarize the attached report.",
+ llm_conn_id="anthropic_default",
+ system_prompt=long_style_guide,
+ cache_prompt=False,
+ )
+
+To choose what is cached or for how long, set the provider's own settings in
+``agent_params["model_settings"]``. Setting any ``anthropic_cache*`` key hands
Anthropic
+caching back to you, and ``cache_prompt`` adds nothing for Anthropic; the same
holds for
+``bedrock_cache*`` and ``openrouter_cache*``. A ``CachePoint`` in the prompt
or the message
+history hands caching back to you for every provider. For example, a mapped
task whose
+instances run further apart than five minutes can keep the system prompt for
an hour, at 2x
+the input price for each write instead of 1.25x:
+
+.. code-block:: python
+
+ AgentOperator.partial(
+ task_id="classify_ticket",
+ llm_conn_id="anthropic_default",
+ system_prompt=long_taxonomy,
+ agent_params={
+ "model_settings": {
+ "anthropic_cache_instructions": "1h",
+ "anthropic_cache_tool_definitions": "1h",
+ }
+ },
+ ).expand(prompt=tickets)
+
+When the provider reports cache activity, the task log shows it under the run
summary:
+
+.. code-block:: text
+
+ LLM run complete: model=claude-sonnet-4-5, requests=2, tool_calls=1,
input_tokens=..., ...
+ LLM prompt cache: cache_read_tokens=..., cache_write_tokens=...
+
+``input_tokens`` includes both counts. With :doc:`../observability` turned on,
each
+request's GenAI span carries them too.
+
Parameters
----------
@@ -344,6 +437,9 @@ Parameters
- ``code_mode``: When ``True``, wraps the agent's tools in a single
``run_code``
tool that the model drives by writing Python, executed in the Monty sandbox.
Requires the ``code-mode`` extra. Default ``False``. See :ref:`code-mode`.
+- ``cache_prompt``: Ask the provider to cache the tool definitions, system
prompt and
+ conversation so later requests read them back at a discount. Default
``True``; a no-op for
+ providers that cache on their own. See :ref:`agent-prompt-caching`.
- ``message_history``: Prior conversation to seed a multi-turn session, as a
list
of pydantic-ai ``ModelMessage`` objects or their JSON form (``str`` /
``bytes``).
When set, the post-run transcript is pushed to XCom under the key
diff --git
a/providers/common/ai/src/airflow/providers/common/ai/durable/fingerprint.py
b/providers/common/ai/src/airflow/providers/common/ai/durable/fingerprint.py
index 6b5a2d490a5..bc4a291ea92 100644
--- a/providers/common/ai/src/airflow/providers/common/ai/durable/fingerprint.py
+++ b/providers/common/ai/src/airflow/providers/common/ai/durable/fingerprint.py
@@ -50,6 +50,8 @@ from pydantic import TypeAdapter
from pydantic_ai.messages import ModelMessagesTypeAdapter
from pydantic_ai.models import ModelRequestParameters
+from airflow.providers.common.ai.utils.prompt_cache import
PROMPT_CACHE_SETTING_NAMES
+
if TYPE_CHECKING:
from pydantic_ai.messages import ModelMessage
from pydantic_ai.settings import ModelSettings
@@ -65,8 +67,10 @@ _VOLATILE_MESSAGE_KEYS = ("timestamp", "run_id",
"conversation_id")
# fingerprint: changing them should not invalidate a cached response, and some
# (``timeout`` can be an ``httpx.Timeout``) are not JSON-serializable, which
# would otherwise force the whole fingerprint to ``None`` and silently disable
-# replay verification for every step.
-_TRANSPORT_ONLY_SETTINGS = frozenset({"timeout"})
+# replay verification for every step. Prompt cache settings only decide what
the
+# provider keeps for the next request, so ``cache_prompt`` can change between
+# attempts without re-running the steps the previous one completed.
+_TRANSPORT_ONLY_SETTINGS = frozenset({"timeout"}) | PROMPT_CACHE_SETTING_NAMES
def _content_settings(model_settings: ModelSettings | None) -> dict[str, Any]
| None:
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 967ab412cff..974c7697cdd 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
@@ -57,6 +57,7 @@ from airflow.providers.common.ai.utils.logging import (
wrap_toolsets_for_logging,
)
from airflow.providers.common.ai.utils.output_type import
rehydrate_pydantic_output
+from airflow.providers.common.ai.utils.prompt_cache import PromptCaching
from airflow.providers.common.ai.utils.toolset_base import ensure_masked
from airflow.providers.common.ai.utils.toolsets import iter_toolsets
from airflow.providers.common.ai.utils.usage import coerce_usage_limits
@@ -366,6 +367,19 @@ class AgentOperator(CancellableAgentRunMixin,
BaseOperator, HITLReviewMixin):
stable per-step call order that code mode does not guarantee), whether
code mode comes from this flag or from a ``CodeMode`` capability.
Default ``False``.
+ :param cache_prompt: When ``True`` (default), asks the provider to cache
the
+ tool definitions, system prompt and conversation so far, so the next
+ request in the run -- and a mapped task's other instances within the
+ cache lifetime -- reads them back at a fraction of the input price
instead
+ of paying for them again. Turns on prompt caching for Anthropic models
and
+ for Bedrock and OpenRouter models that support it; a no-op for OpenAI
and
+ Gemini, which cache long prompts on their own. A provider's own cache
+ settings in ``agent_params["model_settings"]`` or a spec file take
+ precedence: setting any ``anthropic_cache*`` key leaves Anthropic
caching
+ entirely to you, and a ``CachePoint`` in the prompt or message history
+ leaves all of it to you. Set ``False`` where a cache write is rarely
read
+ back, such as a single long request that is not mapped. See
+ :ref:`agent-prompt-caching` for when caching costs more than it saves.
:param message_history: Prior conversation to seed the run with, for
multi-turn sessions that span task runs. Accepts a ``list`` of
pydantic-ai ``ModelMessage`` objects, or their JSON form as ``str`` /
@@ -476,6 +490,7 @@ class AgentOperator(CancellableAgentRunMixin, BaseOperator,
HITLReviewMixin):
usage_limits: UsageLimits | dict[str, Any] | None = None,
durable: bool = False,
code_mode: bool = False,
+ cache_prompt: bool = True,
message_history: list[ModelMessage] | str | bytes | None = None,
# Agent feedback parameters
enable_hitl_review: bool = False,
@@ -510,6 +525,7 @@ class AgentOperator(CancellableAgentRunMixin, BaseOperator,
HITLReviewMixin):
self.durable = durable
self.code_mode = code_mode
+ self.cache_prompt = cache_prompt
# Populated per run in ``execute`` when durable=True. Declared here so
# ``_build_agent`` -- also reached via ``regenerate_with_feedback``
@@ -750,6 +766,8 @@ class AgentOperator(CancellableAgentRunMixin, BaseOperator,
HITLReviewMixin):
capabilities = self._build_durable_capabilities(capabilities,
storage, counter)
if self.code_mode:
capabilities.append(_build_code_mode())
+ if self.cache_prompt:
+ capabilities.append(PromptCaching())
if capabilities:
extra_kwargs["capabilities"] = capabilities
return self.llm_hook.create_agent(
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 6aba0947941..51edfbf8f05 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
@@ -57,10 +57,7 @@ def log_run_summary(
usage.output_tokens,
usage.total_tokens,
)
- if usage.cost is not None:
- # %s on a small Decimal renders scientific notation (e.g. "7.5E-7");
format as
- # plain decimal so cheap runs show a readable dollar amount.
- logger.info("LLM run cost: $%s (USD, best-effort)", format(usage.cost,
"f"))
+ _log_cache_and_cost(logger, usage)
if tool_names := _extract_tool_sequence(result):
logger.info("Tool call sequence: %s", " -> ".join(tool_names))
@@ -80,7 +77,20 @@ def log_run_usage(logger: Logger | logging.Logger, usage:
RunUsage, *, outcome:
usage.output_tokens,
usage.total_tokens,
)
+ _log_cache_and_cost(logger, usage)
+
+
+def _log_cache_and_cost(logger: Logger | logging.Logger, usage: RunUsage) ->
None:
+ if usage.cache_read_tokens or usage.cache_write_tokens:
+ # Part of input_tokens, broken out so the effect of ``cache_prompt``
shows up in the log.
+ logger.info(
+ "LLM prompt cache: cache_read_tokens=%s, cache_write_tokens=%s",
+ usage.cache_read_tokens,
+ usage.cache_write_tokens,
+ )
if usage.cost is not None:
+ # %s on a small Decimal renders scientific notation (e.g. "7.5E-7");
format as
+ # plain decimal so cheap runs show a readable dollar amount.
logger.info("LLM run cost: $%s (USD, best-effort)", format(usage.cost,
"f"))
diff --git
a/providers/common/ai/src/airflow/providers/common/ai/utils/prompt_cache.py
b/providers/common/ai/src/airflow/providers/common/ai/utils/prompt_cache.py
new file mode 100644
index 00000000000..291e1d0e0d8
--- /dev/null
+++ b/providers/common/ai/src/airflow/providers/common/ai/utils/prompt_cache.py
@@ -0,0 +1,140 @@
+# Licensed to the Apache Software Foundation (ASF) under one
+# or more contributor license agreements. See the NOTICE file
+# distributed with this work for additional information
+# regarding copyright ownership. The ASF licenses this file
+# to you under the Apache License, Version 2.0 (the
+# "License"); you may not use this file except in compliance
+# with the License. You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing,
+# software distributed under the License is distributed on an
+# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, either express or implied. See the License for the
+# specific language governing permissions and limitations
+# under the License.
+"""Prompt caching that works the same way whichever provider the connection
resolves to."""
+
+from __future__ import annotations
+
+from dataclasses import KW_ONLY, dataclass
+from typing import TYPE_CHECKING, Any
+
+from pydantic_ai.capabilities import AbstractCapability
+from pydantic_ai.messages import CachePoint, ModelRequest, UserPromptPart
+from pydantic_ai.settings import ModelSettings, merge_model_settings
+
+if TYPE_CHECKING:
+ from pydantic_ai import RunContext
+ from pydantic_ai.agent import AgentModelSettings
+ from pydantic_ai.models.anthropic import AnthropicModelSettings
+ from pydantic_ai.models.bedrock import BedrockModelSettings
+ from pydantic_ai.models.openrouter import OpenRouterModelSettings
+
+# Anthropic writes a cache entry only at a breakpoint, so each of these marks
a prefix a later
+# request can read back: the tool schemas, then the system prompt (the one
that survives a
+# different user prompt per mapped task), then the conversation so far, which
pydantic-ai
+# re-marks on every request of the run. ``anthropic_cache_messages`` rather
than the
+# top-level ``anthropic_cache``: both put the moving breakpoint on the last
message, but
+# pydantic-ai documents the per-block form as the one for Anthropic-compatible
gateways that
+# lack the top-level parameter, and Bedrock and Vertex fall back to the
per-block form anyway.
+_ANTHROPIC_CACHE_SETTINGS: AnthropicModelSettings = {
+ "anthropic_cache_tool_definitions": True,
+ "anthropic_cache_instructions": True,
+ "anthropic_cache_messages": True,
+}
+# The same three breakpoints for Bedrock's Converse API and for OpenRouter.
pydantic-ai adds
+# them only for models whose profile supports prompt caching, so they are a
no-op for the
+# rest of either catalog.
+_BEDROCK_CACHE_SETTINGS: BedrockModelSettings = {
+ "bedrock_cache_tool_definitions": True,
+ "bedrock_cache_instructions": True,
+ "bedrock_cache_messages": True,
+}
+_OPENROUTER_CACHE_SETTINGS: OpenRouterModelSettings = {
+ "openrouter_cache_tool_definitions": True,
+ "openrouter_cache_instructions": True,
+ "openrouter_cache_messages": True,
+}
+# Each provider family's settings, keyed by the prefix every one of its cache
settings shares.
+# OpenAI and Gemini cache long prompts on their own, so they need nothing
here, and a model
+# ignores settings meant for another provider, which is what lets one set
cover a fallback
+# chain that spans providers.
+_CACHE_SETTINGS_BY_PREFIX: tuple[tuple[str, ModelSettings], ...] = (
+ ("anthropic_cache", _ANTHROPIC_CACHE_SETTINGS),
+ ("bedrock_cache", _BEDROCK_CACHE_SETTINGS),
+ ("openrouter_cache", _OPENROUTER_CACHE_SETTINGS),
+)
+
+# Every cache setting of those three families, plus the ``anthropic_cache`` a
caller may set in
+# place of ours. They change what the provider keeps, never what the model
answers.
+PROMPT_CACHE_SETTING_NAMES = frozenset({"anthropic_cache"}).union(
+ *(settings.keys() for _, settings in _CACHE_SETTINGS_BY_PREFIX)
+)
+
+
+def _has_cache_point(ctx: RunContext[Any]) -> bool:
+ """Whether the prompt or the conversation so far carries a ``CachePoint``
marker."""
+ contents = [ctx.prompt] if ctx.prompt is not None else []
+ contents.extend(
+ part.content
+ for message in ctx.messages
+ if isinstance(message, ModelRequest)
+ for part in message.parts
+ if isinstance(part, UserPromptPart)
+ )
+ return any(
+ isinstance(item, CachePoint)
+ for content in contents
+ if not isinstance(content, str)
+ for item in content
+ )
+
+
+def _fill_cache_settings(ctx: RunContext[Any]) -> ModelSettings:
+ """
+ Return the cache settings for each provider family the run has not
configured itself.
+
+ ``ctx.model_settings`` already holds the model's own settings and the
agent's, so a
+ cache setting from either, including one read from a spec file, leaves
that provider's
+ caching entirely to the caller. Taking over the whole family rather than
filling
+ individual keys matters for Anthropic: ``anthropic_cache`` and
+ ``anthropic_cache_messages`` cannot be combined, and a caller who set one
of them
+ would otherwise get a request pydantic-ai refuses to send.
+
+ A ``CachePoint`` in the prompt or the message history leaves caching to
the caller for
+ every provider. Its lifetime may be longer than ours, and Anthropic
requires a
+ longer-lived cache entry to come before a shorter-lived one, which the
tool definitions
+ and system prompt, marked ahead of any message, would break.
+ """
+ if _has_cache_point(ctx):
+ return ModelSettings()
+ configured = ctx.model_settings or {}
+ settings: ModelSettings | None = None
+ for prefix, defaults in _CACHE_SETTINGS_BY_PREFIX:
+ if not any(name.startswith(prefix) for name in configured):
+ settings = merge_model_settings(settings, defaults)
+ return settings or ModelSettings()
+
+
+@dataclass
+class PromptCaching(AbstractCapability[Any]):
+ """
+ Ask the model's provider to cache the repeated prefix of each request.
+
+ Backs ``AgentOperator(cache_prompt=True)``. A provider-agnostic setting
does not exist in
+ pydantic-ai, so this turns on each provider's own: Anthropic models
(direct, Bedrock,
+ Vertex), and Bedrock Converse and OpenRouter models that support caching,
are marked;
+ OpenAI and Gemini already cache automatically. Settings the agent or its
model already
+ carry for a provider win over these, and a ``CachePoint`` in the prompt or
history turns
+ them all off.
+ """
+
+ _: KW_ONLY
+ id: str | None = "prompt_caching"
+
+ def get_model_settings(self) -> AgentModelSettings[Any]:
+ # A callable so it sees the settings merged before it: capability
settings otherwise
+ # override the agent's, and the agent's are what a caller configures.
+ return _fill_cache_settings
diff --git a/providers/common/ai/tests/unit/common/ai/conftest.py
b/providers/common/ai/tests/unit/common/ai/conftest.py
index 18bc46a35d0..c7dc1ce28d3 100644
--- a/providers/common/ai/tests/unit/common/ai/conftest.py
+++ b/providers/common/ai/tests/unit/common/ai/conftest.py
@@ -19,6 +19,7 @@ from __future__ import annotations
from unittest.mock import MagicMock
import pytest
+from pydantic_ai.usage import RunUsage
from tests_common.test_utils.version_compat import AIRFLOW_V_3_1_PLUS
@@ -83,7 +84,15 @@ def make_mock_run_result():
mock_result = MagicMock()
mock_result.output = output
mock_result.usage = MagicMock(
- requests=1, tool_calls=0, input_tokens=0, output_tokens=0,
total_tokens=0, cost=cost
+ spec=RunUsage,
+ requests=1,
+ tool_calls=0,
+ input_tokens=0,
+ output_tokens=0,
+ total_tokens=0,
+ cache_read_tokens=0,
+ cache_write_tokens=0,
+ cost=cost,
)
mock_result.response = MagicMock(model_name="test-model")
mock_result.all_messages.return_value = []
diff --git a/providers/common/ai/tests/unit/common/ai/decorators/test_agent.py
b/providers/common/ai/tests/unit/common/ai/decorators/test_agent.py
index 65190ab6105..a04ec481c4b 100644
--- a/providers/common/ai/tests/unit/common/ai/decorators/test_agent.py
+++ b/providers/common/ai/tests/unit/common/ai/decorators/test_agent.py
@@ -26,6 +26,7 @@ from pydantic_ai.toolsets.function import FunctionToolset
from airflow.providers.common.ai.decorators.agent import
_AgentDecoratedOperator
from airflow.providers.common.ai.toolsets.logging import LoggingToolset
+from airflow.providers.common.ai.utils.prompt_cache import PromptCaching
from airflow.providers.common.ai.utils.toolset_base import MaskingToolset
try:
@@ -196,7 +197,7 @@ class TestAgentDecoratedOperator:
op.execute(context=_make_context())
create_call =
mock_hook_cls.get_hook.return_value.create_agent.call_args
- assert create_call.kwargs["capabilities"] == [thinking]
+ assert create_call.kwargs["capabilities"] == [thinking,
PromptCaching()]
@requires_typed_xcom
@patch("airflow.providers.common.ai.operators.agent.PydanticAIHook",
autospec=True)
diff --git
a/providers/common/ai/tests/unit/common/ai/decorators/test_llm_file_analysis.py
b/providers/common/ai/tests/unit/common/ai/decorators/test_llm_file_analysis.py
index b4bfb040bfa..aa45df1458b 100644
---
a/providers/common/ai/tests/unit/common/ai/decorators/test_llm_file_analysis.py
+++
b/providers/common/ai/tests/unit/common/ai/decorators/test_llm_file_analysis.py
@@ -19,6 +19,7 @@ from __future__ import annotations
from unittest.mock import ANY, MagicMock, patch
import pytest
+from pydantic_ai.usage import RunUsage
from airflow.providers.common.ai.decorators.llm_file_analysis import
_LLMFileAnalysisDecoratedOperator
from airflow.providers.common.ai.utils.file_analysis import FileAnalysisRequest
@@ -28,12 +29,14 @@ def _make_mock_run_result(output):
mock_result = MagicMock(spec=["output", "usage", "response",
"all_messages"])
mock_result.output = output
mock_result.usage = MagicMock(
- spec=["requests", "tool_calls", "input_tokens", "output_tokens",
"total_tokens", "cost"],
+ spec=RunUsage,
requests=1,
tool_calls=0,
input_tokens=0,
output_tokens=0,
total_tokens=0,
+ cache_read_tokens=0,
+ cache_write_tokens=0,
cost=None,
)
mock_result.response = MagicMock(spec=["model_name"],
model_name="test-model")
diff --git
a/providers/common/ai/tests/unit/common/ai/durable/test_fingerprint.py
b/providers/common/ai/tests/unit/common/ai/durable/test_fingerprint.py
index 94e517d1cc3..d555336d955 100644
--- a/providers/common/ai/tests/unit/common/ai/durable/test_fingerprint.py
+++ b/providers/common/ai/tests/unit/common/ai/durable/test_fingerprint.py
@@ -173,6 +173,26 @@ class TestModelRequestFingerprint:
assert no_timeout == float_timeout == httpx_timeout
+ def test_prompt_cache_settings_excluded_from_fingerprint(self):
+ """Turning ``cache_prompt`` on or off between attempts must not re-run
cached steps."""
+ uncached = fingerprint_model_request(
+ "m", make_messages(), {"temperature": 0.2},
ModelRequestParameters()
+ )
+ cached = fingerprint_model_request(
+ "m",
+ make_messages(),
+ {
+ "temperature": 0.2,
+ "anthropic_cache": True,
+ "anthropic_cache_messages": True,
+ "bedrock_cache_instructions": "1h",
+ "openrouter_cache_tool_definitions": True,
+ },
+ ModelRequestParameters(),
+ )
+
+ assert uncached == cached
+
def test_content_settings_still_count_when_timeout_present(self):
"""Stripping timeout must not drop content settings sharing the
dict."""
low = fingerprint_model_request(
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 7d7baedbcb2..be384d21a05 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
@@ -24,6 +24,7 @@ from contextlib import nullcontext
from datetime import timedelta
from decimal import Decimal
from types import ModuleType, SimpleNamespace
+from typing import Any
from unittest.mock import ANY, MagicMock, PropertyMock, call, patch
import pytest
@@ -76,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.prompt_cache import PromptCaching
from airflow.providers.common.ai.utils.toolset_base import MaskingToolset
from airflow.providers.common.ai.utils.toolsets import find_toolset
from airflow.providers.common.ai.utils.usage_budget import (
@@ -765,7 +767,7 @@ class TestAgentOperatorExecute:
# On 3.3+ the agent may also end on a tool call awaiting approval.
expected_output_type = [str, DeferredToolRequests] if
AIRFLOW_V_3_3_PLUS else str
mock_hook_cls.get_hook.return_value.create_agent.assert_called_once_with(
- output_type=expected_output_type, instructions="You are helpful."
+ output_type=expected_output_type, instructions="You are helpful.",
capabilities=[PromptCaching()]
)
mock_agent.run_sync.assert_called_once_with(
"What is the answer?", usage_limits=None, run_id="ti-1",
cancellation_token=ANY, usage=ANY
@@ -835,13 +837,19 @@ class TestAgentOperatorExecute:
@patch("airflow.providers.common.ai.operators.agent.PydanticAIHook",
autospec=True)
def test_code_mode_default_off_no_capabilities(self, mock_hook_cls,
make_mock_run_result):
- """code_mode defaults to False, so no capabilities are injected."""
+ """code_mode defaults to False, so with cache_prompt off no
capabilities are injected."""
mock_hook_cls.get_hook.return_value.create_agent.return_value =
_make_mock_agent(
"ok", make_mock_run_result
)
- op = AgentOperator(task_id="t", prompt="hi", llm_conn_id="my_llm",
toolsets=[MagicMock()])
- op.execute(context=MagicMock())
+ op = AgentOperator(
+ task_id="t",
+ prompt="hi",
+ llm_conn_id="my_llm",
+ toolsets=[MagicMock(spec=AbstractToolset)],
+ cache_prompt=False,
+ )
+ op.execute(context=_make_context())
create_call =
mock_hook_cls.get_hook.return_value.create_agent.call_args
assert "capabilities" not in create_call[1]
@@ -855,9 +863,14 @@ class TestAgentOperatorExecute:
)
op = AgentOperator(
- task_id="t", prompt="hi", llm_conn_id="my_llm",
toolsets=[MagicMock()], code_mode=True
+ task_id="t",
+ prompt="hi",
+ llm_conn_id="my_llm",
+ toolsets=[MagicMock(spec=AbstractToolset)],
+ code_mode=True,
+ cache_prompt=False,
)
- op.execute(context=MagicMock())
+ op.execute(context=_make_context())
create_call =
mock_hook_cls.get_hook.return_value.create_agent.call_args
assert create_call[1]["capabilities"] == ["CM"]
@@ -878,9 +891,10 @@ class TestAgentOperatorExecute:
prompt="hi",
llm_conn_id="my_llm",
code_mode=True,
+ cache_prompt=False,
agent_params={"capabilities": ["existing"]},
)
- op.execute(context=MagicMock())
+ op.execute(context=_make_context())
create_call =
mock_hook_cls.get_hook.return_value.create_agent.call_args
assert create_call[1]["capabilities"] == ["existing", "CM"]
@@ -917,6 +931,53 @@ class TestAgentOperatorExecute:
with pytest.raises(ValueError, match="durable=True and
code_mode=True"):
AgentOperator(task_id="t", prompt="hi", llm_conn_id="my_llm",
durable=True, code_mode=True)
+ @patch("airflow.providers.common.ai.operators.agent.PydanticAIHook",
autospec=True)
+ def test_cache_prompt_default_on_appends_capability_last(self,
mock_hook_cls, make_mock_run_result):
+ """cache_prompt defaults to True and adds PromptCaching after any user
capability."""
+ mock_hook_cls.get_hook.return_value.create_agent.return_value =
_make_mock_agent(
+ "ok", make_mock_run_result
+ )
+
+ op = AgentOperator(
+ task_id="t", prompt="hi", llm_conn_id="my_llm",
agent_params={"capabilities": ["existing"]}
+ )
+ op.execute(context=_make_context())
+
+ create_call =
mock_hook_cls.get_hook.return_value.create_agent.call_args
+ assert create_call[1]["capabilities"] == ["existing", PromptCaching()]
+ assert op.agent_params["capabilities"] == ["existing"]
+
+ @pytest.mark.parametrize(
+ ("cache_prompt", "expected_anthropic_messages"),
+ [pytest.param(True, True, id="on"), pytest.param(False, None,
id="off")],
+ )
+ @patch("airflow.providers.common.ai.operators.agent.PydanticAIHook",
autospec=True)
+ def test_cache_prompt_reaches_the_model_request(
+ self, mock_hook_cls, cache_prompt, expected_anthropic_messages
+ ):
+ """The cache settings arrive on the model request, alongside the
caller's own settings."""
+ seen: dict[str, Any] = {}
+
+ def respond(messages: list[ModelMessage], info: AgentInfo) ->
ModelResponse:
+ seen.update(info.model_settings or {})
+ return ModelResponse(parts=[TextPart("ok")])
+
+ mock_hook_cls.get_hook.return_value.create_agent.side_effect = lambda
**kw: Agent(
+ FunctionModel(respond), **kw
+ )
+
+ op = AgentOperator(
+ task_id="t",
+ prompt="hi",
+ llm_conn_id="my_llm",
+ cache_prompt=cache_prompt,
+ agent_params={"model_settings": {"temperature": 0}},
+ )
+
op.execute(context=_make_context(task_state_store=_make_task_state_store_accessor()))
+
+ assert seen["temperature"] == 0
+ assert seen.get("anthropic_cache_messages") is
expected_anthropic_messages
+
@requires_typed_xcom
@patch("airflow.providers.common.ai.operators.agent.PydanticAIHook",
autospec=True)
def test_execute_structured_output(self, mock_hook_cls,
make_mock_run_result):
@@ -1202,7 +1263,7 @@ class TestAgentOperatorCapabilities:
op.execute(context=_make_context())
create_call =
mock_hook_cls.get_hook.return_value.create_agent.call_args
- assert create_call.kwargs["capabilities"] == [thinking, search]
+ assert create_call.kwargs["capabilities"] == [thinking, search,
PromptCaching()]
@patch("airflow.providers.common.ai.operators.agent.PydanticAIHook",
autospec=True)
def test_capabilities_in_both_places_are_refused(self, mock_hook_cls):
diff --git
a/providers/common/ai/tests/unit/common/ai/operators/test_llm_file_analysis.py
b/providers/common/ai/tests/unit/common/ai/operators/test_llm_file_analysis.py
index ffff1374f65..ccd29d0de53 100644
---
a/providers/common/ai/tests/unit/common/ai/operators/test_llm_file_analysis.py
+++
b/providers/common/ai/tests/unit/common/ai/operators/test_llm_file_analysis.py
@@ -23,6 +23,7 @@ from uuid import uuid4
import pytest
from pydantic import BaseModel
+from pydantic_ai.usage import RunUsage
from airflow.providers.common.ai.operators.llm_file_analysis import
LLMFileAnalysisOperator
from airflow.providers.common.ai.utils.file_analysis import FileAnalysisRequest
@@ -57,12 +58,14 @@ def _make_mock_run_result(output):
mock_result = MagicMock(spec=["output", "usage", "response",
"all_messages"])
mock_result.output = output
mock_result.usage = MagicMock(
- spec=["requests", "tool_calls", "input_tokens", "output_tokens",
"total_tokens", "cost"],
+ spec=RunUsage,
requests=1,
tool_calls=0,
input_tokens=0,
output_tokens=0,
total_tokens=0,
+ cache_read_tokens=0,
+ cache_write_tokens=0,
cost=None,
)
mock_result.response = MagicMock(spec=["model_name"],
model_name="test-model")
diff --git a/providers/common/ai/tests/unit/common/ai/utils/test_logging.py
b/providers/common/ai/tests/unit/common/ai/utils/test_logging.py
index 241500bcf27..828574de4e7 100644
--- a/providers/common/ai/tests/unit/common/ai/utils/test_logging.py
+++ b/providers/common/ai/tests/unit/common/ai/utils/test_logging.py
@@ -20,6 +20,7 @@ import logging
from decimal import Decimal
from unittest.mock import MagicMock
+import pytest
from pydantic import BaseModel
from pydantic_ai import Agent
from pydantic_ai.exceptions import ModelAPIError
@@ -59,7 +60,9 @@ def _make_mock_result(model_name="gpt-5", tool_names=None,
usage_kwargs=None, co
"total_tokens": 3359,
}
result = MagicMock()
- result.usage = MagicMock(cost=cost, **usage_kwargs)
+ result.usage = MagicMock(
+ spec=RunUsage, cost=cost, **{"cache_read_tokens": 0,
"cache_write_tokens": 0, **usage_kwargs}
+ )
result.response = MagicMock(model_name=model_name)
messages: list = []
@@ -147,6 +150,46 @@ class TestLogRunSummary:
records = [r for r in caplog.records if r.name ==
"test.log_run_summary"]
assert not any("LLM run cost" in r.message for r in records)
+ def test_no_cache_tokens_skips_cache_line(self, caplog):
+ logger = logging.getLogger("test.log_run_summary")
+ result = _make_mock_result()
+
+ with caplog.at_level(logging.INFO, logger="test.log_run_summary"):
+ log_run_summary(logger, result)
+
+ records = [r for r in caplog.records if r.name ==
"test.log_run_summary"]
+ assert not any("prompt cache" in r.message for r in records)
+
+ @pytest.mark.parametrize(
+ ("cache_read_tokens", "cache_write_tokens"),
+ [
+ pytest.param(2900, 0, id="read-only"),
+ pytest.param(0, 3000, id="write-only"),
+ pytest.param(2900, 3000, id="read-and-write"),
+ ],
+ )
+ def test_cache_tokens_logged_after_the_usage_line(self, caplog,
cache_read_tokens, cache_write_tokens):
+ logger = logging.getLogger("test.log_run_summary")
+ result = _make_mock_result(
+ usage_kwargs={
+ "requests": 2,
+ "tool_calls": 1,
+ "input_tokens": 6000,
+ "output_tokens": 40,
+ "total_tokens": 6040,
+ "cache_read_tokens": cache_read_tokens,
+ "cache_write_tokens": cache_write_tokens,
+ }
+ )
+
+ with caplog.at_level(logging.INFO, logger="test.log_run_summary"):
+ log_run_summary(logger, result)
+
+ records = [r for r in caplog.records if r.name ==
"test.log_run_summary"]
+ assert records[1].message == (
+ f"LLM prompt cache: cache_read_tokens={cache_read_tokens},
cache_write_tokens={cache_write_tokens}"
+ )
+
def test_cost_set_logs_cost_line_with_value(self, caplog):
logger = logging.getLogger("test.log_run_summary")
result = _make_mock_result(cost=Decimal("0.0123"))
@@ -199,6 +242,16 @@ class TestLogRunSummaryUsageOverride:
class TestLogRunUsage:
+ def test_logs_cache_tokens_on_the_failure_path(self, caplog):
+ logger = logging.getLogger("test.log_run_usage")
+ usage = RunUsage(requests=1, input_tokens=5000,
cache_write_tokens=4000)
+
+ with caplog.at_level(logging.INFO, logger="test.log_run_usage"):
+ log_run_usage(logger, usage, outcome="failed")
+
+ records = [r for r in caplog.records if r.name == "test.log_run_usage"]
+ assert records[1].message == "LLM prompt cache: cache_read_tokens=0,
cache_write_tokens=4000"
+
def test_logs_usage_fields_and_outcome(self):
logger = MagicMock(spec=logging.Logger)
usage = RunUsage(requests=2, tool_calls=1, input_tokens=10,
output_tokens=5)
diff --git
a/providers/common/ai/tests/unit/common/ai/utils/test_prompt_cache.py
b/providers/common/ai/tests/unit/common/ai/utils/test_prompt_cache.py
new file mode 100644
index 00000000000..e6c37ad98cd
--- /dev/null
+++ b/providers/common/ai/tests/unit/common/ai/utils/test_prompt_cache.py
@@ -0,0 +1,209 @@
+# Licensed to the Apache Software Foundation (ASF) under one
+# or more contributor license agreements. See the NOTICE file
+# distributed with this work for additional information
+# regarding copyright ownership. The ASF licenses this file
+# to you under the Apache License, Version 2.0 (the
+# "License"); you may not use this file except in compliance
+# with the License. You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing,
+# software distributed under the License is distributed on an
+# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, either express or implied. See the License for the
+# specific language governing permissions and limitations
+# under the License.
+from __future__ import annotations
+
+import json
+from typing import Any
+
+import httpx2
+import pytest
+from anthropic import AsyncAnthropic
+from pydantic_ai import Agent
+from pydantic_ai.messages import (
+ CachePoint,
+ ModelMessage,
+ ModelRequest,
+ ModelResponse,
+ TextPart,
+ UserPromptPart,
+)
+from pydantic_ai.models.anthropic import AnthropicModel
+from pydantic_ai.models.function import AgentInfo, FunctionModel
+from pydantic_ai.providers.anthropic import AnthropicProvider
+
+from airflow.providers.common.ai.utils.prompt_cache import
PROMPT_CACHE_SETTING_NAMES, PromptCaching
+
+ALL_DEFAULTS = {
+ "anthropic_cache_tool_definitions": True,
+ "anthropic_cache_instructions": True,
+ "anthropic_cache_messages": True,
+ "bedrock_cache_tool_definitions": True,
+ "bedrock_cache_instructions": True,
+ "bedrock_cache_messages": True,
+ "openrouter_cache_tool_definitions": True,
+ "openrouter_cache_instructions": True,
+ "openrouter_cache_messages": True,
+}
+
+
+def _run_and_capture_settings(
+ *,
+ model_settings: Any = None,
+ run_model_settings: Any = None,
+ model_own_settings: Any = None,
+ capabilities: list[Any] | None = None,
+ prompt: Any = "hi",
+ message_history: list[ModelMessage] | None = None,
+) -> dict[str, Any]:
+ """Run a real agent against a FunctionModel and return the settings its
request carried."""
+ seen: dict[str, Any] = {}
+
+ def respond(messages: list[ModelMessage], info: AgentInfo) ->
ModelResponse:
+ seen.update(info.model_settings or {})
+ return ModelResponse(parts=[TextPart("ok")])
+
+ agent = Agent(
+ FunctionModel(respond, settings=model_own_settings),
+ instructions="Be brief.",
+ model_settings=model_settings,
+ capabilities=[PromptCaching()] if capabilities is None else
capabilities,
+ )
+ agent.run_sync(prompt, model_settings=run_model_settings,
message_history=message_history)
+ return seen
+
+
+class TestPromptCaching:
+ def test_turns_on_every_provider_family_by_default(self):
+ assert _run_and_capture_settings() == ALL_DEFAULTS
+
+ def test_keeps_unrelated_agent_settings(self):
+ settings = _run_and_capture_settings(model_settings={"temperature": 0})
+
+ assert settings == {"temperature": 0, **ALL_DEFAULTS}
+
+ @pytest.mark.parametrize(
+ ("caller_settings", "skipped_prefix"),
+ [
+ pytest.param({"anthropic_cache": True}, "anthropic_cache",
id="anthropic-automatic"),
+ pytest.param({"anthropic_cache_instructions": "1h"},
"anthropic_cache", id="anthropic-ttl"),
+ pytest.param({"anthropic_cache_messages": False},
"anthropic_cache", id="anthropic-off"),
+ pytest.param({"bedrock_cache_messages": False}, "bedrock_cache",
id="bedrock-off"),
+ pytest.param({"openrouter_cache_messages": "1h"},
"openrouter_cache", id="openrouter-ttl"),
+ ],
+ )
+ def test_a_caller_cache_setting_takes_over_that_provider_family(self,
caller_settings, skipped_prefix):
+ settings = _run_and_capture_settings(model_settings=caller_settings)
+
+ untouched = {k: v for k, v in ALL_DEFAULTS.items() if not
k.startswith(skipped_prefix)}
+ assert settings == {**caller_settings, **untouched}
+
+ def test_the_model_own_cache_settings_take_over_that_provider_family(self):
+ settings =
_run_and_capture_settings(model_own_settings={"anthropic_cache": "1h"})
+
+ assert settings["anthropic_cache"] == "1h"
+ assert "anthropic_cache_messages" not in settings
+ assert settings["bedrock_cache_messages"] is True
+
+ def test_callable_agent_settings_are_seen(self):
+ settings = _run_and_capture_settings(model_settings=lambda ctx:
{"anthropic_cache": True})
+
+ assert settings["anthropic_cache"] is True
+ assert "anthropic_cache_messages" not in settings
+
+ def test_run_settings_still_win(self):
+ settings =
_run_and_capture_settings(run_model_settings={"anthropic_cache_messages":
False})
+
+ assert settings["anthropic_cache_messages"] is False
+
+ def test_no_capability_means_no_cache_settings(self):
+ assert _run_and_capture_settings(capabilities=[]) == {}
+
+ def test_a_cache_point_in_the_prompt_leaves_caching_to_the_caller(self):
+ settings = _run_and_capture_settings(prompt=["long document",
CachePoint(ttl="1h"), "question"])
+
+ assert settings == {}
+
+ def test_a_cache_point_in_the_history_leaves_caching_to_the_caller(self):
+ history: list[ModelMessage] = [
+ ModelRequest(parts=[UserPromptPart(["long document",
CachePoint(ttl="1h")])]),
+ ModelResponse(parts=[TextPart("noted")]),
+ ]
+
+ settings = _run_and_capture_settings(prompt="question",
message_history=history)
+
+ assert settings == {}
+
+ def
test_setting_names_are_the_defaults_plus_automatic_anthropic_caching(self):
+ assert {*ALL_DEFAULTS, "anthropic_cache"} == PROMPT_CACHE_SETTING_NAMES
+
+
+class TestPromptCachingOnTheAnthropicWire:
+ """What an Anthropic model actually sends, captured at the HTTP
transport."""
+
+ @staticmethod
+ def _send(*, model_settings: Any = None, prompt: Any = "hi") -> dict[str,
Any]:
+ bodies: list[dict[str, Any]] = []
+
+ def handler(request: httpx2.Request) -> httpx2.Response:
+ bodies.append(json.loads(request.content))
+ return httpx2.Response(
+ 200,
+ json={
+ "id": "msg_1",
+ "type": "message",
+ "role": "assistant",
+ "model": "claude-sonnet-4-5",
+ "content": [{"type": "text", "text": "ok"}],
+ "stop_reason": "end_turn",
+ "stop_sequence": None,
+ "usage": {"input_tokens": 10, "output_tokens": 1},
+ },
+ )
+
+ client = AsyncAnthropic(
+ api_key="test",
http_client=httpx2.AsyncClient(transport=httpx2.MockTransport(handler))
+ )
+ model = AnthropicModel("claude-sonnet-4-5",
provider=AnthropicProvider(anthropic_client=client))
+ agent = Agent(
+ model, instructions="Be brief.", model_settings=model_settings,
capabilities=[PromptCaching()]
+ )
+
+ @agent.tool_plain
+ def lookup(key: str) -> str:
+ """Look a key up."""
+ return key
+
+ agent.run_sync(prompt)
+ (body,) = bodies
+ return body
+
+ def test_marks_tools_system_prompt_and_last_message(self):
+ body = self._send()
+
+ assert body["tools"][-1]["cache_control"] == {"type": "ephemeral",
"ttl": "5m"}
+ assert body["system"][-1]["cache_control"] == {"type": "ephemeral",
"ttl": "5m"}
+ assert body["messages"][-1]["content"][-1]["cache_control"] ==
{"type": "ephemeral", "ttl": "5m"}
+ assert "cache_control" not in body
+
+ def test_caller_automatic_caching_is_sent_instead_of_rejected(self):
+ """``anthropic_cache`` cannot be combined with
``anthropic_cache_messages``, so ours step aside."""
+ body = self._send(model_settings={"anthropic_cache": True})
+
+ assert body["cache_control"] == {"type": "ephemeral", "ttl": "5m"}
+ assert "cache_control" not in body["system"][-1]
+
+ def
test_caller_long_lived_cache_point_is_not_preceded_by_shorter_ones(self):
+ """Anthropic needs a longer-lived entry ahead of a shorter one, so
ours stay out."""
+ body = self._send(prompt=["long document", CachePoint(ttl="1h"),
"question"])
+
+ marked = [
+ block["cache_control"]
+ for section in (body["tools"], body["system"], *(m["content"] for
m in body["messages"]))
+ for block in section
+ if "cache_control" in block
+ ]
+ assert marked == [{"type": "ephemeral", "ttl": "1h"}]