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
3 changes: 2 additions & 1 deletion burr/core/action.py
Original file line number Diff line number Diff line change
Expand Up @@ -1507,7 +1507,7 @@ def pydantic(
writes: List[str],
state_input_type: Type["BaseModel"],
state_output_type: Type["BaseModel"],
stream_type: Union[Type["BaseModel"], Type[dict]],
stream_type: Union[Type["BaseModel"], Type[dict], Any],
tags: Optional[List[str]] = None,
) -> Callable:
"""Creates a streaming action that uses pydantic models.
Expand Down Expand Up @@ -1607,3 +1607,4 @@ def create_action(action_: Union[Callable, ActionT], name: str) -> ActionT:
f"Object {action_} is not a valid action. Have you decorated it with @action or @streaming_action?"
)
return action_.with_name(name)

9 changes: 7 additions & 2 deletions burr/integrations/pydantic.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@

import copy
import inspect
import sys
import types
import typing
from typing import (
Expand Down Expand Up @@ -269,7 +270,10 @@ async def async_action_function(state: State, **kwargs) -> State:
return decorator


PartialType = Union[Type[pydantic.BaseModel], Type[dict]]
if sys.version_info >= (3, 10):
PartialType = Union[Type[pydantic.BaseModel], Type[dict], types.UnionType] # noqa: E501
else:
PartialType = Union[Type[pydantic.BaseModel], Type[dict]]

PydanticStreamingActionFunctionSync = Callable[
..., Generator[Tuple[Union[pydantic.BaseModel, dict], Optional[pydantic.BaseModel]], None, None]
Expand All @@ -290,7 +294,7 @@ async def async_action_function(state: State, **kwargs) -> State:

def _validate_and_extract_signature_types_streaming(
fn: PydanticStreamingActionFunction,
stream_type: Optional[Union[Type[pydantic.BaseModel], Type[dict]]],
stream_type: Optional[PartialType],
state_input_type: Optional[Type[pydantic.BaseModel]] = None,
state_output_type: Optional[Type[pydantic.BaseModel]] = None,
) -> Tuple[
Expand Down Expand Up @@ -421,3 +425,4 @@ def construct_data(self, state: State) -> StateModel:

def construct_state(self, data: StateModel) -> State:
return State(model_to_dict(data))

111 changes: 111 additions & 0 deletions tests/integrations/test_burr_pydantic.py
Original file line number Diff line number Diff line change
Expand Up @@ -779,3 +779,114 @@ async def final_result_streamed(
assert state.data.count == 20
assert isinstance(result, IntermediateModel)
assert result.result == 20


# --- Tests for stream_type union support (Issue #607) ---


import sys


@pytest.mark.skipif(
sys.version_info < (3, 10),
reason="Union type syntax with | requires Python 3.10+",
)
def test_validate_streaming_signature_accepts_union_stream_type():
"""Regression test: _validate_and_extract_signature_types_streaming should accept
a union stream_type (e.g. ModelA | ModelB) on Python 3.10+.

Before the fix, PartialType did not include types.UnionType, causing a TypeError
when a union type was passed as stream_type.
"""
from burr.integrations.pydantic import _validate_and_extract_signature_types_streaming

class ModelA(BaseModel):
value: int

class ModelB(BaseModel):
label: str

def dummy_fn(
state: AppStateModel,
) -> Generator[Tuple[ModelA, Optional[AppStateModel]], None, None]:
yield ModelA(value=1), state

union_type = ModelA | ModelB # This is types.UnionType on Python 3.10+

state_in, state_out, resolved_stream_type = _validate_and_extract_signature_types_streaming(
dummy_fn,
stream_type=union_type,
state_input_type=AppStateModel,
state_output_type=AppStateModel,
)
assert state_in is AppStateModel
assert state_out is AppStateModel
assert resolved_stream_type is union_type


@pytest.mark.skipif(
sys.version_info < (3, 10),
reason="Union type syntax with | requires Python 3.10+",
)
def test_pydantic_streaming_action_accepts_union_stream_type():
"""Regression test: @pydantic_streaming_action should accept a union stream_type
on Python 3.10+.

This is the user-facing decorator that was broken by issue #607.
"""

class ResultA(BaseModel):
value: int

class ResultB(BaseModel):
label: str

@pydantic_streaming_action(
reads=["count", "times_called"],
writes=["count", "times_called"],
stream_type=ResultA | ResultB,
state_input_type=AppStateModel,
state_output_type=AppStateModel,
)
def act(
state: AppStateModel,
) -> Generator[Tuple[ResultA, Optional[AppStateModel]], None, None]:
state.count += 1
yield ResultA(value=state.count), state

assert hasattr(act, "bind")
action_fn = getattr(act, FunctionBasedAction.ACTION_FUNCTION, None)
assert action_fn is not None


@pytest.mark.skipif(
sys.version_info < (3, 10),
reason="Union type syntax with | requires Python 3.10+",
)
def test_streaming_action_pydantic_decorator_accepts_union_stream_type():
"""Regression test: @streaming_action.pydantic should accept a union stream_type
on Python 3.10+, matching pydantic_streaming_action behavior.
"""

class ResultA(BaseModel):
value: int

class ResultB(BaseModel):
label: str

@streaming_action.pydantic(
reads=["count", "times_called"],
writes=["count", "times_called"],
stream_type=ResultA | ResultB,
state_input_type=AppStateModel,
state_output_type=AppStateModel,
)
def act(
state: AppStateModel,
) -> Generator[Tuple[ResultA, Optional[AppStateModel]], None, None]:
state.count += 1
yield ResultA(value=state.count), state

assert hasattr(act, "bind")
assert getattr(act, FunctionBasedAction.ACTION_FUNCTION, None) is not None

Loading