Skip to content
Merged
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
18 changes: 16 additions & 2 deletions python/packages/core/agent_framework/_workflows/_agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -64,7 +64,18 @@ def to_dict(self) -> dict[str, Any]:
return {"request_id": self.request_id, "request_event": self.request_event.to_dict()}

@classmethod
def from_dict(cls, payload: dict[str, Any]) -> WorkflowAgent.RequestInfoFunctionArgs:
def from_dict(
cls,
payload: dict[str, Any],
*,
allowed_types: Mapping[str, type[Any]] | None = None,
) -> WorkflowAgent.RequestInfoFunctionArgs:
"""Create request-info function arguments from a dictionary.

Args:
payload: Serialized request-info function arguments.
allowed_types: Optional exact mapping of serialized names to trusted custom types.
"""
if "request_id" not in payload or "request_event" not in payload:
raise ValueError(
"Invalid payload for RequestInfoFunctionArgs. 'request_id' and 'request_event' are required."
Expand All @@ -74,7 +85,10 @@ def from_dict(cls, payload: dict[str, Any]) -> WorkflowAgent.RequestInfoFunction

return cls(
request_id=payload.get("request_id", ""),
request_event=WorkflowEvent.from_dict(payload.get("request_event", {})),
request_event=WorkflowEvent.from_dict(
payload.get("request_event", {}),
allowed_types=allowed_types,
),
)

def __init__(
Expand Down
20 changes: 15 additions & 5 deletions python/packages/core/agent_framework/_workflows/_events.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,7 @@
import sys
import traceback as _traceback
import warnings
from collections.abc import Generator
from collections.abc import Generator, Mapping
from contextlib import contextmanager
from contextvars import ContextVar
from dataclasses import dataclass
Expand Down Expand Up @@ -426,14 +426,24 @@ def to_dict(self) -> dict[str, Any]:
}

@classmethod
def from_dict(cls, data: dict[str, Any]) -> WorkflowEvent[Any]:
"""Create a REQUEST_INFO event from a dictionary."""
def from_dict(
cls,
data: dict[str, Any],
*,
allowed_types: Mapping[str, builtins.type[Any]] | None = None,
) -> WorkflowEvent[Any]:
"""Create a request-info event from a dictionary.

Args:
data: Serialized request-info event fields.
allowed_types: Optional exact mapping of serialized names to trusted custom types.
"""
for prop in ["data", "request_id", "source_executor_id", "request_type", "response_type"]:
if prop not in data:
raise KeyError(f"Missing '{prop}' field in WorkflowEvent dictionary.")

request_data = data["data"]
request_type = deserialize_type(data["request_type"])
request_type = deserialize_type(data["request_type"], allowed_types=allowed_types)

if request_type is not type(request_data):
raise TypeError(
Expand All @@ -444,5 +454,5 @@ def from_dict(cls, data: dict[str, Any]) -> WorkflowEvent[Any]:
request_id=data["request_id"],
source_executor_id=data["source_executor_id"],
request_data=cast(Any, request_data), # type: ignore
response_type=deserialize_type(data["response_type"]),
response_type=deserialize_type(data["response_type"], allowed_types=allowed_types),
)
Original file line number Diff line number Diff line change
@@ -1,7 +1,9 @@
# Copyright (c) Microsoft. All rights reserved.

import sys
import typing
from types import UnionType
from collections.abc import Mapping
from types import ModuleType, UnionType
from typing import Any, TypeGuard, Union, cast, get_args, get_origin

import typing_extensions
Expand All @@ -14,6 +16,17 @@
_TYPEVAR_TYPES: tuple[type, ...] = (type(typing.TypeVar("_T")), type(typing_extensions.TypeVar("_T"))) # pyright: ignore[reportUnknownVariableType]


def _is_runtime_type(value: object) -> TypeGuard[type[Any]]:
if not isinstance(value, type):
return False
try:
type.__getattribute__(value, "__module__")
type.__getattribute__(value, "__qualname__")
except TypeError:
return False
return True


def is_typevar(x: Any) -> bool:
"""Check if x is an unresolved TypeVar instance (from typing or typing_extensions).

Expand Down Expand Up @@ -274,19 +287,46 @@ def serialize_type(t: type) -> str:
return f"{t.__module__}.{t.__qualname__}"


def deserialize_type(serialized_type_string: str) -> type:
def deserialize_type(
serialized_type_string: str,
*,
allowed_types: Mapping[str, type[Any]] | None = None,
) -> type:
"""Deserialize a serialized type string.

Resolution is limited to exact caller-supplied types or types already present
in loaded module namespaces. This function never imports a module selected by
the serialized value.

Args:
serialized_type_string: Fully qualified serialized type name.
allowed_types: Optional exact mapping of serialized names to trusted types.

For example,

deserialize_type("builtins.int") => int
"""
import importlib
if allowed_types is not None and serialized_type_string in allowed_types:
resolved = allowed_types[serialized_type_string]
if not _is_runtime_type(resolved):
raise TypeError(f"allowed_types entry {serialized_type_string!r} must be a type.")
if serialize_type(resolved) != serialized_type_string:
raise ValueError(f"allowed_types entry {serialized_type_string!r} does not match the supplied type.")
return resolved

module_name, _, type_name = serialized_type_string.rpartition(".")
module = importlib.import_module(module_name)

return cast(type, getattr(module, type_name))
module = sys.modules.get(module_name)
if not isinstance(module, ModuleType):
raise ModuleNotFoundError(f"No module named {module_name!r}", name=module_name)

namespace = ModuleType.__getattribute__(module, "__dict__")
if type_name not in namespace:
raise AttributeError(f"{module_name!r} has no attribute {type_name!r}")

resolved = namespace[type_name]
if not _is_runtime_type(resolved):
raise TypeError(f"{serialized_type_string!r} does not resolve to a type.")
return resolved


def is_type_compatible(source_type: type | UnionType | Any, target_type: type | UnionType | Any) -> bool:
Expand Down
Original file line number Diff line number Diff line change
@@ -1,8 +1,12 @@
# Copyright (c) Microsoft. All rights reserved.

import json
import sys
from dataclasses import dataclass, field
from datetime import datetime, timezone
from pathlib import Path

import pytest

from agent_framework import (
FileCheckpointStorage,
Expand Down Expand Up @@ -47,6 +51,91 @@ class TimedApproval:
issued_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))


def test_workflow_event_from_dict_accepts_explicit_allowed_types() -> None:
"""Request-info reconstruction accepts exact trusted custom types."""

@dataclass
class ExplicitRequest:
prompt: str

class ExplicitResponse:
pass

request_type_name = f"{ExplicitRequest.__module__}.{ExplicitRequest.__qualname__}"
response_type_name = f"{ExplicitResponse.__module__}.{ExplicitResponse.__qualname__}"
event = WorkflowEvent.from_dict(
{
"type": "request_info",
"data": ExplicitRequest(prompt="Approve?"),
"request_id": "request-123",
"source_executor_id": "review_gateway",
"request_type": request_type_name,
"response_type": response_type_name,
},
allowed_types={
request_type_name: ExplicitRequest,
response_type_name: ExplicitResponse,
},
)

assert type(event.data) is ExplicitRequest
assert event.request_type is ExplicitRequest
assert event.response_type is ExplicitResponse


def _write_observable_type_module(tmp_path: Path, module_name: str) -> Path:
marker_path = tmp_path / f"{module_name}.imported"
module_path = tmp_path / f"{module_name}.py"
module_path.write_text(f"from pathlib import Path\nPath({str(marker_path)!r}).touch()\nclass Attack:\n pass\n")
return marker_path


def test_workflow_event_from_dict_does_not_import_request_type(
tmp_path: Path,
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""A request-type name cannot cause its module to be imported."""
module_name = "_request_info_untrusted_request_type"
marker_path = _write_observable_type_module(tmp_path, module_name)
monkeypatch.syspath_prepend(str(tmp_path))

with pytest.raises(ModuleNotFoundError, match=f"No module named '{module_name}'"):
WorkflowEvent.from_dict({
"type": "request_info",
"data": "Approve?",
"request_id": "request-123",
"source_executor_id": "review_gateway",
"request_type": f"{module_name}.Attack",
"response_type": "builtins.bool",
})

assert module_name not in sys.modules
assert not marker_path.exists()


def test_workflow_event_from_dict_does_not_import_response_type(
tmp_path: Path,
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""A response-type name cannot cause its module to be imported."""
module_name = "_request_info_untrusted_response_type"
marker_path = _write_observable_type_module(tmp_path, module_name)
monkeypatch.syspath_prepend(str(tmp_path))

with pytest.raises(ModuleNotFoundError, match=f"No module named '{module_name}'"):
WorkflowEvent.from_dict({
"type": "request_info",
"data": "Approve?",
"request_id": "request-123",
"source_executor_id": "review_gateway",
"request_type": "builtins.str",
"response_type": f"{module_name}.Attack",
})

assert module_name not in sys.modules
assert not marker_path.exists()


async def test_rehydrate_request_info_event() -> None:
"""Rehydration should succeed for valid request info events."""
request_info_event = WorkflowEvent.request_info(
Expand Down
50 changes: 50 additions & 0 deletions python/packages/core/tests/workflow/test_typing_utils.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,11 @@
# Copyright (c) Microsoft. All rights reserved.

import importlib
import sys
from dataclasses import dataclass
from types import ModuleType
from typing import Any, Generic, Optional, TypeVar, Union
from unittest.mock import Mock

import pytest

Expand Down Expand Up @@ -337,6 +341,52 @@ def test_deserialize_type_error_handling() -> None:
deserialize_type("builtins.NonExistentType")


def test_deserialize_type_does_not_import_unknown_module(monkeypatch: pytest.MonkeyPatch) -> None:
"""Unknown serialized types fail without importing payload-selected modules."""
imported_modules: list[str] = []

def fail_import(module_name: str) -> None:
imported_modules.append(module_name)
raise AssertionError("deserialize_type must not import payload-selected modules")

monkeypatch.setattr(importlib, "import_module", fail_import)

with pytest.raises(ModuleNotFoundError, match="No module named 'untrusted_request_info_payload'"):
deserialize_type("untrusted_request_info_payload.Attack")

assert imported_modules == []


def test_deserialize_type_accepts_explicit_allowed_type() -> None:
"""Callers can resolve an exact trusted custom type without importing its module."""

class ExplicitType:
pass

serialized_name = f"{ExplicitType.__module__}.{ExplicitType.__qualname__}"

assert deserialize_type(serialized_name, allowed_types={serialized_name: ExplicitType}) is ExplicitType


def test_deserialize_type_rejects_spoofed_allowed_type() -> None:
"""Allowed values must be actual class objects, not objects spoofing ``type``."""
spoofed_type = Mock(spec=type)

with pytest.raises(TypeError, match="must be a type"):
deserialize_type("spoofed.Type", allowed_types={"spoofed.Type": spoofed_type}) # type: ignore[dict-item]


def test_deserialize_type_rejects_spoofed_loaded_type(monkeypatch: pytest.MonkeyPatch) -> None:
"""Loaded namespace values must be actual class objects."""
module_name = "_spoofed_request_info_type"
module = ModuleType(module_name)
module.__dict__["Spoofed"] = Mock(spec=type)
monkeypatch.setitem(sys.modules, module_name, module)

with pytest.raises(TypeError, match="does not resolve to a type"):
deserialize_type(f"{module_name}.Spoofed")


def test_type_compatibility_basic() -> None:
"""Test basic type compatibility scenarios."""
# Exact type match
Expand Down
26 changes: 26 additions & 0 deletions python/packages/core/tests/workflow/test_workflow_agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -335,6 +335,32 @@ def test_request_info_dataclass_arguments_are_serialized_when_content_is_created
assert deserialized_args.request_event.data == HandoffRequest(target_agent="helper", reason="overflow")
assert deserialized_args.request_event.response_type is str

def test_request_info_function_args_from_dict_accepts_explicit_allowed_types(self) -> None:
"""Envelope reconstruction forwards exact trusted custom types."""

@dataclass
class ExplicitRequest:
prompt: str

serialized_name = f"{ExplicitRequest.__module__}.{ExplicitRequest.__qualname__}"
args = WorkflowAgent.RequestInfoFunctionArgs.from_dict(
{
"request_id": "request-123",
"request_event": {
"type": "request_info",
"data": ExplicitRequest(prompt="Approve?"),
"request_id": "request-123",
"source_executor_id": "review_gateway",
"request_type": serialized_name,
"response_type": "builtins.bool",
},
},
allowed_types={serialized_name: ExplicitRequest},
)

assert type(args.request_event.data) is ExplicitRequest
assert args.request_event.response_type is bool

def test_process_request_info_event_passes_through_function_approval_request(self) -> None:
"""If the event data is already a function approval request, it is forwarded unchanged.

Expand Down
Loading