Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,8 @@

from __future__ import annotations

from typing import Any
import json
from typing import Any, get_origin

from pydantic import BaseModel, TypeAdapter, ValidationError

Expand All @@ -39,12 +40,26 @@ def rehydrate_pydantic_output(
pydantic ``TypeAdapter``. When validation fails (reviewer edited the string
into something the type rejects), returns ``raw`` unchanged.

pydantic-ai also accepts an ``output_type`` that is neither a class nor a
generic alias: a list of output types, an output marker such as
``ToolOutput``, or an output function. For these output types, ``raw`` is
parsed with ``json.loads`` instead, and ``raw`` is returned unchanged when
it is not JSON.

When ``serialize_output`` is ``True``, returns the model dumped to a
``dict`` -- matches the operator's ``serialize_output=True`` opt-in for
consumers that want the dict shape.
"""
if output_type is str:
return raw
if not isinstance(output_type, type) and get_origin(output_type) is None:
# TypeAdapter raises for a list of output types or for a marker. For an output function, TypeAdapter
# builds a schema for the function's arguments, so validating the reviewed text would call the
# function again with that text as its arguments.
try:
return json.loads(raw)
except (ValueError, TypeError):
return raw
try:
rehydrated = TypeAdapter(output_type).validate_json(raw)
except (ValidationError, ValueError, TypeError):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -1133,6 +1133,31 @@ def test_execute_with_hitl_returns_the_approved_output_as_output_type(
assert result == expected
assert type(result) is type(expected)

@pytest.mark.skipif(
not AIRFLOW_V_3_1_PLUS, reason="Human in the loop is only compatible with Airflow >= 3.1.0"
)
@patch("airflow.providers.common.ai.operators.agent.AgentOperator.run_hitl_review", autospec=True)
@patch("airflow.providers.common.ai.operators.agent.PydanticAIHook", autospec=True)
def test_execute_with_hitl_parses_the_approved_output_for_a_list_of_output_types(
self, mock_hook_cls, mock_run_hitl, make_mock_run_result
):
mock_agent = MagicMock(spec=["run_sync", "instrument"])
mock_agent.run_sync.return_value = make_mock_run_result(Summary(text="ok", score=1.0))
mock_hook_cls.get_hook.return_value.create_agent.return_value = mock_agent
mock_run_hitl.return_value = '{"text":"ok","score":1.0}'
op = AgentOperator(
task_id="test",
prompt="Summarize",
llm_conn_id="my_llm",
output_type=[Summary, int],
enable_hitl_review=True,
hitl_timeout=timedelta(minutes=5),
)

result = op.execute(context=MagicMock())

assert result == {"text": "ok", "score": 1.0}

@pytest.mark.skipif(
not AIRFLOW_V_3_1_PLUS, reason="Human in the loop is only compatible with Airflow >= 3.1.0"
)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@

import pytest
from pydantic import BaseModel
from pydantic_ai import ToolOutput

from airflow.providers.common.ai.utils.output_type import rehydrate_pydantic_output

Expand All @@ -26,6 +27,10 @@ class A(BaseModel):
x: int


class B(BaseModel):
y: str


class TestRehydratePydanticOutput:
def test_returns_model_instance(self):
result = rehydrate_pydantic_output(A, '{"x": 7}', serialize_output=False)
Expand Down Expand Up @@ -60,3 +65,26 @@ def test_returns_raw_on_schema_mismatch(self):
# ``A`` requires ``x: int`` -- this payload should fail validation
result = rehydrate_pydantic_output(A, '{"y": "no-x-field"}', serialize_output=False)
assert result == '{"y": "no-x-field"}'

@pytest.mark.parametrize(
("output_type", "raw", "expected"),
[
pytest.param([A, B], '{"x": 7}', {"x": 7}, id="list-of-output-types"),
pytest.param(ToolOutput(A), '{"x": 7}', {"x": 7}, id="output-marker"),
pytest.param([A, B], "not-json", "not-json", id="not-json"),
],
)
def test_parses_output_type_that_is_not_a_type_as_json(self, output_type, raw, expected):
assert rehydrate_pydantic_output(output_type, raw, serialize_output=False) == expected

def test_does_not_call_an_output_function_again(self):
calls = []

def make_a(x: int) -> A:
calls.append(x)
return A(x=x)

result = rehydrate_pydantic_output(make_a, '{"x": 7}', serialize_output=False)

assert result == {"x": 7}
assert calls == []