fix(python): coerce JSON workflow resume payloads (#7684)

AG-UI clients send plain JSON, but structured response types were only
accepted as already-built instances, and core's coercion stopped at the
outer object, letting raw dicts sit inside typed fields. Coercion now
walks declared annotations and returns the input untouched whenever it
cannot satisfy them.
This commit is contained in:
NekoPunch
2026-08-19 04:20:24 +00:00
committed by GitHub
parent 1f738cdeb7
commit e74ac4613c
5 changed files with 548 additions and 16 deletions
@@ -10,7 +10,8 @@ import logging
import uuid
from collections.abc import AsyncGenerator
from functools import partial
from typing import Any, cast, get_args, get_origin
from types import UnionType
from typing import Any, Union, cast, get_args, get_origin, get_type_hints
from ag_ui.core import (
ActivitySnapshotEvent,
@@ -34,6 +35,10 @@ from agent_framework import (
Workflow,
WorkflowRunState,
)
from agent_framework._workflows._typing_utils import ( # pyright: ignore[reportPrivateUsage]
is_instance_of,
try_coerce_to_type,
)
from agent_framework.observability import (
_use_telemetry_conversation_id, # pyright: ignore[reportPrivateUsage]
)
@@ -535,6 +540,48 @@ def _coerce_compact_approval_response(request_data: Content, candidate: dict[str
return response
def _without_optional(annotation: Any) -> Any:
"""Unwrap ``X | None`` so optional fields normalize like their plain counterpart."""
if get_origin(annotation) not in (Union, UnionType):
return annotation
members = [member for member in get_args(annotation) if member is not type(None)]
return members[0] if len(members) == 1 else annotation
def _normalize_agui_message_fields(response_type: Any, candidate: Any) -> Any:
"""Convert AG-UI message payloads for fields the response type declares as Message.
Core coercion only understands the canonical ``contents`` form, so AG-UI wire shapes
(``{"role": ..., "content": ...}`` and bare strings) are translated here.
"""
if not isinstance(candidate, dict):
return candidate
try:
field_types = get_type_hints(response_type)
except Exception:
return candidate
normalized = dict(cast(dict[str, Any], candidate))
for name, annotation in field_types.items():
if name not in normalized:
continue
annotation = _without_optional(annotation)
target_type = get_origin(annotation) or annotation
if target_type is Message:
message = _coerce_message(normalized[name])
if message is not None:
normalized[name] = message
elif target_type is list and get_args(annotation)[:1] == (Message,):
items = normalized[name]
if not isinstance(items, list):
continue
messages = [_coerce_message(item) for item in cast(list[Any], items)]
if all(message is not None for message in messages):
normalized[name] = messages
return normalized
def _coerce_response_for_request(request_event: Any, value: Any) -> Any | None:
"""Coerce a candidate value into the request's expected response type."""
response_type = getattr(request_event, "response_type", None)
@@ -598,7 +645,8 @@ def _coerce_response_for_request(request_event: Any, value: Any) -> Any | None:
if target_type is float:
return candidate if isinstance(candidate, (int, float)) and not isinstance(candidate, bool) else None
if isinstance(target_type, type):
return candidate if isinstance(candidate, target_type) else None
coerced = try_coerce_to_type(_normalize_agui_message_fields(response_type, candidate), response_type)
return coerced if is_instance_of(coerced, response_type) else None
# Unknown typing metadata: preserve value as-is.
return candidate
@@ -8,6 +8,7 @@ import subprocess
import sys
from collections import Counter
from collections.abc import AsyncIterator, Callable
from dataclasses import dataclass
from inspect import signature
from typing import Any, cast
@@ -3776,6 +3777,63 @@ def _build_workflow_request_info_app(
return app
async def test_endpoint_workflow_request_info_resumes_dataclass_response_from_json():
"""Dataclass response types resume from plain JSON payloads, as AG-UI clients send them."""
@dataclass
class PlanReview:
review: list[Message]
class PlanReviewExecutor(Executor):
def __init__(self) -> None:
super().__init__(id="plan_review")
@handler
async def start(self, message: Any, ctx: WorkflowContext[Any, Any]) -> None:
del message
await ctx.request_info({"plan": "ship it"}, PlanReview, request_id="plan-review")
@response_handler
async def handle_review(
self, original_request: dict[str, Any], response: PlanReview, ctx: WorkflowContext[Any, Any]
) -> None:
del original_request
verdict = "approved" if not response.review else response.review[0].text
await ctx.yield_output(f"Plan {verdict}") # type: ignore[arg-type] # pyrefly: ignore[bad-argument-type]
app = FastAPI()
add_agent_framework_fastapi_endpoint(
app, WorkflowBuilder(start_executor=PlanReviewExecutor()).build(), path="/workflow"
)
with TestClient(app) as client:
pause_response = client.post(
"/workflow",
json={
"runId": "run-pause",
"threadId": "thread-plan",
"messages": [{"role": "user", "content": "Draft a plan"}],
},
)
assert pause_response.status_code == 200
resume_response = client.post(
"/workflow",
json={
"runId": "run-resume",
"threadId": "thread-plan",
"messages": [],
"resume": [{"interruptId": "plan-review", "status": "resolved", "payload": {"review": []}}],
},
)
assert resume_response.status_code == 200
resume_events = _decode_sse_events(resume_response)
assert not [event for event in resume_events if event.get("type") == "RUN_ERROR"]
text_deltas = [event["delta"] for event in resume_events if event.get("type") == "TEXT_MESSAGE_CONTENT"]
assert "Plan approved" in text_deltas
async def test_endpoint_workflow_request_info_emits_canonical_interrupt_and_resumes():
"""Workflow request_info pauses and resumes through canonical AG-UI interrupt payloads."""
app = _build_workflow_request_info_app()
@@ -4,6 +4,7 @@
import json
from collections.abc import AsyncIterator
from dataclasses import dataclass, make_dataclass
from enum import Enum
from types import SimpleNamespace
from typing import Any, cast
@@ -28,7 +29,9 @@ from agent_framework import (
response_handler,
tool,
)
from agent_framework.orchestrations import MagenticPlanReviewResponse
from conftest import StreamingChatClientStub # pyrefly: ignore[missing-import] # pyright: ignore[reportMissingImports]
from pydantic import BaseModel
from agent_framework_ag_ui._workflow_run import (
_coerce_content,
@@ -1314,6 +1317,93 @@ def test_coerce_response_for_request_bool_int_float_and_mismatch() -> None:
assert _coerce_response_for_request(dict_request, "[1,2,3]") is None
def test_coerce_response_for_request_builds_dataclass_from_json() -> None:
"""JSON objects should map onto dataclass response types such as plan review."""
request = SimpleNamespace(response_type=MagenticPlanReviewResponse)
approved = _coerce_response_for_request(request, {"review": []})
assert isinstance(approved, MagenticPlanReviewResponse)
assert approved.review == []
revised = _coerce_response_for_request(request, '{"review": [{"role": "user", "content": "add tests"}]}')
assert isinstance(revised, MagenticPlanReviewResponse)
assert len(revised.review) == 1
assert revised.review[0].text == "add tests"
from_strings = _coerce_response_for_request(request, {"review": ["add tests"]})
assert isinstance(from_strings, MagenticPlanReviewResponse)
assert from_strings.review[0].text == "add tests"
assert _coerce_response_for_request(request, {"unknown": 1}) is None
assert _coerce_response_for_request(request, "approve") is None
def test_coerce_response_for_request_leaves_non_message_fields_untouched() -> None:
"""Message normalization must follow field annotations, not payload shape."""
@dataclass
class Tagged:
note: Message
revision: Message | None
metadata: dict[str, str]
request = SimpleNamespace(response_type=Tagged)
message_shaped = {"role": "admin", "content": "keep raw"}
tagged = _coerce_response_for_request(
request,
{
"note": {"role": "user", "content": "translate me"},
"revision": {"role": "user", "content": "optional fields too"},
"metadata": message_shaped,
},
)
assert isinstance(tagged, Tagged)
assert tagged.note.text == "translate me"
assert tagged.revision is not None
assert tagged.revision.text == "optional fields too"
assert tagged.metadata == message_shaped
def test_coerce_response_for_request_rejects_malformed_message_field_payloads() -> None:
"""A Message-typed field that is not message-shaped must fail, not crash normalization."""
@dataclass
class Review:
notes: list[Message]
request = SimpleNamespace(response_type=Review)
assert _coerce_response_for_request(request, {"notes": "not-a-list"}) is None
assert _coerce_response_for_request(request, {"notes": [42]}) is None
def test_coerce_response_for_request_skips_normalization_without_resolvable_hints() -> None:
"""Unresolvable annotations skip message normalization instead of failing the resume."""
mystery_type = make_dataclass("Mystery", [("value", "DoesNotExist")])
request = SimpleNamespace(response_type=mystery_type)
mystery = _coerce_response_for_request(request, {"value": 1})
assert type(mystery) is mystery_type
assert vars(mystery) == {"value": 1}
def test_coerce_response_for_request_builds_pydantic_model_from_json() -> None:
"""JSON objects should validate into pydantic response types."""
class ReviewDecision(BaseModel):
approved: bool
request = SimpleNamespace(response_type=ReviewDecision)
decision = _coerce_response_for_request(request, {"approved": True})
assert isinstance(decision, ReviewDecision)
assert decision.approved is True
assert _coerce_response_for_request(request, {"approved": "not-a-bool"}) is None
async def test_workflow_run_emits_run_error_when_stream_raises() -> None:
"""Unexpected stream exceptions should be converted into RUN_ERROR events."""
@@ -3,8 +3,9 @@
import sys
import typing
from collections.abc import Mapping
from dataclasses import InitVar, is_dataclass
from types import ModuleType, UnionType
from typing import Any, TypeGuard, Union, cast, get_args, get_origin
from typing import Any, Literal, TypeGuard, Union, cast, get_args, get_origin, get_type_hints
import typing_extensions
@@ -225,6 +226,23 @@ def is_instance_of(data: Any, target_type: type | UnionType | Any) -> bool:
return isinstance(data, target_type)
def _matches_annotation(data: Any, annotation: Any) -> bool:
"""Check an annotation that may not be runtime-checkable, treating unchecked ones as a match."""
if get_origin(annotation) is Literal:
return any(data == member and type(data) is type(member) for member in get_args(annotation))
try:
return is_instance_of(data, annotation)
except TypeError:
# Annotated, NewType and friends cannot reach isinstance; rejecting them would break
# payloads that were accepted before field-level validation existed.
return True
def _coerced_or_original(coerced: Any, original: Any, target_type: Any) -> Any:
"""Keep a coercion only when it satisfies the annotation, so failures return the input."""
return coerced if _matches_annotation(coerced, target_type) else original
def try_coerce_to_type(data: Any, target_type: type | UnionType | Any) -> Any:
"""Try to coerce data to the target type.
@@ -244,39 +262,122 @@ def try_coerce_to_type(data: Any, target_type: type | UnionType | Any) -> Any:
original_data = data
# If already the right type, return as-is
if is_instance_of(data, target_type):
if _matches_annotation(data, target_type):
return data
# Can't coerce to non-concrete targets (Union, generic, etc.)
origin = get_origin(target_type)
if origin in (UnionType, Union):
for member_type in get_args(target_type):
coerced_member = try_coerce_to_type(data, member_type)
if _matches_annotation(coerced_member, member_type):
return coerced_member
return original_data
if origin is list and isinstance(data, list):
item_types = get_args(target_type) or (Any,)
coerced_list = [try_coerce_to_type(item, item_types[0]) for item in cast(list[Any], data)]
return _coerced_or_original(coerced_list, original_data, target_type)
if origin is dict and isinstance(data, dict):
key_type, value_type = get_args(target_type) or (Any, Any)
coerced_mapping = {
try_coerce_to_type(key, key_type): try_coerce_to_type(value, value_type)
for key, value in cast(dict[Any, Any], data).items()
}
return _coerced_or_original(coerced_mapping, original_data, target_type)
# JSON has no tuples or sets, so sequence payloads have to be rebuilt into them.
if origin in (set, frozenset) and isinstance(data, list):
item_types = get_args(target_type)
item_type = item_types[0] if item_types else Any
coerced_items = [try_coerce_to_type(item, item_type) for item in cast(list[Any], data)]
if not all(_matches_annotation(item, item_type) for item in coerced_items):
return original_data
try:
coerced_set = frozenset(coerced_items) if origin is frozenset else set(coerced_items)
except TypeError:
# An item that stayed unhashable never belonged in this set.
return original_data
return _coerced_or_original(coerced_set, original_data, target_type)
if origin is tuple and isinstance(data, list):
items = cast(list[Any], data)
item_types = get_args(target_type)
if len(item_types) == 2 and item_types[1] is Ellipsis:
item_types = (item_types[0],) * len(items)
if len(item_types) != len(items):
return original_data
coerced_tuple = tuple(
try_coerce_to_type(item, item_type) for item, item_type in zip(items, item_types, strict=True)
)
return _coerced_or_original(coerced_tuple, original_data, target_type)
# Can't coerce to non-concrete targets (generic aliases, etc.)
if not isinstance(target_type, type):
return original_data
target_cls: type[Any] = target_type
# int -> float (JSON integers for float fields)
if isinstance(data, int) and target_cls is float:
if isinstance(data, int) and not isinstance(data, bool) and target_cls is float:
return float(data)
# dict -> dataclass or pydantic model
# dict -> dataclass, pydantic model, or SerializationMixin type
if isinstance(data, dict):
from dataclasses import is_dataclass
payload = cast(dict[str, Any], data)
if is_dataclass(target_cls):
try:
return target_cls(**data)
except (TypeError, ValueError):
return original_data
return _coerce_dict_to_dataclass(payload, target_cls)
model_validate = getattr(target_cls, "model_validate", None)
if callable(model_validate):
try:
return model_validate(data)
return model_validate(payload)
except Exception:
return original_data
from_dict = getattr(target_cls, "from_dict", None)
if callable(from_dict):
try:
return from_dict(payload)
except Exception:
return original_data
return original_data
def _coerce_dict_to_dataclass(data: dict[str, Any], target_cls: type[Any]) -> Any:
"""Build a dataclass from a JSON-like dict, coercing each field to its annotation."""
try:
field_types = get_type_hints(target_cls)
except Exception:
field_types = {}
# __dataclass_fields__ keeps InitVar pseudo-fields that fields() drops, and the constructor takes them.
init_field_names = {name for name, field in target_cls.__dataclass_fields__.items() if field.init}
coerced_fields: dict[str, Any] = {}
for name, value in data.items():
if name not in init_field_names:
return data
annotation = field_types.get(name, Any)
if type(annotation) is InitVar:
annotation = cast(Any, annotation).type
coerced_value = try_coerce_to_type(value, annotation)
# Callers validate the outer type only, so a mismatched field has to fail the whole build.
if not _matches_annotation(coerced_value, annotation):
return data
coerced_fields[name] = coerced_value
try:
return target_cls(**coerced_fields)
except Exception:
# Constructor validation of any kind fails the build; it must not escape to callers
# that have no coercion error path.
return data
def serialize_type(t: type) -> str:
"""Serialize a type to a string.
@@ -2,14 +2,16 @@
import importlib
import sys
from dataclasses import dataclass
from dataclasses import InitVar, dataclass, make_dataclass
from dataclasses import field as dataclass_field
from types import ModuleType
from typing import Any, Generic, Optional, TypeVar, Union
from typing import Annotated, Any, Generic, Literal, NewType, Optional, TypeVar, Union
from unittest.mock import Mock
import pytest
from pydantic import BaseModel
from agent_framework import WorkflowEvent
from agent_framework import Message, WorkflowEvent
from agent_framework._workflows._typing_utils import (
deserialize_type,
is_instance_of,
@@ -533,6 +535,239 @@ def test_coerce_unrelated_types_returns_original() -> None:
assert try_coerce_to_type([1, 2], dict) == [1, 2]
def test_coerce_dict_to_dataclass_coerces_nested_fields() -> None:
"""Dataclass fields should be coerced to their annotations, not left as raw JSON."""
@dataclass
class Point:
x: int
y: int
@dataclass
class Path:
points: list[Point]
label: str | None
result = try_coerce_to_type({"points": [{"x": 1, "y": 2}], "label": None}, Path)
assert isinstance(result, Path)
assert result.points == [Point(x=1, y=2)]
assert result.label is None
def test_coerce_dict_to_dataclass_rejects_bad_field_values() -> None:
"""Field values that do not match their annotation must not build a half-typed object."""
@dataclass
class Point:
x: int
y: int
for payload in ({"x": 1, "y": "two"}, {"x": None, "y": 2}):
assert try_coerce_to_type(payload, Point) is payload
def test_coerce_dict_to_dataclass_respects_post_init_validation() -> None:
"""Constructor-side validation must fail the build instead of escaping to the caller."""
@dataclass
class Decision:
scores: list[int]
def __post_init__(self) -> None:
if self.scores and min(self.scores) < 0:
raise ValueError("scores must be non-negative")
assert try_coerce_to_type({"scores": [1, 2]}, Decision) == Decision(scores=[1, 2])
for payload in ({"scores": [-1]}, {"scores": ["high"]}):
assert try_coerce_to_type(payload, Decision) is payload
def test_coerce_dict_to_dataclass_survives_arbitrary_constructor_errors() -> None:
"""Constructors raising anything at all must fail the build, not the caller."""
class DomainError(Exception): ...
@dataclass
class Strict:
value: int
def __post_init__(self) -> None:
raise DomainError("never valid")
payload = {"value": 1}
assert try_coerce_to_type(payload, Strict) is payload
def test_coerce_tuple_target_rejects_length_mismatch() -> None:
"""A list that cannot fill a fixed-length tuple must come back untouched."""
assert try_coerce_to_type([1, 2, 3], tuple[int, str]) == [1, 2, 3]
assert try_coerce_to_type([1, "a"], tuple[int, str]) == (1, "a")
def test_coerce_dict_to_dataclass_allows_annotations_isinstance_cannot_check() -> None:
"""NewType and Annotated fields cannot reach isinstance, so they must not be rejected."""
UserId = NewType("UserId", int)
@dataclass
class Owner:
user_id: UserId
label: Annotated[str, "display"]
assert try_coerce_to_type({"user_id": 7, "label": "root"}, Owner) == Owner(user_id=UserId(7), label="root")
def test_coerce_dict_to_dataclass_falls_back_when_hints_do_not_resolve() -> None:
"""Unresolvable annotations degrade to no field coercion rather than failing the build."""
mystery = make_dataclass("Mystery", [("value", "DoesNotExist")])
coerced = try_coerce_to_type({"value": 1}, mystery)
assert type(coerced) is mystery
assert vars(coerced) == {"value": 1}
def test_coerce_variadic_tuple_target() -> None:
"""A variadic tuple annotation coerces every item to the single declared type."""
assert try_coerce_to_type([1, 2, 3], tuple[float, ...]) == (1.0, 2.0, 3.0)
assert try_coerce_to_type([], tuple[int, ...]) == ()
items = [1, "x"]
assert try_coerce_to_type(items, tuple[int, ...]) is items
def test_coerce_dict_to_pydantic_model() -> None:
"""Pydantic response types are built through their own validation."""
class Point(BaseModel):
x: int
assert try_coerce_to_type({"x": 1}, Point) == Point(x=1)
payload = {"x": "not-a-number"}
assert try_coerce_to_type(payload, Point) is payload
def test_coerce_container_targets_return_input_on_item_mismatch() -> None:
"""A container whose items cannot match must come back as the very same object."""
items = [1, "x"]
assert try_coerce_to_type(items, list[int]) is items
mapping = {"a": "x"}
assert try_coerce_to_type(mapping, dict[str, int]) is mapping
def test_coerce_dict_to_dataclass_accepts_init_only_variables() -> None:
"""InitVar parameters are constructor arguments, so their payload keys must be accepted."""
@dataclass
class Scaled:
value: int
scale: InitVar[int] = 1
tag: str = dataclass_field(init=False, default="fixed")
def __post_init__(self, scale: int) -> None:
self.value *= scale
assert try_coerce_to_type({"value": 2, "scale": 3}, Scaled) == Scaled(value=6)
for payload in ({"value": 2, "scale": "big"}, {"value": 2, "tag": "injected"}):
assert try_coerce_to_type(payload, Scaled) is payload
def test_coerce_frozenset_target_checks_item_type() -> None:
"""frozenset members are not checked by is_instance_of, so the branch validates them itself."""
payload = ["x"]
assert try_coerce_to_type(payload, frozenset[int]) is payload
assert try_coerce_to_type([1, 2], frozenset[int]) == frozenset({1, 2})
def test_coerce_set_target_rejects_unhashable_items() -> None:
"""Items that stay unhashable must fail the coercion instead of raising."""
payload: list[Any] = [{}]
assert try_coerce_to_type(payload, set[int]) is payload
assert try_coerce_to_type([1, 2], set[int]) == {1, 2}
def test_coerce_dict_to_dataclass_rejects_non_init_fields() -> None:
"""Payload keys that are not constructor parameters must be rejected, not silently dropped."""
@dataclass
class Tagged:
value: int
tag: str = dataclass_field(init=False, default="unset")
payload = {"value": 1, "tag": "injected"}
assert try_coerce_to_type(payload, Tagged) is payload
def test_coerce_dict_to_dataclass_checks_literal_fields() -> None:
"""Literal fields cannot reach isinstance, but their allowed values still have to hold."""
@dataclass
class Verdict:
decision: Literal["approve", "revise"]
assert try_coerce_to_type({"decision": "approve"}, Verdict) == Verdict(decision="approve")
payload = {"decision": "delete-everything"}
assert try_coerce_to_type(payload, Verdict) is payload
def test_coerce_dict_to_dataclass_coerces_container_fields() -> None:
"""Mapping, set and tuple fields arrive as JSON objects/arrays and must be rebuilt."""
@dataclass
class Config:
weights: dict[str, float]
tags: set[str]
bounds: tuple[int, str]
result = try_coerce_to_type({"weights": {"x": 1}, "tags": ["a"], "bounds": [1, "high"]}, Config)
assert isinstance(result, Config)
assert result.weights == {"x": 1.0}
assert result.tags == {"a"}
assert result.bounds == (1, "high")
def test_coerce_bool_is_not_treated_as_float() -> None:
"""JSON booleans must not slip into float fields as 1.0/0.0."""
assert try_coerce_to_type(True, float) is True
def test_coerce_dict_to_serialization_mixin_type() -> None:
"""Types exposing from_dict (Message, Content) should be built from JSON objects."""
result = try_coerce_to_type({"role": "user", "contents": [{"type": "text", "text": "hi"}]}, Message)
assert isinstance(result, Message)
assert result.text == "hi"
payload = {"contents": 5}
assert try_coerce_to_type(payload, Message) is payload
def test_coerce_union_target_picks_matching_member() -> None:
"""Union targets should coerce into the first member that accepts the value."""
@dataclass
class Point:
x: int
y: int
result = try_coerce_to_type({"x": 1, "y": 2}, Point | None)
assert result == Point(x=1, y=2)
def test_coerce_list_target_coerces_items() -> None:
"""Typed list targets should coerce their items."""
@dataclass
class Point:
x: int
y: int
assert try_coerce_to_type([{"x": 1, "y": 2}], list[Point]) == [Point(x=1, y=2)]
def test_coerce_any_returns_original() -> None:
"""Any target type should accept any value without coercion."""
assert try_coerce_to_type(42, Any) == 42