diff --git a/python/packages/core/agent_framework/_workflows/_agent.py b/python/packages/core/agent_framework/_workflows/_agent.py index f0b51969c07..01bd70d300d 100644 --- a/python/packages/core/agent_framework/_workflows/_agent.py +++ b/python/packages/core/agent_framework/_workflows/_agent.py @@ -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." @@ -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__( diff --git a/python/packages/core/agent_framework/_workflows/_events.py b/python/packages/core/agent_framework/_workflows/_events.py index 8788599e652..669887a3c29 100644 --- a/python/packages/core/agent_framework/_workflows/_events.py +++ b/python/packages/core/agent_framework/_workflows/_events.py @@ -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 @@ -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( @@ -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), ) diff --git a/python/packages/core/agent_framework/_workflows/_typing_utils.py b/python/packages/core/agent_framework/_workflows/_typing_utils.py index 404b83f7092..6dff6ba7a48 100644 --- a/python/packages/core/agent_framework/_workflows/_typing_utils.py +++ b/python/packages/core/agent_framework/_workflows/_typing_utils.py @@ -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 @@ -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). @@ -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: diff --git a/python/packages/core/tests/workflow/test_request_info_event_rehydrate.py b/python/packages/core/tests/workflow/test_request_info_event_rehydrate.py index c05ab66984e..1be0a75f58e 100644 --- a/python/packages/core/tests/workflow/test_request_info_event_rehydrate.py +++ b/python/packages/core/tests/workflow/test_request_info_event_rehydrate.py @@ -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, @@ -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( diff --git a/python/packages/core/tests/workflow/test_typing_utils.py b/python/packages/core/tests/workflow/test_typing_utils.py index cc3e7bf4ea3..6e54cd64c7c 100644 --- a/python/packages/core/tests/workflow/test_typing_utils.py +++ b/python/packages/core/tests/workflow/test_typing_utils.py @@ -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 @@ -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 diff --git a/python/packages/core/tests/workflow/test_workflow_agent.py b/python/packages/core/tests/workflow/test_workflow_agent.py index 0cf69aa9414..d5f0c1a6861 100644 --- a/python/packages/core/tests/workflow/test_workflow_agent.py +++ b/python/packages/core/tests/workflow/test_workflow_agent.py @@ -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.