From cd7d399d497cecd00872e46fc7e9b04856a9822d Mon Sep 17 00:00:00 2001 From: Wei Lee Date: Tue, 18 Aug 2026 21:24:12 +0800 Subject: [PATCH] Record OpenAI Responses token usage and response id in XCom OpenAIResponseOperator returned only the aggregated output text, so the response id needed to chain a follow-up call and the token counts that are the run's only cost signal were discarded on every task. Both are now available to downstream tasks without dropping to the hook. The counts are tokens, not money: the OpenAI response carries no cost field. --- .../airflow/providers/openai/hooks/openai.py | 17 ++++++++++ .../providers/openai/operators/openai.py | 16 ++++++++-- .../tests/unit/openai/hooks/test_openai.py | 31 ++++++++++++++++++- .../unit/openai/operators/test_openai.py | 15 ++++++--- 4 files changed, 71 insertions(+), 8 deletions(-) diff --git a/providers/openai/src/airflow/providers/openai/hooks/openai.py b/providers/openai/src/airflow/providers/openai/hooks/openai.py index 97dcc4ceec8dc..b76e89837d56c 100644 --- a/providers/openai/src/airflow/providers/openai/hooks/openai.py +++ b/providers/openai/src/airflow/providers/openai/hooks/openai.py @@ -272,6 +272,23 @@ def create_response(self, input: Any, model: str = "gpt-4o-mini", **kwargs: Any) """ return self.conn.responses.create(model=model, input=input, **kwargs) + @staticmethod + def summarize_response_usage(response: Response) -> dict[str, Any] | None: + """ + Flatten a Responses API call's token usage into an XCom-safe mapping. + + Dumps the model rather than copying a fixed list of fields, so a token-usage + dimension the API adds later is not silently dropped. ``mode="json"`` keeps the + result JSON-serializable for XCom. There is no cost field to flatten: the OpenAI + response reports token counts only, never a monetary amount. + + :param response: A response returned by :meth:`create_response`. ``usage`` is + optional on the SDK model, so ``None`` is returned when it is absent. + """ + if response.usage is None: + return None + return response.usage.model_dump(mode="json") + def get_response(self, response_id: str, **kwargs: Any) -> Response: """ Retrieve a previously created model response. diff --git a/providers/openai/src/airflow/providers/openai/operators/openai.py b/providers/openai/src/airflow/providers/openai/operators/openai.py index dfed6e48d51a9..5556a37a99d3a 100644 --- a/providers/openai/src/airflow/providers/openai/operators/openai.py +++ b/providers/openai/src/airflow/providers/openai/operators/openai.py @@ -83,9 +83,11 @@ class OpenAIResponseOperator(BaseOperator): """ Operator that generates a model response using the OpenAI Responses API. - The operator is synchronous and returns the response's aggregated output text. For - ``previous_response_id`` chaining, ``background=True`` responses, or access to the full - structured response, use :class:`~airflow.providers.openai.hooks.openai.OpenAIHook` directly. + The operator is synchronous and returns the response's aggregated output text; the + response id is also pushed to XCom (see below), so a downstream task can pick it up + for ``previous_response_id`` chaining without going through the hook. For + ``background=True`` responses, or access to the full structured response, use + :class:`~airflow.providers.openai.hooks.openai.OpenAIHook` directly. :param conn_id: The OpenAI connection ID to use. :param input_text: The input prompt for the model. This can be a string or a structured list of @@ -99,6 +101,12 @@ class OpenAIResponseOperator(BaseOperator): :ref:`howto/operator:OpenAIResponseOperator` For possible options, see: https://platform.openai.com/docs/api-reference/responses/create + + ``execute`` also pushes two XCom keys: ``response_id`` (the response's ID, usable as + a downstream call's ``previous_response_id``) and ``usage`` (the flattened response + usage, or ``None`` when the API omits it). ``usage`` reports token counts only -- + the OpenAI response carries no cost field, so pricing a run means multiplying these + counts by your own per-token rate. """ template_fields: Sequence[str] = ("input_text",) @@ -131,6 +139,8 @@ def execute(self, context: Context) -> str: response.status, ) self.log.info("Generated response %s", response.id) + context["ti"].xcom_push(key="response_id", value=response.id) + context["ti"].xcom_push(key="usage", value=self.hook.summarize_response_usage(response)) return response.output_text diff --git a/providers/openai/tests/unit/openai/hooks/test_openai.py b/providers/openai/tests/unit/openai/hooks/test_openai.py index d132a6a3d9582..74514a0aadcda 100644 --- a/providers/openai/tests/unit/openai/hooks/test_openai.py +++ b/providers/openai/tests/unit/openai/hooks/test_openai.py @@ -17,7 +17,7 @@ from __future__ import annotations import os -from unittest.mock import MagicMock, mock_open, patch +from unittest.mock import MagicMock, Mock, mock_open, patch import pytest from openai import OpenAI @@ -34,6 +34,8 @@ from openai.types.beta import Assistant, AssistantDeleted, Thread, ThreadDeleted from openai.types.beta.threads import Message, Run from openai.types.chat import ChatCompletion +from openai.types.responses import Response, ResponseUsage +from openai.types.responses.response_usage import InputTokensDetails, OutputTokensDetails from openai.types.vector_stores import VectorStoreFile, VectorStoreFileBatch, VectorStoreFileDeleted from airflow.exceptions import AirflowProviderDeprecationWarning @@ -338,6 +340,33 @@ def test_cancel_response(mock_openai_hook): assert result is expected +def test_summarize_response_usage(): + usage = ResponseUsage( + input_tokens=10, + input_tokens_details=InputTokensDetails(cache_write_tokens=1, cached_tokens=2), + output_tokens=20, + output_tokens_details=OutputTokensDetails(reasoning_tokens=3), + total_tokens=30, + ) + response = Mock(spec=Response, usage=usage) + + result = OpenAIHook.summarize_response_usage(response) + + assert result == { + "input_tokens": 10, + "input_tokens_details": {"cache_write_tokens": 1, "cached_tokens": 2}, + "output_tokens": 20, + "output_tokens_details": {"reasoning_tokens": 3}, + "total_tokens": 30, + } + + +def test_summarize_response_usage_without_usage(): + response = Mock(spec=Response, usage=None) + + assert OpenAIHook.summarize_response_usage(response) is None + + def test_create_conversation(mock_openai_hook): expected = mock_openai_hook.conn.conversations.create.return_value result = mock_openai_hook.create_conversation(metadata={"topic": "demo"}) diff --git a/providers/openai/tests/unit/openai/operators/test_openai.py b/providers/openai/tests/unit/openai/operators/test_openai.py index 0ec78bb182c51..20f50bd1bcba8 100644 --- a/providers/openai/tests/unit/openai/operators/test_openai.py +++ b/providers/openai/tests/unit/openai/operators/test_openai.py @@ -89,13 +89,17 @@ def test_openai_response_operator_execute(): response_kwargs={"instructions": "Be concise.", "previous_response_id": "resp_prev"}, ) mock_hook_instance = Mock(spec=OpenAIHook) - mock_hook_instance.create_response.return_value = Mock( - spec=Response, output_text="haiku text", id="resp_123", status="completed" - ) + mock_response = Mock(spec=Response, output_text="haiku text", id="resp_123", status="completed") + mock_hook_instance.create_response.return_value = mock_response + mock_hook_instance.summarize_response_usage.return_value = {"input_tokens": 5, "output_tokens": 7} operator.hook = mock_hook_instance - result = operator.execute(Context()) + context = Context() + context["ti"] = Mock() + result = operator.execute(context) + # Backward compat: the return value is still the aggregated output text, unchanged + # by the new XCom pushes below. assert result == "haiku text" mock_hook_instance.create_response.assert_called_once_with( input="Write a haiku.", @@ -103,6 +107,9 @@ def test_openai_response_operator_execute(): instructions="Be concise.", previous_response_id="resp_prev", ) + context["ti"].xcom_push.assert_any_call(key="response_id", value="resp_123") + context["ti"].xcom_push.assert_any_call(key="usage", value={"input_tokens": 5, "output_tokens": 7}) + mock_hook_instance.summarize_response_usage.assert_called_once_with(mock_response) @pytest.mark.parametrize("wait_for_completion", [True, False])