From ccf22ac9635980d696ecdd6ed4a51fe78ed61f17 Mon Sep 17 00:00:00 2001 From: Peter Ibekwe Date: Thu, 30 Apr 2026 15:19:58 -0700 Subject: [PATCH 1/8] Add Python parity for HttpRequestAction in declarative workflow --- .../agent_framework/declarative/__init__.py | 5 + .../agent_framework/declarative/__init__.pyi | 10 + python/packages/declarative/AGENTS.md | 3 +- .../agent_framework_declarative/__init__.py | 10 + .../_workflows/__init__.py | 19 +- .../_workflows/_declarative_builder.py | 23 + .../_workflows/_errors.py | 42 ++ .../_workflows/_executors_http.py | 432 +++++++++++ .../_workflows/_factory.py | 19 +- .../_workflows/_http_handler.py | 229 ++++++ python/packages/declarative/pyproject.toml | 1 + .../test_default_http_request_handler.py | 350 +++++++++ .../tests/test_http_request_executor.py | 673 ++++++++++++++++++ .../test_http_request_yaml_integration.py | 111 +++ .../tests/workflows/http_request.yaml | 29 + .../declarative/invoke_http_request/main.py | 97 +++ .../invoke_http_request/workflow.yaml | 57 ++ python/uv.lock | 2 + 18 files changed, 2105 insertions(+), 7 deletions(-) create mode 100644 python/packages/declarative/agent_framework_declarative/_workflows/_errors.py create mode 100644 python/packages/declarative/agent_framework_declarative/_workflows/_executors_http.py create mode 100644 python/packages/declarative/agent_framework_declarative/_workflows/_http_handler.py create mode 100644 python/packages/declarative/tests/test_default_http_request_handler.py create mode 100644 python/packages/declarative/tests/test_http_request_executor.py create mode 100644 python/packages/declarative/tests/test_http_request_yaml_integration.py create mode 100644 python/packages/declarative/tests/workflows/http_request.yaml create mode 100644 python/samples/03-workflows/declarative/invoke_http_request/main.py create mode 100644 python/samples/03-workflows/declarative/invoke_http_request/workflow.yaml diff --git a/python/packages/core/agent_framework/declarative/__init__.py b/python/packages/core/agent_framework/declarative/__init__.py index 7b7737dd47d..ba88e6a0a98 100644 --- a/python/packages/core/agent_framework/declarative/__init__.py +++ b/python/packages/core/agent_framework/declarative/__init__.py @@ -21,10 +21,15 @@ "AgentFactory", "AgentExternalInputRequest", "AgentExternalInputResponse", + "DeclarativeActionError", "DeclarativeLoaderError", "DeclarativeWorkflowError", + "DefaultHttpRequestHandler", "ExternalInputRequest", "ExternalInputResponse", + "HttpRequestHandler", + "HttpRequestInfo", + "HttpRequestResult", "ProviderLookupError", "ProviderTypeMapping", "WorkflowFactory", diff --git a/python/packages/core/agent_framework/declarative/__init__.pyi b/python/packages/core/agent_framework/declarative/__init__.pyi index 92da0da682e..f18be22f50c 100644 --- a/python/packages/core/agent_framework/declarative/__init__.pyi +++ b/python/packages/core/agent_framework/declarative/__init__.pyi @@ -4,10 +4,15 @@ from agent_framework_declarative import ( AgentExternalInputRequest, AgentExternalInputResponse, AgentFactory, + DeclarativeActionError, DeclarativeLoaderError, DeclarativeWorkflowError, + DefaultHttpRequestHandler, ExternalInputRequest, ExternalInputResponse, + HttpRequestHandler, + HttpRequestInfo, + HttpRequestResult, ProviderLookupError, ProviderTypeMapping, WorkflowFactory, @@ -18,10 +23,15 @@ __all__ = [ "AgentExternalInputRequest", "AgentExternalInputResponse", "AgentFactory", + "DeclarativeActionError", "DeclarativeLoaderError", "DeclarativeWorkflowError", + "DefaultHttpRequestHandler", "ExternalInputRequest", "ExternalInputResponse", + "HttpRequestHandler", + "HttpRequestInfo", + "HttpRequestResult", "ProviderLookupError", "ProviderTypeMapping", "WorkflowFactory", diff --git a/python/packages/declarative/AGENTS.md b/python/packages/declarative/AGENTS.md index 72ac14860c3..1add6146015 100644 --- a/python/packages/declarative/AGENTS.md +++ b/python/packages/declarative/AGENTS.md @@ -8,7 +8,8 @@ YAML/JSON-based declarative agent and workflow definitions. - **`WorkflowFactory`** - Creates workflows from declarative definitions - **`WorkflowState`** - State management for declarative workflows - **`ProviderTypeMapping`** - Maps provider types to implementations -- **`DeclarativeLoaderError`** / **`ProviderLookupError`** - Error types +- **`HttpRequestHandler`** / **`DefaultHttpRequestHandler`** - Pluggable HTTP transport for the `HttpRequestAction` declarative action (configured via `WorkflowFactory(http_request_handler=...)`) +- **`DeclarativeLoaderError`** / **`ProviderLookupError`** / **`DeclarativeWorkflowError`** / **`DeclarativeActionError`** - Error types ## External Input Handling diff --git a/python/packages/declarative/agent_framework_declarative/__init__.py b/python/packages/declarative/agent_framework_declarative/__init__.py index 8200dd42e7d..6afcb3c7912 100644 --- a/python/packages/declarative/agent_framework_declarative/__init__.py +++ b/python/packages/declarative/agent_framework_declarative/__init__.py @@ -6,9 +6,14 @@ from ._workflows import ( AgentExternalInputRequest, AgentExternalInputResponse, + DeclarativeActionError, DeclarativeWorkflowError, + DefaultHttpRequestHandler, ExternalInputRequest, ExternalInputResponse, + HttpRequestHandler, + HttpRequestInfo, + HttpRequestResult, WorkflowFactory, WorkflowState, ) @@ -22,10 +27,15 @@ "AgentExternalInputRequest", "AgentExternalInputResponse", "AgentFactory", + "DeclarativeActionError", "DeclarativeLoaderError", "DeclarativeWorkflowError", + "DefaultHttpRequestHandler", "ExternalInputRequest", "ExternalInputResponse", + "HttpRequestHandler", + "HttpRequestInfo", + "HttpRequestResult", "ProviderLookupError", "ProviderTypeMapping", "WorkflowFactory", diff --git a/python/packages/declarative/agent_framework_declarative/_workflows/__init__.py b/python/packages/declarative/agent_framework_declarative/_workflows/__init__.py index 2968f2b3f92..a1d098a014e 100644 --- a/python/packages/declarative/agent_framework_declarative/_workflows/__init__.py +++ b/python/packages/declarative/agent_framework_declarative/_workflows/__init__.py @@ -67,6 +67,10 @@ RequestExternalInputExecutor, WaitForInputExecutor, ) +from ._executors_http import ( + HTTP_ACTION_EXECUTORS, + HttpRequestActionExecutor, +) from ._executors_tools import ( FUNCTION_TOOL_REGISTRY_KEY, TOOL_ACTION_EXECUTORS, @@ -78,7 +82,13 @@ ToolApprovalState, ToolInvocationResult, ) -from ._factory import DeclarativeWorkflowError, WorkflowFactory +from ._factory import DeclarativeActionError, DeclarativeWorkflowError, WorkflowFactory +from ._http_handler import ( + DefaultHttpRequestHandler, + HttpRequestHandler, + HttpRequestInfo, + HttpRequestResult, +) from ._state import WorkflowState __all__ = [ @@ -90,6 +100,7 @@ "DECLARATIVE_STATE_KEY", "EXTERNAL_INPUT_EXECUTORS", "FUNCTION_TOOL_REGISTRY_KEY", + "HTTP_ACTION_EXECUTORS", "TOOL_ACTION_EXECUTORS", "TOOL_APPROVAL_STATE_KEY", "TOOL_REGISTRY_KEY", @@ -106,12 +117,14 @@ "ContinueLoopExecutor", "ConversationData", "CreateConversationExecutor", + "DeclarativeActionError", "DeclarativeActionExecutor", "DeclarativeMessage", "DeclarativeStateData", "DeclarativeWorkflowBuilder", "DeclarativeWorkflowError", "DeclarativeWorkflowState", + "DefaultHttpRequestHandler", "EmitEventExecutor", "EndConversationExecutor", "EndWorkflowExecutor", @@ -120,6 +133,10 @@ "ExternalLoopState", "ForeachInitExecutor", "ForeachNextExecutor", + "HttpRequestActionExecutor", + "HttpRequestHandler", + "HttpRequestInfo", + "HttpRequestResult", "InvokeAzureAgentExecutor", "InvokeFunctionToolExecutor", "JoinExecutor", diff --git a/python/packages/declarative/agent_framework_declarative/_workflows/_declarative_builder.py b/python/packages/declarative/agent_framework_declarative/_workflows/_declarative_builder.py index 6843c5bd92d..fb5dcb88f87 100644 --- a/python/packages/declarative/agent_framework_declarative/_workflows/_declarative_builder.py +++ b/python/packages/declarative/agent_framework_declarative/_workflows/_declarative_builder.py @@ -26,6 +26,7 @@ DeclarativeActionExecutor, LoopIterationResult, ) +from ._errors import DeclarativeWorkflowError from ._executors_agents import AGENT_ACTION_EXECUTORS, InvokeAzureAgentExecutor from ._executors_basic import BASIC_ACTION_EXECUTORS from ._executors_control_flow import ( @@ -39,7 +40,9 @@ SwitchEvaluatorExecutor, ) from ._executors_external_input import EXTERNAL_INPUT_EXECUTORS +from ._executors_http import HTTP_ACTION_EXECUTORS, HttpRequestActionExecutor from ._executors_tools import TOOL_ACTION_EXECUTORS, InvokeFunctionToolExecutor +from ._http_handler import HttpRequestHandler logger = logging.getLogger(__name__) @@ -51,6 +54,7 @@ **AGENT_ACTION_EXECUTORS, **EXTERNAL_INPUT_EXECUTORS, **TOOL_ACTION_EXECUTORS, + **HTTP_ACTION_EXECUTORS, } # Action kinds that terminate control flow (no fall-through to successor) @@ -85,6 +89,7 @@ "WaitForHumanInput": ["variable"], "EmitEvent": ["event"], "InvokeFunctionTool": ["functionName"], + "HttpRequestAction": ["url"], } # Alternate field names that satisfy required field requirements @@ -129,6 +134,7 @@ def __init__( checkpoint_storage: Any | None = None, validate: bool = True, max_iterations: int | None = None, + http_request_handler: HttpRequestHandler | None = None, ): """Initialize the builder. @@ -141,6 +147,9 @@ def __init__( validate: Whether to validate the workflow definition before building (default: True) max_iterations: Maximum runner supersteps. Falls back to the YAML ``maxTurns`` field, then to the core default (100). + http_request_handler: Handler used to dispatch HttpRequestAction requests. + Must be supplied when the workflow contains any HttpRequestAction; + otherwise build raises ``DeclarativeWorkflowError``. """ self._yaml_def = yaml_definition self._workflow_id = workflow_id or yaml_definition.get("name", "declarative_workflow") @@ -152,6 +161,7 @@ def __init__( self._pending_gotos: list[tuple[Any, str]] = [] # (goto_executor, target_id) self._validate = validate self._seen_explicit_ids: set[str] = set() # Track explicit IDs for duplicate detection + self._http_request_handler = http_request_handler # Resolve max_iterations: explicit arg > YAML maxTurns > core default resolved = max_iterations if max_iterations is not None else yaml_definition.get("maxTurns") if resolved is not None and (not isinstance(resolved, int) or resolved <= 0): @@ -458,6 +468,19 @@ def _create_executor_for_action( executor = InvokeAzureAgentExecutor(action_def, id=action_id, agents=self._agents) elif kind == "InvokeFunctionTool": executor = InvokeFunctionToolExecutor(action_def, id=action_id, tools=self._tools) + elif kind == "HttpRequestAction": + if self._http_request_handler is None: + raise DeclarativeWorkflowError( + f"Workflow defines HttpRequestAction '{action_id}' but no " + "http_request_handler was supplied to WorkflowFactory. Pass " + "http_request_handler=DefaultHttpRequestHandler() (or a custom " + "implementation) to enable HTTP requests." + ) + executor = HttpRequestActionExecutor( + action_def, + id=action_id, + http_request_handler=self._http_request_handler, + ) else: executor = executor_class(action_def, id=action_id) self._executors[action_id] = executor diff --git a/python/packages/declarative/agent_framework_declarative/_workflows/_errors.py b/python/packages/declarative/agent_framework_declarative/_workflows/_errors.py new file mode 100644 index 00000000000..452203268f4 --- /dev/null +++ b/python/packages/declarative/agent_framework_declarative/_workflows/_errors.py @@ -0,0 +1,42 @@ +# Copyright (c) Microsoft. All rights reserved. + +"""Error types for declarative workflow executor modules. + +This module exists so that executor modules and the builder (e.g. +``_executors_http``, ``_declarative_builder``) can raise declarative-specific +exceptions without importing from ``_factory``. ``_factory`` imports +``_declarative_builder`` which imports the executor modules; pulling +``DeclarativeWorkflowError`` from ``_factory`` into an executor or builder +module would therefore introduce a circular import. + +``_factory`` re-exports :class:`DeclarativeWorkflowError` (and +:class:`DeclarativeActionError`) for convenience so existing import paths keep +working. +""" + +from __future__ import annotations + +from agent_framework.exceptions import WorkflowException + + +class DeclarativeWorkflowError(WorkflowException): + """Raised for build-time / factory-level declarative workflow errors. + + Used for YAML parsing/validation issues, missing configuration (e.g. an + HTTP request handler not supplied for a workflow that contains an + ``HttpRequestAction``), and other errors detected before workflow + execution begins. + """ + + pass + + +class DeclarativeActionError(WorkflowException): + """Raised when a declarative action fails at run time. + + Used by executor modules for runtime failures (e.g. transport errors, + non-2xx responses from :class:`HttpRequestActionExecutor`). Build-time and + factory-level errors continue to use :class:`DeclarativeWorkflowError`. + """ + + pass diff --git a/python/packages/declarative/agent_framework_declarative/_workflows/_executors_http.py b/python/packages/declarative/agent_framework_declarative/_workflows/_executors_http.py new file mode 100644 index 00000000000..e041a40c3b8 --- /dev/null +++ b/python/packages/declarative/agent_framework_declarative/_workflows/_executors_http.py @@ -0,0 +1,432 @@ +# Copyright (c) Microsoft. All rights reserved. + +"""Executor for the ``HttpRequestAction`` declarative action. + +Mirrors the .NET ``HttpRequestExecutor``: dispatches an HTTP request through the +configured :class:`HttpRequestHandler`, parses the response body, and assigns +the parsed body and response headers to the declared state paths. + +Security note: response bodies can echo secrets and may be very large. Diagnostic +messages produced for non-2xx responses truncate the body to 256 characters and +collapse CR/LF/TAB to spaces (parity with .NET ``FormatBodyForDiagnostics``). +""" + +from __future__ import annotations + +import json +import logging +from collections.abc import Mapping +from typing import Any + +import httpx +from agent_framework import ( + Message, + WorkflowContext, + handler, +) + +from ._declarative_base import ( + ActionComplete, + DeclarativeActionExecutor, + DeclarativeWorkflowState, +) +from ._errors import DeclarativeActionError +from ._http_handler import HttpRequestHandler, HttpRequestInfo, HttpRequestResult + +__all__ = [ + "HTTP_ACTION_EXECUTORS", + "HttpRequestActionExecutor", +] + +logger = logging.getLogger(__name__) + +_MAX_BODY_DIAGNOSTIC_LENGTH = 256 +_BODY_TRUNCATION_SUFFIX = " \u2026 [truncated]" + + +# Body discriminator aliases. Long forms match the .NET object-model type +# names so YAML produced by .NET round-trips. Short forms are the .NET YAML +# convention used in test fixtures. +_BODY_KIND_JSON = {"json", "JsonRequestContent"} +_BODY_KIND_RAW = {"raw", "RawRequestContent"} +_BODY_KIND_NONE = {"none", "NoRequestContent"} + + +def _get_path(action_def: Mapping[str, Any], key: str) -> str | None: + """Extract a state path from ``response``/``responseHeaders`` field. + + Supports two YAML shapes (matches .NET serialization round-trips): + + - ``response: Local.MyVar`` (plain string). + - ``response: { path: Local.MyVar }`` (object form). + """ + value = action_def.get(key) + if isinstance(value, str): + return value or None + if isinstance(value, Mapping): + path = value.get("path") + return path if isinstance(path, str) and path else None + return None + + +def _format_body_for_diagnostics(body: str | None) -> str: + """Truncate and sanitise a response body for inclusion in error messages. + + Mirrors the .NET ``FormatBodyForDiagnostics`` helper: + + - Empty/None -> empty string. + - Replaces CR/LF/TAB with spaces. + - Truncates to 256 chars with a unicode-ellipsis ``[truncated]`` suffix. + """ + if not body: + return "" + + truncated = len(body) > _MAX_BODY_DIAGNOSTIC_LENGTH + head = body[:_MAX_BODY_DIAGNOSTIC_LENGTH] if truncated else body + sanitized = head.replace("\r", " ").replace("\n", " ").replace("\t", " ") + return sanitized + _BODY_TRUNCATION_SUFFIX if truncated else sanitized + + +def _parse_response_body(body: str | None) -> Any: + """Parse an HTTP response body the same way the .NET executor does. + + JSON-first: if the body parses as JSON, the parsed value is returned. Other + bodies are returned as the raw string. Empty/None bodies return ``None``. + """ + if body is None or body == "": + return None + try: + return json.loads(body) + except (ValueError, json.JSONDecodeError): + return body + + +def _format_query_value(value: Any) -> str | None: + """Format a query-parameter value for URL inclusion. + + Mirrors .NET ``FormatQueryValue``: ``None`` is dropped, ``bool`` becomes + lower-case ``"true"``/``"false"``, numerics use invariant ``str()``, and + other values fall through to ``str()``. + """ + if value is None: + return None + if isinstance(value, bool): + return "true" if value else "false" + if isinstance(value, str): + return value + return str(value) + + +def _get_messages_path(state: DeclarativeWorkflowState, conversation_id_expr: str | None) -> str | None: + """Mirror ``InvokeAzureAgentExecutor._get_conversation_messages_path`` semantics. + + Returns a single state path (no double-mirroring): ``Conversation.messages`` + when no conversation id is provided, otherwise + ``System.conversations.{evaluated_id}.messages``. Returns ``None`` if the + expression evaluates to an empty string (matches .NET ``GetConversationId`` + behaviour where empty becomes ``null`` and the response is not appended). + """ + if not conversation_id_expr: + return None + evaluated = state.eval_if_expression(conversation_id_expr) + if evaluated is None or (isinstance(evaluated, str) and not evaluated): + return None + return f"System.conversations.{evaluated}.messages" + + +class HttpRequestActionExecutor(DeclarativeActionExecutor): + """Executor for the ``HttpRequestAction`` declarative action. + + Dispatches through the supplied :class:`HttpRequestHandler` and: + + - Parses the response body (JSON-first, raw string fall-back). + - Assigns the parsed body to ``response`` path (if configured). + - Folds multi-value response headers (comma-joined) and assigns them to + ``responseHeaders`` path (if configured). + - On 2xx with non-empty body and a configured ``conversationId``, appends + an Assistant :class:`agent_framework.Message` to + ``System.conversations.{id}.messages``. + - On non-2xx, still publishes ``responseHeaders`` (diagnostic) and raises + :class:`DeclarativeActionError` with a status-coded message containing a + truncated/sanitised body preview. + + Transport errors (``httpx.TimeoutException``, ``TimeoutError``, + ``httpx.HTTPError``) become :class:`DeclarativeActionError`. ``CancelledError`` + is intentionally NOT caught so that workflow cancellation propagates. + """ + + def __init__( + self, + action_def: dict[str, Any], + *, + id: str | None = None, + http_request_handler: HttpRequestHandler, + ) -> None: + """Create an HTTP request action executor. + + Args: + action_def: Parsed ``HttpRequestAction`` YAML dict. + id: Optional executor id (defaults to action id or generated). + http_request_handler: Handler used to dispatch HTTP requests. + Required: the builder enforces presence at workflow-build time. + """ + super().__init__(action_def, id=id) + self._http_request_handler = http_request_handler + + @handler + async def handle_action( + self, + trigger: Any, + ctx: WorkflowContext[ActionComplete], + ) -> None: + """Execute the HTTP request action.""" + state = await self._ensure_state_initialized(ctx, trigger) + + method = self._get_method(state) + url = self._get_url(state) + headers = self._get_headers(state) + query_parameters = self._get_query_parameters(state) + body, body_content_type = self._get_body(state) + timeout_ms = self._get_timeout_ms(state) + conversation_id_expr = self._action_def.get("conversationId") + connection_name = self._get_connection_name(state) + + info = HttpRequestInfo( + method=method, + url=url, + headers=headers or {}, + query_parameters=query_parameters or {}, + body=body, + body_content_type=body_content_type, + timeout_ms=timeout_ms, + connection_name=connection_name, + ) + + try: + result = await self._http_request_handler.send(info) + except (httpx.TimeoutException, TimeoutError) as exc: + raise DeclarativeActionError(f"HTTP request to '{url}' timed out.") from exc + except DeclarativeActionError: + raise + except httpx.HTTPError as exc: + raise DeclarativeActionError( + f"HTTP request to '{url}' failed: {type(exc).__name__}" + ) from exc + except Exception as exc: + # Custom HttpRequestHandler implementations may raise arbitrary + # exception types. Wrap them in DeclarativeActionError so workflow + # error handling stays uniform regardless of transport. Note that + # ``asyncio.CancelledError`` is a ``BaseException`` (not + # ``Exception``) and so still propagates unmodified, preserving + # workflow-cancellation semantics. + raise DeclarativeActionError( + f"HTTP request to '{url}' failed: {type(exc).__name__}" + ) from exc + + if result.is_success_status_code: + self._assign_response(state, result) + self._assign_response_headers(state, result) + self._append_response_to_conversation(state, conversation_id_expr, result.body) + await ctx.send_message(ActionComplete()) + return + + # Non-success path: still publish headers diagnostically, then raise. + self._assign_response_headers(state, result) + body_preview = _format_body_for_diagnostics(result.body) + if body_preview: + message = ( + f"HTTP request to '{url}' failed with status code {result.status_code}. " + f"Body: '{body_preview}'" + ) + else: + message = f"HTTP request to '{url}' failed with status code {result.status_code}." + raise DeclarativeActionError(message) + + # ----- Field resolution ---------------------------------------------------- + + def _get_method(self, state: DeclarativeWorkflowState) -> str: + method = self._action_def.get("method") + evaluated = state.eval_if_expression(method) if method is not None else None + if not evaluated: + return "GET" + return str(evaluated).upper() + + def _get_url(self, state: DeclarativeWorkflowState) -> str: + raw = self._action_def.get("url") + if raw is None: + raise ValueError("HttpRequestAction requires a 'url' field.") + evaluated = state.eval_if_expression(raw) + if not isinstance(evaluated, str) or not evaluated: + raise ValueError("HttpRequestAction 'url' evaluated to an empty value.") + return evaluated + + def _get_headers(self, state: DeclarativeWorkflowState) -> dict[str, str] | None: + raw_headers = self._action_def.get("headers") + if not isinstance(raw_headers, Mapping) or not raw_headers: + return None + result: dict[str, str] = {} + for key, value in raw_headers.items(): + if not isinstance(key, str) or not key: + continue + evaluated = state.eval_if_expression(value) + if evaluated is None: + continue + text = str(evaluated) + if not text: + continue + result[key] = text + return result or None + + def _get_query_parameters(self, state: DeclarativeWorkflowState) -> dict[str, str] | None: + raw_params = self._action_def.get("queryParameters") + if not isinstance(raw_params, Mapping) or not raw_params: + return None + result: dict[str, str] = {} + for key, value in raw_params.items(): + if not isinstance(key, str) or not key or value is None: + continue + evaluated = state.eval_if_expression(value) + formatted = _format_query_value(evaluated) + if formatted is not None: + result[key] = formatted + return result or None + + def _get_body(self, state: DeclarativeWorkflowState) -> tuple[str | None, str | None]: + raw_body = self._action_def.get("body") + if raw_body is None: + return None, None + if not isinstance(raw_body, Mapping): + raise ValueError( + "HttpRequestAction 'body' must be a mapping with a 'kind' field " + "(json, raw) or omitted entirely." + ) + + kind_value = raw_body.get("kind") or raw_body.get("$kind") + if kind_value is None: + raise ValueError( + "HttpRequestAction 'body' is missing 'kind'. Use 'json', 'raw', " + "or omit 'body' for no request body." + ) + if not isinstance(kind_value, str): + raise ValueError( + f"HttpRequestAction 'body.kind' must be a string, got {type(kind_value).__name__}." + ) + + if kind_value in _BODY_KIND_NONE: + return None, None + + if kind_value in _BODY_KIND_JSON: + content_expr = raw_body.get("content") + if content_expr is None: + return None, None + evaluated = state.eval_if_expression(content_expr) + try: + body_text = json.dumps(evaluated, default=str) + except (TypeError, ValueError) as exc: + raise ValueError( + f"HttpRequestAction 'body.content' could not be serialised as JSON: {exc}" + ) from exc + return body_text, "application/json" + + if kind_value in _BODY_KIND_RAW: + content_expr = raw_body.get("content") + content_type_expr = raw_body.get("contentType") + content: str | None = None + if content_expr is not None: + evaluated = state.eval_if_expression(content_expr) + content = None if evaluated is None else str(evaluated) + content_type: str | None = None + if content_type_expr is not None: + ct_eval = state.eval_if_expression(content_type_expr) + ct_text = None if ct_eval is None else str(ct_eval) + content_type = ct_text or None + # Match .NET RawRequestContent semantics: when a raw body is sent + # without an explicit content type, default to text/plain so the + # request is interpretable by servers. + if content is not None and not content_type: + content_type = "text/plain" + return content, content_type + + raise ValueError( + f"HttpRequestAction 'body.kind' has unsupported value '{kind_value}'. " + "Expected one of: json, raw, JsonRequestContent, RawRequestContent, " + "NoRequestContent." + ) + + def _get_timeout_ms(self, state: DeclarativeWorkflowState) -> int | None: + raw = self._action_def.get("requestTimeoutInMilliseconds") + if raw is None: + return None + evaluated = state.eval_if_expression(raw) + if evaluated is None: + return None + try: + value = int(evaluated) + except (TypeError, ValueError): + logger.debug( + "HttpRequestAction: ignoring non-numeric requestTimeoutInMilliseconds=%r", + evaluated, + ) + return None + return value if value > 0 else None + + def _get_connection_name(self, state: DeclarativeWorkflowState) -> str | None: + connection = self._action_def.get("connection") + if not isinstance(connection, Mapping): + return None + name_expr = connection.get("name") + if name_expr is None: + return None + evaluated = state.eval_if_expression(name_expr) + if evaluated is None: + return None + text = str(evaluated) + return text or None + + # ----- Result handling ----------------------------------------------------- + + def _assign_response(self, state: DeclarativeWorkflowState, result: HttpRequestResult) -> None: + path = _get_path(self._action_def, "response") + if path is None: + return + state.set(path, _parse_response_body(result.body)) + + def _assign_response_headers(self, state: DeclarativeWorkflowState, result: HttpRequestResult) -> None: + path = _get_path(self._action_def, "responseHeaders") + if path is None: + return + if not result.headers: + state.set(path, None) + return + # Fold multi-value headers with commas (standard HTTP folding) only at + # assignment time. The raw multi-value dict on HttpRequestResult.headers + # is left untouched so callers/tests can inspect duplicates. + flattened: dict[str, str] = {} + for key, values in result.headers.items(): + flattened[key] = ",".join(values) + state.set(path, flattened) + + def _append_response_to_conversation( + self, + state: DeclarativeWorkflowState, + conversation_id_expr: str | None, + body: str, + ) -> None: + if not body: + return + messages_path = _get_messages_path(state, conversation_id_expr) + if messages_path is None: + return + # Ensure the conversation entry exists so downstream agents can reference it. + evaluated_id = messages_path.split(".", 2)[2].rsplit(".", 1)[0] + conversations: dict[str, Any] = state.get("System.conversations") or {} + if evaluated_id not in conversations: + conversations[evaluated_id] = {"id": evaluated_id, "messages": []} + state.set("System.conversations", conversations) + message = Message(role="assistant", contents=[body]) + state.append(messages_path, message) + + +HTTP_ACTION_EXECUTORS: dict[str, type[DeclarativeActionExecutor]] = { + "HttpRequestAction": HttpRequestActionExecutor, +} diff --git a/python/packages/declarative/agent_framework_declarative/_workflows/_factory.py b/python/packages/declarative/agent_framework_declarative/_workflows/_factory.py index 9ba4cb84de6..36f088b05fd 100644 --- a/python/packages/declarative/agent_framework_declarative/_workflows/_factory.py +++ b/python/packages/declarative/agent_framework_declarative/_workflows/_factory.py @@ -24,18 +24,18 @@ SupportsAgentRun, Workflow, ) -from agent_framework.exceptions import WorkflowException from .._loader import AgentFactory from ._declarative_builder import DeclarativeWorkflowBuilder +from ._errors import DeclarativeActionError, DeclarativeWorkflowError +from ._http_handler import HttpRequestHandler logger = logging.getLogger("agent_framework.declarative") -class DeclarativeWorkflowError(WorkflowException): - """Exception raised for errors in declarative workflow processing.""" - - pass +# Re-export DeclarativeWorkflowError (now defined in _errors) so existing +# import paths (`from .._factory import DeclarativeWorkflowError`) keep working. +__all__ = ["DeclarativeActionError", "DeclarativeWorkflowError", "WorkflowFactory"] class WorkflowFactory: @@ -92,6 +92,7 @@ def __init__( env_file: str | None = None, checkpoint_storage: CheckpointStorage | None = None, max_iterations: int | None = None, + http_request_handler: HttpRequestHandler | None = None, ) -> None: """Initialize the workflow factory. @@ -105,6 +106,12 @@ def __init__( max_iterations: Optional maximum runner supersteps. Overrides the YAML ``maxTurns`` field and the core default (100). Workflows with ``GotoAction`` loops (e.g. DeepResearch) typically need a higher value. + http_request_handler: Optional handler used to dispatch HTTP requests for + ``HttpRequestAction``. Required if the workflow contains any + ``HttpRequestAction``; build will fail with :class:`DeclarativeWorkflowError` + otherwise. Use :class:`agent_framework.declarative.DefaultHttpRequestHandler` + for a no-policy ``httpx``-based default, or supply your own implementation + to enforce SSRF guards, allowlisting, or auth resolution. Examples: .. code-block:: python @@ -144,6 +151,7 @@ def __init__( self._tools: dict[str, Any] = {} # Tool registry for InvokeFunctionTool actions self._checkpoint_storage = checkpoint_storage self._max_iterations = max_iterations + self._http_request_handler = http_request_handler def create_workflow_from_yaml_path( self, @@ -387,6 +395,7 @@ def _create_workflow( tools=self._tools, checkpoint_storage=self._checkpoint_storage, max_iterations=self._max_iterations, + http_request_handler=self._http_request_handler, ) workflow = graph_builder.build() except ValueError as e: diff --git a/python/packages/declarative/agent_framework_declarative/_workflows/_http_handler.py b/python/packages/declarative/agent_framework_declarative/_workflows/_http_handler.py new file mode 100644 index 00000000000..75231fe91d2 --- /dev/null +++ b/python/packages/declarative/agent_framework_declarative/_workflows/_http_handler.py @@ -0,0 +1,229 @@ +# Copyright (c) Microsoft. All rights reserved. + +"""HTTP request handler abstraction for declarative workflows. + +Mirrors the .NET ``IHttpRequestHandler`` / ``DefaultHttpRequestHandler`` pair from +``Microsoft.Agents.AI.Workflows.Declarative``. Provides: + +- :class:`HttpRequestInfo` — request input data passed from the executor. +- :class:`HttpRequestResult` — response data returned to the executor. +- :class:`HttpRequestHandler` — :class:`typing.Protocol` callers implement to plug + in custom transports (e.g. with allowlisting, mTLS, retries, etc.). +- :class:`DefaultHttpRequestHandler` — production-grade default backed by + ``httpx.AsyncClient``. + +Security note: :class:`DefaultHttpRequestHandler` performs **no** URL filtering +or SSRF protection. Production deployments should supply a custom handler that +enforces an allowlist or DNS-rebinding-resistant policy. This split mirrors the +.NET design. +""" + +from __future__ import annotations + +import asyncio +from collections.abc import Awaitable, Callable, Mapping +from dataclasses import dataclass, field +from typing import Any, Protocol, runtime_checkable + +import httpx + +__all__ = [ + "DefaultHttpRequestHandler", + "HttpRequestHandler", + "HttpRequestInfo", + "HttpRequestResult", +] + + +@dataclass +class HttpRequestInfo: + """Description of an HTTP request to be dispatched by a :class:`HttpRequestHandler`. + + Mirrors the .NET ``HttpRequestInfo`` record. Field semantics: + + - ``method``: HTTP method (``GET``, ``POST``, etc.). Already upper-cased by the executor. + - ``url``: Absolute URL. Already evaluated from the YAML expression. + - ``headers``: Single-value header map (case-insensitive keys per HTTP semantics + but stored as authored). Empty values are skipped by the executor. + - ``query_parameters``: String key/value pairs appended to the URL. + - ``body``: Request body bytes/text, or ``None`` for no body. + - ``body_content_type``: Content type to send (e.g. ``application/json``). + Ignored when ``body`` is ``None``. + - ``timeout_ms``: Per-request timeout in milliseconds. ``None`` => use the + handler's default. + - ``connection_name``: Optional Foundry connection name for handlers that + resolve auth/credentials by connection. + """ + + method: str + url: str + headers: dict[str, str] = field(default_factory=dict) + query_parameters: dict[str, str] = field(default_factory=dict) + body: str | None = None + body_content_type: str | None = None + timeout_ms: int | None = None + connection_name: str | None = None + + +@dataclass +class HttpRequestResult: + """Response returned by a :class:`HttpRequestHandler`. + + Mirrors the .NET ``HttpRequestResult`` record. ``headers`` preserves + multi-value response headers (e.g. multiple ``Set-Cookie`` headers) as a + ``dict[str, list[str]]``. The executor folds duplicates into a single + comma-joined string only at the point it assigns ``responseHeaders`` to + workflow state. + """ + + status_code: int + is_success_status_code: bool + body: str + headers: dict[str, list[str]] = field(default_factory=dict) + + +@runtime_checkable +class HttpRequestHandler(Protocol): + """Protocol for HTTP request handlers used by ``HttpRequestAction``. + + Implementations must be safe to call concurrently from multiple workflow + runs. Implementations are responsible for any URL allowlisting, SSRF + guards, retry policies, auth resolution, and other policies that the + workflow author wants applied. + """ + + async def send(self, info: HttpRequestInfo) -> HttpRequestResult: + """Dispatch ``info`` and return the response result. + + Args: + info: Description of the request to send. + + Returns: + The response. Implementations should NOT raise on non-2xx status + codes; instead, set ``is_success_status_code`` accordingly. They + SHOULD raise on transport-level failures (connection refused, + DNS errors, timeouts). + """ + ... + + +ClientProvider = Callable[[HttpRequestInfo], Awaitable["httpx.AsyncClient | None"]] + + +class DefaultHttpRequestHandler: + """Default :class:`HttpRequestHandler` backed by :class:`httpx.AsyncClient`. + + Construction modes: + + 1. ``DefaultHttpRequestHandler()`` — owns an internal client created lazily + on first ``send()``. Closed by :meth:`aclose`. + 2. ``DefaultHttpRequestHandler(client=existing)`` — caller-owned client. + Not closed by :meth:`aclose`. + 3. ``DefaultHttpRequestHandler(client_provider=cb)`` — per-request client + lookup (parity with .NET's ``httpClientProvider`` callback). The + provider may return ``None`` to fall back to the owned/default client. + + .. warning:: + + This handler performs **no** URL filtering or SSRF protection. Wrap or + replace it with a custom handler in production. + """ + + def __init__( + self, + *, + client: "httpx.AsyncClient | None" = None, + client_provider: ClientProvider | None = None, + ) -> None: + self._owned_client: httpx.AsyncClient | None = None + self._caller_client = client + self._client_provider = client_provider + # Guards lazy creation of ``_owned_client`` against concurrent first + # ``send()`` calls leaking duplicate clients. + self._owned_client_lock = asyncio.Lock() + + async def send(self, info: HttpRequestInfo) -> HttpRequestResult: + """Dispatch the request and return the parsed result.""" + if not info.url: + raise ValueError("HttpRequestInfo.url must be a non-empty string.") + if not info.method: + raise ValueError("HttpRequestInfo.method must be a non-empty string.") + + client = await self._resolve_client(info) + + timeout: httpx.Timeout | object + if info.timeout_ms is not None and info.timeout_ms > 0: + timeout = httpx.Timeout(info.timeout_ms / 1000.0) + else: + timeout = httpx.USE_CLIENT_DEFAULT + + headers = dict(info.headers) + content: bytes | str | None = None + if info.body is not None: + content = info.body + if not _has_header(headers, "content-type"): + # Match .NET DefaultHttpRequestHandler: when a body is sent + # without an explicit content type, default to ``text/plain`` + # so the request is interpretable by servers and direct + # callers (not just the YAML executor) get sensible defaults. + headers["Content-Type"] = info.body_content_type or "text/plain" + + params: Mapping[str, str] | None = info.query_parameters or None + + response = await client.request( + method=info.method, + url=info.url, + params=params, + headers=headers or None, + content=content, + timeout=timeout, # type: ignore[arg-type] + ) + + # Preserve multi-value headers (e.g. multiple Set-Cookie) as list[str]. + result_headers: dict[str, list[str]] = {} + for key, value in response.headers.multi_items(): + result_headers.setdefault(key, []).append(value) + + body_text = response.text + + return HttpRequestResult( + status_code=response.status_code, + is_success_status_code=200 <= response.status_code < 300, + body=body_text, + headers=result_headers, + ) + + async def aclose(self) -> None: + """Release the owned client, if any. Caller-owned clients are NOT closed.""" + if self._owned_client is not None: + await self._owned_client.aclose() + self._owned_client = None + + async def _resolve_client(self, info: HttpRequestInfo) -> httpx.AsyncClient: + """Pick a client for this request: provider → caller → lazily-owned.""" + if self._client_provider is not None: + provided = await self._client_provider(info) + if provided is not None: + return provided + if self._caller_client is not None: + return self._caller_client + if self._owned_client is None: + # Double-checked locking under asyncio.Lock so concurrent first + # callers don't each create a fresh httpx.AsyncClient and orphan + # one of them. + async with self._owned_client_lock: + if self._owned_client is None: + self._owned_client = httpx.AsyncClient() + return self._owned_client + + async def __aenter__(self) -> "DefaultHttpRequestHandler": + return self + + async def __aexit__(self, exc_type: Any, exc: Any, tb: Any) -> None: + await self.aclose() + + +def _has_header(headers: Mapping[str, str], name: str) -> bool: + """Case-insensitive header presence check.""" + needle = name.lower() + return any(key.lower() == needle for key in headers) diff --git a/python/packages/declarative/pyproject.toml b/python/packages/declarative/pyproject.toml index a26981505c9..a03928ac821 100644 --- a/python/packages/declarative/pyproject.toml +++ b/python/packages/declarative/pyproject.toml @@ -23,6 +23,7 @@ classifiers = [ ] dependencies = [ "agent-framework-core>=1.2.2,<2", + "httpx>=0.27,<1", "powerfx>=0.0.32,<0.0.35; python_version < '3.14'", "pyyaml>=6.0,<7.0", ] diff --git a/python/packages/declarative/tests/test_default_http_request_handler.py b/python/packages/declarative/tests/test_default_http_request_handler.py new file mode 100644 index 00000000000..5391b90adfb --- /dev/null +++ b/python/packages/declarative/tests/test_default_http_request_handler.py @@ -0,0 +1,350 @@ +# Copyright (c) Microsoft. All rights reserved. + +"""Tests for ``DefaultHttpRequestHandler``. + +These tests exercise the real handler against ``httpx.MockTransport`` (no real +network) to cover the parts of the handler not exercisable through the executor +stub: query-param URL composition, content-type forwarding, per-request +timeout overrides, multi-value response header preservation, and client +ownership semantics. +""" + +from __future__ import annotations + +import sys + +import httpx +import pytest + +try: + import powerfx # noqa: F401 + + _powerfx_available = True +except (ImportError, RuntimeError): + _powerfx_available = False + +# These tests don't actually need PowerFx, but the rest of the suite gates on +# Python versions and we keep behaviour consistent. +pytestmark = pytest.mark.skipif( + sys.version_info >= (3, 14), + reason="Skipped on Python 3.14+ to keep parity with rest of declarative suite", +) + +from agent_framework_declarative._workflows._http_handler import ( # noqa: E402 + DefaultHttpRequestHandler, + HttpRequestInfo, +) + + +def _make_handler(transport: httpx.MockTransport) -> DefaultHttpRequestHandler: + """Return a handler with a MockTransport-backed caller-owned client.""" + client = httpx.AsyncClient(transport=transport) + return DefaultHttpRequestHandler(client=client) + + +class TestRequestComposition: + @pytest.mark.asyncio + async def test_query_parameters_merged_into_url(self) -> None: + captured: dict[str, httpx.Request] = {} + + def respond(request: httpx.Request) -> httpx.Response: + captured["req"] = request + return httpx.Response(200, text="ok") + + handler = _make_handler(httpx.MockTransport(respond)) + try: + await handler.send( + HttpRequestInfo( + method="GET", + url="https://api.example.test/items", + query_parameters={"q": "alpha", "limit": "5"}, + ) + ) + finally: + await handler.aclose() + + req = captured["req"] + # httpx exposes the merged URL with QS appended + assert req.url.params.get("q") == "alpha" + assert req.url.params.get("limit") == "5" + + @pytest.mark.asyncio + async def test_body_content_type_forwarded(self) -> None: + captured: dict[str, httpx.Request] = {} + + def respond(request: httpx.Request) -> httpx.Response: + captured["req"] = request + return httpx.Response(204) + + handler = _make_handler(httpx.MockTransport(respond)) + try: + await handler.send( + HttpRequestInfo( + method="POST", + url="https://api.example.test/items", + body='{"k":"v"}', + body_content_type="application/json", + ) + ) + finally: + await handler.aclose() + + req = captured["req"] + assert req.headers.get("content-type") == "application/json" + assert req.content == b'{"k":"v"}' + + @pytest.mark.asyncio + async def test_existing_content_type_header_not_overwritten(self) -> None: + captured: dict[str, httpx.Request] = {} + + def respond(request: httpx.Request) -> httpx.Response: + captured["req"] = request + return httpx.Response(200, text="ok") + + handler = _make_handler(httpx.MockTransport(respond)) + try: + await handler.send( + HttpRequestInfo( + method="POST", + url="https://api.example.test/items", + headers={"Content-Type": "application/xml"}, # caller wins + body="", + body_content_type="application/json", + ) + ) + finally: + await handler.aclose() + + req = captured["req"] + assert req.headers.get("content-type") == "application/xml" + + @pytest.mark.asyncio + async def test_body_without_content_type_defaults_to_text_plain(self) -> None: + """Match .NET DefaultHttpRequestHandler: body without explicit content type → ``text/plain``.""" + captured: dict[str, httpx.Request] = {} + + def respond(request: httpx.Request) -> httpx.Response: + captured["req"] = request + return httpx.Response(204) + + handler = _make_handler(httpx.MockTransport(respond)) + try: + await handler.send( + HttpRequestInfo( + method="POST", + url="https://api.example.test/items", + body="hello", + # No body_content_type and no Content-Type header. + ) + ) + finally: + await handler.aclose() + + req = captured["req"] + assert req.headers.get("content-type") == "text/plain" + assert req.content == b"hello" + + +class TestTimeout: + @pytest.mark.asyncio + async def test_per_request_timeout_surfaces_as_timeout_exception(self) -> None: + def respond(request: httpx.Request) -> httpx.Response: + raise httpx.TimeoutException("simulated timeout", request=request) + + handler = _make_handler(httpx.MockTransport(respond)) + try: + with pytest.raises(httpx.TimeoutException): + await handler.send( + HttpRequestInfo( + method="GET", + url="https://api.example.test/slow", + timeout_ms=50, + ) + ) + finally: + await handler.aclose() + + +class TestResponseHeaders: + @pytest.mark.asyncio + async def test_multi_value_headers_preserved(self) -> None: + def respond(request: httpx.Request) -> httpx.Response: + return httpx.Response( + 200, + text="ok", + headers=[ + ("Content-Type", "application/json"), + ("Set-Cookie", "a=1"), + ("Set-Cookie", "b=2"), + ], + ) + + handler = _make_handler(httpx.MockTransport(respond)) + try: + result = await handler.send( + HttpRequestInfo(method="GET", url="https://api.example.test/x") + ) + finally: + await handler.aclose() + + assert result.is_success_status_code + # The handler keeps multi-value headers as list[str]. + assert result.headers.get("set-cookie") == ["a=1", "b=2"] + assert result.headers.get("content-type") == ["application/json"] + + +class TestClientOwnership: + @pytest.mark.asyncio + async def test_owned_client_is_closed_on_aclose(self) -> None: + handler = DefaultHttpRequestHandler() + # Inject a MockTransport-backed client into the owned slot and verify + # aclose() releases it. Avoids real network access. + owned = httpx.AsyncClient( + transport=httpx.MockTransport(lambda r: httpx.Response(200, text="ok")) + ) + handler._owned_client = owned + assert not owned.is_closed + await handler.aclose() + assert owned.is_closed + + @pytest.mark.asyncio + async def test_caller_owned_client_is_not_closed(self) -> None: + client = httpx.AsyncClient( + transport=httpx.MockTransport(lambda r: httpx.Response(200, text="ok")) + ) + handler = DefaultHttpRequestHandler(client=client) + await handler.send( + HttpRequestInfo(method="GET", url="https://api.example.test/x") + ) + await handler.aclose() + assert not client.is_closed + await client.aclose() # cleanup + + @pytest.mark.asyncio + async def test_concurrent_first_send_creates_single_owned_client(self) -> None: + """Concurrent first-send calls must not race-leak duplicate clients. + + Without the lock, two concurrent calls on a fresh handler would each + observe ``_owned_client is None`` and create their own + ``httpx.AsyncClient``, orphaning one. Verify that lazy initialization + is serialized: all concurrent sends end up using the same client and + ``aclose()`` cleanly closes it. + """ + import asyncio + + # Patch httpx.AsyncClient to count constructions, but only when called + # from inside _resolve_client (no transport=) so we don't break the + # MockTransport-backed clients used elsewhere. + original_ctor = httpx.AsyncClient + construction_count = 0 + + def counting_ctor(*args, **kwargs): # type: ignore[no-untyped-def] + nonlocal construction_count + if not args and not kwargs: + construction_count += 1 + return original_ctor( + transport=httpx.MockTransport( + lambda r: httpx.Response(200, text="ok") + ) + ) + return original_ctor(*args, **kwargs) + + import agent_framework_declarative._workflows._http_handler as hh + + hh.httpx.AsyncClient = counting_ctor # type: ignore[assignment] + try: + handler = DefaultHttpRequestHandler() + try: + await asyncio.gather(*[ + handler.send( + HttpRequestInfo(method="GET", url="https://api.example.test/x") + ) + for _ in range(8) + ]) + finally: + await handler.aclose() + finally: + hh.httpx.AsyncClient = original_ctor # type: ignore[assignment] + + assert construction_count == 1, ( + f"Expected exactly 1 owned client to be lazily created but got {construction_count}" + ) + + +class TestClientProvider: + @pytest.mark.asyncio + async def test_client_provider_overrides_default(self) -> None: + captured: dict[str, str] = {} + + def primary(request: httpx.Request) -> httpx.Response: + captured["transport"] = "primary" + return httpx.Response(200, text="primary") + + def provided(request: httpx.Request) -> httpx.Response: + captured["transport"] = "provided" + return httpx.Response(200, text="provided") + + primary_client = httpx.AsyncClient(transport=httpx.MockTransport(primary)) + provided_client = httpx.AsyncClient(transport=httpx.MockTransport(provided)) + + async def provider(info: HttpRequestInfo) -> httpx.AsyncClient: + return provided_client + + handler = DefaultHttpRequestHandler(client=primary_client, client_provider=provider) + try: + result = await handler.send( + HttpRequestInfo(method="GET", url="https://api.example.test/x") + ) + assert result.body == "provided" + assert captured["transport"] == "provided" + finally: + await handler.aclose() + await primary_client.aclose() + await provided_client.aclose() + + @pytest.mark.asyncio + async def test_client_provider_returning_none_falls_back(self) -> None: + captured: dict[str, str] = {} + + def primary(request: httpx.Request) -> httpx.Response: + captured["transport"] = "primary" + return httpx.Response(200, text="primary") + + async def provider(info: HttpRequestInfo) -> httpx.AsyncClient | None: + return None + + primary_client = httpx.AsyncClient(transport=httpx.MockTransport(primary)) + handler = DefaultHttpRequestHandler(client=primary_client, client_provider=provider) + try: + result = await handler.send( + HttpRequestInfo(method="GET", url="https://api.example.test/x") + ) + assert result.body == "primary" + finally: + await handler.aclose() + await primary_client.aclose() + + +class TestValidation: + @pytest.mark.asyncio + async def test_empty_url_raises(self) -> None: + handler = DefaultHttpRequestHandler() + with pytest.raises(ValueError): + await handler.send(HttpRequestInfo(method="GET", url="")) + + @pytest.mark.asyncio + async def test_empty_method_raises(self) -> None: + handler = DefaultHttpRequestHandler() + with pytest.raises(ValueError): + await handler.send(HttpRequestInfo(method="", url="https://x.test/")) + + +class TestAsyncContextManager: + @pytest.mark.asyncio + async def test_context_manager_closes_owned_client(self) -> None: + async with DefaultHttpRequestHandler() as handler: + owned = httpx.AsyncClient( + transport=httpx.MockTransport(lambda r: httpx.Response(200, text="ok")) + ) + handler._owned_client = owned + assert owned.is_closed diff --git a/python/packages/declarative/tests/test_http_request_executor.py b/python/packages/declarative/tests/test_http_request_executor.py new file mode 100644 index 00000000000..effdc5b298c --- /dev/null +++ b/python/packages/declarative/tests/test_http_request_executor.py @@ -0,0 +1,673 @@ +# Copyright (c) Microsoft. All rights reserved. + +"""Tests for HttpRequestActionExecutor. + +These tests use a stub HttpRequestHandler that returns canned HttpRequestResults. +No real network or httpx transports are exercised. See +test_default_http_request_handler.py for tests that exercise the real +DefaultHttpRequestHandler against httpx.MockTransport. +""" + +from __future__ import annotations + +import asyncio +import sys +from typing import Any + +import httpx +import pytest + +try: + import powerfx # noqa: F401 + + _powerfx_available = True +except (ImportError, RuntimeError): + _powerfx_available = False + +pytestmark = pytest.mark.skipif( + not _powerfx_available or sys.version_info >= (3, 14), + reason="PowerFx engine not available (requires dotnet runtime)", +) + +from agent_framework_declarative._workflows import ( # noqa: E402 + DECLARATIVE_STATE_KEY, + DeclarativeActionError, + DeclarativeWorkflowError, + HttpRequestHandler, + HttpRequestInfo, + HttpRequestResult, + WorkflowFactory, +) + + +class StubHandler: + """Test stub that records the last call and returns a canned result.""" + + def __init__( + self, + result: HttpRequestResult | None = None, + *, + raise_exc: BaseException | None = None, + ) -> None: + self.result = result + self.raise_exc = raise_exc + self.last_info: HttpRequestInfo | None = None + self.call_count = 0 + + async def send(self, info: HttpRequestInfo) -> HttpRequestResult: + self.call_count += 1 + self.last_info = info + if self.raise_exc is not None: + raise self.raise_exc + assert self.result is not None + return self.result + + +def _ok(body: str = "", headers: dict[str, list[str]] | None = None) -> HttpRequestResult: + return HttpRequestResult( + status_code=200, + is_success_status_code=True, + body=body, + headers=headers or {}, + ) + + +def _err( + status: int = 500, body: str = "", headers: dict[str, list[str]] | None = None +) -> HttpRequestResult: + return HttpRequestResult( + status_code=status, + is_success_status_code=False, + body=body, + headers=headers or {}, + ) + + +async def _run(yaml_def: dict[str, Any], handler: HttpRequestHandler) -> Any: + """Build & run a workflow, returning final WorkflowState.""" + factory = WorkflowFactory(http_request_handler=handler) + workflow = factory.create_workflow_from_definition(yaml_def) + return await workflow.run({}) + + +def _state(workflow: Any, events: Any) -> dict[str, Any]: + """Read declarative state out of the workflow after run completes.""" + return workflow._state.get(DECLARATIVE_STATE_KEY) or {} + + +# Helper used by parametrised path tests +_TEST_URL = "https://api.example.test/items" + + +def _action( + *, + method: str | None = None, + url: str = _TEST_URL, + headers: dict[str, Any] | None = None, + query_parameters: dict[str, Any] | None = None, + body: dict[str, Any] | None = None, + response: Any = None, + response_headers: Any = None, + conversation_id: str | None = None, + request_timeout_ms: int | None = None, + connection: dict[str, Any] | None = None, +) -> dict[str, Any]: + action: dict[str, Any] = { + "kind": "HttpRequestAction", + "id": "http_action", + "url": url, + } + if method is not None: + action["method"] = method + if headers is not None: + action["headers"] = headers + if query_parameters is not None: + action["queryParameters"] = query_parameters + if body is not None: + action["body"] = body + if response is not None: + action["response"] = response + if response_headers is not None: + action["responseHeaders"] = response_headers + if conversation_id is not None: + action["conversationId"] = conversation_id + if request_timeout_ms is not None: + action["requestTimeoutInMilliseconds"] = request_timeout_ms + if connection is not None: + action["connection"] = connection + return action + + +def _yaml(action: dict[str, Any]) -> dict[str, Any]: + return {"name": "http_test", "actions": [action]} + + +# ---------- Success path: response parsing ---------------------------------- + + +class TestSuccessPath: + @pytest.mark.asyncio + async def test_get_parses_json_object(self) -> None: + handler = StubHandler(_ok('{"key":"value","number":42}')) + factory = WorkflowFactory(http_request_handler=handler) + workflow = factory.create_workflow_from_definition( + _yaml(_action(method="GET", response="Local.Result")) + ) + await workflow.run({}) + + decl = workflow._state.get(DECLARATIVE_STATE_KEY) + assert decl["Local"]["Result"] == {"key": "value", "number": 42} + assert handler.last_info is not None + assert handler.last_info.method == "GET" + assert handler.last_info.url == _TEST_URL + + @pytest.mark.asyncio + async def test_get_parses_plain_string(self) -> None: + handler = StubHandler(_ok("not-json content")) + factory = WorkflowFactory(http_request_handler=handler) + workflow = factory.create_workflow_from_definition( + _yaml(_action(response="Local.Result")) + ) + await workflow.run({}) + + decl = workflow._state.get(DECLARATIVE_STATE_KEY) + assert decl["Local"]["Result"] == "not-json content" + + @pytest.mark.asyncio + async def test_get_empty_body_yields_none(self) -> None: + handler = StubHandler(_ok("")) + factory = WorkflowFactory(http_request_handler=handler) + workflow = factory.create_workflow_from_definition( + _yaml(_action(response="Local.Result")) + ) + await workflow.run({}) + + decl = workflow._state.get(DECLARATIVE_STATE_KEY) + assert decl["Local"]["Result"] is None + + @pytest.mark.asyncio + async def test_response_object_form_path(self) -> None: + handler = StubHandler(_ok('{"x":1}')) + factory = WorkflowFactory(http_request_handler=handler) + workflow = factory.create_workflow_from_definition( + _yaml(_action(response={"path": "Local.Result"})) + ) + await workflow.run({}) + + decl = workflow._state.get(DECLARATIVE_STATE_KEY) + assert decl["Local"]["Result"] == {"x": 1} + + @pytest.mark.asyncio + async def test_no_response_path_does_not_assign(self) -> None: + handler = StubHandler(_ok('{"x":1}')) + factory = WorkflowFactory(http_request_handler=handler) + workflow = factory.create_workflow_from_definition(_yaml(_action())) + # Should complete without error and without writing anything + await workflow.run({}) + + +# ---------- Method / headers / query params -------------------------------- + + +class TestRequestComposition: + @pytest.mark.asyncio + async def test_default_method_is_get(self) -> None: + handler = StubHandler(_ok()) + factory = WorkflowFactory(http_request_handler=handler) + workflow = factory.create_workflow_from_definition(_yaml(_action())) + await workflow.run({}) + + assert handler.last_info is not None + assert handler.last_info.method == "GET" + + @pytest.mark.asyncio + async def test_method_uppercased(self) -> None: + handler = StubHandler(_ok()) + factory = WorkflowFactory(http_request_handler=handler) + workflow = factory.create_workflow_from_definition(_yaml(_action(method="post"))) + await workflow.run({}) + + assert handler.last_info is not None + assert handler.last_info.method == "POST" + + @pytest.mark.asyncio + async def test_headers_are_forwarded_and_empty_skipped(self) -> None: + handler = StubHandler(_ok()) + factory = WorkflowFactory(http_request_handler=handler) + workflow = factory.create_workflow_from_definition( + _yaml( + _action( + headers={ + "Accept": "application/json", + "X-Empty": "", + "Authorization": "Bearer token", + } + ) + ) + ) + await workflow.run({}) + + assert handler.last_info is not None + assert handler.last_info.headers == { + "Accept": "application/json", + "Authorization": "Bearer token", + } + + @pytest.mark.asyncio + async def test_query_parameters_stringified(self) -> None: + handler = StubHandler(_ok()) + factory = WorkflowFactory(http_request_handler=handler) + workflow = factory.create_workflow_from_definition( + _yaml( + _action( + query_parameters={ + "name": "alpha", + "limit": 10, + "active": True, + "ratio": 0.5, + "missing": None, # dropped + } + ) + ) + ) + await workflow.run({}) + + assert handler.last_info is not None + assert handler.last_info.query_parameters == { + "name": "alpha", + "limit": "10", + "active": "true", + "ratio": "0.5", + } + + +# ---------- Body composition ------------------------------------------------ + + +class TestBody: + @pytest.mark.asyncio + async def test_post_json_body_sets_content_type_and_serialises(self) -> None: + handler = StubHandler(_ok()) + factory = WorkflowFactory(http_request_handler=handler) + workflow = factory.create_workflow_from_definition( + _yaml( + _action( + method="POST", + body={"kind": "json", "content": {"k": "v", "n": 1}}, + ) + ) + ) + await workflow.run({}) + + info = handler.last_info + assert info is not None + assert info.body_content_type == "application/json" + assert info.body is not None + # JSON serialized, key order may vary + import json + + assert json.loads(info.body) == {"k": "v", "n": 1} + + @pytest.mark.asyncio + async def test_post_raw_body_uses_declared_content_type(self) -> None: + handler = StubHandler(_ok()) + factory = WorkflowFactory(http_request_handler=handler) + workflow = factory.create_workflow_from_definition( + _yaml( + _action( + method="POST", + body={ + "kind": "raw", + "content": "raw body text", + "contentType": "text/plain", + }, + ) + ) + ) + await workflow.run({}) + + info = handler.last_info + assert info is not None + assert info.body == "raw body text" + assert info.body_content_type == "text/plain" + + @pytest.mark.asyncio + async def test_post_raw_body_without_content_type_defaults_to_text_plain(self) -> None: + """Match .NET RawRequestContent: no contentType => default text/plain. + + Otherwise the request is sent without a Content-Type header which most + servers will treat as application/octet-stream and fail to parse. + """ + handler = StubHandler(_ok()) + factory = WorkflowFactory(http_request_handler=handler) + workflow = factory.create_workflow_from_definition( + _yaml( + _action( + method="POST", + body={"kind": "raw", "content": "plain body"}, + ) + ) + ) + await workflow.run({}) + + info = handler.last_info + assert info is not None + assert info.body == "plain body" + assert info.body_content_type == "text/plain" + + @pytest.mark.asyncio + async def test_long_form_body_kinds_accepted(self) -> None: + handler = StubHandler(_ok()) + factory = WorkflowFactory(http_request_handler=handler) + workflow = factory.create_workflow_from_definition( + _yaml( + _action( + method="POST", + body={"kind": "JsonRequestContent", "content": {"k": 1}}, + ) + ) + ) + await workflow.run({}) + info = handler.last_info + assert info is not None + assert info.body_content_type == "application/json" + + @pytest.mark.asyncio + async def test_unknown_body_kind_raises(self) -> None: + handler = StubHandler(_ok()) + factory = WorkflowFactory(http_request_handler=handler) + workflow = factory.create_workflow_from_definition( + _yaml(_action(body={"kind": "weirdform", "content": "x"})) + ) + with pytest.raises(Exception) as excinfo: + await workflow.run({}) + # Should surface as ValueError (potentially wrapped by runner) + msg = str(excinfo.value) + assert "weirdform" in msg or "unsupported value" in msg + + @pytest.mark.asyncio + async def test_no_body_omitted(self) -> None: + handler = StubHandler(_ok()) + factory = WorkflowFactory(http_request_handler=handler) + workflow = factory.create_workflow_from_definition(_yaml(_action())) + await workflow.run({}) + info = handler.last_info + assert info is not None + assert info.body is None + assert info.body_content_type is None + + +# ---------- Non-2xx and error handling ------------------------------------- + + +class TestErrorHandling: + @pytest.mark.asyncio + async def test_non_2xx_raises_declarative_action_error(self) -> None: + handler = StubHandler(_err(status=500, body="server exploded")) + factory = WorkflowFactory(http_request_handler=handler) + workflow = factory.create_workflow_from_definition(_yaml(_action())) + with pytest.raises(DeclarativeActionError) as excinfo: + await workflow.run({}) + msg = str(excinfo.value) + assert "500" in msg + assert "server exploded" in msg + + @pytest.mark.asyncio + async def test_non_2xx_long_body_truncated(self) -> None: + big_body = "A" * 1000 + handler = StubHandler(_err(status=500, body=big_body)) + factory = WorkflowFactory(http_request_handler=handler) + workflow = factory.create_workflow_from_definition(_yaml(_action())) + with pytest.raises(DeclarativeActionError) as excinfo: + await workflow.run({}) + msg = str(excinfo.value) + assert "[truncated]" in msg + assert len(msg) < 512 + # Should NOT contain the full 1000-char body + assert big_body not in msg + + @pytest.mark.asyncio + async def test_non_2xx_empty_body_omits_body_section(self) -> None: + handler = StubHandler(_err(status=404, body="")) + factory = WorkflowFactory(http_request_handler=handler) + workflow = factory.create_workflow_from_definition(_yaml(_action())) + with pytest.raises(DeclarativeActionError) as excinfo: + await workflow.run({}) + msg = str(excinfo.value) + assert "404" in msg + assert "Body:" not in msg + + @pytest.mark.asyncio + async def test_non_2xx_control_chars_collapsed(self) -> None: + handler = StubHandler(_err(status=500, body="line1\r\nline2\tlong")) + factory = WorkflowFactory(http_request_handler=handler) + workflow = factory.create_workflow_from_definition(_yaml(_action())) + with pytest.raises(DeclarativeActionError) as excinfo: + await workflow.run({}) + msg = str(excinfo.value) + assert "\r" not in msg + assert "\n" not in msg + assert "\t" not in msg + assert "line1 line2 long" in msg + + @pytest.mark.asyncio + async def test_timeout_exception_becomes_declarative_action_error(self) -> None: + handler = StubHandler(raise_exc=httpx.TimeoutException("timeout")) + factory = WorkflowFactory(http_request_handler=handler) + workflow = factory.create_workflow_from_definition(_yaml(_action())) + with pytest.raises(DeclarativeActionError) as excinfo: + await workflow.run({}) + assert "timed out" in str(excinfo.value) + + @pytest.mark.asyncio + async def test_stdlib_timeout_error_becomes_declarative_action_error(self) -> None: + handler = StubHandler(raise_exc=TimeoutError("clock")) + factory = WorkflowFactory(http_request_handler=handler) + workflow = factory.create_workflow_from_definition(_yaml(_action())) + with pytest.raises(DeclarativeActionError) as excinfo: + await workflow.run({}) + assert "timed out" in str(excinfo.value) + + @pytest.mark.asyncio + async def test_transport_error_becomes_declarative_action_error(self) -> None: + handler = StubHandler(raise_exc=httpx.ConnectError("dns failure")) + factory = WorkflowFactory(http_request_handler=handler) + workflow = factory.create_workflow_from_definition(_yaml(_action())) + with pytest.raises(DeclarativeActionError) as excinfo: + await workflow.run({}) + msg = str(excinfo.value) + assert "failed" in msg + assert _TEST_URL in msg + + @pytest.mark.asyncio + async def test_cancelled_error_propagates_unchanged(self) -> None: + """CancelledError from the handler must propagate so cancellation works.""" + handler = StubHandler(raise_exc=asyncio.CancelledError()) + factory = WorkflowFactory(http_request_handler=handler) + workflow = factory.create_workflow_from_definition(_yaml(_action())) + # CancelledError is allowed to surface as either CancelledError or as + # the runner's wrapped form, but it MUST NOT be DeclarativeActionError. + with pytest.raises(BaseException) as excinfo: + await workflow.run({}) + assert not isinstance(excinfo.value, DeclarativeActionError) + + @pytest.mark.asyncio + async def test_generic_exception_from_custom_handler_wrapped(self) -> None: + """A custom handler raising a non-httpx Exception must be wrapped. + + Authors can plug in custom HttpRequestHandler implementations that use + any transport (requests-like clients, gRPC bridges, mock test doubles, + etc.). The executor must wrap arbitrary Exception subclasses uniformly + so that workflow error handling stays consistent across transports. + """ + handler = StubHandler(raise_exc=RuntimeError("custom transport blew up")) + factory = WorkflowFactory(http_request_handler=handler) + workflow = factory.create_workflow_from_definition(_yaml(_action())) + with pytest.raises(DeclarativeActionError) as excinfo: + await workflow.run({}) + msg = str(excinfo.value) + assert "failed" in msg + assert "RuntimeError" in msg + assert _TEST_URL in msg + + +# ---------- Response headers ------------------------------------------------ + + +class TestResponseHeaders: + @pytest.mark.asyncio + async def test_response_headers_folded_with_commas(self) -> None: + handler = StubHandler( + _ok( + "ok", + headers={ + "Content-Type": ["application/json"], + "Set-Cookie": ["a=1", "b=2"], + }, + ) + ) + factory = WorkflowFactory(http_request_handler=handler) + workflow = factory.create_workflow_from_definition( + _yaml(_action(response_headers="Local.H")) + ) + await workflow.run({}) + decl = workflow._state.get(DECLARATIVE_STATE_KEY) + h = decl["Local"]["H"] + assert h["Content-Type"] == "application/json" + assert h["Set-Cookie"] == "a=1,b=2" + + @pytest.mark.asyncio + async def test_response_headers_empty_assigned_none(self) -> None: + handler = StubHandler(_ok("ok", headers={})) + factory = WorkflowFactory(http_request_handler=handler) + workflow = factory.create_workflow_from_definition( + _yaml(_action(response_headers="Local.H")) + ) + await workflow.run({}) + decl = workflow._state.get(DECLARATIVE_STATE_KEY) + assert decl["Local"]["H"] is None + + @pytest.mark.asyncio + async def test_non_2xx_still_publishes_headers(self) -> None: + handler = StubHandler(_err(status=500, body="boom", headers={"X-Trace": ["abc"]})) + factory = WorkflowFactory(http_request_handler=handler) + workflow = factory.create_workflow_from_definition( + _yaml(_action(response_headers="Local.H")) + ) + with pytest.raises(DeclarativeActionError): + await workflow.run({}) + decl = workflow._state.get(DECLARATIVE_STATE_KEY) + assert decl["Local"]["H"] == {"X-Trace": "abc"} + + +# ---------- ConversationId append ------------------------------------------- + + +class TestConversationAppend: + @pytest.mark.asyncio + async def test_conversation_id_appends_message(self) -> None: + handler = StubHandler(_ok('{"answer":"hello"}')) + factory = WorkflowFactory(http_request_handler=handler) + workflow = factory.create_workflow_from_definition( + _yaml( + _action( + response="Local.Result", + conversation_id="conv-test-1", + ) + ) + ) + await workflow.run({}) + decl = workflow._state.get(DECLARATIVE_STATE_KEY) + conv = decl["System"]["conversations"].get("conv-test-1") + assert conv is not None + assert len(conv["messages"]) == 1 + + @pytest.mark.asyncio + async def test_empty_conversation_id_does_not_append(self) -> None: + handler = StubHandler(_ok('{"answer":"hello"}')) + factory = WorkflowFactory(http_request_handler=handler) + workflow = factory.create_workflow_from_definition( + _yaml(_action(response="Local.Result", conversation_id="")) + ) + await workflow.run({}) + decl = workflow._state.get(DECLARATIVE_STATE_KEY) + # Auto-init creates an entry for the System.ConversationId conversation, + # but it should NOT have HTTP-appended messages from us. + for _cid, conv in decl["System"]["conversations"].items(): + assert conv["messages"] == [] + + @pytest.mark.asyncio + async def test_empty_body_skips_conversation_append(self) -> None: + handler = StubHandler(_ok("")) + factory = WorkflowFactory(http_request_handler=handler) + workflow = factory.create_workflow_from_definition( + _yaml(_action(conversation_id="conv-test-1")) + ) + await workflow.run({}) + decl = workflow._state.get(DECLARATIVE_STATE_KEY) + # No conversation entry should have been created either. + assert "conv-test-1" not in decl["System"]["conversations"] + + +# ---------- Connection name ------------------------------------------------- + + +class TestConnection: + @pytest.mark.asyncio + async def test_connection_name_forwarded(self) -> None: + handler = StubHandler(_ok()) + factory = WorkflowFactory(http_request_handler=handler) + workflow = factory.create_workflow_from_definition( + _yaml(_action(connection={"name": "my-connection"})) + ) + await workflow.run({}) + assert handler.last_info is not None + assert handler.last_info.connection_name == "my-connection" + + +# ---------- Build-time validation ------------------------------------------- + + +class TestBuildTimeValidation: + def test_missing_url_fails_validation(self) -> None: + handler = StubHandler(_ok()) + factory = WorkflowFactory(http_request_handler=handler) + bad = { + "name": "no_url", + "actions": [{"kind": "HttpRequestAction", "id": "x"}], + } + with pytest.raises(DeclarativeWorkflowError): + factory.create_workflow_from_definition(bad) + + def test_missing_handler_fails_at_build(self) -> None: + factory = WorkflowFactory() # no handler + with pytest.raises(DeclarativeWorkflowError) as excinfo: + factory.create_workflow_from_definition(_yaml(_action())) + assert "http_request_handler" in str(excinfo.value) + + +# ---------- Timeout forwarding ---------------------------------------------- + + +class TestTimeout: + @pytest.mark.asyncio + async def test_timeout_ms_forwarded(self) -> None: + handler = StubHandler(_ok()) + factory = WorkflowFactory(http_request_handler=handler) + workflow = factory.create_workflow_from_definition( + _yaml(_action(request_timeout_ms=2500)) + ) + await workflow.run({}) + assert handler.last_info is not None + assert handler.last_info.timeout_ms == 2500 + + @pytest.mark.asyncio + async def test_timeout_ms_zero_treated_as_unset(self) -> None: + handler = StubHandler(_ok()) + factory = WorkflowFactory(http_request_handler=handler) + workflow = factory.create_workflow_from_definition( + _yaml(_action(request_timeout_ms=0)) + ) + await workflow.run({}) + assert handler.last_info is not None + assert handler.last_info.timeout_ms is None diff --git a/python/packages/declarative/tests/test_http_request_yaml_integration.py b/python/packages/declarative/tests/test_http_request_yaml_integration.py new file mode 100644 index 00000000000..49cd0d15e83 --- /dev/null +++ b/python/packages/declarative/tests/test_http_request_yaml_integration.py @@ -0,0 +1,111 @@ +# Copyright (c) Microsoft. All rights reserved. + +"""End-to-end YAML integration test for ``HttpRequestAction``. + +Loads the ``tests/workflows/http_request.yaml`` fixture (parity with the .NET +integration fixture) through ``WorkflowFactory.create_workflow_from_yaml_path`` +with a stub :class:`HttpRequestHandler` and asserts state is populated. +""" + +from __future__ import annotations + +import sys +from pathlib import Path +from typing import Any + +import pytest + +try: + import powerfx # noqa: F401 + + _powerfx_available = True +except (ImportError, RuntimeError): + _powerfx_available = False + +pytestmark = [ + pytest.mark.skipif( + not _powerfx_available, + reason="powerfx not available — declarative workflows require it.", + ), + pytest.mark.skipif( + sys.version_info >= (3, 14), + reason="Skipped on Python 3.14+ to keep parity with declarative suite.", + ), +] + +from agent_framework_declarative import WorkflowFactory # noqa: E402 +from agent_framework_declarative._workflows import DECLARATIVE_STATE_KEY # noqa: E402 +from agent_framework_declarative._workflows._http_handler import ( # noqa: E402 + HttpRequestInfo, + HttpRequestResult, +) + +FIXTURE_PATH = Path(__file__).parent / "workflows" / "http_request.yaml" + + +class _StubHandler: + """Test double that records requests and returns a canned response.""" + + def __init__(self, result: HttpRequestResult) -> None: + self._result = result + self.received: list[HttpRequestInfo] = [] + + async def send(self, info: HttpRequestInfo) -> HttpRequestResult: + self.received.append(info) + return self._result + + +@pytest.mark.asyncio +async def test_http_request_yaml_roundtrip() -> None: + handler = _StubHandler( + HttpRequestResult( + status_code=200, + is_success_status_code=True, + body='{"name": "runtime", "visibility": "public", "stars": 12345}', + headers={ + "content-type": ["application/json"], + "x-ratelimit-remaining": ["59"], + }, + ) + ) + + factory = WorkflowFactory(http_request_handler=handler) + workflow = factory.create_workflow_from_yaml_path(FIXTURE_PATH) + await workflow.run({}) + + decl: dict[str, Any] = workflow._state.get(DECLARATIVE_STATE_KEY) or {} + local = decl.get("Local") or {} + + assert local.get("RepoOwner") == "dotnet" + repo_info = local.get("RepoInfo") + assert isinstance(repo_info, dict), f"Expected dict body, got {type(repo_info)!r}" + assert repo_info["name"] == "runtime" + assert repo_info["visibility"] == "public" + assert repo_info["stars"] == 12345 + + repo_headers = local.get("RepoHeaders") + assert isinstance(repo_headers, dict) + # Single-value header surfaces as plain string. + assert repo_headers.get("content-type") == "application/json" + assert repo_headers.get("x-ratelimit-remaining") == "59" + + # Stub got the right call. + assert len(handler.received) == 1 + sent = handler.received[0] + assert sent.method == "GET" + assert sent.url == "https://api.github.com/repos/dotnet/runtime" + assert sent.headers["Accept"] == "application/vnd.github+json" + assert sent.headers["User-Agent"] == "agent-framework-integration-test" + + +@pytest.mark.asyncio +async def test_http_request_yaml_missing_handler_fails_at_build_time() -> None: + """Without an http_request_handler, building the workflow must raise.""" + from agent_framework_declarative._workflows._errors import DeclarativeWorkflowError + + factory = WorkflowFactory() # no handler configured + with pytest.raises(DeclarativeWorkflowError) as excinfo: + factory.create_workflow_from_yaml_path(FIXTURE_PATH) + msg = str(excinfo.value) + assert "HttpRequestAction" in msg + assert "http_request_handler" in msg diff --git a/python/packages/declarative/tests/workflows/http_request.yaml b/python/packages/declarative/tests/workflows/http_request.yaml new file mode 100644 index 00000000000..382d3fafe20 --- /dev/null +++ b/python/packages/declarative/tests/workflows/http_request.yaml @@ -0,0 +1,29 @@ +# +# Integration fixture: end-to-end HttpRequestAction round-trip using a +# stub HttpRequestHandler. Mirrors the .NET integration fixture in +# dotnet/tests/.../Workflows/HttpRequest.yaml. +# +kind: Workflow +trigger: + + kind: OnConversationStart + id: workflow_http_request_test + actions: + + # Set the repo owner used to form the request URL. + - kind: SetVariable + id: set_repo_owner + variable: Local.RepoOwner + value: dotnet + + # Invoke the (stubbed) GitHub repo API. + - kind: HttpRequestAction + id: fetch_repo_info + conversationId: =System.ConversationId + method: GET + url: =Concatenate("https://api.github.com/repos/", Local.RepoOwner, "/runtime") + headers: + Accept: application/vnd.github+json + User-Agent: agent-framework-integration-test + response: Local.RepoInfo + responseHeaders: Local.RepoHeaders diff --git a/python/samples/03-workflows/declarative/invoke_http_request/main.py b/python/samples/03-workflows/declarative/invoke_http_request/main.py new file mode 100644 index 00000000000..ebcbcc0a164 --- /dev/null +++ b/python/samples/03-workflows/declarative/invoke_http_request/main.py @@ -0,0 +1,97 @@ +# Copyright (c) Microsoft. All rights reserved. + +"""Invoke HTTP Request sample - demonstrates the HttpRequestAction declarative action. + +This sample shows how to: + 1. Configure a ``WorkflowFactory`` with a ``HttpRequestHandler`` so the YAML + ``HttpRequestAction`` can dispatch real HTTP calls. + 2. Fetch JSON from a public REST endpoint (the GitHub repository API) and + bind the parsed response to a workflow variable. + 3. Mirror the response body into the conversation via ``conversationId`` so + a downstream Foundry agent can answer questions about it using only that + conversation context. + +Security note: + ``DefaultHttpRequestHandler`` issues HTTP calls to whatever URL the + workflow author specifies and performs **no** allowlisting or SSRF + guards. For production use, replace it with a custom handler that + enforces an allowlist or DNS-rebinding-resistant policy and adds any + required authentication headers per call. + +Run with: + python -m samples.03-workflows.declarative.invoke_http_request.main +""" + +import asyncio +import os +from pathlib import Path + +from agent_framework import Agent +from agent_framework.declarative import ( + DefaultHttpRequestHandler, + WorkflowFactory, +) +from agent_framework.foundry import FoundryChatClient +from azure.identity import AzureCliCredential + +GITHUB_REPO_INFO_AGENT_INSTRUCTIONS = """\ +You answer the user's question about a GitHub repository using ONLY the JSON +data already present in the conversation history. If the answer is not +contained in the conversation, say so plainly rather than guessing. Be concise +and helpful. +""" + + +async def main() -> None: + """Run the invoke HTTP request workflow.""" + chat_client = FoundryChatClient( + project_endpoint=os.environ["FOUNDRY_PROJECT_ENDPOINT"], + model=os.environ["FOUNDRY_MODEL"], + credential=AzureCliCredential(), + ) + + # The agent has no tools — it answers the question about the GitHub + # repository using only the JSON data that ``HttpRequestAction`` adds to + # the conversation. + github_repo_info_agent = Agent( + client=chat_client, + name="GitHubRepoInfoAgent", + instructions=GITHUB_REPO_INFO_AGENT_INSTRUCTIONS, + ) + + agents = {"GitHubRepoInfoAgent": github_repo_info_agent} + + # The default HttpRequestHandler is sufficient for this sample because + # the GitHub REST endpoint used here does not require authentication. + # For authenticated endpoints, supply a custom client_provider callback + # to DefaultHttpRequestHandler so each request can be routed through a + # pre-configured httpx.AsyncClient with the appropriate credentials. + async with DefaultHttpRequestHandler() as http_handler: + factory = WorkflowFactory( + agents=agents, + http_request_handler=http_handler, + ) + + workflow_path = Path(__file__).parent / "workflow.yaml" + workflow = factory.create_workflow_from_yaml_path(workflow_path) + + print("=" * 60) + print("Invoke HTTP Request Workflow Demo") + print("=" * 60) + print() + print("Ask one question about the microsoft/agent-framework repo.") + print() + + user_input = input("You: ").strip() # noqa: ASYNC250 + if not user_input: + user_input = "Please summarize the repository." + + print("\nAgent: ", end="", flush=True) + async for event in workflow.run(user_input, stream=True): + if event.type == "output" and isinstance(event.data, str): + print(event.data, end="", flush=True) + print() + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/python/samples/03-workflows/declarative/invoke_http_request/workflow.yaml b/python/samples/03-workflows/declarative/invoke_http_request/workflow.yaml new file mode 100644 index 00000000000..f1bfd223bb2 --- /dev/null +++ b/python/samples/03-workflows/declarative/invoke_http_request/workflow.yaml @@ -0,0 +1,57 @@ +# +# This workflow demonstrates the HttpRequestAction declarative action. +# +# HttpRequestAction lets a workflow author issue an HTTP call directly from +# YAML without writing any Python glue. It can: +# +# - fetch data from external REST endpoints, +# - store the parsed response in a workflow variable, and +# - add the response body to the conversation so a downstream agent can +# answer questions based on it. +# +# This sample fetches public metadata for the microsoft/agent-framework +# repository from the GitHub REST API (no authentication required) and uses +# a Foundry agent to answer a single question about it. +# +# Example input: +# How many open issues does the repository have? +# +kind: Workflow +trigger: + + kind: OnConversationStart + id: workflow_invoke_http_request_demo + actions: + + # Set the repository org/name used to form the request URL. + - kind: SetVariable + id: set_repo_name + variable: Local.RepoName + value: microsoft/agent-framework + + # Invoke the GitHub repo API. The response body is parsed into + # Local.RepoInfo and also added to the conversation (via conversationId) + # so the agent below can answer questions based on it. + - kind: HttpRequestAction + id: fetch_repo_info + conversationId: =System.ConversationId + method: GET + url: =Concatenate("https://api.github.com/repos/", Local.RepoName) + headers: + Accept: application/vnd.github+json + User-Agent: agent-framework-sample + response: Local.RepoInfo + + # Use the agent to answer the user's question using the conversation + # context (which now contains the GitHub JSON response). The user's + # original message is already in the conversation as System.LastMessage, + # and the executor's input fallback chain extracts its ``Text`` field + # automatically when ``input.messages`` is omitted. + - kind: InvokeAzureAgent + id: answer_question + conversationId: =System.ConversationId + agent: + name: GitHubRepoInfoAgent + output: + autoSend: true + messages: Local.AgentResponse diff --git a/python/uv.lock b/python/uv.lock index 12c5f39b275..85a968174a3 100644 --- a/python/uv.lock +++ b/python/uv.lock @@ -428,6 +428,7 @@ version = "1.0.0b260429" source = { editable = "packages/declarative" } dependencies = [ { name = "agent-framework-core", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, + { name = "httpx", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, { name = "powerfx", marker = "(python_full_version < '3.14' and sys_platform == 'darwin') or (python_full_version < '3.14' and sys_platform == 'linux') or (python_full_version < '3.14' and sys_platform == 'win32')" }, { name = "pyyaml", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, ] @@ -440,6 +441,7 @@ dev = [ [package.metadata] requires-dist = [ { name = "agent-framework-core", editable = "packages/core" }, + { name = "httpx", specifier = ">=0.27,<1" }, { name = "powerfx", marker = "python_full_version < '3.14'", specifier = ">=0.0.32,<0.0.35" }, { name = "pyyaml", specifier = ">=6.0,<7.0" }, ] From 7999bf3c2dd766d9a26f376298b7227e8fbb8afc Mon Sep 17 00:00:00 2001 From: Peter Ibekwe Date: Thu, 30 Apr 2026 16:26:38 -0700 Subject: [PATCH 2/8] Ran pyupgrade and pright to fix CI issues --- .../_workflows/_executors_http.py | 60 ++++++++----------- .../_workflows/_http_handler.py | 20 +++++-- .../test_default_http_request_handler.py | 39 +++--------- .../tests/test_http_request_executor.py | 56 +++++------------ 4 files changed, 61 insertions(+), 114 deletions(-) diff --git a/python/packages/declarative/agent_framework_declarative/_workflows/_executors_http.py b/python/packages/declarative/agent_framework_declarative/_workflows/_executors_http.py index e041a40c3b8..3df7a2bca18 100644 --- a/python/packages/declarative/agent_framework_declarative/_workflows/_executors_http.py +++ b/python/packages/declarative/agent_framework_declarative/_workflows/_executors_http.py @@ -64,7 +64,7 @@ def _get_path(action_def: Mapping[str, Any], key: str) -> str | None: if isinstance(value, str): return value or None if isinstance(value, Mapping): - path = value.get("path") + path = value.get("path") # type: ignore[reportUnknownMemberType, reportUnknownVariableType] return path if isinstance(path, str) and path else None return None @@ -97,7 +97,7 @@ def _parse_response_body(body: str | None) -> Any: return None try: return json.loads(body) - except (ValueError, json.JSONDecodeError): + except json.JSONDecodeError: return body @@ -118,13 +118,14 @@ def _format_query_value(value: Any) -> str | None: def _get_messages_path(state: DeclarativeWorkflowState, conversation_id_expr: str | None) -> str | None: - """Mirror ``InvokeAzureAgentExecutor._get_conversation_messages_path`` semantics. - - Returns a single state path (no double-mirroring): ``Conversation.messages`` - when no conversation id is provided, otherwise - ``System.conversations.{evaluated_id}.messages``. Returns ``None`` if the - expression evaluates to an empty string (matches .NET ``GetConversationId`` - behaviour where empty becomes ``null`` and the response is not appended). + """Return the configured conversation messages path, if any. + + Returns ``System.conversations.{evaluated_id}.messages`` when a + ``conversation_id_expr`` is configured and evaluates to a non-empty value. + Returns ``None`` when no conversation id expression is configured or when + the expression evaluates to ``None`` or an empty string (matches .NET + ``GetConversationId`` behaviour where empty becomes ``null`` and the + response is not appended). """ if not conversation_id_expr: return None @@ -209,9 +210,7 @@ async def handle_action( except DeclarativeActionError: raise except httpx.HTTPError as exc: - raise DeclarativeActionError( - f"HTTP request to '{url}' failed: {type(exc).__name__}" - ) from exc + raise DeclarativeActionError(f"HTTP request to '{url}' failed: {type(exc).__name__}") from exc except Exception as exc: # Custom HttpRequestHandler implementations may raise arbitrary # exception types. Wrap them in DeclarativeActionError so workflow @@ -219,9 +218,7 @@ async def handle_action( # ``asyncio.CancelledError`` is a ``BaseException`` (not # ``Exception``) and so still propagates unmodified, preserving # workflow-cancellation semantics. - raise DeclarativeActionError( - f"HTTP request to '{url}' failed: {type(exc).__name__}" - ) from exc + raise DeclarativeActionError(f"HTTP request to '{url}' failed: {type(exc).__name__}") from exc if result.is_success_status_code: self._assign_response(state, result) @@ -234,10 +231,7 @@ async def handle_action( self._assign_response_headers(state, result) body_preview = _format_body_for_diagnostics(result.body) if body_preview: - message = ( - f"HTTP request to '{url}' failed with status code {result.status_code}. " - f"Body: '{body_preview}'" - ) + message = f"HTTP request to '{url}' failed with status code {result.status_code}. Body: '{body_preview}'" else: message = f"HTTP request to '{url}' failed with status code {result.status_code}." raise DeclarativeActionError(message) @@ -265,7 +259,7 @@ def _get_headers(self, state: DeclarativeWorkflowState) -> dict[str, str] | None if not isinstance(raw_headers, Mapping) or not raw_headers: return None result: dict[str, str] = {} - for key, value in raw_headers.items(): + for key, value in raw_headers.items(): # type: ignore[reportUnknownVariableType] if not isinstance(key, str) or not key: continue evaluated = state.eval_if_expression(value) @@ -282,7 +276,7 @@ def _get_query_parameters(self, state: DeclarativeWorkflowState) -> dict[str, st if not isinstance(raw_params, Mapping) or not raw_params: return None result: dict[str, str] = {} - for key, value in raw_params.items(): + for key, value in raw_params.items(): # type: ignore[reportUnknownVariableType] if not isinstance(key, str) or not key or value is None: continue evaluated = state.eval_if_expression(value) @@ -297,40 +291,34 @@ def _get_body(self, state: DeclarativeWorkflowState) -> tuple[str | None, str | return None, None if not isinstance(raw_body, Mapping): raise ValueError( - "HttpRequestAction 'body' must be a mapping with a 'kind' field " - "(json, raw) or omitted entirely." + "HttpRequestAction 'body' must be a mapping with a 'kind' field (json, raw) or omitted entirely." ) - kind_value = raw_body.get("kind") or raw_body.get("$kind") + kind_value: Any = raw_body.get("kind") or raw_body.get("$kind") # type: ignore[reportUnknownMemberType] if kind_value is None: raise ValueError( - "HttpRequestAction 'body' is missing 'kind'. Use 'json', 'raw', " - "or omit 'body' for no request body." + "HttpRequestAction 'body' is missing 'kind'. Use 'json', 'raw', or omit 'body' for no request body." ) if not isinstance(kind_value, str): - raise ValueError( - f"HttpRequestAction 'body.kind' must be a string, got {type(kind_value).__name__}." - ) + raise ValueError(f"HttpRequestAction 'body.kind' must be a string, got {kind_value!r}.") if kind_value in _BODY_KIND_NONE: return None, None if kind_value in _BODY_KIND_JSON: - content_expr = raw_body.get("content") + content_expr: Any = raw_body.get("content") # type: ignore[reportUnknownMemberType] if content_expr is None: return None, None evaluated = state.eval_if_expression(content_expr) try: body_text = json.dumps(evaluated, default=str) except (TypeError, ValueError) as exc: - raise ValueError( - f"HttpRequestAction 'body.content' could not be serialised as JSON: {exc}" - ) from exc + raise ValueError(f"HttpRequestAction 'body.content' could not be serialised as JSON: {exc}") from exc return body_text, "application/json" if kind_value in _BODY_KIND_RAW: - content_expr = raw_body.get("content") - content_type_expr = raw_body.get("contentType") + content_expr = raw_body.get("content") # type: ignore[reportUnknownMemberType] + content_type_expr: Any = raw_body.get("contentType") # type: ignore[reportUnknownMemberType] content: str | None = None if content_expr is not None: evaluated = state.eval_if_expression(content_expr) @@ -374,7 +362,7 @@ def _get_connection_name(self, state: DeclarativeWorkflowState) -> str | None: connection = self._action_def.get("connection") if not isinstance(connection, Mapping): return None - name_expr = connection.get("name") + name_expr: Any = connection.get("name") # type: ignore[reportUnknownMemberType] if name_expr is None: return None evaluated = state.eval_if_expression(name_expr) diff --git a/python/packages/declarative/agent_framework_declarative/_workflows/_http_handler.py b/python/packages/declarative/agent_framework_declarative/_workflows/_http_handler.py index 75231fe91d2..90ff5b87b4c 100644 --- a/python/packages/declarative/agent_framework_declarative/_workflows/_http_handler.py +++ b/python/packages/declarative/agent_framework_declarative/_workflows/_http_handler.py @@ -57,8 +57,8 @@ class HttpRequestInfo: method: str url: str - headers: dict[str, str] = field(default_factory=dict) - query_parameters: dict[str, str] = field(default_factory=dict) + headers: dict[str, str] = field(default_factory=dict) # type: ignore[reportUnknownVariableType] + query_parameters: dict[str, str] = field(default_factory=dict) # type: ignore[reportUnknownVariableType] body: str | None = None body_content_type: str | None = None timeout_ms: int | None = None @@ -74,12 +74,17 @@ class HttpRequestResult: ``dict[str, list[str]]``. The executor folds duplicates into a single comma-joined string only at the point it assigns ``responseHeaders`` to workflow state. + + Header keys are normalized to lowercase so that lookups are consistent + regardless of the server's transmitted casing (HTTP headers are + case-insensitive per RFC 7230 §3.2). Custom :class:`HttpRequestHandler` + implementations should follow the same convention. """ status_code: int is_success_status_code: bool body: str - headers: dict[str, list[str]] = field(default_factory=dict) + headers: dict[str, list[str]] = field(default_factory=dict) # type: ignore[reportUnknownVariableType] @runtime_checkable @@ -132,7 +137,7 @@ class DefaultHttpRequestHandler: def __init__( self, *, - client: "httpx.AsyncClient | None" = None, + client: httpx.AsyncClient | None = None, client_provider: ClientProvider | None = None, ) -> None: self._owned_client: httpx.AsyncClient | None = None @@ -180,9 +185,12 @@ async def send(self, info: HttpRequestInfo) -> HttpRequestResult: ) # Preserve multi-value headers (e.g. multiple Set-Cookie) as list[str]. + # Normalize names to lowercase so lookups are consistent and case + # variations from the transport do not create duplicate logical keys + # (HTTP headers are case-insensitive per RFC 7230 §3.2). result_headers: dict[str, list[str]] = {} for key, value in response.headers.multi_items(): - result_headers.setdefault(key, []).append(value) + result_headers.setdefault(key.lower(), []).append(value) body_text = response.text @@ -216,7 +224,7 @@ async def _resolve_client(self, info: HttpRequestInfo) -> httpx.AsyncClient: self._owned_client = httpx.AsyncClient() return self._owned_client - async def __aenter__(self) -> "DefaultHttpRequestHandler": + async def __aenter__(self) -> DefaultHttpRequestHandler: return self async def __aexit__(self, exc_type: Any, exc: Any, tb: Any) -> None: diff --git a/python/packages/declarative/tests/test_default_http_request_handler.py b/python/packages/declarative/tests/test_default_http_request_handler.py index 5391b90adfb..ecdce3d7ff0 100644 --- a/python/packages/declarative/tests/test_default_http_request_handler.py +++ b/python/packages/declarative/tests/test_default_http_request_handler.py @@ -181,9 +181,7 @@ def respond(request: httpx.Request) -> httpx.Response: handler = _make_handler(httpx.MockTransport(respond)) try: - result = await handler.send( - HttpRequestInfo(method="GET", url="https://api.example.test/x") - ) + result = await handler.send(HttpRequestInfo(method="GET", url="https://api.example.test/x")) finally: await handler.aclose() @@ -199,9 +197,7 @@ async def test_owned_client_is_closed_on_aclose(self) -> None: handler = DefaultHttpRequestHandler() # Inject a MockTransport-backed client into the owned slot and verify # aclose() releases it. Avoids real network access. - owned = httpx.AsyncClient( - transport=httpx.MockTransport(lambda r: httpx.Response(200, text="ok")) - ) + owned = httpx.AsyncClient(transport=httpx.MockTransport(lambda r: httpx.Response(200, text="ok"))) handler._owned_client = owned assert not owned.is_closed await handler.aclose() @@ -209,13 +205,9 @@ async def test_owned_client_is_closed_on_aclose(self) -> None: @pytest.mark.asyncio async def test_caller_owned_client_is_not_closed(self) -> None: - client = httpx.AsyncClient( - transport=httpx.MockTransport(lambda r: httpx.Response(200, text="ok")) - ) + client = httpx.AsyncClient(transport=httpx.MockTransport(lambda r: httpx.Response(200, text="ok"))) handler = DefaultHttpRequestHandler(client=client) - await handler.send( - HttpRequestInfo(method="GET", url="https://api.example.test/x") - ) + await handler.send(HttpRequestInfo(method="GET", url="https://api.example.test/x")) await handler.aclose() assert not client.is_closed await client.aclose() # cleanup @@ -242,11 +234,7 @@ def counting_ctor(*args, **kwargs): # type: ignore[no-untyped-def] nonlocal construction_count if not args and not kwargs: construction_count += 1 - return original_ctor( - transport=httpx.MockTransport( - lambda r: httpx.Response(200, text="ok") - ) - ) + return original_ctor(transport=httpx.MockTransport(lambda r: httpx.Response(200, text="ok"))) return original_ctor(*args, **kwargs) import agent_framework_declarative._workflows._http_handler as hh @@ -256,10 +244,7 @@ def counting_ctor(*args, **kwargs): # type: ignore[no-untyped-def] handler = DefaultHttpRequestHandler() try: await asyncio.gather(*[ - handler.send( - HttpRequestInfo(method="GET", url="https://api.example.test/x") - ) - for _ in range(8) + handler.send(HttpRequestInfo(method="GET", url="https://api.example.test/x")) for _ in range(8) ]) finally: await handler.aclose() @@ -292,9 +277,7 @@ async def provider(info: HttpRequestInfo) -> httpx.AsyncClient: handler = DefaultHttpRequestHandler(client=primary_client, client_provider=provider) try: - result = await handler.send( - HttpRequestInfo(method="GET", url="https://api.example.test/x") - ) + result = await handler.send(HttpRequestInfo(method="GET", url="https://api.example.test/x")) assert result.body == "provided" assert captured["transport"] == "provided" finally: @@ -316,9 +299,7 @@ async def provider(info: HttpRequestInfo) -> httpx.AsyncClient | None: primary_client = httpx.AsyncClient(transport=httpx.MockTransport(primary)) handler = DefaultHttpRequestHandler(client=primary_client, client_provider=provider) try: - result = await handler.send( - HttpRequestInfo(method="GET", url="https://api.example.test/x") - ) + result = await handler.send(HttpRequestInfo(method="GET", url="https://api.example.test/x")) assert result.body == "primary" finally: await handler.aclose() @@ -343,8 +324,6 @@ class TestAsyncContextManager: @pytest.mark.asyncio async def test_context_manager_closes_owned_client(self) -> None: async with DefaultHttpRequestHandler() as handler: - owned = httpx.AsyncClient( - transport=httpx.MockTransport(lambda r: httpx.Response(200, text="ok")) - ) + owned = httpx.AsyncClient(transport=httpx.MockTransport(lambda r: httpx.Response(200, text="ok"))) handler._owned_client = owned assert owned.is_closed diff --git a/python/packages/declarative/tests/test_http_request_executor.py b/python/packages/declarative/tests/test_http_request_executor.py index effdc5b298c..4030cf42946 100644 --- a/python/packages/declarative/tests/test_http_request_executor.py +++ b/python/packages/declarative/tests/test_http_request_executor.py @@ -72,9 +72,7 @@ def _ok(body: str = "", headers: dict[str, list[str]] | None = None) -> HttpRequ ) -def _err( - status: int = 500, body: str = "", headers: dict[str, list[str]] | None = None -) -> HttpRequestResult: +def _err(status: int = 500, body: str = "", headers: dict[str, list[str]] | None = None) -> HttpRequestResult: return HttpRequestResult( status_code=status, is_success_status_code=False, @@ -150,9 +148,7 @@ class TestSuccessPath: async def test_get_parses_json_object(self) -> None: handler = StubHandler(_ok('{"key":"value","number":42}')) factory = WorkflowFactory(http_request_handler=handler) - workflow = factory.create_workflow_from_definition( - _yaml(_action(method="GET", response="Local.Result")) - ) + workflow = factory.create_workflow_from_definition(_yaml(_action(method="GET", response="Local.Result"))) await workflow.run({}) decl = workflow._state.get(DECLARATIVE_STATE_KEY) @@ -165,9 +161,7 @@ async def test_get_parses_json_object(self) -> None: async def test_get_parses_plain_string(self) -> None: handler = StubHandler(_ok("not-json content")) factory = WorkflowFactory(http_request_handler=handler) - workflow = factory.create_workflow_from_definition( - _yaml(_action(response="Local.Result")) - ) + workflow = factory.create_workflow_from_definition(_yaml(_action(response="Local.Result"))) await workflow.run({}) decl = workflow._state.get(DECLARATIVE_STATE_KEY) @@ -177,9 +171,7 @@ async def test_get_parses_plain_string(self) -> None: async def test_get_empty_body_yields_none(self) -> None: handler = StubHandler(_ok("")) factory = WorkflowFactory(http_request_handler=handler) - workflow = factory.create_workflow_from_definition( - _yaml(_action(response="Local.Result")) - ) + workflow = factory.create_workflow_from_definition(_yaml(_action(response="Local.Result"))) await workflow.run({}) decl = workflow._state.get(DECLARATIVE_STATE_KEY) @@ -189,9 +181,7 @@ async def test_get_empty_body_yields_none(self) -> None: async def test_response_object_form_path(self) -> None: handler = StubHandler(_ok('{"x":1}')) factory = WorkflowFactory(http_request_handler=handler) - workflow = factory.create_workflow_from_definition( - _yaml(_action(response={"path": "Local.Result"})) - ) + workflow = factory.create_workflow_from_definition(_yaml(_action(response={"path": "Local.Result"}))) await workflow.run({}) decl = workflow._state.get(DECLARATIVE_STATE_KEY) @@ -376,9 +366,7 @@ async def test_long_form_body_kinds_accepted(self) -> None: async def test_unknown_body_kind_raises(self) -> None: handler = StubHandler(_ok()) factory = WorkflowFactory(http_request_handler=handler) - workflow = factory.create_workflow_from_definition( - _yaml(_action(body={"kind": "weirdform", "content": "x"})) - ) + workflow = factory.create_workflow_from_definition(_yaml(_action(body={"kind": "weirdform", "content": "x"}))) with pytest.raises(Exception) as excinfo: await workflow.run({}) # Should surface as ValueError (potentially wrapped by runner) @@ -527,9 +515,7 @@ async def test_response_headers_folded_with_commas(self) -> None: ) ) factory = WorkflowFactory(http_request_handler=handler) - workflow = factory.create_workflow_from_definition( - _yaml(_action(response_headers="Local.H")) - ) + workflow = factory.create_workflow_from_definition(_yaml(_action(response_headers="Local.H"))) await workflow.run({}) decl = workflow._state.get(DECLARATIVE_STATE_KEY) h = decl["Local"]["H"] @@ -540,9 +526,7 @@ async def test_response_headers_folded_with_commas(self) -> None: async def test_response_headers_empty_assigned_none(self) -> None: handler = StubHandler(_ok("ok", headers={})) factory = WorkflowFactory(http_request_handler=handler) - workflow = factory.create_workflow_from_definition( - _yaml(_action(response_headers="Local.H")) - ) + workflow = factory.create_workflow_from_definition(_yaml(_action(response_headers="Local.H"))) await workflow.run({}) decl = workflow._state.get(DECLARATIVE_STATE_KEY) assert decl["Local"]["H"] is None @@ -551,9 +535,7 @@ async def test_response_headers_empty_assigned_none(self) -> None: async def test_non_2xx_still_publishes_headers(self) -> None: handler = StubHandler(_err(status=500, body="boom", headers={"X-Trace": ["abc"]})) factory = WorkflowFactory(http_request_handler=handler) - workflow = factory.create_workflow_from_definition( - _yaml(_action(response_headers="Local.H")) - ) + workflow = factory.create_workflow_from_definition(_yaml(_action(response_headers="Local.H"))) with pytest.raises(DeclarativeActionError): await workflow.run({}) decl = workflow._state.get(DECLARATIVE_STATE_KEY) @@ -586,9 +568,7 @@ async def test_conversation_id_appends_message(self) -> None: async def test_empty_conversation_id_does_not_append(self) -> None: handler = StubHandler(_ok('{"answer":"hello"}')) factory = WorkflowFactory(http_request_handler=handler) - workflow = factory.create_workflow_from_definition( - _yaml(_action(response="Local.Result", conversation_id="")) - ) + workflow = factory.create_workflow_from_definition(_yaml(_action(response="Local.Result", conversation_id=""))) await workflow.run({}) decl = workflow._state.get(DECLARATIVE_STATE_KEY) # Auto-init creates an entry for the System.ConversationId conversation, @@ -600,9 +580,7 @@ async def test_empty_conversation_id_does_not_append(self) -> None: async def test_empty_body_skips_conversation_append(self) -> None: handler = StubHandler(_ok("")) factory = WorkflowFactory(http_request_handler=handler) - workflow = factory.create_workflow_from_definition( - _yaml(_action(conversation_id="conv-test-1")) - ) + workflow = factory.create_workflow_from_definition(_yaml(_action(conversation_id="conv-test-1"))) await workflow.run({}) decl = workflow._state.get(DECLARATIVE_STATE_KEY) # No conversation entry should have been created either. @@ -617,9 +595,7 @@ class TestConnection: async def test_connection_name_forwarded(self) -> None: handler = StubHandler(_ok()) factory = WorkflowFactory(http_request_handler=handler) - workflow = factory.create_workflow_from_definition( - _yaml(_action(connection={"name": "my-connection"})) - ) + workflow = factory.create_workflow_from_definition(_yaml(_action(connection={"name": "my-connection"}))) await workflow.run({}) assert handler.last_info is not None assert handler.last_info.connection_name == "my-connection" @@ -654,9 +630,7 @@ class TestTimeout: async def test_timeout_ms_forwarded(self) -> None: handler = StubHandler(_ok()) factory = WorkflowFactory(http_request_handler=handler) - workflow = factory.create_workflow_from_definition( - _yaml(_action(request_timeout_ms=2500)) - ) + workflow = factory.create_workflow_from_definition(_yaml(_action(request_timeout_ms=2500))) await workflow.run({}) assert handler.last_info is not None assert handler.last_info.timeout_ms == 2500 @@ -665,9 +639,7 @@ async def test_timeout_ms_forwarded(self) -> None: async def test_timeout_ms_zero_treated_as_unset(self) -> None: handler = StubHandler(_ok()) factory = WorkflowFactory(http_request_handler=handler) - workflow = factory.create_workflow_from_definition( - _yaml(_action(request_timeout_ms=0)) - ) + workflow = factory.create_workflow_from_definition(_yaml(_action(request_timeout_ms=0))) await workflow.run({}) assert handler.last_info is not None assert handler.last_info.timeout_ms is None From 3c91ba4050e43941dbbd3fd053cf022960622fcb Mon Sep 17 00:00:00 2001 From: Peter Ibekwe Date: Thu, 30 Apr 2026 19:00:25 -0700 Subject: [PATCH 3/8] Fix conversation ID dot parsing for http executor --- .../_workflows/_executors_http.py | 9 +++------ 1 file changed, 3 insertions(+), 6 deletions(-) diff --git a/python/packages/declarative/agent_framework_declarative/_workflows/_executors_http.py b/python/packages/declarative/agent_framework_declarative/_workflows/_executors_http.py index 3df7a2bca18..0a91748c6ba 100644 --- a/python/packages/declarative/agent_framework_declarative/_workflows/_executors_http.py +++ b/python/packages/declarative/agent_framework_declarative/_workflows/_executors_http.py @@ -405,12 +405,9 @@ def _append_response_to_conversation( messages_path = _get_messages_path(state, conversation_id_expr) if messages_path is None: return - # Ensure the conversation entry exists so downstream agents can reference it. - evaluated_id = messages_path.split(".", 2)[2].rsplit(".", 1)[0] - conversations: dict[str, Any] = state.get("System.conversations") or {} - if evaluated_id not in conversations: - conversations[evaluated_id] = {"id": evaluated_id, "messages": []} - state.set("System.conversations", conversations) + # Mirrors InvokeAzureAgentExecutor: rely on state.append to lazily + # create the conversation entry. Avoids re-parsing the id back out + # of the dotted path string. message = Message(role="assistant", contents=[body]) state.append(messages_path, message) From dca9dc081ba612cbb153c13e003b8ff8e69384d2 Mon Sep 17 00:00:00 2001 From: Peter Ibekwe Date: Fri, 1 May 2026 15:38:49 -0700 Subject: [PATCH 4/8] Removed unnecessary export command --- .../agent_framework_declarative/_workflows/__init__.py | 3 ++- .../agent_framework_declarative/_workflows/_errors.py | 8 ++------ .../agent_framework_declarative/_workflows/_factory.py | 6 ++---- .../packages/declarative/tests/test_workflow_factory.py | 6 ++---- 4 files changed, 8 insertions(+), 15 deletions(-) diff --git a/python/packages/declarative/agent_framework_declarative/_workflows/__init__.py b/python/packages/declarative/agent_framework_declarative/_workflows/__init__.py index a1d098a014e..c199e4551b7 100644 --- a/python/packages/declarative/agent_framework_declarative/_workflows/__init__.py +++ b/python/packages/declarative/agent_framework_declarative/_workflows/__init__.py @@ -25,6 +25,7 @@ LoopIterationResult, ) from ._declarative_builder import ALL_ACTION_EXECUTORS, DeclarativeWorkflowBuilder +from ._errors import DeclarativeActionError, DeclarativeWorkflowError from ._executors_agents import ( AGENT_ACTION_EXECUTORS, AGENT_REGISTRY_KEY, @@ -82,7 +83,7 @@ ToolApprovalState, ToolInvocationResult, ) -from ._factory import DeclarativeActionError, DeclarativeWorkflowError, WorkflowFactory +from ._factory import WorkflowFactory from ._http_handler import ( DefaultHttpRequestHandler, HttpRequestHandler, diff --git a/python/packages/declarative/agent_framework_declarative/_workflows/_errors.py b/python/packages/declarative/agent_framework_declarative/_workflows/_errors.py index 452203268f4..e3372ebf06b 100644 --- a/python/packages/declarative/agent_framework_declarative/_workflows/_errors.py +++ b/python/packages/declarative/agent_framework_declarative/_workflows/_errors.py @@ -6,12 +6,8 @@ ``_executors_http``, ``_declarative_builder``) can raise declarative-specific exceptions without importing from ``_factory``. ``_factory`` imports ``_declarative_builder`` which imports the executor modules; pulling -``DeclarativeWorkflowError`` from ``_factory`` into an executor or builder -module would therefore introduce a circular import. - -``_factory`` re-exports :class:`DeclarativeWorkflowError` (and -:class:`DeclarativeActionError`) for convenience so existing import paths keep -working. +:class:`DeclarativeWorkflowError` from ``_factory`` into an executor or +builder module would therefore introduce a circular import. """ from __future__ import annotations diff --git a/python/packages/declarative/agent_framework_declarative/_workflows/_factory.py b/python/packages/declarative/agent_framework_declarative/_workflows/_factory.py index 36f088b05fd..d1e21d76e98 100644 --- a/python/packages/declarative/agent_framework_declarative/_workflows/_factory.py +++ b/python/packages/declarative/agent_framework_declarative/_workflows/_factory.py @@ -27,15 +27,13 @@ from .._loader import AgentFactory from ._declarative_builder import DeclarativeWorkflowBuilder -from ._errors import DeclarativeActionError, DeclarativeWorkflowError +from ._errors import DeclarativeWorkflowError from ._http_handler import HttpRequestHandler logger = logging.getLogger("agent_framework.declarative") -# Re-export DeclarativeWorkflowError (now defined in _errors) so existing -# import paths (`from .._factory import DeclarativeWorkflowError`) keep working. -__all__ = ["DeclarativeActionError", "DeclarativeWorkflowError", "WorkflowFactory"] +__all__ = ["WorkflowFactory"] class WorkflowFactory: diff --git a/python/packages/declarative/tests/test_workflow_factory.py b/python/packages/declarative/tests/test_workflow_factory.py index e313f78799d..720bca3498b 100644 --- a/python/packages/declarative/tests/test_workflow_factory.py +++ b/python/packages/declarative/tests/test_workflow_factory.py @@ -4,10 +4,8 @@ import pytest -from agent_framework_declarative._workflows._factory import ( - DeclarativeWorkflowError, - WorkflowFactory, -) +from agent_framework_declarative._workflows._errors import DeclarativeWorkflowError +from agent_framework_declarative._workflows._factory import WorkflowFactory try: import powerfx # noqa: F401 From 55e665ade00cc3bdb2abf26f6adde660c9a96fc5 Mon Sep 17 00:00:00 2001 From: Peter Ibekwe Date: Mon, 4 May 2026 13:08:21 -0700 Subject: [PATCH 5/8] Initial implementation of invoke mcp tool in python --- .../agent_framework/declarative/__init__.py | 5 + .../agent_framework/declarative/__init__.pyi | 10 + python/packages/declarative/AGENTS.md | 1 + .../agent_framework_declarative/__init__.py | 10 + .../_workflows/__init__.py | 18 + .../_workflows/_declarative_builder.py | 22 + .../_workflows/_executors_mcp.py | 611 +++++++++++++++++ .../_workflows/_factory.py | 11 + .../_workflows/_mcp_handler.py | 420 ++++++++++++ .../tests/test_default_mcp_tool_handler.py | 415 ++++++++++++ .../tests/test_invoke_mcp_tool_executor.py | 631 ++++++++++++++++++ .../declarative/invoke_mcp_tool/main.py | 100 +++ .../declarative/invoke_mcp_tool/workflow.yaml | 64 ++ 13 files changed, 2318 insertions(+) create mode 100644 python/packages/declarative/agent_framework_declarative/_workflows/_executors_mcp.py create mode 100644 python/packages/declarative/agent_framework_declarative/_workflows/_mcp_handler.py create mode 100644 python/packages/declarative/tests/test_default_mcp_tool_handler.py create mode 100644 python/packages/declarative/tests/test_invoke_mcp_tool_executor.py create mode 100644 python/samples/03-workflows/declarative/invoke_mcp_tool/main.py create mode 100644 python/samples/03-workflows/declarative/invoke_mcp_tool/workflow.yaml diff --git a/python/packages/core/agent_framework/declarative/__init__.py b/python/packages/core/agent_framework/declarative/__init__.py index ba88e6a0a98..90c73ef8bd5 100644 --- a/python/packages/core/agent_framework/declarative/__init__.py +++ b/python/packages/core/agent_framework/declarative/__init__.py @@ -25,11 +25,16 @@ "DeclarativeLoaderError", "DeclarativeWorkflowError", "DefaultHttpRequestHandler", + "DefaultMCPToolHandler", "ExternalInputRequest", "ExternalInputResponse", "HttpRequestHandler", "HttpRequestInfo", "HttpRequestResult", + "MCPToolApprovalRequest", + "MCPToolHandler", + "MCPToolInvocation", + "MCPToolResult", "ProviderLookupError", "ProviderTypeMapping", "WorkflowFactory", diff --git a/python/packages/core/agent_framework/declarative/__init__.pyi b/python/packages/core/agent_framework/declarative/__init__.pyi index f18be22f50c..bd6bf73fba6 100644 --- a/python/packages/core/agent_framework/declarative/__init__.pyi +++ b/python/packages/core/agent_framework/declarative/__init__.pyi @@ -8,11 +8,16 @@ from agent_framework_declarative import ( DeclarativeLoaderError, DeclarativeWorkflowError, DefaultHttpRequestHandler, + DefaultMCPToolHandler, ExternalInputRequest, ExternalInputResponse, HttpRequestHandler, HttpRequestInfo, HttpRequestResult, + MCPToolApprovalRequest, + MCPToolHandler, + MCPToolInvocation, + MCPToolResult, ProviderLookupError, ProviderTypeMapping, WorkflowFactory, @@ -27,11 +32,16 @@ __all__ = [ "DeclarativeLoaderError", "DeclarativeWorkflowError", "DefaultHttpRequestHandler", + "DefaultMCPToolHandler", "ExternalInputRequest", "ExternalInputResponse", "HttpRequestHandler", "HttpRequestInfo", "HttpRequestResult", + "MCPToolApprovalRequest", + "MCPToolHandler", + "MCPToolInvocation", + "MCPToolResult", "ProviderLookupError", "ProviderTypeMapping", "WorkflowFactory", diff --git a/python/packages/declarative/AGENTS.md b/python/packages/declarative/AGENTS.md index 1add6146015..3c9402fc4e5 100644 --- a/python/packages/declarative/AGENTS.md +++ b/python/packages/declarative/AGENTS.md @@ -9,6 +9,7 @@ YAML/JSON-based declarative agent and workflow definitions. - **`WorkflowState`** - State management for declarative workflows - **`ProviderTypeMapping`** - Maps provider types to implementations - **`HttpRequestHandler`** / **`DefaultHttpRequestHandler`** - Pluggable HTTP transport for the `HttpRequestAction` declarative action (configured via `WorkflowFactory(http_request_handler=...)`) +- **`MCPToolHandler`** / **`DefaultMCPToolHandler`** - Pluggable MCP transport for the `InvokeMcpTool` declarative action (configured via `WorkflowFactory(mcp_tool_handler=...)`) - **`DeclarativeLoaderError`** / **`ProviderLookupError`** / **`DeclarativeWorkflowError`** / **`DeclarativeActionError`** - Error types ## External Input Handling diff --git a/python/packages/declarative/agent_framework_declarative/__init__.py b/python/packages/declarative/agent_framework_declarative/__init__.py index 6afcb3c7912..ad639fb5217 100644 --- a/python/packages/declarative/agent_framework_declarative/__init__.py +++ b/python/packages/declarative/agent_framework_declarative/__init__.py @@ -9,11 +9,16 @@ DeclarativeActionError, DeclarativeWorkflowError, DefaultHttpRequestHandler, + DefaultMCPToolHandler, ExternalInputRequest, ExternalInputResponse, HttpRequestHandler, HttpRequestInfo, HttpRequestResult, + MCPToolApprovalRequest, + MCPToolHandler, + MCPToolInvocation, + MCPToolResult, WorkflowFactory, WorkflowState, ) @@ -31,11 +36,16 @@ "DeclarativeLoaderError", "DeclarativeWorkflowError", "DefaultHttpRequestHandler", + "DefaultMCPToolHandler", "ExternalInputRequest", "ExternalInputResponse", "HttpRequestHandler", "HttpRequestInfo", "HttpRequestResult", + "MCPToolApprovalRequest", + "MCPToolHandler", + "MCPToolInvocation", + "MCPToolResult", "ProviderLookupError", "ProviderTypeMapping", "WorkflowFactory", diff --git a/python/packages/declarative/agent_framework_declarative/_workflows/__init__.py b/python/packages/declarative/agent_framework_declarative/_workflows/__init__.py index c199e4551b7..d06fdeba175 100644 --- a/python/packages/declarative/agent_framework_declarative/_workflows/__init__.py +++ b/python/packages/declarative/agent_framework_declarative/_workflows/__init__.py @@ -72,6 +72,11 @@ HTTP_ACTION_EXECUTORS, HttpRequestActionExecutor, ) +from ._executors_mcp import ( + MCP_ACTION_EXECUTORS, + InvokeMcpToolActionExecutor, + MCPToolApprovalRequest, +) from ._executors_tools import ( FUNCTION_TOOL_REGISTRY_KEY, TOOL_ACTION_EXECUTORS, @@ -90,6 +95,12 @@ HttpRequestInfo, HttpRequestResult, ) +from ._mcp_handler import ( + DefaultMCPToolHandler, + MCPToolHandler, + MCPToolInvocation, + MCPToolResult, +) from ._state import WorkflowState __all__ = [ @@ -102,6 +113,7 @@ "EXTERNAL_INPUT_EXECUTORS", "FUNCTION_TOOL_REGISTRY_KEY", "HTTP_ACTION_EXECUTORS", + "MCP_ACTION_EXECUTORS", "TOOL_ACTION_EXECUTORS", "TOOL_APPROVAL_STATE_KEY", "TOOL_REGISTRY_KEY", @@ -126,6 +138,7 @@ "DeclarativeWorkflowError", "DeclarativeWorkflowState", "DefaultHttpRequestHandler", + "DefaultMCPToolHandler", "EmitEventExecutor", "EndConversationExecutor", "EndWorkflowExecutor", @@ -140,9 +153,14 @@ "HttpRequestResult", "InvokeAzureAgentExecutor", "InvokeFunctionToolExecutor", + "InvokeMcpToolActionExecutor", "JoinExecutor", "LoopControl", "LoopIterationResult", + "MCPToolApprovalRequest", + "MCPToolHandler", + "MCPToolInvocation", + "MCPToolResult", "QuestionExecutor", "RequestExternalInputExecutor", "ResetVariableExecutor", diff --git a/python/packages/declarative/agent_framework_declarative/_workflows/_declarative_builder.py b/python/packages/declarative/agent_framework_declarative/_workflows/_declarative_builder.py index fb5dcb88f87..67b4a582734 100644 --- a/python/packages/declarative/agent_framework_declarative/_workflows/_declarative_builder.py +++ b/python/packages/declarative/agent_framework_declarative/_workflows/_declarative_builder.py @@ -41,8 +41,10 @@ ) from ._executors_external_input import EXTERNAL_INPUT_EXECUTORS from ._executors_http import HTTP_ACTION_EXECUTORS, HttpRequestActionExecutor +from ._executors_mcp import MCP_ACTION_EXECUTORS, InvokeMcpToolActionExecutor from ._executors_tools import TOOL_ACTION_EXECUTORS, InvokeFunctionToolExecutor from ._http_handler import HttpRequestHandler +from ._mcp_handler import MCPToolHandler logger = logging.getLogger(__name__) @@ -55,6 +57,7 @@ **EXTERNAL_INPUT_EXECUTORS, **TOOL_ACTION_EXECUTORS, **HTTP_ACTION_EXECUTORS, + **MCP_ACTION_EXECUTORS, } # Action kinds that terminate control flow (no fall-through to successor) @@ -90,6 +93,7 @@ "EmitEvent": ["event"], "InvokeFunctionTool": ["functionName"], "HttpRequestAction": ["url"], + "InvokeMcpTool": ["serverUrl", "toolName"], } # Alternate field names that satisfy required field requirements @@ -135,6 +139,7 @@ def __init__( validate: bool = True, max_iterations: int | None = None, http_request_handler: HttpRequestHandler | None = None, + mcp_tool_handler: MCPToolHandler | None = None, ): """Initialize the builder. @@ -150,6 +155,9 @@ def __init__( http_request_handler: Handler used to dispatch HttpRequestAction requests. Must be supplied when the workflow contains any HttpRequestAction; otherwise build raises ``DeclarativeWorkflowError``. + mcp_tool_handler: Handler used to dispatch InvokeMcpTool calls. + Must be supplied when the workflow contains any InvokeMcpTool; + otherwise build raises ``DeclarativeWorkflowError``. """ self._yaml_def = yaml_definition self._workflow_id = workflow_id or yaml_definition.get("name", "declarative_workflow") @@ -162,6 +170,7 @@ def __init__( self._validate = validate self._seen_explicit_ids: set[str] = set() # Track explicit IDs for duplicate detection self._http_request_handler = http_request_handler + self._mcp_tool_handler = mcp_tool_handler # Resolve max_iterations: explicit arg > YAML maxTurns > core default resolved = max_iterations if max_iterations is not None else yaml_definition.get("maxTurns") if resolved is not None and (not isinstance(resolved, int) or resolved <= 0): @@ -481,6 +490,19 @@ def _create_executor_for_action( id=action_id, http_request_handler=self._http_request_handler, ) + elif kind == "InvokeMcpTool": + if self._mcp_tool_handler is None: + raise DeclarativeWorkflowError( + f"Workflow defines InvokeMcpTool '{action_id}' but no " + "mcp_tool_handler was supplied to WorkflowFactory. Pass " + "mcp_tool_handler=DefaultMCPToolHandler() (or a custom " + "implementation) to enable MCP tool invocations." + ) + executor = InvokeMcpToolActionExecutor( + action_def, + id=action_id, + mcp_tool_handler=self._mcp_tool_handler, + ) else: executor = executor_class(action_def, id=action_id) self._executors[action_id] = executor diff --git a/python/packages/declarative/agent_framework_declarative/_workflows/_executors_mcp.py b/python/packages/declarative/agent_framework_declarative/_workflows/_executors_mcp.py new file mode 100644 index 00000000000..a1501a3cb16 --- /dev/null +++ b/python/packages/declarative/agent_framework_declarative/_workflows/_executors_mcp.py @@ -0,0 +1,611 @@ +# Copyright (c) Microsoft. All rights reserved. + +"""Executor for the ``InvokeMcpTool`` declarative action. + +Mirrors the .NET ``InvokeMcpToolExecutor``: dispatches an MCP tool call through +the configured :class:`MCPToolHandler`, parses tool outputs, and routes +results to the configured ``output.{result, messages, autoSend}`` paths and +optional conversation history. Supports a human-in-loop approval flow via +``ctx.request_info()`` / :func:`@response_handler` for ``requireApproval=true``. + +Security notes: + +- The executor never echoes header VALUES (auth tokens, API keys) into the + approval request — only header NAMES are surfaced to the caller. This + matches the security posture of :mod:`._executors_http` (which never logs + request headers either) and prevents secrets from leaking through workflow + events that are typically observable to operators / UIs. +- ``_MCPToolApprovalState`` snapshots the EVALUATED values for non-secret + fields (server URL, tool name, arguments) at approval-request time so that + subsequent state mutations cannot make the executor "approve X then call + Y". Headers are stored as the raw expression strings (not evaluated values) + so secrets are not persisted in the workflow's checkpoint state. They are + re-evaluated on resume. +- Tool outputs flow back into agent conversations through ``conversationId`` + and through Tool-role messages emitted to ``output.messages``. They share + the same prompt-injection risk surface as ``HttpRequestAction``: workflow + authors must trust the MCP server they invoke. +""" + +import json +import logging +import uuid +from collections.abc import Mapping +from dataclasses import dataclass, field +from typing import Any + +import httpx +from agent_framework import ( + Content, + Message, + WorkflowContext, + handler, + response_handler, +) +from agent_framework.exceptions import ToolExecutionException + +from ._declarative_base import ( + ActionComplete, + DeclarativeActionExecutor, + DeclarativeWorkflowState, +) +from ._executors_tools import ToolApprovalResponse +from ._mcp_handler import MCPToolHandler, MCPToolInvocation, MCPToolResult + +__all__ = [ + "MCP_ACTION_EXECUTORS", + "InvokeMcpToolActionExecutor", + "MCPToolApprovalRequest", +] + +logger = logging.getLogger(__name__) + +_MCP_APPROVAL_STATE_KEY = "_mcp_tool_approval_state" + + +# --------------------------------------------------------------------------- +# Request / state types +# --------------------------------------------------------------------------- + + +@dataclass +class MCPToolApprovalRequest: + """Approval request emitted before invoking an MCP tool. + + Mirrors :class:`agent_framework_declarative.ToolApprovalRequest` but for + MCP-style invocations. Only header NAMES are surfaced — header values are + intentionally omitted because they typically carry authentication + secrets. + + Attributes: + request_id: Unique identifier for this approval request. Matches the + id workflow event-emitters use. + tool_name: Evaluated name of the tool to be invoked. + server_url: Evaluated MCP server URL. + server_label: Optional human-readable label for diagnostics. + arguments: Evaluated arguments to be forwarded to the tool. + header_names: Sorted list of outbound header names (no values). Empty + when no headers are configured. + """ + + request_id: str + tool_name: str + server_url: str + server_label: str | None + arguments: dict[str, Any] + header_names: list[str] = field(default_factory=lambda: []) + + +@dataclass +class _MCPToolApprovalState: + """Internal state saved during the approval yield for resumption. + + Stores **evaluated** values for non-secret fields to prevent + "approve X / execute Y" attacks. Stores the raw expression string for + ``headers`` so that secret values are NOT persisted in checkpoint state; + the expressions are re-evaluated against current state on resume. + """ + + server_url: str + tool_name: str + server_label: str | None + arguments: dict[str, Any] + connection_name: str | None + headers_def: Any + auto_send: bool + conversation_id_expr: str | None + output_messages_path: str | None + output_result_path: str | None + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + + +def _get_messages_path(state: DeclarativeWorkflowState, conversation_id_expr: str | None) -> str | None: + """Return the configured conversation messages path, if any. + + Returns ``System.conversations.{evaluated_id}.messages`` when a + ``conversation_id_expr`` is configured and evaluates to a non-empty value. + Returns ``None`` when no conversation id expression is configured or when + the expression evaluates to ``None`` or an empty string (mirrors .NET + ``GetConversationId`` behaviour). + """ + if not conversation_id_expr: + return None + evaluated = state.eval_if_expression(conversation_id_expr) + if evaluated is None or (isinstance(evaluated, str) and not evaluated): + return None + return f"System.conversations.{evaluated}.messages" + + +def _get_output_path(action_def: Mapping[str, Any], key: str) -> str | None: + """Extract a state path from ``output.{key}`` field. + + Supports two YAML shapes: + + - ``output: { result: Local.MyVar }`` — plain string. + - ``output: { result: { path: Local.MyVar } }`` — object form. + """ + output: Any = action_def.get("output") + if not isinstance(output, Mapping): + return None + value: Any = output.get(key) # type: ignore[reportUnknownMemberType] + if isinstance(value, str): + return value or None + if isinstance(value, Mapping): + path: Any = value.get("path") # type: ignore[reportUnknownMemberType] + return path if isinstance(path, str) and path else None + return None + + +def _format_outputs_for_send(parsed_results: list[Any]) -> str: + """Render parsed MCP outputs to a string for ``ctx.yield_output(...)``. + + - All-string list → newline-joined. + - Single dict / list → JSON. + - Empty / mixed → JSON-dump the whole list. + """ + if not parsed_results: + return "" + if all(isinstance(item, str) for item in parsed_results): + return "\n".join(parsed_results) # type: ignore[arg-type] + if len(parsed_results) == 1 and isinstance(parsed_results[0], (dict, list)): + return json.dumps(parsed_results[0], ensure_ascii=False) + return json.dumps(parsed_results, ensure_ascii=False) + + +# --------------------------------------------------------------------------- +# Executor +# --------------------------------------------------------------------------- + + +class InvokeMcpToolActionExecutor(DeclarativeActionExecutor): + """Executor for the ``InvokeMcpTool`` declarative action. + + Dispatches through the supplied :class:`MCPToolHandler` and: + + - Evaluates ``serverUrl`` / ``toolName`` / ``serverLabel`` / ``arguments`` + / ``headers`` / ``connection.name`` from the action definition. + - When ``requireApproval=true``: emits a :class:`MCPToolApprovalRequest` + via ``ctx.request_info()`` and yields. On resume, the response is + checked; on rejection, ``output.result`` is set to ``"Error: ..."`` and + no tool call is made. + - On success: parses each :class:`agent_framework.Content` output (text → + JSON-first / data / uri → URI string) and assigns the parsed list to + ``output.result``. Builds a single Tool-role :class:`Message` + containing all output contents and assigns it to ``output.messages``. + When ``output.autoSend`` is true (default), emits the rendered string + via ``ctx.yield_output(...)``. When ``conversationId`` is configured, + appends an Assistant-role :class:`Message` with the same contents to + ``System.conversations.{id}.messages``. + - On error returned by the handler (``is_error=True``): assigns + ``"Error: "`` to ``output.result`` and completes normally + (parity with .NET ``AssignErrorAsync``). + + .. note:: + + ``output.messages`` receives a SINGLE Tool-role :class:`Message` + (containing the full tool output as ``contents``), unlike + :class:`agent_framework_declarative.InvokeFunctionToolExecutor` which + writes a list of two messages (assistant call + tool result). This + matches the .NET ``InvokeMcpToolExecutor`` output contract. + """ + + def __init__( + self, + action_def: dict[str, Any], + *, + id: str | None = None, + mcp_tool_handler: MCPToolHandler, + ) -> None: + """Create an MCP tool action executor. + + Args: + action_def: Parsed ``InvokeMcpTool`` YAML dict. + id: Optional executor id (defaults to action id or generated). + mcp_tool_handler: Handler used to dispatch MCP tool calls. + Required: the builder enforces presence at workflow-build + time. + """ + super().__init__(action_def, id=id) + self._mcp_tool_handler = mcp_tool_handler + + # ----- Main handler -------------------------------------------------------- + + @handler + async def handle_action( + self, + trigger: Any, + ctx: WorkflowContext[ActionComplete, str], + ) -> None: + """Execute the MCP tool action.""" + state = await self._ensure_state_initialized(ctx, trigger) + + server_url = self._get_server_url(state) + tool_name = self._get_tool_name(state) + server_label = self._get_server_label(state) + arguments = self._get_arguments(state) + headers = self._get_headers(state) + connection_name = self._get_connection_name(state) + require_approval = self._get_require_approval(state) + auto_send = self._get_auto_send(state) + conversation_id_expr = self._action_def.get("conversationId") + output_messages_path = _get_output_path(self._action_def, "messages") + output_result_path = _get_output_path(self._action_def, "result") + + if require_approval: + request_id = str(uuid.uuid4()) + approval_state = _MCPToolApprovalState( + server_url=server_url, + tool_name=tool_name, + server_label=server_label, + arguments=arguments, + connection_name=connection_name, + headers_def=self._action_def.get("headers"), + auto_send=auto_send, + conversation_id_expr=conversation_id_expr if isinstance(conversation_id_expr, str) else None, + output_messages_path=output_messages_path, + output_result_path=output_result_path, + ) + ctx.state.set(self._approval_key(), approval_state) + + request = MCPToolApprovalRequest( + request_id=request_id, + tool_name=tool_name, + server_url=server_url, + server_label=server_label, + arguments=arguments, + header_names=sorted(headers.keys()), + ) + logger.info( + "%s: requesting approval for MCP tool '%s' on '%s'", + self.__class__.__name__, + tool_name, + server_url, + ) + await ctx.request_info(request, ToolApprovalResponse, request_id=request_id) + # Workflow yields here — resume in handle_approval_response. + return + + # No approval required - invoke directly. + invocation = MCPToolInvocation( + server_url=server_url, + tool_name=tool_name, + server_label=server_label, + arguments=arguments, + headers=headers, + connection_name=connection_name, + ) + result = await self._invoke_with_narrow_catch(invocation) + await self._process_result( + ctx=ctx, + state=state, + result=result, + auto_send=auto_send, + conversation_id_expr=conversation_id_expr if isinstance(conversation_id_expr, str) else None, + output_messages_path=output_messages_path, + output_result_path=output_result_path, + ) + await ctx.send_message(ActionComplete()) + + # ----- Approval response handler ------------------------------------------ + + @response_handler + async def handle_approval_response( + self, + original_request: MCPToolApprovalRequest, + response: ToolApprovalResponse, + ctx: WorkflowContext[ActionComplete, str], + ) -> None: + """Resume after the workflow yielded for an approval request.""" + state = self._get_state(ctx.state) + approval_key = self._approval_key() + + try: + approval_state: _MCPToolApprovalState = ctx.state.get(approval_key) + except KeyError: + logger.error("%s: approval state missing for executor '%s'", self.__class__.__name__, self.id) + await ctx.send_message(ActionComplete()) + return + try: + ctx.state.delete(approval_key) + except KeyError: + logger.warning("%s: approval state already deleted for '%s'", self.__class__.__name__, self.id) + + if not response.approved: + logger.info( + "%s: MCP tool '%s' rejected: %s", + self.__class__.__name__, + approval_state.tool_name, + response.reason, + ) + self._assign_error( + state, approval_state.output_result_path, "MCP tool invocation was not approved by user." + ) + await ctx.send_message(ActionComplete()) + return + + # Approved — re-evaluate headers (not stored at approval time for security). + headers = self._evaluate_headers(state, approval_state.headers_def) + + invocation = MCPToolInvocation( + server_url=approval_state.server_url, + tool_name=approval_state.tool_name, + server_label=approval_state.server_label, + arguments=approval_state.arguments, + headers=headers, + connection_name=approval_state.connection_name, + ) + result = await self._invoke_with_narrow_catch(invocation) + await self._process_result( + ctx=ctx, + state=state, + result=result, + auto_send=approval_state.auto_send, + conversation_id_expr=approval_state.conversation_id_expr, + output_messages_path=approval_state.output_messages_path, + output_result_path=approval_state.output_result_path, + ) + await ctx.send_message(ActionComplete()) + + # ----- Field resolution ---------------------------------------------------- + + def _get_server_url(self, state: DeclarativeWorkflowState) -> str: + raw = self._action_def.get("serverUrl") + if raw is None: + raise ValueError("InvokeMcpTool requires a 'serverUrl' field.") + evaluated = state.eval_if_expression(raw) + if not isinstance(evaluated, str) or not evaluated: + raise ValueError("InvokeMcpTool 'serverUrl' evaluated to an empty value.") + return evaluated + + def _get_tool_name(self, state: DeclarativeWorkflowState) -> str: + raw = self._action_def.get("toolName") + if raw is None: + raise ValueError("InvokeMcpTool requires a 'toolName' field.") + evaluated = state.eval_if_expression(raw) + if not isinstance(evaluated, str) or not evaluated: + raise ValueError("InvokeMcpTool 'toolName' evaluated to an empty value.") + return evaluated + + def _get_server_label(self, state: DeclarativeWorkflowState) -> str | None: + raw = self._action_def.get("serverLabel") + if raw is None: + return None + evaluated = state.eval_if_expression(raw) + if evaluated is None: + return None + text = str(evaluated) + return text or None + + def _get_arguments(self, state: DeclarativeWorkflowState) -> dict[str, Any]: + """Evaluate ``arguments`` map. Preserves ``None`` values (parity with .NET).""" + raw = self._action_def.get("arguments") + if raw is None: + return {} + if not isinstance(raw, Mapping) or not raw: + return {} + result: dict[str, Any] = {} + for key, value in raw.items(): # type: ignore[reportUnknownVariableType] + if not isinstance(key, str) or not key: + continue + result[key] = state.eval_if_expression(value) + return result + + def _get_headers(self, state: DeclarativeWorkflowState) -> dict[str, str]: + return self._evaluate_headers(state, self._action_def.get("headers")) + + @staticmethod + def _evaluate_headers(state: DeclarativeWorkflowState, headers_def: Any) -> dict[str, str]: + """Evaluate the ``headers`` map. Empty string values are skipped.""" + if not isinstance(headers_def, Mapping) or not headers_def: + return {} + result: dict[str, str] = {} + for key, value in headers_def.items(): # type: ignore[reportUnknownVariableType] + if not isinstance(key, str) or not key: + continue + evaluated = state.eval_if_expression(value) + if evaluated is None: + continue + text = str(evaluated) + if not text: + continue + result[key] = text + return result + + def _get_connection_name(self, state: DeclarativeWorkflowState) -> str | None: + connection = self._action_def.get("connection") + if not isinstance(connection, Mapping): + return None + name_expr: Any = connection.get("name") # type: ignore[reportUnknownMemberType] + if name_expr is None: + return None + evaluated = state.eval_if_expression(name_expr) + if evaluated is None: + return None + text = str(evaluated) + return text or None + + def _get_require_approval(self, state: DeclarativeWorkflowState) -> bool: + raw = self._action_def.get("requireApproval") + if raw is None: + return False + evaluated = state.eval_if_expression(raw) + if isinstance(evaluated, bool): + return evaluated + if isinstance(evaluated, str): + return evaluated.strip().lower() in {"true", "1", "yes"} + return bool(evaluated) + + def _get_auto_send(self, state: DeclarativeWorkflowState) -> bool: + output: Any = self._action_def.get("output") + if not isinstance(output, Mapping): + return True + raw: Any = output.get("autoSend") # type: ignore[reportUnknownMemberType] + if raw is None: + return True + evaluated = state.eval_if_expression(raw) + if isinstance(evaluated, bool): + return evaluated + if isinstance(evaluated, str): + return evaluated.strip().lower() in {"true", "1", "yes"} + return bool(evaluated) + + # ----- Invocation + error handling ---------------------------------------- + + async def _invoke_with_narrow_catch(self, invocation: MCPToolInvocation) -> MCPToolResult: + """Invoke the handler with a narrow exception catch. + + Only known transport / tool exceptions are normalised to an error + result. Programmer bugs (TypeError, ValueError from misuse, etc.) + propagate so they fail loudly. + + ``asyncio.CancelledError`` is a ``BaseException``, not ``Exception``, + so it is not caught here and propagates unchanged for workflow + cancellation. + """ + try: + return await self._mcp_tool_handler.invoke_tool(invocation) + except ToolExecutionException as exc: + message = str(exc) or type(exc).__name__ + return MCPToolResult( + outputs=[Content.from_text(f"Error: {message}")], + is_error=True, + error_message=message, + ) + except httpx.HTTPError as exc: + message = f"{type(exc).__name__}: {exc}" if str(exc) else type(exc).__name__ + return MCPToolResult( + outputs=[Content.from_text(f"Error: {message}")], + is_error=True, + error_message=message, + ) + except Exception as exc: + try: + from mcp.shared.exceptions import McpError + except ImportError: # pragma: no cover - mcp is a hard dep + raise + if isinstance(exc, McpError): + message = str(exc) or type(exc).__name__ + return MCPToolResult( + outputs=[Content.from_text(f"Error: {message}")], + is_error=True, + error_message=message, + ) + raise + + # ----- Result handling ----------------------------------------------------- + + async def _process_result( + self, + *, + ctx: WorkflowContext[ActionComplete, str], + state: DeclarativeWorkflowState, + result: MCPToolResult, + auto_send: bool, + conversation_id_expr: str | None, + output_messages_path: str | None, + output_result_path: str | None, + ) -> None: + """Apply ``result`` to workflow state per the configured output paths.""" + if result.is_error: + # Error path mirrors .NET ``AssignErrorAsync`` — only the result + # path is touched; messages / autoSend / conversation are not. + self._assign_error( + state, + output_result_path, + result.error_message or "MCP tool invocation failed.", + ) + return + + parsed_results = _parse_outputs(result.outputs) + if output_result_path is not None and parsed_results: + state.set(output_result_path, parsed_results) + + # Single Tool-role message (matches .NET line 178 contract). Differs + # from InvokeFunctionTool's two-message [assistant call, tool result] + # convention. + tool_message = Message(role="tool", contents=list(result.outputs)) + if output_messages_path is not None: + state.set(output_messages_path, tool_message) + + if auto_send and parsed_results: + await ctx.yield_output(_format_outputs_for_send(parsed_results)) + + if conversation_id_expr: + messages_path = _get_messages_path(state, conversation_id_expr) + if messages_path is not None: + # Mirrors .NET: conversation gets ASSISTANT-role message with + # the same outputs (so chat history reads it as the agent's + # contribution). + assistant_message = Message(role="assistant", contents=list(result.outputs)) + state.append(messages_path, assistant_message) + + @staticmethod + def _assign_error( + state: DeclarativeWorkflowState, + output_result_path: str | None, + error_message: str, + ) -> None: + """Mirror .NET ``AssignErrorAsync``: store ``"Error: "`` at the result path.""" + if output_result_path is None: + return + state.set(output_result_path, f"Error: {error_message}") + + def _approval_key(self) -> str: + return f"{_MCP_APPROVAL_STATE_KEY}_{self.id}" + + +def _parse_outputs(outputs: list[Content]) -> list[Any]: + """Parse :class:`Content` outputs into Python values for ``output.result``. + + Mirrors .NET ``AssignResultAsync``: + + - ``TextContent`` → JSON-parse text; on failure use the raw text. + - ``DataContent`` / ``UriContent`` → ``content.uri``. + - Other content kinds → ``str(content)``. + """ + parsed: list[Any] = [] + for content in outputs: + kind = getattr(content, "type", None) + if kind == "text": + text_value = getattr(content, "text", None) + text_str = "" if text_value is None else str(text_value) + try: + parsed.append(json.loads(text_str)) + except (json.JSONDecodeError, ValueError): + parsed.append(text_str) + continue + if kind in ("data", "uri"): + uri_value = getattr(content, "uri", None) + parsed.append("" if uri_value is None else str(uri_value)) + continue + parsed.append(str(content)) + return parsed + + +MCP_ACTION_EXECUTORS: dict[str, type[DeclarativeActionExecutor]] = { + "InvokeMcpTool": InvokeMcpToolActionExecutor, +} diff --git a/python/packages/declarative/agent_framework_declarative/_workflows/_factory.py b/python/packages/declarative/agent_framework_declarative/_workflows/_factory.py index d1e21d76e98..221dfec3cc4 100644 --- a/python/packages/declarative/agent_framework_declarative/_workflows/_factory.py +++ b/python/packages/declarative/agent_framework_declarative/_workflows/_factory.py @@ -29,6 +29,7 @@ from ._declarative_builder import DeclarativeWorkflowBuilder from ._errors import DeclarativeWorkflowError from ._http_handler import HttpRequestHandler +from ._mcp_handler import MCPToolHandler logger = logging.getLogger("agent_framework.declarative") @@ -91,6 +92,7 @@ def __init__( checkpoint_storage: CheckpointStorage | None = None, max_iterations: int | None = None, http_request_handler: HttpRequestHandler | None = None, + mcp_tool_handler: MCPToolHandler | None = None, ) -> None: """Initialize the workflow factory. @@ -110,6 +112,13 @@ def __init__( otherwise. Use :class:`agent_framework.declarative.DefaultHttpRequestHandler` for a no-policy ``httpx``-based default, or supply your own implementation to enforce SSRF guards, allowlisting, or auth resolution. + mcp_tool_handler: Optional handler used to dispatch MCP tool calls for + ``InvokeMcpTool``. Required if the workflow contains any + ``InvokeMcpTool``; build will fail with :class:`DeclarativeWorkflowError` + otherwise. Use :class:`agent_framework.declarative.DefaultMCPToolHandler` + for a default backed by :class:`agent_framework.MCPStreamableHTTPTool`, + or supply your own implementation to enforce SSRF guards, allowlisting, + or auth/connection resolution. Examples: .. code-block:: python @@ -150,6 +159,7 @@ def __init__( self._checkpoint_storage = checkpoint_storage self._max_iterations = max_iterations self._http_request_handler = http_request_handler + self._mcp_tool_handler = mcp_tool_handler def create_workflow_from_yaml_path( self, @@ -394,6 +404,7 @@ def _create_workflow( checkpoint_storage=self._checkpoint_storage, max_iterations=self._max_iterations, http_request_handler=self._http_request_handler, + mcp_tool_handler=self._mcp_tool_handler, ) workflow = graph_builder.build() except ValueError as e: diff --git a/python/packages/declarative/agent_framework_declarative/_workflows/_mcp_handler.py b/python/packages/declarative/agent_framework_declarative/_workflows/_mcp_handler.py new file mode 100644 index 00000000000..5ea17aabc1b --- /dev/null +++ b/python/packages/declarative/agent_framework_declarative/_workflows/_mcp_handler.py @@ -0,0 +1,420 @@ +# Copyright (c) Microsoft. All rights reserved. + +"""MCP tool handler abstraction for declarative workflows. + +Mirrors the .NET ``IMcpToolHandler`` / ``DefaultMcpToolHandler`` pair from +``Microsoft.Agents.AI.Workflows.Declarative.Mcp``. Provides: + +- :class:`MCPToolInvocation` — request input data passed from the executor. +- :class:`MCPToolResult` — response data returned to the executor. +- :class:`MCPToolHandler` — :class:`typing.Protocol` callers implement to plug + in custom transports (e.g. with allowlisting, Foundry connection resolution, + per-server auth, etc.). +- :class:`DefaultMCPToolHandler` — production-grade default backed by + :class:`agent_framework.MCPStreamableHTTPTool`. + +Security note: :class:`DefaultMCPToolHandler` performs **no** URL filtering or +SSRF protection. Production deployments should supply a custom handler that +enforces an allowlist or DNS-rebinding-resistant policy. This split mirrors the +.NET design. + +Prompt-injection note: MCP tool outputs flow back into agent conversations +(via ``conversationId`` and Tool-role messages emitted by the executor) so +they share the same risk surface as ``HttpRequestAction``. Workflow authors +must trust the MCP server they invoke. +""" + +from __future__ import annotations + +import asyncio +import hashlib +import json +import logging +from collections import OrderedDict +from collections.abc import Awaitable, Callable +from dataclasses import dataclass, field +from typing import TYPE_CHECKING, Any, Protocol, cast, runtime_checkable + +import httpx + +if TYPE_CHECKING: + from agent_framework import Content + +__all__ = [ + "ClientProvider", + "DefaultMCPToolHandler", + "MCPToolHandler", + "MCPToolInvocation", + "MCPToolResult", +] + +logger = logging.getLogger(__name__) + +_DEFAULT_CACHE_MAX_SIZE = 32 + + +@dataclass +class MCPToolInvocation: + """Description of an MCP tool call to be dispatched by a :class:`MCPToolHandler`. + + Mirrors the input parameters of the .NET ``IMcpToolHandler.InvokeToolAsync`` + method. Field semantics: + + - ``server_url``: Absolute URL of the MCP server. Already evaluated from + the YAML expression. + - ``server_label``: Optional human-readable label used for diagnostics + and as the underlying ``MCPStreamableHTTPTool`` name. + - ``tool_name``: Name of the tool to invoke on the MCP server. + - ``arguments``: Tool arguments. Already evaluated; values may be any + JSON-serialisable Python object (str, int, bool, dict, list, None). + - ``headers``: Outbound HTTP headers (e.g. authentication). Empty values + are skipped by the executor before construction. + - ``connection_name``: Optional Foundry connection name forwarded for + handlers that resolve auth/credentials by connection. The default + handler does not consume this field. + """ + + server_url: str + tool_name: str + server_label: str | None = None + arguments: dict[str, Any] = field(default_factory=dict) # type: ignore[reportUnknownVariableType] + headers: dict[str, str] = field(default_factory=dict) # type: ignore[reportUnknownVariableType] + connection_name: str | None = None + + +def _empty_outputs() -> list[Any]: + """Default factory for ``MCPToolResult.outputs``. + + Typed as ``list[Any]`` here to keep the dataclass field's runtime + factory simple; the public type on :class:`MCPToolResult` is + ``list[Content]``. + """ + return [] + + +@dataclass +class MCPToolResult: + """Response returned by an :class:`MCPToolHandler`. + + Mirrors the .NET ``McpServerToolResultContent`` shape. ``outputs`` is a + list of :class:`agent_framework.Content` items as parsed by the MCP + transport (TextContent / DataContent / UriContent / etc.). + + On error, ``is_error`` is ``True``, ``error_message`` carries a human + readable description, and ``outputs`` typically contains a single + ``Content.from_text("Error: ...")`` entry for downstream display. + """ + + outputs: list[Content] = field(default_factory=_empty_outputs) + is_error: bool = False + error_message: str | None = None + + +@runtime_checkable +class MCPToolHandler(Protocol): + """Protocol for MCP tool handlers used by ``InvokeMcpTool``. + + Mirrors :class:`HttpRequestHandler` — declares ONLY the invocation method. + Lifecycle methods (``aclose`` / ``__aenter__`` / ``__aexit__``) are NOT + part of the Protocol; concrete implementations may add them as + appropriate. + + Implementations must be safe to call concurrently from multiple workflow + runs. Implementations are responsible for any URL allowlisting, SSRF + guards, retry policies, auth resolution, and other policies the workflow + author wants applied. + """ + + async def invoke_tool(self, invocation: MCPToolInvocation) -> MCPToolResult: + """Dispatch ``invocation`` and return the result. + + Args: + invocation: Description of the MCP tool call to perform. + + Returns: + The :class:`MCPToolResult` carrying the parsed outputs (or an + error flag if the tool raised). Implementations SHOULD return a + result with ``is_error=True`` rather than raising for transport + or tool-level failures, so the workflow can store the message in + ``output.result`` (matching .NET ``AssignErrorAsync`` behaviour). + They MAY raise on unexpected programming errors — these will be + propagated unchanged by the executor so they fail loudly. + """ + ... + + +ClientProvider = Callable[[MCPToolInvocation], Awaitable["httpx.AsyncClient | None"]] + + +@dataclass +class _CacheEntry: + """Internal record stored in the LRU cache.""" + + tool: Any # MCPStreamableHTTPTool — typed Any to avoid import at module load + owned_httpx_client: httpx.AsyncClient | None + + +class DefaultMCPToolHandler: + """Default :class:`MCPToolHandler` backed by :class:`agent_framework.MCPStreamableHTTPTool`. + + Caches one :class:`agent_framework.MCPStreamableHTTPTool` instance per + ``(server_url, headers_hash)`` in a bounded LRU. The cache prevents + re-establishing an MCP session for every invocation while ensuring + different header sets (auth tokens) cannot share a session — matches the + .NET design intent while bounding cardinality. + + Construction modes: + + 1. ``DefaultMCPToolHandler()`` — owns its own ``httpx.AsyncClient`` + instances created lazily per cache entry. Closed by :meth:`aclose`. + 2. ``DefaultMCPToolHandler(client_provider=cb)`` — per-server client + lookup (parity with .NET ``httpClientProvider`` callback). The + callback receives the full :class:`MCPToolInvocation` so it can + dispatch on ``server_url`` / ``connection_name`` / ``server_label``. + Returning ``None`` falls back to an internally-created client. Caller + supplied clients are NOT closed by :meth:`aclose`. + + .. warning:: + + This handler performs **no** URL filtering or SSRF protection. Wrap + or replace it with a custom handler in production deployments. + + Args: + client_provider: Optional per-server ``httpx.AsyncClient`` provider. + cache_max_size: Maximum number of cached MCP clients. When exceeded, + the least-recently-used entry is evicted and its client closed + (only owned clients are closed; caller-supplied ones are not). + Defaults to ``32``. + """ + + def __init__( + self, + *, + client_provider: ClientProvider | None = None, + cache_max_size: int = _DEFAULT_CACHE_MAX_SIZE, + ) -> None: + if cache_max_size <= 0: + raise ValueError(f"cache_max_size must be positive, got {cache_max_size}") + self._client_provider = client_provider + self._cache_max_size = cache_max_size + self._cache: OrderedDict[tuple[str, str], _CacheEntry] = OrderedDict() + # Outer lock guards the cache + in-flight-future map only — never + # held across network I/O. + self._cache_lock = asyncio.Lock() + # Per-key in-flight futures: while one task is connecting, other + # tasks awaiting the same key will await the same future and share + # the resulting cache entry. + self._inflight: dict[tuple[str, str], asyncio.Future[_CacheEntry]] = {} + + async def invoke_tool(self, invocation: MCPToolInvocation) -> MCPToolResult: + """Invoke ``invocation.tool_name`` on the cached MCP client for the server.""" + from agent_framework import Content + from agent_framework.exceptions import ToolExecutionException + + try: + entry = await self._get_or_create_entry(invocation) + except Exception as exc: + # Connect / cache lookup failures surface as tool errors so the + # workflow can store them at output.result without crashing. + logger.warning( + "DefaultMCPToolHandler: failed to obtain MCP client for url=%s tool=%s: %s", + invocation.server_url, + invocation.tool_name, + exc, + ) + message = f"Failed to connect to MCP server: {type(exc).__name__}: {exc}".rstrip(": ") + return MCPToolResult( + outputs=[Content.from_text(f"Error: {message}")], + is_error=True, + error_message=message, + ) + + try: + raw = await entry.tool.call_tool(invocation.tool_name, **invocation.arguments) + except ToolExecutionException as exc: + logger.info( + "DefaultMCPToolHandler: tool '%s' on '%s' raised ToolExecutionException", + invocation.tool_name, + invocation.server_url, + ) + message = str(exc) or type(exc).__name__ + return MCPToolResult( + outputs=[Content.from_text(f"Error: {message}")], + is_error=True, + error_message=message, + ) + except httpx.HTTPError as exc: + message = f"{type(exc).__name__}: {exc}" if str(exc) else type(exc).__name__ + return MCPToolResult( + outputs=[Content.from_text(f"Error: {message}")], + is_error=True, + error_message=message, + ) + except Exception as exc: + # Be defensive about MCP errors that may bubble up without being + # wrapped in ToolExecutionException by custom parsers. + try: + from mcp.shared.exceptions import McpError + except ImportError: # pragma: no cover - mcp is a hard dep but stay defensive + raise + if isinstance(exc, McpError): + message = str(exc) or type(exc).__name__ + return MCPToolResult( + outputs=[Content.from_text(f"Error: {message}")], + is_error=True, + error_message=message, + ) + raise + + # Defensive normalisation: call_tool is typed ``str | list[Content]``. + # Default parser returns list, but custom parse_tool_results may return str. + if isinstance(raw, str): + outputs: list[Content] = [Content.from_text(raw)] + else: + outputs = list(raw) + return MCPToolResult(outputs=outputs) + + async def aclose(self) -> None: + """Close all cached MCP clients and the owned httpx clients. + + Caller-supplied :class:`httpx.AsyncClient` instances (returned by the + ``client_provider`` callback) are NOT closed. + """ + async with self._cache_lock: + entries = list(self._cache.values()) + self._cache.clear() + for entry in entries: + await self._close_entry(entry) + + async def __aenter__(self) -> DefaultMCPToolHandler: + return self + + async def __aexit__(self, exc_type: Any, exc: Any, tb: Any) -> None: + await self.aclose() + + # ------------------------------------------------------------------ + # Internal helpers + # ------------------------------------------------------------------ + + async def _get_or_create_entry(self, invocation: MCPToolInvocation) -> _CacheEntry: + """Look up (or create) the cached MCP client for this invocation.""" + key = self._cache_key(invocation.server_url, invocation.headers) + + # Phase 1: check the cache and either claim creation or wait for an + # already in-flight creation. + creating = False + async with self._cache_lock: + existing = self._cache.get(key) + if existing is not None: + self._cache.move_to_end(key) + return existing + inflight = self._inflight.get(key) + if inflight is None: + inflight = asyncio.get_running_loop().create_future() + self._inflight[key] = inflight + creating = True + + if not creating: + return await inflight + + # Phase 2: we own creation. Build the entry outside the lock. + try: + entry = await self._create_entry(invocation) + except BaseException as exc: + async with self._cache_lock: + self._inflight.pop(key, None) + if not inflight.done(): + inflight.set_exception(exc if isinstance(exc, BaseException) else RuntimeError(str(exc))) + # Mark the exception retrieved to suppress noisy "Future exception + # was never retrieved" warnings when there are no other awaiters + # (other awaiters still see the exception through their ``await``). + inflight.exception() + raise + + # Phase 3: insert with LRU eviction; resolve the in-flight future. + evicted: _CacheEntry | None = None + duplicate: _CacheEntry | None = None + async with self._cache_lock: + self._inflight.pop(key, None) + existing = self._cache.get(key) + if existing is not None: + # Another writer beat us; prefer the existing entry and + # discard ours after the lock is released. + self._cache.move_to_end(key) + duplicate = entry + entry = existing + else: + self._cache[key] = entry + self._cache.move_to_end(key) + if len(self._cache) > self._cache_max_size: + _evicted_key, evicted = self._cache.popitem(last=False) + if not inflight.done(): + inflight.set_result(entry) + + if duplicate is not None: + await self._close_entry(duplicate) + if evicted is not None: + await self._close_entry(evicted) + return entry + + async def _create_entry(self, invocation: MCPToolInvocation) -> _CacheEntry: + """Construct (and connect) a fresh MCP client for ``invocation``.""" + from agent_framework import MCPStreamableHTTPTool + + provided_client: httpx.AsyncClient | None = None + if self._client_provider is not None: + provided_client = await self._client_provider(invocation) + # Capture headers for this cache entry so the header_provider closure + # always returns the same set, regardless of the runtime kwargs. + captured_headers = dict(invocation.headers) + + def _header_provider(_kwargs: dict[str, Any]) -> dict[str, str]: + return captured_headers + + tool: Any = MCPStreamableHTTPTool( + name=invocation.server_label or "McpClient", + url=invocation.server_url, + load_prompts=False, + http_client=provided_client, + header_provider=_header_provider if captured_headers else None, + ) + try: + await tool.connect() + except BaseException: + try: + await tool.close() + except Exception: # pragma: no cover - best effort + logger.debug("DefaultMCPToolHandler: error closing tool after failed connect", exc_info=True) + raise + + # ``MCPStreamableHTTPTool.get_mcp_client`` lazily creates an + # ``httpx.AsyncClient`` when no caller client was provided AND a + # ``header_provider`` was set. We treat any client allocated this + # way as owned (closed by the handler). When the caller supplies + # one, we never close it. + owned_client: httpx.AsyncClient | None = None + if provided_client is None: + owned_client = cast("httpx.AsyncClient | None", getattr(tool, "_httpx_client", None)) + return _CacheEntry(tool=tool, owned_httpx_client=owned_client) + + async def _close_entry(self, entry: _CacheEntry) -> None: + """Close the MCP tool and any owned httpx client.""" + try: + await entry.tool.close() + except Exception: # pragma: no cover - best effort + logger.debug("DefaultMCPToolHandler: error closing MCP tool", exc_info=True) + if entry.owned_httpx_client is not None: + try: + await entry.owned_httpx_client.aclose() + except Exception: # pragma: no cover - best effort + logger.debug("DefaultMCPToolHandler: error closing owned httpx client", exc_info=True) + + @staticmethod + def _cache_key(server_url: str, headers: dict[str, str] | None) -> tuple[str, str]: + """Build an order-independent cache key for ``(server_url, headers)``.""" + if not headers: + headers_hash = "0" + else: + payload = json.dumps(sorted(headers.items()), ensure_ascii=False) + headers_hash = hashlib.sha256(payload.encode("utf-8")).hexdigest() + return (server_url, headers_hash) diff --git a/python/packages/declarative/tests/test_default_mcp_tool_handler.py b/python/packages/declarative/tests/test_default_mcp_tool_handler.py new file mode 100644 index 00000000000..c40d275e807 --- /dev/null +++ b/python/packages/declarative/tests/test_default_mcp_tool_handler.py @@ -0,0 +1,415 @@ +# Copyright (c) Microsoft. All rights reserved. + +"""Tests for ``DefaultMCPToolHandler``. + +These tests exercise the real handler against a fake ``MCPStreamableHTTPTool`` +(no real MCP server, no real network) to cover the parts of the handler not +exercisable through the executor stub: cache hit/miss/eviction, concurrent +connect via in-flight futures, header isolation across cache keys, +string-result normalisation, ``load_prompts=False`` verification, and +owned-vs-caller httpx close semantics. +""" + +from __future__ import annotations + +import asyncio +import sys +from typing import Any +from unittest.mock import patch + +import httpx +import pytest +from agent_framework import Content +from agent_framework.exceptions import ToolExecutionException + +from agent_framework_declarative._workflows._mcp_handler import ( + DefaultMCPToolHandler, + MCPToolInvocation, +) + +pytestmark = pytest.mark.skipif( + sys.version_info >= (3, 14), + reason="Skipped on Python 3.14+ to keep parity with rest of declarative suite", +) + + +class FakeTool: + """Stand-in for ``MCPStreamableHTTPTool``. + + Records constructor kwargs, tracks connect/close lifecycle, and dispatches + ``call_tool`` to a per-instance handler. + """ + + instances: list[FakeTool] = [] + + def __init__(self, **kwargs: Any) -> None: + self.kwargs = kwargs + self.connect_count = 0 + self.close_count = 0 + self.connect_delay: float = 0.0 + self.connect_error: BaseException | None = None + self.call_handler: Any = lambda **_a: [Content.from_text("ok")] + self._httpx_client: httpx.AsyncClient | None = None + # Mimic MCPStreamableHTTPTool: when no caller client AND header_provider + # is set, lazily allocate an owned httpx client during connect. + FakeTool.instances.append(self) + + async def connect(self) -> None: + if self.connect_delay: + await asyncio.sleep(self.connect_delay) + if self.connect_error is not None: + raise self.connect_error + self.connect_count += 1 + # Mimic lazy httpx allocation when no client provided AND header_provider set. + if self.kwargs.get("http_client") is None and self.kwargs.get("header_provider") is not None: + self._httpx_client = httpx.AsyncClient() + + async def close(self) -> None: + self.close_count += 1 + + async def call_tool(self, tool_name: str, **arguments: Any) -> Any: + return self.call_handler(tool_name=tool_name, **arguments) + + +@pytest.fixture(autouse=True) +def _clear_fake_instances() -> None: + FakeTool.instances.clear() + + +def _patch_tool() -> Any: + """Patch the lazy import inside ``_create_entry`` to substitute FakeTool.""" + import agent_framework + + return patch.object(agent_framework, "MCPStreamableHTTPTool", FakeTool) + + +def _invocation( + *, server_url: str = "https://mcp.example/api", tool_name: str = "search", **overrides: Any +) -> MCPToolInvocation: + return MCPToolInvocation( + server_url=server_url, + tool_name=tool_name, + **overrides, + ) + + +# ---------- Construction --------------------------------------------------- + + +class TestConstruction: + def test_invalid_cache_size_raises(self) -> None: + with pytest.raises(ValueError): + DefaultMCPToolHandler(cache_max_size=0) + with pytest.raises(ValueError): + DefaultMCPToolHandler(cache_max_size=-3) + + +# ---------- Tool kwargs ---------------------------------------------------- + + +class TestToolKwargs: + @pytest.mark.asyncio + async def test_load_prompts_false_passed_to_tool(self) -> None: + handler = DefaultMCPToolHandler() + with _patch_tool(): + await handler.invoke_tool(_invocation()) + assert len(FakeTool.instances) == 1 + assert FakeTool.instances[0].kwargs["load_prompts"] is False + + @pytest.mark.asyncio + async def test_server_label_used_as_tool_name(self) -> None: + handler = DefaultMCPToolHandler() + with _patch_tool(): + await handler.invoke_tool(_invocation(server_label="MyMcp")) + assert FakeTool.instances[0].kwargs["name"] == "MyMcp" + + @pytest.mark.asyncio + async def test_default_tool_name_when_no_label(self) -> None: + handler = DefaultMCPToolHandler() + with _patch_tool(): + await handler.invoke_tool(_invocation(server_label=None)) + assert FakeTool.instances[0].kwargs["name"] == "McpClient" + + @pytest.mark.asyncio + async def test_no_header_provider_when_no_headers(self) -> None: + handler = DefaultMCPToolHandler() + with _patch_tool(): + await handler.invoke_tool(_invocation(headers={})) + assert FakeTool.instances[0].kwargs["header_provider"] is None + + @pytest.mark.asyncio + async def test_header_provider_returns_captured_headers(self) -> None: + handler = DefaultMCPToolHandler() + with _patch_tool(): + await handler.invoke_tool(_invocation(headers={"Authorization": "Bearer T"})) + provider = FakeTool.instances[0].kwargs["header_provider"] + assert provider({}) == {"Authorization": "Bearer T"} + # Even if runtime kwargs change, captured headers stay the same. + assert provider({"foo": "bar"}) == {"Authorization": "Bearer T"} + + +# ---------- Cache behaviour ------------------------------------------------ + + +class TestCache: + @pytest.mark.asyncio + async def test_same_url_and_headers_hit_cache(self) -> None: + handler = DefaultMCPToolHandler() + with _patch_tool(): + await handler.invoke_tool(_invocation(headers={"X": "1"})) + await handler.invoke_tool(_invocation(headers={"X": "1"})) + # One tool created, connect called once. + assert len(FakeTool.instances) == 1 + assert FakeTool.instances[0].connect_count == 1 + + @pytest.mark.asyncio + async def test_different_headers_create_separate_entries(self) -> None: + handler = DefaultMCPToolHandler() + with _patch_tool(): + await handler.invoke_tool(_invocation(headers={"Authorization": "tk-A"})) + await handler.invoke_tool(_invocation(headers={"Authorization": "tk-B"})) + assert len(FakeTool.instances) == 2 + + @pytest.mark.asyncio + async def test_different_urls_create_separate_entries(self) -> None: + handler = DefaultMCPToolHandler() + with _patch_tool(): + await handler.invoke_tool(_invocation(server_url="https://mcp.a/api")) + await handler.invoke_tool(_invocation(server_url="https://mcp.b/api")) + assert len(FakeTool.instances) == 2 + + @pytest.mark.asyncio + async def test_lru_eviction_closes_old_entry(self) -> None: + handler = DefaultMCPToolHandler(cache_max_size=2) + with _patch_tool(): + await handler.invoke_tool(_invocation(server_url="https://a/")) + await handler.invoke_tool(_invocation(server_url="https://b/")) + # Inserting a third evicts the LRU entry (the first one). + await handler.invoke_tool(_invocation(server_url="https://c/")) + assert len(FakeTool.instances) == 3 + # First instance (https://a/) was evicted → close() called. + assert FakeTool.instances[0].kwargs["url"] == "https://a/" + assert FakeTool.instances[0].close_count == 1 + # Other two remain in cache → not closed. + assert FakeTool.instances[1].close_count == 0 + assert FakeTool.instances[2].close_count == 0 + + @pytest.mark.asyncio + async def test_repeated_use_keeps_lru_alive(self) -> None: + handler = DefaultMCPToolHandler(cache_max_size=2) + with _patch_tool(): + await handler.invoke_tool(_invocation(server_url="https://a/")) + await handler.invoke_tool(_invocation(server_url="https://b/")) + # Touch a → b becomes LRU. + await handler.invoke_tool(_invocation(server_url="https://a/")) + # Insert c → b is evicted. + await handler.invoke_tool(_invocation(server_url="https://c/")) + # b was evicted. + b = FakeTool.instances[1] + assert b.kwargs["url"] == "https://b/" + assert b.close_count == 1 + # a survived. + a = FakeTool.instances[0] + assert a.kwargs["url"] == "https://a/" + assert a.close_count == 0 + + @pytest.mark.asyncio + async def test_concurrent_connect_shares_one_entry(self) -> None: + """Multiple concurrent invocations with the same key must share one tool.""" + handler = DefaultMCPToolHandler() + + # Slow down connect so concurrency window is observable. + original_connect = FakeTool.connect + + async def slow_connect(self: FakeTool) -> None: + self.connect_delay = 0.05 + await original_connect(self) + + with _patch_tool(), patch.object(FakeTool, "connect", slow_connect): + results = await asyncio.gather( + handler.invoke_tool(_invocation(headers={"X": "1"})), + handler.invoke_tool(_invocation(headers={"X": "1"})), + handler.invoke_tool(_invocation(headers={"X": "1"})), + handler.invoke_tool(_invocation(headers={"X": "1"})), + ) + assert all(not r.is_error for r in results) + # Only one tool was created and connected, despite 4 concurrent calls. + assert len(FakeTool.instances) == 1 + assert FakeTool.instances[0].connect_count == 1 + + +# ---------- Aclose semantics ---------------------------------------------- + + +class TestAclose: + @pytest.mark.asyncio + async def test_aclose_closes_owned_clients(self) -> None: + handler = DefaultMCPToolHandler() + with _patch_tool(): + await handler.invoke_tool(_invocation(headers={"X": "1"})) + tool = FakeTool.instances[0] + owned = tool._httpx_client + assert owned is not None + await handler.aclose() + assert tool.close_count == 1 + assert owned.is_closed + + @pytest.mark.asyncio + async def test_aclose_does_not_close_caller_supplied_client(self) -> None: + caller_client = httpx.AsyncClient() + + async def provider(_inv: MCPToolInvocation) -> httpx.AsyncClient: + return caller_client + + handler = DefaultMCPToolHandler(client_provider=provider) + try: + with _patch_tool(): + await handler.invoke_tool(_invocation(headers={"X": "1"})) + await handler.aclose() + assert FakeTool.instances[0].close_count == 1 + # Caller client must still be usable. + assert not caller_client.is_closed + finally: + await caller_client.aclose() + + @pytest.mark.asyncio + async def test_async_context_manager(self) -> None: + with _patch_tool(): + async with DefaultMCPToolHandler() as handler: + await handler.invoke_tool(_invocation()) + tool = FakeTool.instances[0] + assert tool.close_count == 1 + + +# ---------- Result normalisation ------------------------------------------ + + +class TestResultNormalisation: + @pytest.mark.asyncio + async def test_string_result_wrapped_in_text_content(self) -> None: + handler = DefaultMCPToolHandler() + with _patch_tool(): + inv = _invocation() + result = await handler.invoke_tool(inv) + # The fake's default already returns a list; replace handler for this test. + FakeTool.instances[0].call_handler = lambda **_a: "raw string body" + result = await handler.invoke_tool(inv) + assert result.is_error is False + assert len(result.outputs) == 1 + assert result.outputs[0].text == "raw string body" # type: ignore[reportAttributeAccessIssue] + + @pytest.mark.asyncio + async def test_list_result_passed_through(self) -> None: + handler = DefaultMCPToolHandler() + custom = [Content.from_text("a"), Content.from_text("b")] + with _patch_tool(): + inv = _invocation() + await handler.invoke_tool(inv) + FakeTool.instances[0].call_handler = lambda **_a: custom + result = await handler.invoke_tool(inv) + assert result.is_error is False + assert len(result.outputs) == 2 + + +# ---------- Error mapping -------------------------------------------------- + + +class TestErrorMapping: + @pytest.mark.asyncio + async def test_tool_execution_exception_returns_error_result(self) -> None: + handler = DefaultMCPToolHandler() + + def boom(**_a: Any) -> Any: + raise ToolExecutionException("server says no") + + with _patch_tool(): + inv = _invocation() + await handler.invoke_tool(inv) + FakeTool.instances[0].call_handler = boom + result = await handler.invoke_tool(inv) + assert result.is_error is True + assert result.error_message == "server says no" + assert result.outputs[0].text.startswith("Error:") # type: ignore[reportAttributeAccessIssue] + + @pytest.mark.asyncio + async def test_httpx_error_returns_error_result(self) -> None: + handler = DefaultMCPToolHandler() + + def boom(**_a: Any) -> Any: + raise httpx.ConnectError("dns failure") + + with _patch_tool(): + inv = _invocation() + await handler.invoke_tool(inv) + FakeTool.instances[0].call_handler = boom + result = await handler.invoke_tool(inv) + assert result.is_error is True + assert "dns failure" in (result.error_message or "") + + @pytest.mark.asyncio + async def test_unexpected_exception_propagates(self) -> None: + """RuntimeError (not in the narrow catch list) must propagate.""" + handler = DefaultMCPToolHandler() + + def boom(**_a: Any) -> Any: + raise RuntimeError("programmer error") + + with _patch_tool(): + inv = _invocation() + await handler.invoke_tool(inv) + FakeTool.instances[0].call_handler = boom + with pytest.raises(RuntimeError, match="programmer error"): + await handler.invoke_tool(inv) + + @pytest.mark.asyncio + async def test_connect_failure_returns_error_result(self) -> None: + handler = DefaultMCPToolHandler() + with ( + _patch_tool(), + patch.object( + FakeTool, + "connect", + lambda self: (_ for _ in ()).throw(httpx.ConnectError("server down")), + ), + ): + result = await handler.invoke_tool(_invocation()) + assert result.is_error is True + assert result.outputs[0].text.startswith("Error:") # type: ignore[reportAttributeAccessIssue] + # Failed connect must clear in-flight + cache entries. + assert handler._inflight == {} + assert len(handler._cache) == 0 + + @pytest.mark.asyncio + async def test_cancelled_error_propagates(self) -> None: + """asyncio.CancelledError is BaseException, must NOT be swallowed.""" + handler = DefaultMCPToolHandler() + + def boom(**_a: Any) -> Any: + raise asyncio.CancelledError + + with _patch_tool(): + inv = _invocation() + await handler.invoke_tool(inv) + FakeTool.instances[0].call_handler = boom + with pytest.raises(asyncio.CancelledError): + await handler.invoke_tool(inv) + + +# ---------- Cache key isolation ------------------------------------------- + + +class TestCacheKey: + def test_key_order_independent(self) -> None: + k1 = DefaultMCPToolHandler._cache_key("https://x/", {"A": "1", "B": "2"}) + k2 = DefaultMCPToolHandler._cache_key("https://x/", {"B": "2", "A": "1"}) + assert k1 == k2 + + def test_key_distinguishes_values(self) -> None: + k1 = DefaultMCPToolHandler._cache_key("https://x/", {"A": "1"}) + k2 = DefaultMCPToolHandler._cache_key("https://x/", {"A": "2"}) + assert k1 != k2 + + def test_empty_headers_use_fixed_hash(self) -> None: + k1 = DefaultMCPToolHandler._cache_key("https://x/", None) + k2 = DefaultMCPToolHandler._cache_key("https://x/", {}) + assert k1 == k2 diff --git a/python/packages/declarative/tests/test_invoke_mcp_tool_executor.py b/python/packages/declarative/tests/test_invoke_mcp_tool_executor.py new file mode 100644 index 00000000000..867b8543bbb --- /dev/null +++ b/python/packages/declarative/tests/test_invoke_mcp_tool_executor.py @@ -0,0 +1,631 @@ +# Copyright (c) Microsoft. All rights reserved. + +"""Tests for ``InvokeMcpToolActionExecutor``. + +Use a stub :class:`MCPToolHandler` that returns canned :class:`MCPToolResult`s. +No real MCP server or network is exercised. See +``test_default_mcp_tool_handler.py`` for tests that exercise the real +``DefaultMCPToolHandler`` against a mocked ``MCPStreamableHTTPTool``. +""" + +import sys +from typing import Any + +import httpx +import pytest + +try: + import powerfx # noqa: F401 + + _powerfx_available = True +except (ImportError, RuntimeError): + _powerfx_available = False + +pytestmark = pytest.mark.skipif( + not _powerfx_available or sys.version_info >= (3, 14), + reason="PowerFx engine not available (requires dotnet runtime)", +) + +from agent_framework import Content, Message # noqa: E402 +from agent_framework.exceptions import ToolExecutionException # noqa: E402 + +from agent_framework_declarative._workflows import ( # noqa: E402 + DECLARATIVE_STATE_KEY, + DeclarativeWorkflowError, + MCPToolHandler, + MCPToolInvocation, + MCPToolResult, + WorkflowFactory, +) + + +class StubMcpHandler: + """Test stub recording the last call and returning a canned result.""" + + def __init__( + self, + result: MCPToolResult | None = None, + *, + raise_exc: BaseException | None = None, + ) -> None: + self.result = result + self.raise_exc = raise_exc + self.last_invocation: MCPToolInvocation | None = None + self.invocations: list[MCPToolInvocation] = [] + self.call_count = 0 + + async def invoke_tool(self, invocation: MCPToolInvocation) -> MCPToolResult: + self.call_count += 1 + self.last_invocation = invocation + self.invocations.append(invocation) + if self.raise_exc is not None: + raise self.raise_exc + assert self.result is not None + return self.result + + +def _ok(outputs: list[Content] | None = None) -> MCPToolResult: + return MCPToolResult(outputs=outputs or [Content.from_text("hello")]) + + +def _err(message: str = "boom") -> MCPToolResult: + return MCPToolResult( + outputs=[Content.from_text(f"Error: {message}")], + is_error=True, + error_message=message, + ) + + +def _action( + *, + server_url: str = "https://mcp.example/api", + tool_name: str = "search", + server_label: str | None = None, + arguments: dict[str, Any] | None = None, + headers: dict[str, Any] | None = None, + require_approval: Any = None, + connection: dict[str, Any] | None = None, + conversation_id: str | None = None, + output: dict[str, Any] | None = None, +) -> dict[str, Any]: + action: dict[str, Any] = { + "kind": "InvokeMcpTool", + "id": "mcp_action", + "serverUrl": server_url, + "toolName": tool_name, + } + if server_label is not None: + action["serverLabel"] = server_label + if arguments is not None: + action["arguments"] = arguments + if headers is not None: + action["headers"] = headers + if require_approval is not None: + action["requireApproval"] = require_approval + if connection is not None: + action["connection"] = connection + if conversation_id is not None: + action["conversationId"] = conversation_id + if output is not None: + action["output"] = output + return action + + +def _yaml(action: dict[str, Any]) -> dict[str, Any]: + return {"name": "mcp_test", "actions": [action]} + + +# ---------- Builder enforcement -------------------------------------------- + + +class TestBuilderEnforcement: + def test_missing_handler_raises_at_build_time(self) -> None: + factory = WorkflowFactory() + with pytest.raises(DeclarativeWorkflowError) as excinfo: + factory.create_workflow_from_definition(_yaml(_action())) + assert "InvokeMcpTool" in str(excinfo.value) + assert "mcp_tool_handler" in str(excinfo.value) + + def test_missing_server_url_fails_validation(self) -> None: + handler = StubMcpHandler(_ok()) + factory = WorkflowFactory(mcp_tool_handler=handler) + action = _action() + del action["serverUrl"] + with pytest.raises(Exception) as excinfo: + factory.create_workflow_from_definition(_yaml(action)) + assert "serverUrl" in str(excinfo.value) + + def test_missing_tool_name_fails_validation(self) -> None: + handler = StubMcpHandler(_ok()) + factory = WorkflowFactory(mcp_tool_handler=handler) + action = _action() + del action["toolName"] + with pytest.raises(Exception) as excinfo: + factory.create_workflow_from_definition(_yaml(action)) + assert "toolName" in str(excinfo.value) + + +# ---------- Field forwarding ---------------------------------------------- + + +class TestFieldForwarding: + @pytest.mark.asyncio + async def test_basic_invocation_forwards_required_fields(self) -> None: + handler = StubMcpHandler(_ok()) + factory = WorkflowFactory(mcp_tool_handler=handler) + workflow = factory.create_workflow_from_definition(_yaml(_action())) + await workflow.run({}) + assert handler.call_count == 1 + inv = handler.last_invocation + assert inv is not None + assert inv.server_url == "https://mcp.example/api" + assert inv.tool_name == "search" + assert inv.server_label is None + assert inv.headers == {} + assert inv.arguments == {} + assert inv.connection_name is None + + @pytest.mark.asyncio + async def test_arguments_evaluated_and_preserves_none(self) -> None: + handler = StubMcpHandler(_ok()) + factory = WorkflowFactory(mcp_tool_handler=handler) + workflow = factory.create_workflow_from_definition( + _yaml( + _action( + arguments={ + "query": "weather today", + "limit": 5, + "fresh": True, + "missing": None, + } + ) + ) + ) + await workflow.run({}) + inv = handler.last_invocation + assert inv is not None + # ``None`` is preserved (parity with .NET) — caller decides. + assert inv.arguments == { + "query": "weather today", + "limit": 5, + "fresh": True, + "missing": None, + } + + @pytest.mark.asyncio + async def test_headers_drop_empty_values(self) -> None: + handler = StubMcpHandler(_ok()) + factory = WorkflowFactory(mcp_tool_handler=handler) + workflow = factory.create_workflow_from_definition( + _yaml( + _action( + headers={ + "Authorization": "Bearer token-123", + "X-Trace": "trace-id", + "X-Empty": "", + } + ) + ) + ) + await workflow.run({}) + inv = handler.last_invocation + assert inv is not None + assert inv.headers == { + "Authorization": "Bearer token-123", + "X-Trace": "trace-id", + } + + @pytest.mark.asyncio + async def test_server_label_and_connection_name_forwarded(self) -> None: + handler = StubMcpHandler(_ok()) + factory = WorkflowFactory(mcp_tool_handler=handler) + workflow = factory.create_workflow_from_definition( + _yaml( + _action( + server_label="docs-mcp", + connection={"name": "azure-conn"}, + ) + ) + ) + await workflow.run({}) + inv = handler.last_invocation + assert inv is not None + assert inv.server_label == "docs-mcp" + assert inv.connection_name == "azure-conn" + + +# ---------- Output handling ------------------------------------------------ + + +class TestOutput: + @pytest.mark.asyncio + async def test_output_result_parses_json_text(self) -> None: + handler = StubMcpHandler(_ok([Content.from_text('{"k":"v","n":1}')])) + factory = WorkflowFactory(mcp_tool_handler=handler) + workflow = factory.create_workflow_from_definition(_yaml(_action(output={"result": "Local.Result"}))) + await workflow.run({}) + decl = workflow._state.get(DECLARATIVE_STATE_KEY) + assert decl["Local"]["Result"] == [{"k": "v", "n": 1}] + + @pytest.mark.asyncio + async def test_output_result_falls_back_to_raw_text(self) -> None: + handler = StubMcpHandler(_ok([Content.from_text("plain text not json")])) + factory = WorkflowFactory(mcp_tool_handler=handler) + workflow = factory.create_workflow_from_definition(_yaml(_action(output={"result": "Local.Result"}))) + await workflow.run({}) + decl = workflow._state.get(DECLARATIVE_STATE_KEY) + assert decl["Local"]["Result"] == ["plain text not json"] + + @pytest.mark.asyncio + async def test_output_messages_writes_single_tool_role_message(self) -> None: + handler = StubMcpHandler(_ok([Content.from_text("hi"), Content.from_text("there")])) + factory = WorkflowFactory(mcp_tool_handler=handler) + workflow = factory.create_workflow_from_definition(_yaml(_action(output={"messages": "Local.Messages"}))) + await workflow.run({}) + decl = workflow._state.get(DECLARATIVE_STATE_KEY) + msg = decl["Local"]["Messages"] + # Single Tool-role message containing both contents (parity with .NET). + assert isinstance(msg, Message) + assert str(msg.role).lower() == "tool" + assert len(msg.contents) == 2 + + @pytest.mark.asyncio + async def test_uri_content_serialised_as_uri_string(self) -> None: + uri_content = Content.from_uri("https://example.com/file.txt", media_type="text/plain") + handler = StubMcpHandler(_ok([uri_content])) + factory = WorkflowFactory(mcp_tool_handler=handler) + workflow = factory.create_workflow_from_definition(_yaml(_action(output={"result": "Local.Result"}))) + await workflow.run({}) + decl = workflow._state.get(DECLARATIVE_STATE_KEY) + assert decl["Local"]["Result"] == ["https://example.com/file.txt"] + + @pytest.mark.asyncio + async def test_output_path_object_form(self) -> None: + handler = StubMcpHandler(_ok([Content.from_text("ok")])) + factory = WorkflowFactory(mcp_tool_handler=handler) + workflow = factory.create_workflow_from_definition(_yaml(_action(output={"result": {"path": "Local.Result"}}))) + await workflow.run({}) + decl = workflow._state.get(DECLARATIVE_STATE_KEY) + assert decl["Local"]["Result"] == ["ok"] + + +# ---------- Conversation append -------------------------------------------- + + +class TestConversation: + @pytest.mark.asyncio + async def test_conversation_id_appends_assistant_message(self) -> None: + handler = StubMcpHandler(_ok([Content.from_text("answer")])) + factory = WorkflowFactory(mcp_tool_handler=handler) + workflow = factory.create_workflow_from_definition( + _yaml( + _action( + conversation_id="conv-42", + output={"result": "Local.Result"}, + ) + ) + ) + await workflow.run({}) + decl = workflow._state.get(DECLARATIVE_STATE_KEY) + conv = decl["System"]["conversations"]["conv-42"] + msgs = conv["messages"] if isinstance(conv, dict) else conv.messages + assert len(msgs) == 1 + appended = msgs[0] + assert str(appended.role).lower() == "assistant" + # Same contents as the tool output. + assert len(appended.contents) == 1 + + @pytest.mark.asyncio + async def test_empty_conversation_id_does_not_append(self) -> None: + handler = StubMcpHandler(_ok([Content.from_text("answer")])) + factory = WorkflowFactory(mcp_tool_handler=handler) + workflow = factory.create_workflow_from_definition( + _yaml( + _action( + conversation_id="", + output={"result": "Local.Result"}, + ) + ) + ) + await workflow.run({}) + decl = workflow._state.get(DECLARATIVE_STATE_KEY) + # Empty conversation id must not produce a `""` entry under System.conversations. + conversations = decl.get("System", {}).get("conversations", {}) + assert "" not in conversations + + +# ---------- Approval flow -------------------------------------------------- + + +@pytest.fixture +def mock_state(): # type: ignore[no-untyped-def] + from unittest.mock import MagicMock + + state = MagicMock() + state._data = {} + + def _get(key: str, default: Any = None) -> Any: + if key not in state._data: + if default is not None: + return default + raise KeyError(key) + return state._data[key] + + def _set(key: str, value: Any) -> None: + state._data[key] = value + + def _delete(key: str) -> None: + if key in state._data: + del state._data[key] + else: + raise KeyError(key) + + state.get = MagicMock(side_effect=_get) + state.set = MagicMock(side_effect=_set) + state.delete = MagicMock(side_effect=_delete) + return state + + +@pytest.fixture +def mock_context(mock_state): # type: ignore[no-untyped-def] + from unittest.mock import AsyncMock, MagicMock + + ctx = MagicMock() + ctx.state = mock_state + ctx.send_message = AsyncMock() + ctx.yield_output = AsyncMock() + ctx.request_info = AsyncMock() + return ctx + + +def _seed_state(mock_state) -> None: # type: ignore[no-untyped-def] + """Pre-seed the declarative state container as the executors expect.""" + from agent_framework_declarative._workflows import DECLARATIVE_STATE_KEY + + mock_state._data[DECLARATIVE_STATE_KEY] = { + "Local": {}, + "Custom": {}, + "Workflow": {}, + "System": { + "ConversationId": "00000000-0000-0000-0000-000000000000", + "LastMessage": {"Id": "", "Text": ""}, + "LastMessageText": "", + "LastMessageId": "", + }, + "Agent": {}, + "Conversation": {"messages": [], "history": []}, + "Inputs": {}, + } + + +class TestApprovalFlow: + @pytest.mark.asyncio + async def test_approval_required_emits_request_and_yields(self, mock_state, mock_context) -> None: # type: ignore[no-untyped-def] + from agent_framework_declarative._workflows._declarative_base import ActionTrigger + from agent_framework_declarative._workflows._executors_mcp import ( + _MCP_APPROVAL_STATE_KEY, + InvokeMcpToolActionExecutor, + MCPToolApprovalRequest, + ) + + _seed_state(mock_state) + handler = StubMcpHandler(_ok()) + executor = InvokeMcpToolActionExecutor( + _action( + require_approval=True, + arguments={"q": "x"}, + headers={"Authorization": "Bearer SECRET"}, + output={"result": "Local.Result"}, + ), + mcp_tool_handler=handler, + ) + await executor.handle_action(ActionTrigger(), mock_context) + + # Approval request emitted. + mock_context.request_info.assert_called_once() + request = mock_context.request_info.call_args[0][0] + assert isinstance(request, MCPToolApprovalRequest) + assert request.tool_name == "search" + assert request.arguments == {"q": "x"} + assert request.header_names == ["Authorization"] + + # NEVER expose the actual auth token in any field of the approval payload. + for value in request.__dict__.values(): + assert "SECRET" not in str(value) + + # Workflow should yield (no ActionComplete sent yet). + mock_context.send_message.assert_not_called() + + # Handler not invoked yet. + assert handler.call_count == 0 + + # Approval state stored. + approval_key = f"{_MCP_APPROVAL_STATE_KEY}_mcp_action" + assert approval_key in mock_state._data + + @pytest.mark.asyncio + async def test_approval_response_approved_invokes_handler(self, mock_state, mock_context) -> None: # type: ignore[no-untyped-def] + from agent_framework_declarative._workflows import ActionComplete, ToolApprovalResponse + from agent_framework_declarative._workflows._executors_mcp import ( + _MCP_APPROVAL_STATE_KEY, + InvokeMcpToolActionExecutor, + MCPToolApprovalRequest, + _MCPToolApprovalState, + ) + + _seed_state(mock_state) + handler = StubMcpHandler(_ok([Content.from_text('{"ok":true}')])) + executor = InvokeMcpToolActionExecutor( + _action( + require_approval=True, + output={"result": "Local.Result"}, + ), + mcp_tool_handler=handler, + ) + # Pre-populate approval state. + approval_key = f"{_MCP_APPROVAL_STATE_KEY}_mcp_action" + mock_state._data[approval_key] = _MCPToolApprovalState( + server_url="https://mcp.example/api", + tool_name="search", + server_label=None, + arguments={"q": "x"}, + connection_name=None, + headers_def={"Authorization": "Bearer tk"}, + auto_send=False, + conversation_id_expr=None, + output_messages_path=None, + output_result_path="Local.Result", + ) + await executor.handle_approval_response( + MCPToolApprovalRequest( + request_id="req-1", + tool_name="search", + server_url="https://mcp.example/api", + server_label=None, + arguments={"q": "x"}, + ), + ToolApprovalResponse(approved=True), + mock_context, + ) + + assert handler.call_count == 1 + inv = handler.last_invocation + assert inv is not None + # Headers are re-evaluated from headers_def. + assert inv.headers == {"Authorization": "Bearer tk"} + # Approval state was cleaned up. + assert approval_key not in mock_state._data + # ActionComplete was sent. + mock_context.send_message.assert_called_once() + sent = mock_context.send_message.call_args[0][0] + assert isinstance(sent, ActionComplete) + + @pytest.mark.asyncio + async def test_approval_response_rejected_assigns_error(self, mock_state, mock_context) -> None: # type: ignore[no-untyped-def] + from agent_framework_declarative._workflows import ToolApprovalResponse + from agent_framework_declarative._workflows._executors_mcp import ( + _MCP_APPROVAL_STATE_KEY, + InvokeMcpToolActionExecutor, + MCPToolApprovalRequest, + _MCPToolApprovalState, + ) + + _seed_state(mock_state) + handler = StubMcpHandler(_ok()) + executor = InvokeMcpToolActionExecutor( + _action( + require_approval=True, + output={"result": "Local.Result"}, + ), + mcp_tool_handler=handler, + ) + approval_key = f"{_MCP_APPROVAL_STATE_KEY}_mcp_action" + mock_state._data[approval_key] = _MCPToolApprovalState( + server_url="https://mcp.example/api", + tool_name="search", + server_label=None, + arguments={}, + connection_name=None, + headers_def=None, + auto_send=True, + conversation_id_expr=None, + output_messages_path=None, + output_result_path="Local.Result", + ) + await executor.handle_approval_response( + MCPToolApprovalRequest( + request_id="req-2", + tool_name="search", + server_url="https://mcp.example/api", + server_label=None, + arguments={}, + ), + ToolApprovalResponse(approved=False, reason="not authorized"), + mock_context, + ) + + assert handler.call_count == 0 + # Error string assigned at output.result. + from agent_framework_declarative._workflows import DECLARATIVE_STATE_KEY + + result = mock_state._data[DECLARATIVE_STATE_KEY]["Local"]["Result"] + assert result == "Error: MCP tool invocation was not approved by user." + + +# ---------- Error handling ------------------------------------------------- + + +class TestErrorHandling: + @pytest.mark.asyncio + async def test_handler_returns_error_result_assigns_error_string(self) -> None: + handler = StubMcpHandler(_err("server down")) + factory = WorkflowFactory(mcp_tool_handler=handler) + workflow = factory.create_workflow_from_definition(_yaml(_action(output={"result": "Local.Result"}))) + await workflow.run({}) + decl = workflow._state.get(DECLARATIVE_STATE_KEY) + assert decl["Local"]["Result"] == "Error: server down" + + @pytest.mark.asyncio + async def test_tool_execution_exception_becomes_error_result(self) -> None: + handler = StubMcpHandler(raise_exc=ToolExecutionException("invalid arguments")) + factory = WorkflowFactory(mcp_tool_handler=handler) + workflow = factory.create_workflow_from_definition(_yaml(_action(output={"result": "Local.Result"}))) + await workflow.run({}) + decl = workflow._state.get(DECLARATIVE_STATE_KEY) + assert decl["Local"]["Result"] == "Error: invalid arguments" + + @pytest.mark.asyncio + async def test_httpx_error_becomes_error_result(self) -> None: + handler = StubMcpHandler(raise_exc=httpx.ConnectError("dns fail")) + factory = WorkflowFactory(mcp_tool_handler=handler) + workflow = factory.create_workflow_from_definition(_yaml(_action(output={"result": "Local.Result"}))) + await workflow.run({}) + decl = workflow._state.get(DECLARATIVE_STATE_KEY) + result = decl["Local"]["Result"] + assert isinstance(result, str) + assert result.startswith("Error:") + assert "ConnectError" in result + + @pytest.mark.asyncio + async def test_unexpected_exception_propagates(self) -> None: + """Programmer bugs (TypeError etc.) must NOT be swallowed.""" + handler = StubMcpHandler(raise_exc=TypeError("bad type")) + factory = WorkflowFactory(mcp_tool_handler=handler) + workflow = factory.create_workflow_from_definition(_yaml(_action())) + with pytest.raises(Exception) as excinfo: + await workflow.run({}) + # Either the TypeError reaches us or it gets wrapped by the runner — + # either way the message must surface. + assert "bad type" in str(excinfo.value) + + +# ---------- autoSend ------------------------------------------------------- + + +class TestAutoSend: + @pytest.mark.asyncio + async def test_auto_send_default_true_yields_output(self) -> None: + handler = StubMcpHandler(_ok([Content.from_text("hello")])) + factory = WorkflowFactory(mcp_tool_handler=handler) + workflow = factory.create_workflow_from_definition(_yaml(_action())) + events = await workflow.run({}) + outputs = events.get_outputs() + assert len(outputs) == 1 + + @pytest.mark.asyncio + async def test_auto_send_false_suppresses_yield(self) -> None: + handler = StubMcpHandler(_ok([Content.from_text("hello")])) + factory = WorkflowFactory(mcp_tool_handler=handler) + workflow = factory.create_workflow_from_definition(_yaml(_action(output={"autoSend": False}))) + events = await workflow.run({}) + outputs = events.get_outputs() + assert outputs == [] + + +# ---------- Protocol structure -------------------------------------------- + + +class TestProtocol: + def test_stub_handler_satisfies_protocol(self) -> None: + handler = StubMcpHandler(_ok()) + assert isinstance(handler, MCPToolHandler) diff --git a/python/samples/03-workflows/declarative/invoke_mcp_tool/main.py b/python/samples/03-workflows/declarative/invoke_mcp_tool/main.py new file mode 100644 index 00000000000..c95b0c46912 --- /dev/null +++ b/python/samples/03-workflows/declarative/invoke_mcp_tool/main.py @@ -0,0 +1,100 @@ +# Copyright (c) Microsoft. All rights reserved. + +"""Invoke MCP Tool sample - demonstrates the InvokeMcpTool declarative action. + +This sample shows how to: + 1. Configure a ``WorkflowFactory`` with a ``MCPToolHandler`` so the YAML + ``InvokeMcpTool`` action can dispatch real MCP tool calls. + 2. Invoke a tool on a public unauthenticated MCP server (the Microsoft + Learn Docs MCP server at ``https://learn.microsoft.com/api/mcp``, + calling ``microsoft_docs_search``). + 3. Bind the parsed tool result to a workflow variable and mirror it into + the conversation via ``conversationId`` so a downstream Foundry agent + can answer questions using only that context. + +Security note: + ``DefaultMCPToolHandler`` connects to whatever MCP server URL the + workflow author specifies and performs **no** allowlisting or SSRF + guards. For production use, replace it with a custom handler that + enforces an allowlist and adds any required authentication headers + per server. MCP tool outputs flow back into agent conversations and + therefore share the same prompt-injection risk surface as + ``HttpRequestAction``: only invoke MCP servers you trust. + +Run with: + python -m samples.03-workflows.declarative.invoke_mcp_tool.main +""" + +import asyncio +import os +from pathlib import Path + +from agent_framework import Agent +from agent_framework.declarative import ( + DefaultMCPToolHandler, + WorkflowFactory, +) +from agent_framework.foundry import FoundryChatClient +from azure.identity import AzureCliCredential + +DOCS_AGENT_INSTRUCTIONS = """\ +You answer the user's question about Microsoft technology using ONLY the +search results already present in the conversation history. If the answer is +not contained in the conversation, say so plainly rather than guessing. Be +concise and cite the relevant document title or URL when possible. +""" + + +async def main() -> None: + """Run the invoke MCP tool workflow.""" + chat_client = FoundryChatClient( + project_endpoint=os.environ["FOUNDRY_PROJECT_ENDPOINT"], + model=os.environ["FOUNDRY_MODEL"], + credential=AzureCliCredential(), + ) + + # The agent has no tools — it answers using only the search results that + # ``InvokeMcpTool`` adds to the conversation. + docs_agent = Agent( + client=chat_client, + name="DocsAgent", + instructions=DOCS_AGENT_INSTRUCTIONS, + ) + + agents = {"DocsAgent": docs_agent} + + # The default MCPToolHandler is sufficient for this sample because the + # Microsoft Learn Docs MCP server is public and unauthenticated. For + # authenticated servers, supply a ``client_provider`` callback to route + # requests through a pre-configured ``httpx.AsyncClient`` carrying the + # appropriate credentials, or wrap the handler with one that injects + # headers per call. + async with DefaultMCPToolHandler() as mcp_handler: + factory = WorkflowFactory( + agents=agents, + mcp_tool_handler=mcp_handler, + ) + + workflow_path = Path(__file__).parent / "workflow.yaml" + workflow = factory.create_workflow_from_yaml_path(workflow_path) + + print("=" * 60) + print("Invoke MCP Tool Workflow Demo") + print("=" * 60) + print() + print("Ask one question that can be answered from the Microsoft Learn docs or provide a keyword to search.") + print() + + user_input = input("You: ").strip() # noqa: ASYNC250 + if not user_input: + user_input = "What is the Agent Framework declarative workflow runtime?" + + print("\nAgent: ", end="", flush=True) + async for event in workflow.run(user_input, stream=True): + if event.type == "output" and isinstance(event.data, str): + print(event.data, end="", flush=True) + print() + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/python/samples/03-workflows/declarative/invoke_mcp_tool/workflow.yaml b/python/samples/03-workflows/declarative/invoke_mcp_tool/workflow.yaml new file mode 100644 index 00000000000..b83dc052ffa --- /dev/null +++ b/python/samples/03-workflows/declarative/invoke_mcp_tool/workflow.yaml @@ -0,0 +1,64 @@ +# +# This workflow demonstrates the InvokeMcpTool declarative action. +# +# InvokeMcpTool lets a workflow author call a tool exposed by a Model Context +# Protocol (MCP) server directly from YAML without writing any Python glue. +# It can: +# +# - dispatch a tool call against an MCP server (with optional auth headers), +# - store the parsed tool result in a workflow variable, and +# - add the result to the conversation so a downstream agent can answer +# questions based on it. +# +# This sample calls ``microsoft_docs_search`` on the public Microsoft Learn +# Docs MCP server (no authentication required) and uses a Foundry agent to +# answer a single question about Microsoft technology using the search +# results. +# +# Example inputs (Choose one or provide yours): +# How do I configure logging in the Agent Framework? +# Gpt-5.4-mini +# +kind: Workflow +trigger: + + kind: OnConversationStart + id: workflow_invoke_mcp_tool_demo + actions: + + # Capture the user's question into a local variable so the MCP tool call + # can pass it as an argument. + - kind: SetVariable + id: capture_query + variable: Local.SearchQuery + value: =System.LastMessage.Text + + # Invoke microsoft_docs_search on the Microsoft Learn Docs MCP server. + # The result is parsed into Local.SearchResults and also added to the + # conversation (via conversationId) so the agent below can answer the + # user's question based on it. + - kind: InvokeMcpTool + id: search_docs + conversationId: =System.ConversationId + serverUrl: https://learn.microsoft.com/api/mcp + serverLabel: MicrosoftLearnDocs + toolName: microsoft_docs_search + arguments: + query: =Local.SearchQuery + output: + autoSend: false + result: Local.SearchResults + + # Use the agent to answer the user's question using the conversation + # context (which now contains the MCP search results). The user's + # original message is already in the conversation as System.LastMessage, + # and the executor's input fallback chain extracts its ``Text`` field + # automatically when ``input.messages`` is omitted. + - kind: InvokeAzureAgent + id: answer_question + conversationId: =System.ConversationId + agent: + name: DocsAgent + output: + autoSend: true + messages: Local.AgentResponse From 92cd194122e9d98999c1f99a2a22e1564c5b1f7b Mon Sep 17 00:00:00 2001 From: Peter Ibekwe Date: Mon, 4 May 2026 13:47:24 -0700 Subject: [PATCH 6/8] Update sample to support require approval to be toggled by environment variable. --- .../agent_framework/declarative/__init__.py | 2 + .../agent_framework/declarative/__init__.pyi | 4 + .../agent_framework_declarative/__init__.py | 4 + .../declarative/invoke_mcp_tool/main.py | 111 +++++++++++++++++- .../declarative/invoke_mcp_tool/workflow.yaml | 21 +++- 5 files changed, 133 insertions(+), 9 deletions(-) diff --git a/python/packages/core/agent_framework/declarative/__init__.py b/python/packages/core/agent_framework/declarative/__init__.py index 90c73ef8bd5..b5e9c9ef9e7 100644 --- a/python/packages/core/agent_framework/declarative/__init__.py +++ b/python/packages/core/agent_framework/declarative/__init__.py @@ -37,6 +37,8 @@ "MCPToolResult", "ProviderLookupError", "ProviderTypeMapping", + "ToolApprovalRequest", + "ToolApprovalResponse", "WorkflowFactory", "WorkflowState", ] diff --git a/python/packages/core/agent_framework/declarative/__init__.pyi b/python/packages/core/agent_framework/declarative/__init__.pyi index bd6bf73fba6..c64e7304419 100644 --- a/python/packages/core/agent_framework/declarative/__init__.pyi +++ b/python/packages/core/agent_framework/declarative/__init__.pyi @@ -20,6 +20,8 @@ from agent_framework_declarative import ( MCPToolResult, ProviderLookupError, ProviderTypeMapping, + ToolApprovalRequest, + ToolApprovalResponse, WorkflowFactory, WorkflowState, ) @@ -44,6 +46,8 @@ __all__ = [ "MCPToolResult", "ProviderLookupError", "ProviderTypeMapping", + "ToolApprovalRequest", + "ToolApprovalResponse", "WorkflowFactory", "WorkflowState", ] diff --git a/python/packages/declarative/agent_framework_declarative/__init__.py b/python/packages/declarative/agent_framework_declarative/__init__.py index ad639fb5217..84bc404d5d8 100644 --- a/python/packages/declarative/agent_framework_declarative/__init__.py +++ b/python/packages/declarative/agent_framework_declarative/__init__.py @@ -19,6 +19,8 @@ MCPToolHandler, MCPToolInvocation, MCPToolResult, + ToolApprovalRequest, + ToolApprovalResponse, WorkflowFactory, WorkflowState, ) @@ -48,6 +50,8 @@ "MCPToolResult", "ProviderLookupError", "ProviderTypeMapping", + "ToolApprovalRequest", + "ToolApprovalResponse", "WorkflowFactory", "WorkflowState", "__version__", diff --git a/python/samples/03-workflows/declarative/invoke_mcp_tool/main.py b/python/samples/03-workflows/declarative/invoke_mcp_tool/main.py index c95b0c46912..5d08cd5bf09 100644 --- a/python/samples/03-workflows/declarative/invoke_mcp_tool/main.py +++ b/python/samples/03-workflows/declarative/invoke_mcp_tool/main.py @@ -11,6 +11,12 @@ 3. Bind the parsed tool result to a workflow variable and mirror it into the conversation via ``conversationId`` so a downstream Foundry agent can answer questions using only that context. + 4. Optionally pause the MCP tool call for human approval. The YAML reads + ``requireApproval`` from ``Workflow.Inputs.requireApproval`` so the + host can flip the behaviour without editing the workflow definition. + Set the ``MCP_REQUIRE_APPROVAL`` environment variable (``1`` / ``true`` + / ``yes``) to enable the approval flow; leave it unset for the + "fire-and-forget" default. Security note: ``DefaultMCPToolHandler`` connects to whatever MCP server URL the @@ -21,8 +27,16 @@ therefore share the same prompt-injection risk surface as ``HttpRequestAction``: only invoke MCP servers you trust. + The approval flow is also a defence-in-depth control: even with a + trusted server, requiring human approval lets a reviewer inspect + tool name, arguments, and outbound header NAMES (never values) + before any network call is made. + Run with: python -m samples.03-workflows.declarative.invoke_mcp_tool.main + +Run with approval prompts: + MCP_REQUIRE_APPROVAL=1 python -m samples.03-workflows.declarative.invoke_mcp_tool.main """ import asyncio @@ -32,6 +46,8 @@ from agent_framework import Agent from agent_framework.declarative import ( DefaultMCPToolHandler, + MCPToolApprovalRequest, + ToolApprovalResponse, WorkflowFactory, ) from agent_framework.foundry import FoundryChatClient @@ -44,6 +60,44 @@ concise and cite the relevant document title or URL when possible. """ +_TRUTHY = {"1", "true", "yes", "on"} + + +def _read_require_approval_flag() -> bool: + """Return True when the MCP_REQUIRE_APPROVAL env var requests approval.""" + return os.environ.get("MCP_REQUIRE_APPROVAL", "").strip().lower() in _TRUTHY + + +def _prompt_for_approval(request: MCPToolApprovalRequest) -> ToolApprovalResponse: + """Render the pending MCP call to stdout and read approve/reject from the user.""" + print() + print("-" * 60) + print("MCP tool approval required") + print("-" * 60) + print(f" tool: {request.tool_name}") + print(f" server label: {request.server_label or '(unset)'}") + print(f" server url: {request.server_url}") + if request.arguments: + print(" arguments:") + for key, value in request.arguments.items(): + print(f" {key}: {value!r}") + if request.header_names: + # Only NAMES are surfaced; values are intentionally withheld because + # they typically carry authentication secrets. + print(f" outbound header names: {', '.join(request.header_names)}") + else: + print(" outbound header names: (none)") + print("-" * 60) + + while True: + answer = input("Approve this MCP call? [y/N] ").strip().lower() # noqa: ASYNC250 + if answer in {"y", "yes"}: + return ToolApprovalResponse(approved=True) + if answer in {"", "n", "no"}: + reason = input("Reason for rejection (optional): ").strip() # noqa: ASYNC250 + return ToolApprovalResponse(approved=False, reason=reason or None) + print("Please answer 'y' or 'n'.") + async def main() -> None: """Run the invoke MCP tool workflow.""" @@ -63,6 +117,8 @@ async def main() -> None: agents = {"DocsAgent": docs_agent} + require_approval = _read_require_approval_flag() + # The default MCPToolHandler is sufficient for this sample because the # Microsoft Learn Docs MCP server is public and unauthenticated. For # authenticated servers, supply a ``client_provider`` callback to route @@ -80,6 +136,10 @@ async def main() -> None: print("=" * 60) print("Invoke MCP Tool Workflow Demo") + if require_approval: + print("(MCP_REQUIRE_APPROVAL is set — you will be prompted before the tool runs)") + else: + print("(set MCP_REQUIRE_APPROVAL=1 to enable the human-approval flow)") print("=" * 60) print() print("Ask one question that can be answered from the Microsoft Learn docs or provide a keyword to search.") @@ -89,11 +149,52 @@ async def main() -> None: if not user_input: user_input = "What is the Agent Framework declarative workflow runtime?" - print("\nAgent: ", end="", flush=True) - async for event in workflow.run(user_input, stream=True): - if event.type == "output" and isinstance(event.data, str): - print(event.data, end="", flush=True) - print() + # Drive the workflow via dict-shaped inputs so the YAML can read + # both the user's question (``Workflow.Inputs.text``) and the + # approval toggle (``Workflow.Inputs.requireApproval``) without + # any Python-side mutation of the workflow definition. + workflow_inputs: dict[str, object] = { + "text": user_input, + "requireApproval": require_approval, + } + + # The request_info loop below handles the MCP approval flow when + # the YAML requests it. When ``requireApproval`` is false the + # workflow never emits an ``MCPToolApprovalRequest`` event, so + # the loop runs exactly once and exits cleanly — both modes share + # the same code path. + pending: tuple[str, MCPToolApprovalRequest] | None = None + produced_output = False + printed_agent_prefix = False + + while True: + if pending is None: + stream = workflow.run(workflow_inputs, stream=True) + else: + pending_id, pending_request = pending + response = _prompt_for_approval(pending_request) + stream = workflow.run(stream=True, responses={pending_id: response}) + pending = None + + async for event in stream: + if event.type == "output" and isinstance(event.data, str): + if not printed_agent_prefix: + print("\nAgent: ", end="", flush=True) + printed_agent_prefix = True + print(event.data, end="", flush=True) + produced_output = True + elif event.type == "request_info" and isinstance(event.data, MCPToolApprovalRequest): + pending = (event.request_id, event.data) + + if pending is None: + if not produced_output: + # Workflow finished without producing any agent output + # (e.g. the user rejected the MCP tool call and the + # downstream agent had nothing to summarise). + print("\n(no response produced)") + else: + print() + break if __name__ == "__main__": diff --git a/python/samples/03-workflows/declarative/invoke_mcp_tool/workflow.yaml b/python/samples/03-workflows/declarative/invoke_mcp_tool/workflow.yaml index b83dc052ffa..55f9f0754d4 100644 --- a/python/samples/03-workflows/declarative/invoke_mcp_tool/workflow.yaml +++ b/python/samples/03-workflows/declarative/invoke_mcp_tool/workflow.yaml @@ -19,6 +19,12 @@ # How do I configure logging in the Agent Framework? # Gpt-5.4-mini # +# Workflow inputs (set by the host via ``workflow.run({...})``): +# text: The user's question (required). +# requireApproval: Optional bool. When true, the MCP tool call pauses for +# human approval before contacting the server. Defaults +# to false when omitted. +# kind: Workflow trigger: @@ -31,18 +37,23 @@ trigger: - kind: SetVariable id: capture_query variable: Local.SearchQuery - value: =System.LastMessage.Text + value: =Workflow.Inputs.text # Invoke microsoft_docs_search on the Microsoft Learn Docs MCP server. # The result is parsed into Local.SearchResults and also added to the # conversation (via conversationId) so the agent below can answer the # user's question based on it. + # + # ``requireApproval`` reads from Workflow.Inputs so the host can toggle + # the human-approval flow without editing this YAML. When the input is + # absent or evaluates to a falsy value, the tool runs without pausing. - kind: InvokeMcpTool id: search_docs conversationId: =System.ConversationId serverUrl: https://learn.microsoft.com/api/mcp serverLabel: MicrosoftLearnDocs toolName: microsoft_docs_search + requireApproval: =Workflow.Inputs.requireApproval arguments: query: =Local.SearchQuery output: @@ -51,14 +62,16 @@ trigger: # Use the agent to answer the user's question using the conversation # context (which now contains the MCP search results). The user's - # original message is already in the conversation as System.LastMessage, - # and the executor's input fallback chain extracts its ``Text`` field - # automatically when ``input.messages`` is omitted. + # question is supplied via ``input.messages`` (sourced from the workflow + # inputs), and the prior conversation history is bound via + # ``conversationId``. - kind: InvokeAzureAgent id: answer_question conversationId: =System.ConversationId agent: name: DocsAgent + input: + messages: =Workflow.Inputs.text output: autoSend: true messages: Local.AgentResponse From ea0b3c12107e128aec7af08b87b1d8dd13762fcf Mon Sep 17 00:00:00 2001 From: Peter Ibekwe Date: Mon, 4 May 2026 15:08:04 -0700 Subject: [PATCH 7/8] Fix cache and PR comments --- .../_workflows/_executors_mcp.py | 9 +- .../_workflows/_mcp_handler.py | 122 ++++++++++++--- .../tests/test_default_mcp_tool_handler.py | 140 +++++++++++++++++- .../tests/test_invoke_mcp_tool_executor.py | 33 +++++ 4 files changed, 271 insertions(+), 33 deletions(-) diff --git a/python/packages/declarative/agent_framework_declarative/_workflows/_executors_mcp.py b/python/packages/declarative/agent_framework_declarative/_workflows/_executors_mcp.py index a1501a3cb16..73b66341ea3 100644 --- a/python/packages/declarative/agent_framework_declarative/_workflows/_executors_mcp.py +++ b/python/packages/declarative/agent_framework_declarative/_workflows/_executors_mcp.py @@ -163,15 +163,18 @@ def _get_output_path(action_def: Mapping[str, Any], key: str) -> str | None: def _format_outputs_for_send(parsed_results: list[Any]) -> str: """Render parsed MCP outputs to a string for ``ctx.yield_output(...)``. + - Empty list → ``""``. - All-string list → newline-joined. - - Single dict / list → JSON. - - Empty / mixed → JSON-dump the whole list. + - Single element (any type — scalar, dict, list) → JSON-dumped element. + This avoids surprising ``"[42]"`` / ``"[true]"`` / ``"[null]"`` when + an MCP tool returns a single scalar JSON value. + - Multi-element non-string list → JSON-dump the whole list. """ if not parsed_results: return "" if all(isinstance(item, str) for item in parsed_results): return "\n".join(parsed_results) # type: ignore[arg-type] - if len(parsed_results) == 1 and isinstance(parsed_results[0], (dict, list)): + if len(parsed_results) == 1: return json.dumps(parsed_results[0], ensure_ascii=False) return json.dumps(parsed_results, ensure_ascii=False) diff --git a/python/packages/declarative/agent_framework_declarative/_workflows/_mcp_handler.py b/python/packages/declarative/agent_framework_declarative/_workflows/_mcp_handler.py index 5ea17aabc1b..658ce42c232 100644 --- a/python/packages/declarative/agent_framework_declarative/_workflows/_mcp_handler.py +++ b/python/packages/declarative/agent_framework_declarative/_workflows/_mcp_handler.py @@ -158,10 +158,17 @@ class DefaultMCPToolHandler: """Default :class:`MCPToolHandler` backed by :class:`agent_framework.MCPStreamableHTTPTool`. Caches one :class:`agent_framework.MCPStreamableHTTPTool` instance per - ``(server_url, headers_hash)`` in a bounded LRU. The cache prevents - re-establishing an MCP session for every invocation while ensuring - different header sets (auth tokens) cannot share a session — matches the - .NET design intent while bounding cardinality. + ``(server_url, server_label, connection_name, headers_hash)`` in a + bounded LRU. The cache prevents re-establishing an MCP session for every + invocation while ensuring different header sets (auth tokens) cannot + share a session — matches the .NET design intent while bounding + cardinality. ``server_label`` and ``connection_name`` participate in + the key so that callers using ``client_provider`` to dispatch on those + fields receive a fresh client per logical connection (see below). + Header *names* are lower-cased inside the hash payload only — the + headers passed on the wire keep the caller's original casing — so two + YAML actions that spell ``Authorization`` differently still share a + cache entry. Construction modes: @@ -197,14 +204,17 @@ def __init__( raise ValueError(f"cache_max_size must be positive, got {cache_max_size}") self._client_provider = client_provider self._cache_max_size = cache_max_size - self._cache: OrderedDict[tuple[str, str], _CacheEntry] = OrderedDict() + self._cache: OrderedDict[tuple[str, str, str, str], _CacheEntry] = OrderedDict() # Outer lock guards the cache + in-flight-future map only — never # held across network I/O. self._cache_lock = asyncio.Lock() # Per-key in-flight futures: while one task is connecting, other # tasks awaiting the same key will await the same future and share # the resulting cache entry. - self._inflight: dict[tuple[str, str], asyncio.Future[_CacheEntry]] = {} + self._inflight: dict[tuple[str, str, str, str], asyncio.Future[_CacheEntry]] = {} + # Set by ``aclose`` to prevent post-close cache insertions and to + # reject new ``invoke_tool`` calls. Once set, never cleared. + self._closed = False async def invoke_tool(self, invocation: MCPToolInvocation) -> MCPToolResult: """Invoke ``invocation.tool_name`` on the cached MCP client for the server.""" @@ -279,10 +289,32 @@ async def aclose(self) -> None: Caller-supplied :class:`httpx.AsyncClient` instances (returned by the ``client_provider`` callback) are NOT closed. + + Idempotent — a second call returns immediately. Drains any in-flight + ``_create_entry`` tasks before returning so their resources are + cleaned up; the in-flight tasks see ``self._closed`` in phase 3 of + :meth:`_get_or_create_entry`, close their own entry, and resolve + their future with ``RuntimeError("DefaultMCPToolHandler is closed")``. """ async with self._cache_lock: + if self._closed: + return + self._closed = True entries = list(self._cache.values()) self._cache.clear() + inflight_futures = list(self._inflight.values()) + + # Wait for in-flight creations to finish their self-cleanup. Each + # in-flight task self-closes its entry under the closed-flag branch + # in phase 3 and resolves its future with ``RuntimeError``; we + # swallow it here because the failure is expected at shutdown. + for fut in inflight_futures: + try: + await fut + except BaseException: + logger.debug("DefaultMCPToolHandler: in-flight future raised during aclose", exc_info=True) + continue + for entry in entries: await self._close_entry(entry) @@ -298,12 +330,19 @@ async def __aexit__(self, exc_type: Any, exc: Any, tb: Any) -> None: async def _get_or_create_entry(self, invocation: MCPToolInvocation) -> _CacheEntry: """Look up (or create) the cached MCP client for this invocation.""" - key = self._cache_key(invocation.server_url, invocation.headers) + key = self._cache_key( + invocation.server_url, + invocation.server_label, + invocation.connection_name, + invocation.headers, + ) # Phase 1: check the cache and either claim creation or wait for an # already in-flight creation. creating = False async with self._cache_lock: + if self._closed: + raise RuntimeError("DefaultMCPToolHandler is closed") existing = self._cache.get(key) if existing is not None: self._cache.move_to_end(key) @@ -332,25 +371,44 @@ async def _get_or_create_entry(self, invocation: MCPToolInvocation) -> _CacheEnt raise # Phase 3: insert with LRU eviction; resolve the in-flight future. + # If ``aclose`` ran while we were connecting, ``_closed`` is now + # True; don't insert into the cache (it has been drained), close + # the just-built entry, and surface the closed-handler error to + # all awaiters of the future. evicted: _CacheEntry | None = None duplicate: _CacheEntry | None = None + handler_closed = False async with self._cache_lock: self._inflight.pop(key, None) - existing = self._cache.get(key) - if existing is not None: - # Another writer beat us; prefer the existing entry and - # discard ours after the lock is released. - self._cache.move_to_end(key) - duplicate = entry - entry = existing + if self._closed: + handler_closed = True else: - self._cache[key] = entry - self._cache.move_to_end(key) - if len(self._cache) > self._cache_max_size: - _evicted_key, evicted = self._cache.popitem(last=False) + existing = self._cache.get(key) + if existing is not None: + # Another writer beat us; prefer the existing entry and + # discard ours after the lock is released. + self._cache.move_to_end(key) + duplicate = entry + entry = existing + else: + self._cache[key] = entry + self._cache.move_to_end(key) + if len(self._cache) > self._cache_max_size: + _evicted_key, evicted = self._cache.popitem(last=False) + if not inflight.done(): + inflight.set_result(entry) + + if handler_closed: + # Close our orphaned entry; resolve the future with a clear + # error so the caller (and any other awaiters) surface a + # consistent "handler is closed" failure rather than receiving + # an entry we are about to close behind their back. + await self._close_entry(entry) + err = RuntimeError("DefaultMCPToolHandler is closed") if not inflight.done(): - inflight.set_result(entry) - + inflight.set_exception(err) + inflight.exception() + raise err if duplicate is not None: await self._close_entry(duplicate) if evicted is not None: @@ -410,11 +468,27 @@ async def _close_entry(self, entry: _CacheEntry) -> None: logger.debug("DefaultMCPToolHandler: error closing owned httpx client", exc_info=True) @staticmethod - def _cache_key(server_url: str, headers: dict[str, str] | None) -> tuple[str, str]: - """Build an order-independent cache key for ``(server_url, headers)``.""" + def _cache_key( + server_url: str, + server_label: str | None, + connection_name: str | None, + headers: dict[str, str] | None, + ) -> tuple[str, str, str, str]: + """Build an order-independent cache key for the invocation identity. + + The key includes ``server_label`` and ``connection_name`` so that + callers using ``client_provider`` to dispatch on those fields + receive a fresh client per logical connection (matches the + documented dispatch contract). + + Header *names* are lower-cased inside the hash payload only so + that ``Authorization`` and ``authorization`` map to the same + cache entry. Header values remain case-sensitive (per RFC 7235). + """ if not headers: headers_hash = "0" else: - payload = json.dumps(sorted(headers.items()), ensure_ascii=False) + normalized = sorted((k.lower(), v) for k, v in headers.items()) + payload = json.dumps(normalized, ensure_ascii=False) headers_hash = hashlib.sha256(payload.encode("utf-8")).hexdigest() - return (server_url, headers_hash) + return (server_url, server_label or "", connection_name or "", headers_hash) diff --git a/python/packages/declarative/tests/test_default_mcp_tool_handler.py b/python/packages/declarative/tests/test_default_mcp_tool_handler.py index c40d275e807..3a5c67e1d63 100644 --- a/python/packages/declarative/tests/test_default_mcp_tool_handler.py +++ b/python/packages/declarative/tests/test_default_mcp_tool_handler.py @@ -237,6 +237,54 @@ async def slow_connect(self: FakeTool) -> None: assert len(FakeTool.instances) == 1 assert FakeTool.instances[0].connect_count == 1 + @pytest.mark.asyncio + async def test_different_connection_names_create_separate_entries(self) -> None: + """Same URL/headers but different ``connection_name`` must dispatch separately.""" + handler = DefaultMCPToolHandler() + with _patch_tool(): + await handler.invoke_tool(_invocation(connection_name="conn-A")) + await handler.invoke_tool(_invocation(connection_name="conn-B")) + assert len(FakeTool.instances) == 2 + + @pytest.mark.asyncio + async def test_different_server_labels_create_separate_entries(self) -> None: + """Same URL/headers but different ``server_label`` must dispatch separately.""" + handler = DefaultMCPToolHandler() + with _patch_tool(): + await handler.invoke_tool(_invocation(server_label="LabelA")) + await handler.invoke_tool(_invocation(server_label="LabelB")) + assert len(FakeTool.instances) == 2 + + @pytest.mark.asyncio + async def test_full_identity_match_hits_cache(self) -> None: + """All four identity components match → single cached entry.""" + handler = DefaultMCPToolHandler() + with _patch_tool(): + await handler.invoke_tool(_invocation(server_label="Lbl", connection_name="C", headers={"X": "1"})) + await handler.invoke_tool(_invocation(server_label="Lbl", connection_name="C", headers={"X": "1"})) + assert len(FakeTool.instances) == 1 + assert FakeTool.instances[0].connect_count == 1 + + @pytest.mark.asyncio + async def test_header_name_case_collapses_to_one_cache_entry(self) -> None: + """Header name spelling differences (case-only) must share a cache entry.""" + handler = DefaultMCPToolHandler() + with _patch_tool(): + await handler.invoke_tool(_invocation(headers={"Authorization": "tk"})) + await handler.invoke_tool(_invocation(headers={"authorization": "tk"})) + await handler.invoke_tool(_invocation(headers={"AUTHORIZATION": "tk"})) + assert len(FakeTool.instances) == 1 + assert FakeTool.instances[0].connect_count == 1 + + @pytest.mark.asyncio + async def test_header_value_case_does_not_collapse(self) -> None: + """Header *values* remain case-sensitive (different tokens → different sessions).""" + handler = DefaultMCPToolHandler() + with _patch_tool(): + await handler.invoke_tool(_invocation(headers={"Authorization": "Bearer-A"})) + await handler.invoke_tool(_invocation(headers={"Authorization": "bearer-a"})) + assert len(FakeTool.instances) == 2 + # ---------- Aclose semantics ---------------------------------------------- @@ -280,6 +328,66 @@ async def test_async_context_manager(self) -> None: tool = FakeTool.instances[0] assert tool.close_count == 1 + @pytest.mark.asyncio + async def test_aclose_is_idempotent(self) -> None: + """A second ``aclose`` is a no-op (no exception, no double-close).""" + handler = DefaultMCPToolHandler() + with _patch_tool(): + await handler.invoke_tool(_invocation(headers={"X": "1"})) + await handler.aclose() + await handler.aclose() + assert FakeTool.instances[0].close_count == 1 + + @pytest.mark.asyncio + async def test_invoke_after_close_returns_error_result(self) -> None: + """Post-close ``invoke_tool`` surfaces a tool error rather than crashing.""" + handler = DefaultMCPToolHandler() + with _patch_tool(): + await handler.aclose() + result = await handler.invoke_tool(_invocation()) + assert result.is_error is True + assert "closed" in (result.error_message or "").lower() + + @pytest.mark.asyncio + async def test_aclose_drains_inflight_creation(self) -> None: + """An in-flight ``_create_entry`` must not leak when ``aclose`` races with it. + + Reproduces the race described in PR #5630 review-comment 3: + task A claims an inflight future and starts a slow connect; task B + runs ``aclose``; task A must self-clean (close its tool + httpx + client) and surface a closed-handler error rather than orphaning + the entry. + """ + handler = DefaultMCPToolHandler() + connect_started = asyncio.Event() + release_connect = asyncio.Event() + original_connect = FakeTool.connect + + async def gated_connect(self: FakeTool) -> None: + connect_started.set() + await release_connect.wait() + await original_connect(self) + + with _patch_tool(), patch.object(FakeTool, "connect", gated_connect): + invoke_task = asyncio.create_task(handler.invoke_tool(_invocation(headers={"X": "1"}))) + # Wait until task A is mid-connect. + await connect_started.wait() + # Race: kick off aclose. It must wait for the in-flight task. + close_task = asyncio.create_task(handler.aclose()) + # Yield once to ensure aclose has set _closed and is awaiting. + await asyncio.sleep(0) + # Allow the connect to complete; phase 3 sees _closed and self-cleans. + release_connect.set() + result = await invoke_task + await close_task + + # Entry was created and then closed by the in-flight task itself. + assert len(FakeTool.instances) == 1 + assert FakeTool.instances[0].close_count == 1 + # The originating invocation surfaces a closed-handler error. + assert result.is_error is True + assert "closed" in (result.error_message or "").lower() + # ---------- Result normalisation ------------------------------------------ @@ -400,16 +508,36 @@ def boom(**_a: Any) -> Any: class TestCacheKey: def test_key_order_independent(self) -> None: - k1 = DefaultMCPToolHandler._cache_key("https://x/", {"A": "1", "B": "2"}) - k2 = DefaultMCPToolHandler._cache_key("https://x/", {"B": "2", "A": "1"}) + k1 = DefaultMCPToolHandler._cache_key("https://x/", None, None, {"A": "1", "B": "2"}) + k2 = DefaultMCPToolHandler._cache_key("https://x/", None, None, {"B": "2", "A": "1"}) assert k1 == k2 def test_key_distinguishes_values(self) -> None: - k1 = DefaultMCPToolHandler._cache_key("https://x/", {"A": "1"}) - k2 = DefaultMCPToolHandler._cache_key("https://x/", {"A": "2"}) + k1 = DefaultMCPToolHandler._cache_key("https://x/", None, None, {"A": "1"}) + k2 = DefaultMCPToolHandler._cache_key("https://x/", None, None, {"A": "2"}) assert k1 != k2 def test_empty_headers_use_fixed_hash(self) -> None: - k1 = DefaultMCPToolHandler._cache_key("https://x/", None) - k2 = DefaultMCPToolHandler._cache_key("https://x/", {}) + k1 = DefaultMCPToolHandler._cache_key("https://x/", None, None, None) + k2 = DefaultMCPToolHandler._cache_key("https://x/", None, None, {}) assert k1 == k2 + + def test_key_distinguishes_connection_name(self) -> None: + k1 = DefaultMCPToolHandler._cache_key("https://x/", None, "conn-A", None) + k2 = DefaultMCPToolHandler._cache_key("https://x/", None, "conn-B", None) + assert k1 != k2 + + def test_key_distinguishes_server_label(self) -> None: + k1 = DefaultMCPToolHandler._cache_key("https://x/", "Lbl-A", None, None) + k2 = DefaultMCPToolHandler._cache_key("https://x/", "Lbl-B", None, None) + assert k1 != k2 + + def test_key_collapses_header_name_case(self) -> None: + k1 = DefaultMCPToolHandler._cache_key("https://x/", None, None, {"Authorization": "tk"}) + k2 = DefaultMCPToolHandler._cache_key("https://x/", None, None, {"authorization": "tk"}) + assert k1 == k2 + + def test_key_keeps_header_value_case(self) -> None: + k1 = DefaultMCPToolHandler._cache_key("https://x/", None, None, {"X": "Bearer-A"}) + k2 = DefaultMCPToolHandler._cache_key("https://x/", None, None, {"X": "bearer-a"}) + assert k1 != k2 diff --git a/python/packages/declarative/tests/test_invoke_mcp_tool_executor.py b/python/packages/declarative/tests/test_invoke_mcp_tool_executor.py index 867b8543bbb..fdee1f7df1d 100644 --- a/python/packages/declarative/tests/test_invoke_mcp_tool_executor.py +++ b/python/packages/declarative/tests/test_invoke_mcp_tool_executor.py @@ -629,3 +629,36 @@ class TestProtocol: def test_stub_handler_satisfies_protocol(self) -> None: handler = StubMcpHandler(_ok()) assert isinstance(handler, MCPToolHandler) + + +# ---------- _format_outputs_for_send -------------------------------------- + + +class TestFormatOutputsForSend: + """Direct tests for the auto-send rendering helper. + + Regression for PR #5630 review-comment 4: a single scalar JSON value + must render bare (e.g. ``"42"``) rather than wrapped (``"[42]"``). + """ + + @pytest.mark.parametrize( + ("parsed", "expected"), + [ + ([], ""), + (["hello"], "hello"), + (["a", "b"], "a\nb"), + ([42], "42"), + ([3.14], "3.14"), + ([True], "true"), + ([False], "false"), + ([None], "null"), + ([{"k": "v"}], '{"k": "v"}'), + ([[1, 2]], "[1, 2]"), + (["hello", 42], '["hello", 42]'), + ([{"a": 1}, {"b": 2}], '[{"a": 1}, {"b": 2}]'), + ], + ) + def test_format_outputs_for_send(self, parsed: list[Any], expected: str) -> None: + from agent_framework_declarative._workflows._executors_mcp import _format_outputs_for_send + + assert _format_outputs_for_send(parsed) == expected From 745f9a717865e9fb7301438d7ac8ae261b9f74ca Mon Sep 17 00:00:00 2001 From: Peter Ibekwe <109177538+peibekwe@users.noreply.github.com> Date: Tue, 5 May 2026 08:09:12 -0700 Subject: [PATCH 8/8] Update python/samples/03-workflows/declarative/invoke_mcp_tool/main.py Co-authored-by: Eduard van Valkenburg --- python/samples/03-workflows/declarative/invoke_mcp_tool/main.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/python/samples/03-workflows/declarative/invoke_mcp_tool/main.py b/python/samples/03-workflows/declarative/invoke_mcp_tool/main.py index 5d08cd5bf09..85b513b5620 100644 --- a/python/samples/03-workflows/declarative/invoke_mcp_tool/main.py +++ b/python/samples/03-workflows/declarative/invoke_mcp_tool/main.py @@ -33,7 +33,7 @@ before any network call is made. Run with: - python -m samples.03-workflows.declarative.invoke_mcp_tool.main + python samples/03-workflows/declarative/invoke_mcp_tool/main.py Run with approval prompts: MCP_REQUIRE_APPROVAL=1 python -m samples.03-workflows.declarative.invoke_mcp_tool.main