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 a4e96624b3e Add `OpenAIAgentSessionOperator` for OpenAI Managed Agents
(#73447)
a4e96624b3e is described below
commit a4e96624b3e7818a1f7266da27d4389752029747
Author: Kaxil Naik <[email protected]>
AuthorDate: Mon Sep 21 13:04:38 2026 +0100
Add `OpenAIAgentSessionOperator` for OpenAI Managed Agents (#73447)
* Add OpenAIAgentSessionOperator for OpenAI Managed Agents
Run one turn in a fresh OpenAI Managed Agents session from Airflow, with
deferrable polling of the first turn, cancellation of the active turn on
timeout, failure or kill, and session, turn and usage IDs pushed to XCom.
The Agents API is feature-detected on the installed SDK so the provider
floor stays compatible with libraries pinning openai<3.
* Clearer errors for failed agent turns and invalid agent kwargs
Format the SDK error code and message when a Managed Agents turn fails,
omit the suffix when a cancelled turn carries no error, and reject a
non-dict session_kwargs['agent'] with a ValueError before creating a
session.
---
providers/openai/docs/operators/openai.rst | 73 +++++
providers/openai/provider.yaml | 2 +
.../src/airflow/providers/openai/exceptions.py | 4 +
.../airflow/providers/openai/get_provider_info.py | 16 +-
.../src/airflow/providers/openai/hooks/openai.py | 67 +++++
.../airflow/providers/openai/operators/agent.py | 179 +++++++++++
.../src/airflow/providers/openai/triggers/agent.py | 89 ++++++
.../tests/system/openai/example_openai_agent.py | 56 ++++
.../openai/tests/unit/openai/hooks/test_openai.py | 9 +
.../tests/unit/openai/operators/test_agent.py | 334 +++++++++++++++++++++
.../tests/unit/openai/triggers/test_agent.py | 77 +++++
11 files changed, 904 insertions(+), 2 deletions(-)
diff --git a/providers/openai/docs/operators/openai.rst
b/providers/openai/docs/operators/openai.rst
index 25f5ea11101..4aa67bef12b 100644
--- a/providers/openai/docs/operators/openai.rst
+++ b/providers/openai/docs/operators/openai.rst
@@ -224,3 +224,76 @@ An example of using the operator:
:language: python
:start-after: [START howto_operator_openai_trigger_operator]
:end-before: [END howto_operator_openai_trigger_operator]
+
+.. _howto/operator:OpenAIAgentSessionOperator:
+
+Managed Agents sessions
+=======================
+
+Use
:class:`~airflow.providers.openai.operators.agent.OpenAIAgentSessionOperator`
+to submit a message to OpenAI's Managed Agents service. The service runs the
agent
+loop. Airflow waits for the first turn to complete, optionally releasing the
worker
+with ``deferrable=True``. This requires OpenAI Python SDK 3.13.0 or newer and
access
+to the beta Agents API on your configured endpoint.
+
+The provider's base dependency still permits older SDKs for other OpenAI APIs.
+Install ``openai>=3.13.0`` on both workers and triggerers to use Managed
Agents.
+Libraries that require ``openai<3`` (including current LlamaIndex OpenAI LLM
+integrations) cannot share that environment.
+
+.. exampleinclude:: /../../openai/tests/system/openai/example_openai_agent.py
+ :language: python
+ :start-after: [START howto_operator_openai_agent]
+ :end-before: [END howto_operator_openai_agent]
+
+Parameters
+^^^^^^^^^^
+
+* ``input``: Initial user message.
+* ``environment``: SDK environment configuration, such as ``{"type": "none"}``,
+ or an environment template reference for a hosted sandbox.
+* ``agent_id``: An existing saved agent. Alternatively, supply an inline agent
+ with a model in ``session_kwargs["agent"]``.
+* ``session_kwargs``: SDK session creation options, including agent overrides,
+ ``vault_ids`` and ``metadata``. The keys ``input``, ``environment``,
``agent_id``
+ and ``stream`` are reserved.
+* ``conn_id``: OpenAI connection, defaulting to ``openai_default``.
+* ``deferrable``: Whether to release the worker while waiting. Defaults to the
+ Airflow ``operators.default_deferrable`` setting.
+* ``poll_interval``: Seconds between checks, defaulting to 10.
+* ``timeout``: Seconds to wait for completion, defaulting to 3600. A shorter
+ ``execution_timeout`` still applies to a deferred task and preempts the
+ cancel-on-timeout path below.
+
+Transient polling failures are retried; three consecutive failures fail the
task.
+
+The operator returns the session ID. When XCom pushing is enabled, it also
writes
+``session_id``, ``turn_id`` and the turn's available token ``usage``. Usage
includes
+the Airflow ``try_number``; it represents the current attempt, not cumulative
spend
+across retries. Full message histories and artifacts are not stored in XCom.
+Retrieve them with ``OpenAIHook().get_conn().beta.agents.sessions.items`` and
+``.artifacts`` using the returned session ID.
+
+Each attempt creates a fresh session. Do not submit additional turns to it
while
+this task is running. An idle session without a visible turn is not treated as
+success. Failed or cancelled turns fail the task. Client-side function tools
are
+not executed by the operator and fail the task when requested; use service-side
+tools instead. A self-hosted environment must have an independently managed
worker.
+
+On timeout or polling failure, the operator requests cancellation of its
session's
+active turn. It retains the session and artifacts for inspection. Cancellation
does
+not delete the environment or guarantee that its resources have been released.
+Killing a synchronous task also requests cancellation. Cancellation of a killed
+deferred task requires Airflow 3.3 or newer; on older versions, cancel it
manually.
+A hard worker termination or Airflow execution timeout can bypass cleanup.
Retrying
+the task creates another session and can repeat external side effects.
+
+Hook methods
+^^^^^^^^^^^^
+
+:class:`~airflow.providers.openai.hooks.openai.OpenAIHook` provides
+``create_agent``, ``create_agent_session``, ``get_agent_session`` and
+``cancel_agent_session``. ``poll_agent_session`` checks the first turn of a
fresh,
+exclusively owned session; it is not a general waiter for reused sessions.
+For other resources, use the SDK client returned by ``get_conn()``. See the
+`OpenAI Agents API reference
<https://developers.openai.com/api/reference/python/resources/beta/subresources/agents>`__.
diff --git a/providers/openai/provider.yaml b/providers/openai/provider.yaml
index 3e203e7e641..6c2ce768e58 100644
--- a/providers/openai/provider.yaml
+++ b/providers/openai/provider.yaml
@@ -87,11 +87,13 @@ operators:
- integration-name: OpenAI
python-modules:
- airflow.providers.openai.operators.openai
+ - airflow.providers.openai.operators.agent
triggers:
- integration-name: OpenAI
python-modules:
- airflow.providers.openai.triggers.openai
+ - airflow.providers.openai.triggers.agent
connection-types:
- hook-class-name: airflow.providers.openai.hooks.openai.OpenAIHook
diff --git a/providers/openai/src/airflow/providers/openai/exceptions.py
b/providers/openai/src/airflow/providers/openai/exceptions.py
index 85f015880c5..09618b9048e 100644
--- a/providers/openai/src/airflow/providers/openai/exceptions.py
+++ b/providers/openai/src/airflow/providers/openai/exceptions.py
@@ -30,3 +30,7 @@ class OpenAIBatchTimeout(AirflowException):
class OpenAITriggerEventError(AirflowException):
"""Raise when a deferred task resumes with a missing or malformed trigger
event."""
+
+
+class OpenAIAgentSessionError(AirflowException):
+ """Raise when a Managed Agents session fails or cannot run."""
diff --git a/providers/openai/src/airflow/providers/openai/get_provider_info.py
b/providers/openai/src/airflow/providers/openai/get_provider_info.py
index 3f9ff71b5d6..7c249555ae6 100644
--- a/providers/openai/src/airflow/providers/openai/get_provider_info.py
+++ b/providers/openai/src/airflow/providers/openai/get_provider_info.py
@@ -48,10 +48,22 @@ def get_provider_info():
{"integration-name": "OpenAI", "python-modules":
["airflow.providers.openai.hooks.openai"]}
],
"operators": [
- {"integration-name": "OpenAI", "python-modules":
["airflow.providers.openai.operators.openai"]}
+ {
+ "integration-name": "OpenAI",
+ "python-modules": [
+ "airflow.providers.openai.operators.openai",
+ "airflow.providers.openai.operators.agent",
+ ],
+ }
],
"triggers": [
- {"integration-name": "OpenAI", "python-modules":
["airflow.providers.openai.triggers.openai"]}
+ {
+ "integration-name": "OpenAI",
+ "python-modules": [
+ "airflow.providers.openai.triggers.openai",
+ "airflow.providers.openai.triggers.agent",
+ ],
+ }
],
"connection-types": [
{
diff --git a/providers/openai/src/airflow/providers/openai/hooks/openai.py
b/providers/openai/src/airflow/providers/openai/hooks/openai.py
index b00fe215015..f9ba07a8497 100644
--- a/providers/openai/src/airflow/providers/openai/hooks/openai.py
+++ b/providers/openai/src/airflow/providers/openai/hooks/openai.py
@@ -56,6 +56,7 @@ from airflow.exceptions import
AirflowProviderDeprecationWarning
from airflow.providers.common.compat.module_loading import import_string
from airflow.providers.common.compat.sdk import BaseHook
from airflow.providers.openai.exceptions import (
+ OpenAIAgentSessionError,
OpenAIBatchJobException,
OpenAIBatchTimeout,
OpenAITriggerEventError,
@@ -714,3 +715,69 @@ class OpenAIHook(BaseHook):
"""
batch = self.conn.batches.cancel(batch_id=batch_id)
return batch
+
+ #: Consecutive ``poll_agent_session`` failures tolerated before an agent
wait gives up.
+ MAX_CONSECUTIVE_POLL_FAILURES = 3
+
+ @cached_property
+ def _managed_agents(self) -> Any:
+ agents = getattr(self.conn.beta, "agents", None)
+ if agents is None:
+ raise OpenAIAgentSessionError(
+ "Managed Agents requires openai>=3.13.0. Upgrade the OpenAI
SDK on workers and triggerers."
+ )
+ return agents
+
+ def create_agent(self, **kwargs: Any) -> Any:
+ """Create a reusable Managed Agent using the SDK's agent configuration
arguments."""
+ return self._managed_agents.create(**kwargs)
+
+ def create_agent_session(self, *, input: str, environment: dict[str, Any],
**kwargs: Any) -> Any:
+ """Create a fresh Managed Agents session and submit its initial
turn."""
+ if "stream" in kwargs:
+ raise ValueError("create_agent_session does not support streaming")
+ return self._managed_agents.sessions.create(
+ input=input, environment=environment, stream=False, **kwargs
+ )
+
+ def get_agent_session(self, session_id: str) -> Any:
+ """Retrieve a Managed Agents session, including required actions and
usage."""
+ return self._managed_agents.sessions.retrieve(session_id)
+
+ def cancel_agent_session(self, session_id: str) -> None:
+ """Request cancellation of the session's active turn, preserving its
history and artifacts."""
+ self._managed_agents.sessions.events.create(
+ session_id, events=[{"type": "agent.session.input.cancel"}]
+ )
+
+ def poll_agent_session(self, session_id: str) -> dict[str, Any] | None:
+ """
+ Check the first turn of a fresh, exclusively owned session.
+
+ Return a terminal result, or ``None`` while waiting. An idle session
without
+ a visible turn is not completion: the submitted input may still be
queued.
+ This helper must not be used to wait for subsequent turns of a reused
session.
+ """
+ session = self.get_agent_session(session_id)
+ turns = self._managed_agents.sessions.turns.list(session_id,
order="asc", limit=1)
+ turn = turns.data[0] if turns.data else None
+ result: dict[str, Any] = {"session_id": session_id}
+ if turn is not None:
+ result["turn_id"] = turn.id
+ result["usage"] = turn.usage.model_dump(mode="json") if turn.usage
is not None else None
+ if turn.status == "completed":
+ return {**result, "status": "success"}
+ if turn.status in {"failed", "cancelled"}:
+ message = f"Agent turn {turn.id} {turn.status}"
+ if turn.error is not None:
+ message += f" ({turn.error.code}): {turn.error.message}"
+ return {**result, "status": "error", "message": message}
+ if session.status == "failed":
+ return {**result, "status": "error", "message": f"Agent session
failed: {session.error}"}
+ if any(action.type == "function_call" for action in
session.required_actions):
+ return {
+ **result,
+ "status": "error",
+ "message": "The agent requested a client-side function tool.
Use server-side tools with this operator.",
+ }
+ return None
diff --git a/providers/openai/src/airflow/providers/openai/operators/agent.py
b/providers/openai/src/airflow/providers/openai/operators/agent.py
new file mode 100644
index 00000000000..b778e9d63e9
--- /dev/null
+++ b/providers/openai/src/airflow/providers/openai/operators/agent.py
@@ -0,0 +1,179 @@
+# 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 math
+import time
+from collections.abc import Sequence
+from datetime import timedelta
+from functools import cached_property
+from typing import TYPE_CHECKING, Any
+
+from airflow.providers.common.compat.sdk import BaseOperator, conf
+from airflow.providers.openai.exceptions import OpenAIAgentSessionError,
OpenAITriggerEventError
+from airflow.providers.openai.hooks.openai import OpenAIHook
+from airflow.providers.openai.triggers.agent import OpenAIAgentSessionTrigger
+
+if TYPE_CHECKING:
+ from airflow.providers.common.compat.sdk import Context
+
+
+class OpenAIAgentSessionOperator(BaseOperator):
+ """
+ Run one turn in a fresh OpenAI Managed Agents session and return its
session ID.
+
+ The session is retained for downstream retrieval of items and artifacts. A
retry
+ creates a new session and can repeat external side effects.
+
+ :param input: Initial user message. (templated)
+ :param environment: SDK environment configuration or template reference.
(templated)
+ :param agent_id: Saved agent ID. Alternatively supply agent.model in
session_kwargs. (templated)
+ :param session_kwargs: Additional SDK session creation arguments, such as
agent,
+ vault_ids and metadata. Must not contain input, environment, agent_id
or stream. (templated)
+ :param conn_id: OpenAI connection ID. (templated)
+ :param deferrable: Release the worker while waiting for completion.
+ :param poll_interval: Seconds between polls.
+ :param timeout: Maximum seconds to wait for the initial turn. A shorter
+ ``execution_timeout`` still applies and preempts the cancel-on-timeout
path.
+ """
+
+ template_fields: Sequence[str] = ("input", "environment", "agent_id",
"session_kwargs", "conn_id")
+ template_fields_renderers = {"environment": "json", "session_kwargs":
"json"}
+
+ def __init__(
+ self,
+ *,
+ input: str,
+ environment: dict[str, Any],
+ agent_id: str | None = None,
+ session_kwargs: dict[str, Any] | None = None,
+ conn_id: str = OpenAIHook.default_conn_name,
+ deferrable: bool = conf.getboolean("operators", "default_deferrable",
fallback=False),
+ poll_interval: float = 10,
+ timeout: float = 3600,
+ **kwargs: Any,
+ ) -> None:
+ super().__init__(**kwargs)
+ for name, value in (("poll_interval", poll_interval), ("timeout",
timeout)):
+ if not math.isfinite(value) or value <= 0:
+ raise ValueError(f"{name} must be finite and positive")
+ self.input = input
+ self.environment = environment
+ self.agent_id = agent_id
+ self.session_kwargs = session_kwargs or {}
+ self.conn_id = conn_id
+ self.deferrable = deferrable
+ self.poll_interval = poll_interval
+ self.timeout = timeout
+ self.session_id: str | None = None
+
+ @cached_property
+ def hook(self) -> OpenAIHook:
+ """Return the connection's OpenAI hook."""
+ return OpenAIHook(conn_id=self.conn_id)
+
+ def execute(self, context: Context) -> str:
+ reserved = {"input", "environment", "agent_id", "stream"} &
self.session_kwargs.keys()
+ if reserved:
+ raise ValueError(f"Reserved session_kwargs: {sorted(reserved)}")
+ agent = self.session_kwargs.get("agent")
+ if agent is not None and not isinstance(agent, dict):
+ raise ValueError("session_kwargs['agent'] must be a dict of SDK
agent fields")
+ if not self.agent_id and not (agent or {}).get("model"):
+ raise ValueError("Supply agent_id or
session_kwargs['agent']['model']")
+ if not self.input:
+ raise ValueError("input must not be empty")
+ create_kwargs = dict(self.session_kwargs)
+ if self.agent_id:
+ create_kwargs["agent_id"] = self.agent_id
+ session = self.hook.create_agent_session(
+ input=self.input, environment=self.environment, **create_kwargs
+ )
+ self.session_id = session.id
+ try:
+ if self.do_xcom_push:
+ context["ti"].xcom_push(key="session_id", value=session.id)
+ if self.deferrable:
+ self.defer(
+ trigger=OpenAIAgentSessionTrigger(
+ conn_id=self.conn_id,
+ session_id=session.id,
+ poll_interval=self.poll_interval,
+ end_time=time.time() + self.timeout,
+ ),
+ method_name="execute_complete",
+ kwargs={"session_id": session.id},
+ timeout=self.execution_timeout
+ or timedelta(seconds=self.timeout + self.poll_interval +
60),
+ )
+ deadline = time.monotonic() + self.timeout
+ consecutive_failures = 0
+ while time.monotonic() < deadline:
+ try:
+ result = self.hook.poll_agent_session(session.id)
+ except Exception as exc:
+ consecutive_failures += 1
+ if consecutive_failures >=
OpenAIHook.MAX_CONSECUTIVE_POLL_FAILURES:
+ raise
+ self.log.warning("Polling agent session %s failed (%s);
retrying.", session.id, exc)
+ else:
+ consecutive_failures = 0
+ if result is not None:
+ break
+ time.sleep(min(self.poll_interval, max(0, deadline -
time.monotonic())))
+ else:
+ raise OpenAIAgentSessionError(f"Agent session {session.id}
timed out")
+ except Exception:
+ self.on_kill()
+ raise
+ return self.execute_complete(context, result, session_id=session.id)
+
+ def execute_complete(self, context: Context, event: Any = None,
session_id: str | None = None) -> str:
+ """Validate completion, record usage, and return the owned session
ID."""
+ self.session_id = session_id or self.session_id
+ if (
+ not isinstance(event, dict)
+ or event.get("status") not in ("success", "error", "timeout")
+ or not self.session_id
+ or event.get("session_id") != self.session_id
+ ):
+ self.on_kill()
+ raise OpenAITriggerEventError("Invalid Managed Agents trigger
event")
+ if event["status"] != "success":
+ self.on_kill()
+ if self.do_xcom_push:
+ try:
+ if event.get("turn_id"):
+ context["ti"].xcom_push(key="turn_id",
value=event["turn_id"])
+ usage = event.get("usage")
+ if usage is not None:
+ context["ti"].xcom_push(
+ key="usage", value={**usage, "try_number":
context["ti"].try_number}
+ )
+ except Exception:
+ self.log.exception("Could not record agent turn usage for
session %s", self.session_id)
+ if event["status"] != "success":
+ raise OpenAIAgentSessionError(event.get("message", "Agent session
failed"))
+ return self.session_id
+
+ def on_kill(self) -> None:
+ """Request cancellation without deleting session history or
artifacts."""
+ if self.session_id:
+ try:
+ self.hook.cancel_agent_session(self.session_id)
+ except Exception:
+ self.log.exception("Could not cancel agent session %s",
self.session_id)
diff --git a/providers/openai/src/airflow/providers/openai/triggers/agent.py
b/providers/openai/src/airflow/providers/openai/triggers/agent.py
new file mode 100644
index 00000000000..3a01b99e921
--- /dev/null
+++ b/providers/openai/src/airflow/providers/openai/triggers/agent.py
@@ -0,0 +1,89 @@
+# 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 asyncio
+import time
+from collections.abc import AsyncIterator
+from typing import Any
+
+from airflow.providers.openai.hooks.openai import OpenAIHook
+from airflow.triggers.base import BaseTrigger, TriggerEvent
+
+
+class OpenAIAgentSessionTrigger(BaseTrigger):
+ """
+ Wait for the first turn of an exclusively owned Managed Agents session.
+
+ :param conn_id: OpenAI connection ID.
+ :param session_id: Fresh session whose first turn is being awaited.
+ :param poll_interval: Seconds between polls.
+ :param end_time: Epoch deadline, preserved across triggerer restarts.
+ """
+
+ def __init__(self, conn_id: str, session_id: str, poll_interval: float,
end_time: float) -> None:
+ super().__init__()
+ self.conn_id = conn_id
+ self.session_id = session_id
+ self.poll_interval = poll_interval
+ self.end_time = end_time
+
+ def serialize(self) -> tuple[str, dict[str, Any]]:
+ """Serialize identifiers and the absolute deadline, without
credentials."""
+ return (
+
"airflow.providers.openai.triggers.agent.OpenAIAgentSessionTrigger",
+ {
+ "conn_id": self.conn_id,
+ "session_id": self.session_id,
+ "poll_interval": self.poll_interval,
+ "end_time": self.end_time,
+ },
+ )
+
+ async def run(self) -> AsyncIterator[TriggerEvent]:
+ """Poll off the event loop and emit one terminal result."""
+ hook = OpenAIHook(conn_id=self.conn_id)
+ consecutive_failures = 0
+ while time.time() < self.end_time:
+ try:
+ result = await asyncio.to_thread(hook.poll_agent_session,
self.session_id)
+ except Exception as exc:
+ # Tolerate transient polling errors rather than cancelling a
live session.
+ consecutive_failures += 1
+ if consecutive_failures >=
OpenAIHook.MAX_CONSECUTIVE_POLL_FAILURES:
+ yield TriggerEvent(
+ {"status": "error", "session_id": self.session_id,
"message": str(exc)}
+ )
+ return
+ self.log.warning("Polling agent session %s failed (%s);
retrying.", self.session_id, exc)
+ else:
+ consecutive_failures = 0
+ if result is not None:
+ yield TriggerEvent(result)
+ return
+ await asyncio.sleep(min(self.poll_interval, max(0, self.end_time -
time.time())))
+ yield TriggerEvent(
+ {"status": "timeout", "session_id": self.session_id, "message":
"Agent session timed out"}
+ )
+
+ async def on_kill(self) -> None:
+ """Cancel a killed deferred task's turn on Airflow versions supporting
trigger cleanup."""
+ hook = OpenAIHook(conn_id=self.conn_id)
+ try:
+ await asyncio.to_thread(hook.cancel_agent_session, self.session_id)
+ except Exception:
+ self.log.exception("Could not cancel agent session %s",
self.session_id)
diff --git a/providers/openai/tests/system/openai/example_openai_agent.py
b/providers/openai/tests/system/openai/example_openai_agent.py
new file mode 100644
index 00000000000..92a48e30700
--- /dev/null
+++ b/providers/openai/tests/system/openai/example_openai_agent.py
@@ -0,0 +1,56 @@
+# 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 pendulum
+
+from airflow.providers.common.compat.sdk import dag
+from airflow.providers.openai.operators.agent import OpenAIAgentSessionOperator
+
+
+@dag(
+ schedule=None,
+ start_date=pendulum.datetime(2026, 1, 1, tz="UTC"),
+ catchup=False,
+ tags=["example", "openai"],
+)
+def example_openai_agent():
+ # [START howto_operator_openai_agent]
+ OpenAIAgentSessionOperator(
+ task_id="run_agent",
+ input="Explain how Airflow retries affect a task that calls an
external API.",
+ environment={"type": "none"},
+ session_kwargs={
+ "agent": {
+ "model": "gpt-6-astra",
+ "instructions": "Give a short explanation suitable for a data
engineer.",
+ }
+ },
+ deferrable=True,
+ poll_interval=10,
+ timeout=600,
+ )
+ # [END howto_operator_openai_agent]
+
+
+example_dag = example_openai_agent()
+
+
+from tests_common.test_utils.system_tests import get_test_run # noqa: E402
+
+# Needed to run the example DAG with pytest (see:
contributing-docs/testing/system_tests.rst)
+test_run = get_test_run(example_dag)
diff --git a/providers/openai/tests/unit/openai/hooks/test_openai.py
b/providers/openai/tests/unit/openai/hooks/test_openai.py
index 81bcdaa3dca..cc370911777 100644
--- a/providers/openai/tests/unit/openai/hooks/test_openai.py
+++ b/providers/openai/tests/unit/openai/hooks/test_openai.py
@@ -17,6 +17,7 @@
from __future__ import annotations
import os
+from types import SimpleNamespace
from unittest.mock import MagicMock, mock_open, patch
import pytest
@@ -39,6 +40,7 @@ from openai.types.vector_stores import VectorStoreFile,
VectorStoreFileBatch, Ve
from airflow.exceptions import AirflowProviderDeprecationWarning
from airflow.models import Connection
from airflow.providers.openai.exceptions import (
+ OpenAIAgentSessionError,
OpenAIBatchJobException,
OpenAIBatchTimeout,
OpenAITriggerEventError,
@@ -953,3 +955,10 @@ class TestValidateTriggerEvent:
)
def test_valid_event_is_returned(self, event):
assert validate_execute_complete_event(event) is event
+
+
+def test_managed_agents_reports_sdk_upgrade_without_breaking_hook():
+ hook = OpenAIHook()
+ hook.conn = SimpleNamespace(beta=SimpleNamespace())
+ with pytest.raises(OpenAIAgentSessionError, match="requires
openai>=3.13.0"):
+ hook.create_agent_session(input="Hello", environment={"type": "none"},
agent_id="agent")
diff --git a/providers/openai/tests/unit/openai/operators/test_agent.py
b/providers/openai/tests/unit/openai/operators/test_agent.py
new file mode 100644
index 00000000000..aa1e5d01fa1
--- /dev/null
+++ b/providers/openai/tests/unit/openai/operators/test_agent.py
@@ -0,0 +1,334 @@
+# 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 datetime import timedelta
+from types import SimpleNamespace
+from unittest.mock import MagicMock, patch
+
+import pytest
+from openai import OpenAI
+
+from airflow.providers.common.compat.sdk import AirflowException, TaskDeferred
+from airflow.providers.openai.hooks.openai import OpenAIHook
+from airflow.providers.openai.operators.agent import OpenAIAgentSessionOperator
+from airflow.providers.openai.triggers.agent import OpenAIAgentSessionTrigger
+
+httpx2 = pytest.importorskip("httpx2")
+AgentSession =
pytest.importorskip("openai.types.beta.agent_session").AgentSession
+Turn = pytest.importorskip("openai.types.beta.agents.sessions.turn").Turn
+SessionTurnError =
pytest.importorskip("openai.types.beta.session_turn_error").SessionTurnError
+
+
[email protected]
+def hook():
+ hook = OpenAIHook()
+ with OpenAI(api_key="test") as client:
+ hook.conn = MagicMock(spec=client)
+ yield hook
+
+
[email protected]("status", ["queued", "in_progress", "waiting",
"completed", "failed", "cancelled"])
+def test_poll_turn_status(hook, status):
+ hook.conn.beta.agents.sessions.retrieve.return_value =
AgentSession.model_construct(
+ status="idle", required_actions=[], error=None
+ )
+ hook.conn.beta.agents.sessions.turns.list.return_value.data = [
+ Turn.model_construct(id="turn", status=status, error=None, usage=None)
+ ]
+ result = hook.poll_agent_session("session")
+ if status in {"queued", "in_progress", "waiting"}:
+ assert result is None
+ else:
+ assert result["status"] == ("success" if status == "completed" else
"error")
+ assert result["turn_id"] == "turn"
+
hook.conn.beta.agents.sessions.turns.list.assert_called_once_with("session",
order="asc", limit=1)
+
+
+def test_idle_without_visible_turn_keeps_waiting(hook):
+ hook.conn.beta.agents.sessions.retrieve.return_value =
AgentSession.model_construct(
+ status="idle", required_actions=[], error=None
+ )
+ hook.conn.beta.agents.sessions.turns.list.return_value.data = []
+ assert hook.poll_agent_session("session") is None
+
+
[email protected](
+ ("status", "actions", "expected"),
+ [
+ ("failed", [], "error"),
+ ("requires_action", [SimpleNamespace(type="function_call")], "error"),
+ ("requires_action", [SimpleNamespace(type="environment_connection")],
None),
+ ],
+)
+def test_poll_session_failure_and_required_actions(hook, status, actions,
expected):
+ hook.conn.beta.agents.sessions.retrieve.return_value =
AgentSession.model_construct(
+ status=status, required_actions=actions, error="Environment failed"
+ )
+ hook.conn.beta.agents.sessions.turns.list.return_value.data = []
+ result = hook.poll_agent_session("session")
+ assert (result["status"] if result else None) == expected
+
+
+def test_create_and_cancel_sdk_payload(hook):
+ hook.create_agent_session(input="Hello", environment={"type": "none"},
agent_id="agent")
+ hook.conn.beta.agents.sessions.create.assert_called_once_with(
+ input="Hello", environment={"type": "none"}, agent_id="agent",
stream=False
+ )
+ hook.cancel_agent_session("session")
+ hook.conn.beta.agents.sessions.events.create.assert_called_once_with(
+ "session", events=[{"type": "agent.session.input.cancel"}]
+ )
+
+
[email protected]
+def operator():
+ operator = OpenAIAgentSessionOperator(
+ task_id="agent",
+ input="Hello",
+ environment={"type": "none"},
+ agent_id="agent",
+ do_xcom_push=False,
+ )
+ operator.hook = MagicMock(spec=OpenAIHook)
+ operator.hook.create_agent_session.return_value =
AgentSession.model_construct(id="session")
+ return operator
+
+
+def test_sync_completion(operator):
+ operator.hook.poll_agent_session.return_value = {"status": "success",
"session_id": "session"}
+ assert operator.execute({}) == "session"
+ operator.hook.cancel_agent_session.assert_not_called()
+
+
+def test_deferral_roundtrip(operator):
+ operator.deferrable = True
+ with pytest.raises(TaskDeferred) as exc:
+ operator.execute({})
+ assert exc.value.kwargs == {"session_id": "session"}
+ path, kwargs = exc.value.trigger.serialize()
+ assert path.endswith(".OpenAIAgentSessionTrigger")
+ assert OpenAIAgentSessionTrigger(**kwargs).serialize() == (path, kwargs)
+ operator.hook.cancel_agent_session.assert_not_called()
+ assert (
+ operator.execute_complete({}, {"status": "success", "session_id":
"session"}, session_id="session")
+ == "session"
+ )
+
+
[email protected]("event", [None, {}, {"status": "success",
"session_id": "other"}])
+def test_invalid_event_cancels_owned_session(operator, event):
+ with pytest.raises(AirflowException, match="Invalid"):
+ operator.execute_complete({}, event, session_id="session")
+ operator.hook.cancel_agent_session.assert_called_once_with("session")
+
+
[email protected]("status", ["error", "timeout"])
+def test_deferred_failure_cancels(operator, status):
+ with pytest.raises(AirflowException, match="Failed"):
+ operator.execute_complete(
+ {}, {"status": status, "session_id": "session", "message":
"Failed"}, session_id="session"
+ )
+ operator.hook.cancel_agent_session.assert_called_once_with("session")
+
+
+def test_sync_poll_exception_cancels_and_preserves_error(operator):
+ operator.hook.poll_agent_session.side_effect = RuntimeError("Read failed")
+ operator.hook.cancel_agent_session.side_effect = RuntimeError("Cancel
failed")
+ with patch("airflow.providers.openai.operators.agent.time.sleep",
autospec=True) as sleep:
+ with pytest.raises(RuntimeError, match="Read failed"):
+ operator.execute({})
+ assert operator.hook.poll_agent_session.call_count ==
OpenAIHook.MAX_CONSECUTIVE_POLL_FAILURES
+ assert sleep.call_count == OpenAIHook.MAX_CONSECUTIVE_POLL_FAILURES - 1
+
+
+def test_sync_transient_poll_failure_recovers_without_cancelling(operator):
+ operator.hook.poll_agent_session.side_effect = [
+ RuntimeError("Read failed"),
+ None,
+ RuntimeError("Read failed again"),
+ {"status": "success", "session_id": "session"},
+ ]
+ with patch("airflow.providers.openai.operators.agent.time.sleep",
autospec=True):
+ assert operator.execute({}) == "session"
+ operator.hook.cancel_agent_session.assert_not_called()
+
+
+def test_deferral_timeout_is_capped_by_execution_timeout(operator):
+ operator.deferrable = True
+ operator.execution_timeout = timedelta(seconds=60)
+ with pytest.raises(TaskDeferred) as exc:
+ operator.execute({})
+ assert exc.value.timeout == timedelta(seconds=60)
+
+
+def test_deferral_timeout_defaults_to_operator_timeout(operator):
+ operator.deferrable = True
+ with pytest.raises(TaskDeferred) as exc:
+ operator.execute({})
+ assert exc.value.timeout == timedelta(seconds=operator.timeout +
operator.poll_interval + 60)
+
+
+def test_sync_timeout_cancels(operator):
+ with patch("airflow.providers.openai.operators.agent.time.monotonic",
side_effect=[0, 3601]):
+ with pytest.raises(AirflowException, match="timed out"):
+ operator.execute({})
+ operator.hook.cancel_agent_session.assert_called_once_with("session")
+
+
[email protected](
+ ("name", "value"), [("timeout", 0), ("poll_interval", -1), ("timeout",
float("inf"))]
+)
+def test_invalid_wait_configuration(name, value):
+ with pytest.raises(ValueError, match=name):
+ OpenAIAgentSessionOperator(task_id="agent", input="Hi",
environment={}, **{name: value})
+
+
[email protected]("reserved", ["input", "environment", "agent_id",
"stream"])
+def test_reserved_kwargs_rejected_before_creation(operator, reserved):
+ operator.session_kwargs = {reserved: "bad"}
+ with pytest.raises(ValueError, match="Reserved"):
+ operator.execute({})
+ operator.hook.create_agent_session.assert_not_called()
+
+
[email protected]("agent", [None, "gpt-6-astra", ["gpt-6-astra"]])
+def test_non_dict_agent_rejected_with_value_error(operator, agent):
+ operator.agent_id = None
+ operator.session_kwargs = {"agent": agent}
+ with pytest.raises(ValueError, match="agent"):
+ operator.execute({})
+ operator.hook.create_agent_session.assert_not_called()
+
+
[email protected](
+ ("status", "error", "expected"),
+ [
+ ("cancelled", None, "Agent turn turn cancelled"),
+ (
+ "failed",
+ SessionTurnError.model_construct(code="usage_limit_exceeded",
message="Add credits."),
+ "Agent turn turn failed (usage_limit_exceeded): Add credits.",
+ ),
+ ],
+)
+def test_failed_turn_message_formats_error(hook, status, error, expected):
+ hook.conn.beta.agents.sessions.retrieve.return_value =
AgentSession.model_construct(
+ status="idle", required_actions=[], error=None
+ )
+ hook.conn.beta.agents.sessions.turns.list.return_value.data = [
+ Turn.model_construct(id="turn", status=status, error=error, usage=None)
+ ]
+ assert hook.poll_agent_session("session")["message"] == expected
+
+
+def test_real_sdk_request_and_response_contract():
+ requests = []
+ turn_reads = 0
+
+ def respond(request):
+ nonlocal turn_reads
+ requests.append(request)
+ assert request.headers["OpenAI-Beta"] == "agents=v1"
+ if request.url.path.endswith("/events"):
+ return httpx2.Response(204)
+ if request.url.path.endswith("/turns"):
+ turn_reads += 1
+ return httpx2.Response(
+ 200,
+ json={
+ "object": "list",
+ "has_more": False,
+ "data": []
+ if turn_reads == 1
+ else [
+ {
+ "id": "turn",
+ "session_id": "session",
+ "agent_id": "agent",
+ "object": "agent.session.turn",
+ "created_at": 0,
+ "status": "completed",
+ "usage": {"input_tokens": 3, "output_tokens": 2,
"total_tokens": 5},
+ }
+ ],
+ },
+ )
+ return httpx2.Response(
+ 200,
+ json={"id": "session", "status": "idle", "required_actions": [],
"error": None},
+ )
+
+ with OpenAI(api_key="test",
http_client=httpx2.Client(transport=httpx2.MockTransport(respond))) as client:
+ hook = OpenAIHook()
+ hook.conn = client
+ session = hook.create_agent_session(input="Hello",
environment={"type": "none"}, agent_id="agent")
+ assert session.id == "session"
+ assert hook.poll_agent_session(session.id) is None
+ result = hook.poll_agent_session(session.id)
+ assert result["status"] == "success"
+ assert result["usage"]["total_tokens"] == 5
+ hook.cancel_agent_session(session.id)
+ assert json.loads(requests[0].content) == {
+ "input": "Hello",
+ "environment": {"type": "none"},
+ "agent_id": "agent",
+ "stream": False,
+ }
+ assert json.loads(requests[-1].content) == {"events": [{"type":
"agent.session.input.cancel"}]}
+ assert dict(requests[2].url.params) == {"order": "asc", "limit": "1"}
+
+
+def test_usage_is_recorded_with_attempt(operator):
+ operator.do_xcom_push = True
+ push = MagicMock(spec=lambda **kwargs: None)
+ context = {"ti": SimpleNamespace(xcom_push=push, try_number=2)}
+ event = {"status": "success", "session_id": "session", "turn_id": "turn",
"usage": {"total_tokens": 5}}
+ assert operator.execute_complete(context, event, session_id="session") ==
"session"
+ push.assert_any_call(key="turn_id", value="turn")
+ push.assert_any_call(key="usage", value={"total_tokens": 5, "try_number":
2})
+ assert event["usage"] == {"total_tokens": 5}
+
+
+def test_usage_failure_does_not_mask_agent_failure(operator):
+ operator.do_xcom_push = True
+ push = MagicMock(spec=lambda **kwargs: None,
side_effect=RuntimeError("XCom unavailable"))
+ context = {"ti": SimpleNamespace(xcom_push=push, try_number=1)}
+ with pytest.raises(AirflowException, match="Agent failed"):
+ operator.execute_complete(
+ context,
+ {
+ "status": "error",
+ "session_id": "session",
+ "turn_id": "turn",
+ "message": "Agent failed",
+ },
+ session_id="session",
+ )
+ operator.hook.cancel_agent_session.assert_called_once_with("session")
+
+
+def test_sync_failed_turn_cancels_once(operator):
+ operator.hook.poll_agent_session.return_value = {
+ "status": "error",
+ "session_id": "session",
+ "message": "Agent failed",
+ }
+ with pytest.raises(AirflowException, match="Agent failed"):
+ operator.execute({})
+ operator.hook.cancel_agent_session.assert_called_once_with("session")
diff --git a/providers/openai/tests/unit/openai/triggers/test_agent.py
b/providers/openai/tests/unit/openai/triggers/test_agent.py
new file mode 100644
index 00000000000..a1979174dfd
--- /dev/null
+++ b/providers/openai/tests/unit/openai/triggers/test_agent.py
@@ -0,0 +1,77 @@
+# 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 time
+from unittest.mock import patch
+
+import pytest
+
+from airflow.providers.openai.hooks.openai import OpenAIHook
+from airflow.providers.openai.triggers.agent import OpenAIAgentSessionTrigger
+
+
[email protected]
+async def test_trigger_completes():
+ trigger = OpenAIAgentSessionTrigger("openai_default", "session", 1,
time.time() + 60)
+ event = {"status": "success", "session_id": "session", "turn_id": "turn",
"usage": None}
+ with patch.object(OpenAIHook, "poll_agent_session", autospec=True,
return_value=event):
+ events = [result.payload async for result in trigger.run()]
+ assert events == [event]
+
+
[email protected]
+async def test_trigger_timeout_survives_serialization():
+ trigger = OpenAIAgentSessionTrigger("openai_default", "session", 1,
time.time() - 1)
+ _, kwargs = trigger.serialize()
+ with patch.object(OpenAIHook, "poll_agent_session", autospec=True) as poll:
+ events = [result.payload async for result in
OpenAIAgentSessionTrigger(**kwargs).run()]
+ assert events[0]["status"] == "timeout"
+ poll.assert_not_called()
+
+
[email protected]
+async def test_trigger_gives_up_after_consecutive_poll_errors():
+ trigger = OpenAIAgentSessionTrigger("openai_default", "session", 0,
time.time() + 60)
+ with patch.object(
+ OpenAIHook, "poll_agent_session", autospec=True,
side_effect=RuntimeError("API failed")
+ ) as poll:
+ events = [result.payload async for result in trigger.run()]
+ assert events == [{"status": "error", "session_id": "session", "message":
"API failed"}]
+ assert poll.call_count == OpenAIHook.MAX_CONSECUTIVE_POLL_FAILURES
+
+
[email protected]
+async def test_trigger_recovers_from_transient_poll_error():
+ trigger = OpenAIAgentSessionTrigger("openai_default", "session", 0,
time.time() + 60)
+ event = {"status": "success", "session_id": "session", "turn_id": "turn",
"usage": None}
+ with patch.object(
+ OpenAIHook,
+ "poll_agent_session",
+ autospec=True,
+ side_effect=[RuntimeError("API failed"), None, RuntimeError("API
failed again"), event],
+ ):
+ events = [result.payload async for result in trigger.run()]
+ assert events == [event]
+
+
[email protected]
+async def test_trigger_kill():
+ trigger = OpenAIAgentSessionTrigger("openai_default", "session", 1,
time.time() + 60)
+ with patch.object(OpenAIHook, "cancel_agent_session", autospec=True) as
cancel:
+ await trigger.on_kill()
+ assert cancel.call_args.args[1] == "session"