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 9081eda445e Support Pydantic structured outputs in
`OpenAIResponseOperator` (#69812)
9081eda445e is described below
commit 9081eda445e9b9cb9167f5d5bf5a356e13a47011
Author: Yash jain <[email protected]>
AuthorDate: Tue Oct 6 05:24:06 2026 +0530
Support Pydantic structured outputs in `OpenAIResponseOperator` (#69812)
* Add OpenAIHook.parse_response for structured Responses API output
Wraps the SDK's responses.parse(), which turns a Pydantic model into a
strict JSON-schema request and returns a ParsedResponse whose
output_parsed is an instance of that model, or None when the response
carries no parsed output. The method is generic over the model type,
mirroring the SDK's own TextFormatT, so callers get output_parsed typed
as that model rather than Any.
* Move OpenAIResponseOperator's XCom metadata push into a helper
The response_id and usage pushes move verbatim into
_push_response_metadata so the structured-output path added next can
record the same metadata. No behaviour change.
* Support Pydantic structured outputs in OpenAIResponseOperator
Pass a Pydantic BaseModel subclass as text_format and the operator calls
OpenAIHook.parse_response instead of create_response, returning the
parsed model's model_dump(mode="json") so enums, dates and other
non-JSON field types are safe to push to XCom.
The structured path fails the task rather than returning partial data. A
response whose status is not "completed" raises even when its partial
JSON validates, as does one with no parsed output (a refusal, or a
tools-only response). The SDK's ValidationError for output it cannot
parse, usually truncation at max_output_tokens, is re-raised as
ValueError naming the model. Errors carry the response id and whatever
the API reported: status, error, incomplete_details, refusal text or
output item types.
text_format is keyword-only and validated when the operator is
constructed, since a pydantic dataclass, which the SDK also accepts,
would only fail after the billed call returned. Token ceilings and
templated response_kwargs go through _build_response_kwargs() as on the
text path, and response_id and usage are pushed before the structured
output is checked, so a rejected response still records what it cost.
* Test OpenAIResponseOperator structured outputs
Responses are real ParsedResponse objects built with model_construct and
real output items rather than Mocks, so an SDK field rename breaks the
tests instead of production. Covers the JSON-mode dump with a plain Enum
(the test fails if mode="json" is reverted), token ceilings reaching
parse_response as ints, invalid ceilings failing before any request,
XCom metadata including for rejected responses, refusal, tools-only,
incomplete-but-valid and failed responses, the ValidationError
conversion, and text_format validation.
* Document structured outputs in the OpenAIResponseOperator guide
Adds a structured outputs section with an example Dag task, lists
parse_response among the hook's Responses methods, and spells out where
the structured path differs from the plain-text one: it fails on an
incomplete response instead of returning truncated output, and records
response_id and usage before checking the output.
* Type the structured-response details helper with ParsedResponse
response: Any turned off type checking for every attribute the helper
reads. With ParsedResponse[BaseModel], mypy narrows on output.type and
content.type and checks each field against the SDK, so a renamed or
misspelled field is a type error instead of passing silently.
* Make text_format keyword-only in OpenAIHook.parse_response
text_format sat where create_response has model, so
parse_response(prompt, "gpt-4o"), written by analogy with
create_response, passed the model id as the format. model is now the
second positional argument, as in create_response, and text_format is
keyword-only, as in the SDK's responses.parse.
* Document that background=True fails a structured request
With text_format set, a background response comes back queued or
in_progress, so the status check raises ValueError while the response
keeps running on OpenAI's side. Say so in the response_kwargs docs and
in the guide's background note.
* Name response_id and usage as reserved text_format fields
With multiple_outputs=True, each top-level field of the structured
result is pushed as its own XCom after execute returns, so a field named
response_id or usage would overwrite the metadata the operator pushes
under those keys. The text_format docstring and the guide now say so.
---
providers/openai/docs/operators/openai.rst | 42 ++-
.../src/airflow/providers/openai/hooks/openai.py | 35 ++-
.../airflow/providers/openai/operators/openai.py | 150 ++++++++--
.../openai/tests/system/openai/example_openai.py | 14 +
.../openai/tests/unit/openai/hooks/test_openai.py | 34 +++
.../tests/unit/openai/operators/test_openai.py | 307 ++++++++++++++++++++-
6 files changed, 548 insertions(+), 34 deletions(-)
diff --git a/providers/openai/docs/operators/openai.rst
b/providers/openai/docs/operators/openai.rst
index 4aa67bef12b..7ebdc6ed9d2 100644
--- a/providers/openai/docs/operators/openai.rst
+++ b/providers/openai/docs/operators/openai.rst
@@ -47,7 +47,8 @@ OpenAIResponseOperator
Use the
:class:`~airflow.providers.openai.operators.openai.OpenAIResponseOperator` to
generate a
model response with the OpenAI Responses API, OpenAI's recommended interface
for text generation and
-tool use. The operator returns the response's aggregated output text. When
``do_xcom_push`` is
+tool use. By default, the operator returns the response's aggregated output
text; it can also return
+Pydantic-validated structured output as a JSON-compatible value (see below).
When ``do_xcom_push`` is
enabled (the default), ``execute`` also pushes two XCom keys: ``response_id``
(the response's ID,
usable as a downstream task's ``previous_response_id`` for chaining) and
``usage`` (the response's
token usage, or ``None`` when the API omits it). ``usage`` is the nested dict
returned by
@@ -172,17 +173,52 @@ know about yet. Options worth knowing about:
before the response finishes. ``OpenAIResponseOperator`` is synchronous:
it makes one
``create_response`` call and returns ``response.output_text`` immediately,
so a response
started with ``background=True`` comes back incomplete, and the operator
logs its own warning
- because ``response.status`` is not yet ``"completed"``. Do not set
``background=True`` on
+ because ``response.status`` is not yet ``"completed"``. With
``text_format`` set, the task
+ raises ``ValueError`` instead, and the background response is left running
on OpenAI's side.
+ Do not set ``background=True`` on
``OpenAIResponseOperator``. If you need a background response, create it
from a ``@task``
using :class:`~airflow.providers.openai.hooks.openai.OpenAIHook`'s
``create_response`` directly.
+Structured outputs (Pydantic models)
+^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
+
+To request a structured response, pass a Pydantic ``BaseModel`` subclass as
``text_format``. The
+operator then calls the Responses API's structured-output path
(``responses.parse``) and returns
+the parsed model's JSON-mode dump (via ``model_dump(mode="json")``). This
renders supported field
+types such as enums and dates as JSON-compatible values before the result is
pushed to XCom. Most
+models produce a ``dict``; a Pydantic custom model serializer may produce
another JSON shape such
+as a list or scalar.
+
+``response_kwargs``, ``max_output_tokens`` and ``max_tool_calls`` are
validated and passed through
+for structured requests exactly as for plain-text ones, and the
``response_id`` and ``usage`` XCom
+keys are pushed the same way -- before the checks below, so a rejected
response still records its
+id and token usage. With ``multiple_outputs=True``, each top-level field of
the result is pushed as
+its own XCom after ``execute`` returns, so ``response_id`` and ``usage`` are
reserved field names: a
+field with either name would overwrite the XCom the operator pushes under that
key.
+
+The operator rejects any response that did not complete -- incomplete, failed,
or still in
+progress -- even if the partial output happens to match the Pydantic model, so
unlike the
+plain-text path, reaching ``max_output_tokens`` fails the task instead of
returning truncated
+output. The resulting ``ValueError`` includes the response id and
+available API details, such as ``status``, ``error``, ``incomplete_details``,
refusal text, or
+output item types. If the SDK cannot parse the model output, it raises
``ValidationError`` before
+returning a response object; the operator converts that to ``ValueError``
naming the requested
+model and notes that reaching ``max_output_tokens`` is a likely cause. In that
case, a response id
+and API details are unavailable, and neither XCom key is pushed.
+
+.. exampleinclude:: /../../openai/tests/system/openai/example_openai.py
+ :language: python
+ :start-after: [START howto_operator_openai_response_structured]
+ :end-before: [END howto_operator_openai_response_structured]
+
Using the OpenAIHook for Responses and Conversations
=====================================================
The :class:`~airflow.providers.openai.hooks.openai.OpenAIHook` exposes the
Responses and
Conversations APIs directly for use inside ``@task`` functions or custom
operators:
-- Responses: ``create_response``, ``get_response``, ``delete_response`` and
``cancel_response``
+- Responses: ``create_response``, ``parse_response`` (structured-output
wrapper),
+ ``get_response``, ``delete_response`` and ``cancel_response``
(the last cancels a response created with ``background=True``).
- Conversations: ``create_conversation``, ``get_conversation``,
``update_conversation`` and
``delete_conversation``. Pass the conversation id to ``create_response``
(via the operator's
diff --git a/providers/openai/src/airflow/providers/openai/hooks/openai.py
b/providers/openai/src/airflow/providers/openai/hooks/openai.py
index 761f29c4fd2..3987999b961 100644
--- a/providers/openai/src/airflow/providers/openai/hooks/openai.py
+++ b/providers/openai/src/airflow/providers/openai/hooks/openai.py
@@ -20,7 +20,7 @@ from __future__ import annotations
import time
from enum import Enum
from functools import cached_property
-from typing import TYPE_CHECKING, Any, BinaryIO, Literal, overload
+from typing import TYPE_CHECKING, Any, BinaryIO, Literal, TypeVar, overload
from deprecated import deprecated
from openai import OpenAI
@@ -50,8 +50,9 @@ if TYPE_CHECKING:
ChatCompletionUserMessageParam,
)
from openai.types.conversations import Conversation,
ConversationDeletedResource
- from openai.types.responses import Response
+ from openai.types.responses import ParsedResponse, Response
from openai.types.vector_stores import VectorStoreFile,
VectorStoreFileBatch, VectorStoreFileDeleted
+ from pydantic import BaseModel
from airflow.exceptions import AirflowProviderDeprecationWarning
from airflow.providers.common.compat.module_loading import import_string
from airflow.providers.common.compat.sdk import AirflowException, BaseHook
@@ -72,6 +73,11 @@ _ASSISTANTS_DEPRECATION_REASON = (
"See https://platform.openai.com/docs/guides/migrate-to-responses."
)
+#: Generic type variable for the Pydantic model used as the ``text_format`` in
structured-output
+#: Responses API calls. Mirrors the SDK's ``TextFormatT`` so
``parse_response`` returns a
+#: ``ParsedResponse[T]`` — callers get ``output_parsed`` typed as ``T | None``.
+_TextFormatT = TypeVar("_TextFormatT", bound="BaseModel")
+
class BatchStatus(str, Enum):
"""Enum for the status of a batch."""
@@ -310,6 +316,31 @@ class OpenAIHook(BaseHook):
"""
return self.conn.responses.create(model=model, input=input, **kwargs)
+ def parse_response(
+ self,
+ input: Any,
+ model: str = "gpt-4o-mini",
+ *,
+ text_format: type[_TextFormatT],
+ **kwargs: Any,
+ ) -> ParsedResponse[_TextFormatT]:
+ """
+ Create a model response and parse it into a Pydantic model via the
Responses API.
+
+ Wraps :py:meth:`openai.resources.responses.Responses.parse`. The SDK
converts
+ ``text_format`` into a JSON schema, sends it as a structured-output
request, and
+ returns a :class:`~openai.types.responses.ParsedResponse` whose
``output_parsed``
+ attribute is an instance of ``text_format``, or ``None`` when the
response carries no
+ parsed output (for example a refusal).
+
+ :param input: Text, image, or file input(s) to the model.
+ :param model: ID of the model to use.
+ :param text_format: A Pydantic ``BaseModel`` subclass describing the
expected
+ structured output. Keyword-only, as in the SDK. The SDK converts
it to a JSON schema
+ and sends the structured-output request.
+ """
+ return self.conn.responses.parse(input=input, model=model,
text_format=text_format, **kwargs)
+
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 fe03fe1bf40..f761442bc37 100644
--- a/providers/openai/src/airflow/providers/openai/operators/openai.py
+++ b/providers/openai/src/airflow/providers/openai/operators/openai.py
@@ -21,6 +21,8 @@ from collections.abc import Sequence
from functools import cached_property
from typing import TYPE_CHECKING, Any, ClassVar
+from pydantic import BaseModel, ValidationError
+
from airflow.providers.common.compat.sdk import BaseOperator, conf
from airflow.providers.openai.hooks.openai import (
OpenAIHook,
@@ -31,9 +33,35 @@ from airflow.providers.openai.hooks.openai import (
from airflow.providers.openai.triggers.openai import OpenAIBatchTrigger
if TYPE_CHECKING:
+ from openai.types.responses import ParsedResponse, Response
+
from airflow.providers.common.compat.sdk import Context
+def _get_structured_response_details(response: ParsedResponse[BaseModel]) ->
str:
+ """Return API-reported context for a structured-response failure."""
+ details = [f"status={response.status!r}"]
+ if response.error is not None:
+ details.append(f"error={response.error!r}")
+ if response.incomplete_details is not None:
+ details.append(f"incomplete_details={response.incomplete_details!r}")
+
+ refusals = [
+ content.refusal
+ for output in response.output
+ if output.type == "message"
+ for content in output.content
+ if content.type == "refusal"
+ ]
+ if refusals:
+ details.append(f"refusal={'; '.join(refusals)!r}")
+ else:
+ output_types = [output.type for output in response.output]
+ if output_types:
+ details.append(f"output_types={output_types!r}")
+ return ", ".join(details)
+
+
class OpenAIEmbeddingOperator(BaseOperator):
"""
Operator that accepts input text to generate OpenAI embeddings using the
specified model.
@@ -92,10 +120,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; 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
+ The operator is synchronous and returns the response's aggregated output
text, or, when
+ ``text_format`` is set, the structured output parsed into that Pydantic
model (see
+ ``text_format`` below). 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 response object, use
:class:`~airflow.providers.openai.hooks.openai.OpenAIHook` directly.
``max_output_tokens`` caps the number of tokens generated for the
response; ``max_tool_calls``
@@ -115,10 +144,13 @@ class OpenAIResponseOperator(BaseOperator):
input items.
:param model: The OpenAI model to use.
:param response_kwargs: Additional keyword arguments to pass to the OpenAI
``create_response``
- method (for example ``instructions``, ``tools``, ``conversation`` or
``previous_response_id``).
- Templated, so values (e.g. ``previous_response_id``) may reference
upstream XCom.
+ method, or ``parse_response`` when ``text_format`` is set (for example
``instructions``,
+ ``tools``, ``conversation`` or ``previous_response_id``). Templated,
so values (e.g.
+ ``previous_response_id``) may reference upstream XCom.
Do not set ``background`` or ``stream`` here: ``background=True``
returns before the response
- completes, so this operator logs a warning and the returned output
text may be empty, while
+ completes, so this operator logs a warning and the returned output
text may be empty (with
+ ``text_format`` set, the parsed response comes back ``queued`` or
``in_progress``, so the task
+ raises ``ValueError`` and the background response is left running on
OpenAI's side), while
``stream=True`` returns an object without ``status`` or
``output_text``, so the task raises
``AttributeError``. See :ref:`howto/operator:OpenAIResponseOperator`
for these and other
options this operator can pass through, such as ``truncation`` and
``metadata``. ``max_output_tokens``
@@ -148,6 +180,16 @@ class OpenAIResponseOperator(BaseOperator):
:param max_tool_calls: Optional upper bound on the number of built-in tool
calls the model may
make while generating the response. Same templating, type, validation,
blank-as-unset, and
mutual-exclusion rules as ``max_output_tokens``.
+ :param text_format: Optional Pydantic ``BaseModel`` subclass describing
the expected structured
+ output. When set, the operator calls ``parse_response`` instead of
``create_response`` and
+ returns the parsed model's ``model_dump(mode="json")``, so enums,
dates and other non-JSON
+ field types reach XCom as their JSON representations. The task fails
with ``ValueError``
+ rather than returning partial data when the response did not complete
(for example because
+ ``max_output_tokens`` was reached), when it carries no parsed output
(for example a
+ refusal), or when the SDK cannot validate the output against the
model. With
+ ``multiple_outputs=True`` each top-level field is pushed as its own
XCom after ``execute``
+ returns, so ``response_id`` and ``usage`` are reserved field names: a
field with either
+ name would overwrite the XCom this operator pushes under that key.
.. seealso::
For more information on how to use this operator, take a look at the
guide:
@@ -161,7 +203,9 @@ class OpenAIResponseOperator(BaseOperator):
not ``None`` it also carries a ``try_number`` key recording which attempt
produced it --
XCom is cleared at the start of every attempt, so this makes it visible
that the value
only reflects the current attempt rather than a silently under-reported
total across
- retries. Both XCom pushes are skipped when ``do_xcom_push=False``.
+ retries. With ``text_format`` set, both keys are pushed before the
structured output is
+ checked, so a response that then fails the task still records its id and
token usage.
+ Both XCom pushes are skipped when ``do_xcom_push=False``.
"""
template_fields: Sequence[str] = (
@@ -183,6 +227,7 @@ class OpenAIResponseOperator(BaseOperator):
*,
max_output_tokens: int | str | None = None,
max_tool_calls: int | str | None = None,
+ text_format: type[BaseModel] | None = None,
**kwargs: Any,
):
super().__init__(**kwargs)
@@ -192,11 +237,29 @@ class OpenAIResponseOperator(BaseOperator):
self.response_kwargs = response_kwargs or {}
self.max_output_tokens = max_output_tokens
self.max_tool_calls = max_tool_calls
+ self.text_format = text_format
self._supplied_ceilings: frozenset[str] = frozenset(
name for name in self._TOKEN_CEILING_PARAM_NAMES if getattr(self,
name) is not None
)
self._validate_no_response_kwargs_conflict()
self._validate_literal_ceiling_values()
+ self._validate_text_format()
+
+ def _validate_text_format(self) -> None:
+ """
+ Reject a ``text_format`` that is not a Pydantic ``BaseModel`` subclass.
+
+ The SDK also accepts other Pydantic-compatible types, such as a
``pydantic.dataclasses``
+ class, but those have no ``model_dump``: they would complete the
billed API call and only
+ then fail. Checking when the operator is constructed rejects them
before any request.
+ """
+ if self.text_format is not None and not (
+ isinstance(self.text_format, type) and
issubclass(self.text_format, BaseModel)
+ ):
+ raise TypeError(
+ f"Task {self.task_id!r}: 'text_format' must be a Pydantic
BaseModel subclass, "
+ f"got {self.text_format!r}."
+ )
def _validate_no_response_kwargs_conflict(self) -> None:
"""Reject a ceiling set both as an operator argument and in
``response_kwargs``."""
@@ -297,10 +360,56 @@ class OpenAIResponseOperator(BaseOperator):
response_kwargs[param_name] =
self._coerce_token_ceiling(param_name, value)
return response_kwargs
- def execute(self, context: Context) -> str:
- response = self.hook.create_response(
- input=self.input_text, model=self.model,
**self._build_response_kwargs()
- )
+ def _push_response_metadata(self, context: Context, response: Response) ->
None:
+ """Push the response id and token usage to XCom when ``do_xcom_push``
is enabled."""
+ if self.do_xcom_push:
+ context["ti"].xcom_push(key="response_id", value=response.id)
+ # model_dump (not a hand-picked field list) keeps a token-usage
dimension
+ # the API adds later from being silently dropped; mode="json"
keeps the
+ # value XCom-serializable.
+ #
+ # XCom is cleared at the start of every attempt, so this key only
ever holds
+ # the last one. Stamping the attempt makes that visible rather
than silently
+ # under-reporting total spend across retries. Built as a new dict
rather than
+ # mutating what model_dump() returned.
+ usage = (
+ {**response.usage.model_dump(mode="json"), "try_number":
context["ti"].try_number}
+ if response.usage is not None
+ else None
+ )
+ context["ti"].xcom_push(key="usage", value=usage)
+
+ def execute(self, context: Context) -> str | dict[str, Any]:
+ response_kwargs = self._build_response_kwargs()
+ if self.text_format is not None:
+ try:
+ parsed = self.hook.parse_response(
+ input=self.input_text,
+ model=self.model,
+ text_format=self.text_format,
+ **response_kwargs,
+ )
+ except ValidationError as exc:
+ # ``responses.parse`` raises ``ValidationError`` when the
model's JSON output
+ # can't be coerced into ``text_format`` — most commonly
because the response
+ # was truncated (e.g. ``max_output_tokens`` hit) mid-JSON.
Convert to a clean
+ # ``ValueError`` so callers see a consistent shape across all
parse failures.
+ raise ValueError(
+ f"OpenAI Responses API returned a payload that does not
match "
+ f"{self.text_format.__name__!r}. The response may have
been truncated because "
+ f"max_output_tokens was reached: {exc}"
+ ) from exc
+
+ self.log.info("Generated response %s", parsed.id)
+ # Pushed before the checks below: a response they reject was still
billed.
+ self._push_response_metadata(context, parsed)
+ details = _get_structured_response_details(parsed)
+ if parsed.status != "completed":
+ raise ValueError(f"Response {parsed.id} did not complete
({details}).")
+ if parsed.output_parsed is None:
+ raise ValueError(f"Response {parsed.id} did not return a
structured output ({details}).")
+ return parsed.output_parsed.model_dump(mode="json")
+ response = self.hook.create_response(input=self.input_text,
model=self.model, **response_kwargs)
if response.status == "incomplete":
reason = response.incomplete_details.reason if
response.incomplete_details else None
if reason and response.output_text:
@@ -333,22 +442,7 @@ class OpenAIResponseOperator(BaseOperator):
response.status,
)
self.log.info("Generated response %s", response.id)
- if self.do_xcom_push:
- context["ti"].xcom_push(key="response_id", value=response.id)
- # model_dump (not a hand-picked field list) keeps a token-usage
dimension
- # the API adds later from being silently dropped; mode="json"
keeps the
- # value XCom-serializable.
- #
- # XCom is cleared at the start of every attempt, so this key only
ever holds
- # the last one. Stamping the attempt makes that visible rather
than silently
- # under-reporting total spend across retries. Built as a new dict
rather than
- # mutating what model_dump() returned.
- usage = (
- {**response.usage.model_dump(mode="json"), "try_number":
context["ti"].try_number}
- if response.usage is not None
- else None
- )
- context["ti"].xcom_push(key="usage", value=usage)
+ self._push_response_metadata(context, response)
return response.output_text
diff --git a/providers/openai/tests/system/openai/example_openai.py
b/providers/openai/tests/system/openai/example_openai.py
index e16818d397f..cfa28392c81 100644
--- a/providers/openai/tests/system/openai/example_openai.py
+++ b/providers/openai/tests/system/openai/example_openai.py
@@ -17,6 +17,7 @@
from __future__ import annotations
import pendulum
+from pydantic import BaseModel
# This example uses common.compat for Airflow 2.x/3.x compatibility.
# If you only need Airflow 3+, you can use: from airflow.sdk import dag, task
@@ -132,6 +133,19 @@ def example_openai_dag():
)
# [END howto_operator_openai_response]
+ # [START howto_operator_openai_response_structured]
+ class Person(BaseModel):
+ name: str
+ age: int
+
+ OpenAIResponseOperator(
+ task_id="openai_response_structured",
+ conn_id="openai_default",
+ input_text="Extract the name and age from: 'Alice is 30 years old'.",
+ text_format=Person,
+ )
+ # [END howto_operator_openai_response_structured]
+
create_embeddings_using_hook()
diff --git a/providers/openai/tests/unit/openai/hooks/test_openai.py
b/providers/openai/tests/unit/openai/hooks/test_openai.py
index 6b56df91f55..0b0222198d7 100644
--- a/providers/openai/tests/unit/openai/hooks/test_openai.py
+++ b/providers/openai/tests/unit/openai/hooks/test_openai.py
@@ -36,6 +36,7 @@ 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.vector_stores import VectorStoreFile, VectorStoreFileBatch,
VectorStoreFileDeleted
+from pydantic import BaseModel
from airflow.exceptions import AirflowProviderDeprecationWarning
from airflow.models import Connection
@@ -322,6 +323,39 @@ def test_create_response(mock_openai_hook):
assert result is expected
+def test_parse_response(mock_openai_hook):
+ class Person(BaseModel):
+ name: str
+
+ expected = mock_openai_hook.conn.responses.parse.return_value
+ result = mock_openai_hook.parse_response(
+ input="Extract: Alice",
+ text_format=Person,
+ model=MODEL,
+ instructions="Be precise.",
+ )
+ mock_openai_hook.conn.responses.parse.assert_called_once_with(
+ model=MODEL,
+ input="Extract: Alice",
+ text_format=Person,
+ instructions="Be precise.",
+ )
+ assert result is expected
+
+
+def
test_parse_response_matches_create_response_positional_order(mock_openai_hook):
+ class Person(BaseModel):
+ name: str
+
+ # model is the second positional argument, as in create_response;
text_format is keyword-only.
+ mock_openai_hook.parse_response("Extract: Alice", MODEL,
text_format=Person)
+ mock_openai_hook.conn.responses.parse.assert_called_once_with(
+ model=MODEL, input="Extract: Alice", text_format=Person
+ )
+ with pytest.raises(TypeError, match="text_format"):
+ mock_openai_hook.parse_response("Extract: Alice", Person)
+
+
def test_get_response(mock_openai_hook):
expected = mock_openai_hook.conn.responses.retrieve.return_value
result = mock_openai_hook.get_response("resp_123")
diff --git a/providers/openai/tests/unit/openai/operators/test_openai.py
b/providers/openai/tests/unit/openai/operators/test_openai.py
index 7e6cdc5c684..dafb984ffc8 100644
--- a/providers/openai/tests/unit/openai/operators/test_openai.py
+++ b/providers/openai/tests/unit/openai/operators/test_openai.py
@@ -18,16 +18,29 @@ from __future__ import annotations
from datetime import datetime
from decimal import Decimal
+from enum import Enum
from fractions import Fraction
+from typing import Any
from unittest import mock
from unittest.mock import Mock
import jinja2
import pytest
from openai.types.batch import Batch
-from openai.types.responses import Response, ResponseUsage
+from openai.types.responses import (
+ ParsedResponse,
+ ParsedResponseOutputMessage,
+ ParsedResponseOutputText,
+ Response,
+ ResponseError,
+ ResponseFunctionToolCall,
+ ResponseOutputRefusal,
+ ResponseUsage,
+)
from openai.types.responses.response import IncompleteDetails
from openai.types.responses.response_usage import InputTokensDetails,
OutputTokensDetails
+from pydantic import BaseModel, ValidationError
+from pydantic.dataclasses import dataclass as pydantic_dataclass
from airflow.providers.common.compat.sdk import DAG, BaseOperator, Context,
TaskDeferred, XComArg
from airflow.providers.openai.exceptions import (
@@ -733,6 +746,298 @@ class TestOpenAIResponseOperatorTokenCeilings:
)
+class _StructuredPerson(BaseModel):
+ """Pydantic model used by the structured-output operator tests."""
+
+ name: str
+
+
+@pydantic_dataclass
+class _StructuredPersonDataclass:
+ name: str
+
+
+class _Priority(Enum):
+ LOW = "low"
+ HIGH = "high"
+
+
+class _StructuredTask(BaseModel):
+ title: str
+ priority: _Priority
+
+
+def _build_usage() -> ResponseUsage:
+ return ResponseUsage(
+ input_tokens=5,
+ input_tokens_details=InputTokensDetails(cached_tokens=1,
cache_write_tokens=0),
+ output_tokens=7,
+ output_tokens_details=OutputTokensDetails(reasoning_tokens=2),
+ total_tokens=12,
+ )
+
+
+def _build_parsed_response(
+ output_parsed: BaseModel | None = None,
+ *,
+ response_id: str = "resp_structured",
+ status: str = "completed",
+ error: ResponseError | None = None,
+ incomplete_details: IncompleteDetails | None = None,
+ refusal: str | None = None,
+ output_items: list[Any] | None = None,
+ usage: ResponseUsage | None = None,
+) -> ParsedResponse:
+ content: list[ParsedResponseOutputText[BaseModel] | ResponseOutputRefusal]
+ if output_items is not None:
+ output = output_items
+ elif output_parsed is not None:
+ content = [
+ ParsedResponseOutputText[BaseModel](
+ annotations=[],
+ text=output_parsed.model_dump_json(),
+ type="output_text",
+ parsed=output_parsed,
+ )
+ ]
+ output = [
+ ParsedResponseOutputMessage[BaseModel](
+ id=f"msg_{response_id}",
+ content=content,
+ role="assistant",
+ status="completed",
+ type="message",
+ )
+ ]
+ elif refusal is not None:
+ content = [ResponseOutputRefusal(refusal=refusal, type="refusal")]
+ output = [
+ ParsedResponseOutputMessage[BaseModel](
+ id=f"msg_{response_id}",
+ content=content,
+ role="assistant",
+ status="completed",
+ type="message",
+ )
+ ]
+ else:
+ output = []
+ return ParsedResponse[BaseModel].model_construct(
+ id=response_id,
+ status=status,
+ output=output,
+ error=error,
+ incomplete_details=incomplete_details,
+ usage=usage,
+ )
+
+
+class TestOpenAIResponseOperatorStructuredOutput:
+ @staticmethod
+ def _operator(**kwargs: Any) -> tuple[OpenAIResponseOperator, Mock]:
+ kwargs.setdefault("text_format", _StructuredPerson)
+ operator = OpenAIResponseOperator(
+ task_id=TASK_ID, conn_id=CONN_ID, input_text="Extract: Alice",
model="test_model", **kwargs
+ )
+ mock_hook_instance = Mock(spec=OpenAIHook)
+ operator.hook = mock_hook_instance
+ return operator, mock_hook_instance
+
+ def test_returns_parsed_model_as_dict(self):
+ operator, hook = self._operator(response_kwargs={"instructions": "Be
precise."})
+ hook.parse_response.return_value =
_build_parsed_response(_StructuredPerson(name="Alice"))
+
+ result = operator.execute(_build_execute_context())
+
+ assert result == {"name": "Alice"}
+ hook.parse_response.assert_called_once_with(
+ input="Extract: Alice",
+ model="test_model",
+ text_format=_StructuredPerson,
+ instructions="Be precise.",
+ )
+ hook.create_response.assert_not_called()
+
+ def test_dumps_enum_field_as_json_value(self):
+ # A plain Enum, with no str mixin: under model_dump()'s default
mode="python" the value
+ # would stay the live _Priority.HIGH member, which neither equals
"high" nor is a str.
+ operator, hook = self._operator(text_format=_StructuredTask)
+ hook.parse_response.return_value = _build_parsed_response(
+ _StructuredTask(title="Deploy", priority=_Priority.HIGH)
+ )
+
+ result = operator.execute(_build_execute_context())
+
+ assert result == {"title": "Deploy", "priority": "high"}
+ assert isinstance(result, dict)
+ assert isinstance(result["priority"], str)
+
+ def test_token_ceilings_apply_to_structured_request(self):
+ operator, hook = self._operator(max_output_tokens="100",
response_kwargs={"max_tool_calls": 5})
+ hook.parse_response.return_value =
_build_parsed_response(_StructuredPerson(name="Alice"))
+
+ operator.execute(_build_execute_context())
+
+ hook.parse_response.assert_called_once_with(
+ input="Extract: Alice",
+ model="test_model",
+ text_format=_StructuredPerson,
+ max_output_tokens=100,
+ max_tool_calls=5,
+ )
+
+ def test_invalid_ceiling_raises_before_structured_request(self):
+ operator, hook = self._operator(max_output_tokens="not-a-number")
+
+ with pytest.raises(ValueError, match="max_output_tokens"):
+ operator.execute(_build_execute_context())
+
+ hook.parse_response.assert_not_called()
+
+ @pytest.mark.parametrize(
+ ("do_xcom_push", "expected_push_count"),
+ [
+ pytest.param(True, 2, id="enabled"),
+ pytest.param(False, 0, id="disabled"),
+ ],
+ )
+ def test_pushes_response_id_and_usage(self, do_xcom_push,
expected_push_count):
+ operator, hook = self._operator(do_xcom_push=do_xcom_push)
+ usage = _build_usage()
+ hook.parse_response.return_value = _build_parsed_response(
+ _StructuredPerson(name="Alice"), response_id="resp_str_1",
usage=usage
+ )
+ context = _build_execute_context(try_number=2)
+
+ operator.execute(context)
+
+ assert context["ti"].xcom_push.call_count == expected_push_count
+ if do_xcom_push:
+ context["ti"].xcom_push.assert_any_call(key="response_id",
value="resp_str_1")
+ context["ti"].xcom_push.assert_any_call(
+ key="usage", value={**usage.model_dump(mode="json"),
"try_number": 2}
+ )
+
+ def test_rejected_response_still_records_id_and_usage(self):
+ # The API call behind a rejected response was still billed, so its id
and token usage
+ # are pushed before the structured output is checked.
+ operator, hook = self._operator()
+ usage = _build_usage()
+ hook.parse_response.return_value = _build_parsed_response(
+ response_id="resp_refused", refusal="I cannot help with that
request.", usage=usage
+ )
+ context = _build_execute_context()
+
+ with pytest.raises(ValueError, match="did not return a structured
output"):
+ operator.execute(context)
+
+ context["ti"].xcom_push.assert_any_call(key="response_id",
value="resp_refused")
+ context["ti"].xcom_push.assert_any_call(
+ key="usage", value={**usage.model_dump(mode="json"), "try_number":
1}
+ )
+
+ def test_refusal_raises_with_refusal_text(self):
+ operator, hook = self._operator()
+ hook.parse_response.return_value = _build_parsed_response(
+ response_id="resp_refused", refusal="I cannot help with that
request."
+ )
+
+ with pytest.raises(ValueError, match="did not return a structured
output") as excinfo:
+ operator.execute(_build_execute_context())
+
+ message = str(excinfo.value)
+ assert "resp_refused" in message
+ assert "status='completed'" in message
+ assert "refusal='I cannot help with that request.'" in message
+
+ def test_tools_only_response_raises_with_output_types(self):
+ operator, hook = self._operator()
+ hook.parse_response.return_value = _build_parsed_response(
+ response_id="resp_tool_call",
+ output_items=[
+ ResponseFunctionToolCall(
+ arguments='{"name": "Alice"}',
+ call_id="call_1",
+ name="extract_person",
+ type="function_call",
+ status="completed",
+ )
+ ],
+ )
+
+ with pytest.raises(ValueError, match="did not return a structured
output") as excinfo:
+ operator.execute(_build_execute_context())
+
+ assert "output_types=['function_call']" in str(excinfo.value)
+
+ def test_incomplete_response_raises_even_with_valid_model(self):
+ operator, hook = self._operator()
+ hook.parse_response.return_value = _build_parsed_response(
+ _StructuredPerson(name="Alice"),
+ response_id="resp_incomplete",
+ status="incomplete",
+ incomplete_details=IncompleteDetails(reason="max_output_tokens"),
+ )
+
+ with pytest.raises(ValueError, match="did not complete") as excinfo:
+ operator.execute(_build_execute_context())
+
+ message = str(excinfo.value)
+ assert "status='incomplete'" in message
+ assert "reason='max_output_tokens'" in message
+
+ def test_failed_response_raises_with_error(self):
+ operator, hook = self._operator()
+ hook.parse_response.return_value = _build_parsed_response(
+ response_id="resp_failed",
+ status="failed",
+ error=ResponseError(code="server_error", message="The model
failed."),
+ )
+
+ with pytest.raises(ValueError, match="did not complete") as excinfo:
+ operator.execute(_build_execute_context())
+
+ message = str(excinfo.value)
+ assert "status='failed'" in message
+ assert "code='server_error'" in message
+ assert "message='The model failed.'" in message
+
+ def test_validation_error_raises_value_error_naming_model(self):
+ # ``responses.parse`` raises ``pydantic.ValidationError`` when the
model's JSON output
+ # can't be coerced into ``text_format`` (e.g. truncated mid-JSON on
``max_output_tokens``).
+ # The operator converts it to a ``ValueError`` so callers see one
exception type across
+ # all parse failures.
+ operator, hook = self._operator()
+ with pytest.raises(ValidationError) as exc_info:
+ _StructuredPerson.model_validate({})
+ hook.parse_response.side_effect = exc_info.value
+ context = _build_execute_context()
+
+ with pytest.raises(ValueError,
match="'_StructuredPerson'.*max_output_tokens") as excinfo:
+ operator.execute(context)
+
+ assert excinfo.value.__cause__ is exc_info.value
+ # parse() raised before returning a response, so there is no id or
usage to record.
+ context["ti"].xcom_push.assert_not_called()
+
+ @pytest.mark.parametrize(
+ "text_format",
+ [
+ pytest.param(_StructuredPersonDataclass, id="pydantic-dataclass"),
+ pytest.param(_StructuredPerson(name="Alice"), id="model-instance"),
+ pytest.param({"type": "object"}, id="json-schema-dict"),
+ ],
+ )
+ def test_rejects_non_base_model_text_format(self, text_format):
+ with pytest.raises(TypeError, match="Pydantic BaseModel subclass"):
+ OpenAIResponseOperator(
+ task_id=TASK_ID,
+ conn_id=CONN_ID,
+ input_text="Extract: Alice",
+ text_format=text_format,
+ )
+
+
@pytest.mark.parametrize("wait_for_completion", [True, False])
def test_openai_trigger_batch_operator_not_deferred(mock_batch,
wait_for_completion):
operator = OpenAITriggerBatchOperator(