diff --git a/mcpgateway/auth.py b/mcpgateway/auth.py index 01ed674ee1..423905c66c 100644 --- a/mcpgateway/auth.py +++ b/mcpgateway/auth.py @@ -84,7 +84,7 @@ from mcpgateway.config import settings from mcpgateway.db import EmailUser, fresh_db_session, SessionLocal from mcpgateway.plugins import get_plugin_manager -from mcpgateway.plugins.utils import build_request_extensions, record_plugin_metrics +from mcpgateway.plugins.utils import build_hook_extensions, record_plugin_metrics from mcpgateway.services.observability_service import current_trace_id from mcpgateway.transports.context import UserContext from mcpgateway.utils.correlation_id import get_correlation_id @@ -1472,8 +1472,8 @@ def _set_trace_for_user(user_obj: EmailUser, *, teams: Any = _UNSET, auth_method context_table = getattr(request.state, "plugin_context_table", None) if request else None - # Invoke custom auth resolution hook - # violations_as_exceptions=True so PluginViolationError is raised for explicit denials + # Dual-write: HttpAuthResolveUserPayload.headers is still required in cpex 0.1.1; + # real data lives on extensions.http.headers. auth_result, context_table_result = await plugin_manager.invoke_hook( HttpHookType.HTTP_AUTH_RESOLVE_USER, payload=HttpAuthResolveUserPayload( @@ -1485,7 +1485,7 @@ def _set_trace_for_user(user_obj: EmailUser, *, teams: Any = _UNSET, auth_method global_context=global_context, local_contexts=context_table, violations_as_exceptions=True, # Raise PluginViolationError for auth denials - extensions=build_request_extensions(), + extensions=build_hook_extensions(headers), ) record_plugin_metrics(current_trace_id.get(), auth_result.metadata) diff --git a/mcpgateway/middleware/http_auth_middleware.py b/mcpgateway/middleware/http_auth_middleware.py index 6aef6d0ac0..2a748207d2 100644 --- a/mcpgateway/middleware/http_auth_middleware.py +++ b/mcpgateway/middleware/http_auth_middleware.py @@ -23,7 +23,7 @@ # First-Party from mcpgateway.config import settings from mcpgateway.plugins import get_plugin_manager -from mcpgateway.plugins.utils import build_request_extensions, record_plugin_metrics +from mcpgateway.plugins.utils import build_hook_extensions, headers_from_modified_extensions, record_plugin_metrics from mcpgateway.services.observability_service import current_trace_id from mcpgateway.utils.correlation_id import generate_correlation_id, get_correlation_id from mcpgateway.utils.verify_credentials import _resolve_auth_header_name @@ -69,6 +69,9 @@ async def run_pre_request_hooks( global_context = GlobalContext(request_id=request_id, server_id=None, tenant_id=None, content_type=content_type) try: + ext_in = build_hook_extensions(headers) + # Dual-write: HttpPreRequestPayload.headers is still required in cpex 0.1.1; + # real data lives on extensions.http.headers. pre_result, context_table = await plugin_manager.invoke_hook( HttpHookType.HTTP_PRE_REQUEST, payload=HttpPreRequestPayload( @@ -81,15 +84,18 @@ async def run_pre_request_hooks( global_context=global_context, local_contexts=None, violations_as_exceptions=False, - extensions=build_request_extensions(), + extensions=ext_in, ) record_plugin_metrics(current_trace_id.get(), pre_result.metadata) - if not pre_result.modified_payload: + plugin_headers = headers_from_modified_extensions(pre_result) + if plugin_headers is not None: + modified_headers_dict = plugin_headers + elif pre_result.modified_payload: + modified_headers_dict = pre_result.modified_payload.root + else: return headers, global_context, context_table - modified_headers_dict = pre_result.modified_payload.root - # Security: prevent plugin hooks from overriding auth-sensitive # headers that were already present on the inbound request. # Plugins MAY create new auth headers (e.g. x-api-key → authorization @@ -251,7 +257,7 @@ async def dispatch(self, request: Request, call_next): global_context=global_context, local_contexts=context_table, violations_as_exceptions=False, - extensions=build_request_extensions(), + extensions=build_hook_extensions(dict(request.headers)), ) record_plugin_metrics(current_trace_id.get(), post_result.metadata) diff --git a/mcpgateway/plugins/utils.py b/mcpgateway/plugins/utils.py index b37f2e2736..43cd3bc9ab 100644 --- a/mcpgateway/plugins/utils.py +++ b/mcpgateway/plugins/utils.py @@ -11,10 +11,11 @@ import logging import math import re +from collections.abc import Mapping from typing import Any, Dict, Optional # Third-Party -from cpex.framework.extensions import Extensions, RequestExtension +from cpex.framework.extensions import Extensions, HttpExtension, RequestExtension logger = logging.getLogger(__name__) @@ -432,6 +433,51 @@ def build_request_extensions() -> Optional[Extensions]: return None +def build_hook_extensions(headers: Mapping[str, str] | None = None) -> Optional[Extensions]: + """Build hook ``Extensions`` with optional HTTP headers merged onto request/trace context. + + Prefer this over putting headers on deprecated payload ``.headers`` fields. + Headers land on ``extensions.http.headers`` (``HttpExtension``). When ``headers`` + is empty/None, behaves like :func:`build_request_extensions`. + + Args: + headers: Optional HTTP headers to expose to plugins via ``extensions.http``. + + Returns: + ``Extensions`` with ``http`` and/or ``request`` populated, or ``None`` when + there is neither a trace context nor headers. + """ + base = build_request_extensions() + if not headers: + return base + http = HttpExtension(headers={str(k): str(v) for k, v in headers.items()}) + if base is None: + return Extensions(http=http) + return base.model_copy(update={"http": http}) + + +def headers_from_modified_extensions(result: Any) -> Optional[Dict[str, str]]: + """Extract ``extensions.http.headers`` from a plugin hook result, if present. + + Returns ``None`` when the result has no real ``modified_extensions.http`` + (including MagicMock stand-ins used in unit tests), so callers can fall back + to legacy ``modified_payload`` paths. + + Args: + result: Plugin hook result (``PluginResult`` or test double). + + Returns: + Header dict, or ``None`` if unavailable. + """ + ext = getattr(result, "modified_extensions", None) + if not isinstance(ext, Extensions): + return None + http = getattr(ext, "http", None) + if not isinstance(http, HttpExtension): + return None + return dict(http.headers) + + def apply_attribute_mapping(attributes: dict[str, Any], mapping: dict[str, str]) -> dict[str, Any]: """Apply attribute name mapping (renaming) to a dictionary of attributes. diff --git a/mcpgateway/services/a2a_service.py b/mcpgateway/services/a2a_service.py index e3edaa480a..bc8e70d762 100644 --- a/mcpgateway/services/a2a_service.py +++ b/mcpgateway/services/a2a_service.py @@ -35,7 +35,7 @@ from mcpgateway.db import fresh_db_session, get_for_update from mcpgateway.db import Tool as DbTool from mcpgateway.observability import create_span, set_span_attribute, set_span_error -from mcpgateway.plugins.utils import build_request_extensions, record_plugin_metrics +from mcpgateway.plugins.utils import build_hook_extensions, build_request_extensions, headers_from_modified_extensions, record_plugin_metrics from mcpgateway.schemas import A2AAgentAggregateMetrics, A2AAgentCreate, A2AAgentMetrics, A2AAgentRead, A2AAgentUpdate from mcpgateway.services.a2a_protocol import prepare_a2a_invocation from mcpgateway.services.base_service import BaseService @@ -2250,7 +2250,6 @@ async def invoke_agent( AgentHookType, AgentPreInvokePayload, GlobalContext, - HttpHeaderPayload, PluginViolationError, ) @@ -2292,44 +2291,44 @@ async def invoke_agent( # Fire pre-invoke hook — can modify parameters, headers, and agent metadata if plugin_manager and plugin_manager.has_hooks_for(AgentHookType.AGENT_PRE_INVOKE): try: + ext_in = build_hook_extensions(plugin_headers) pre_result, context_table = await plugin_manager.invoke_hook( AgentHookType.AGENT_PRE_INVOKE, payload=AgentPreInvokePayload( agent_id=agent_id, messages=[{"role": "user", "content": parameters}] if parameters else [], - headers=HttpHeaderPayload(root=plugin_headers), parameters=parameters if isinstance(parameters, dict) else {}, ), global_context=global_context, local_contexts=context_table, violations_as_exceptions=True, - extensions=build_request_extensions(), + extensions=ext_in, ) record_plugin_metrics(current_trace_id.get(), pre_result.metadata) if pre_result.modified_payload: if pre_result.modified_payload.parameters is not None: parameters = pre_result.modified_payload.parameters - if pre_result.modified_payload.headers is not None: - # Security: Re-filter plugin-returned headers to prevent malicious - # plugins from injecting sensitive headers into downstream requests - # (PR #5183 review fix) - plugin_returned = pre_result.modified_payload.headers.model_dump() - safe_headers = self._refilter_plugin_headers( - plugin_headers=plugin_returned, - agent=agent, - feature_flag_enabled=settings.enable_sensitive_header_passthrough, + plugin_returned = headers_from_modified_extensions(pre_result) + if plugin_returned is not None: + # Security: Re-filter plugin-returned headers to prevent malicious + # plugins from injecting sensitive headers into downstream requests + # (PR #5183 review fix) + safe_headers = self._refilter_plugin_headers( + plugin_headers=plugin_returned, + agent=agent, + feature_flag_enabled=settings.enable_sensitive_header_passthrough, + ) + prepared.headers.update(safe_headers) + + # Log security-blocked headers for forensic awareness + if plugin_returned.keys() - safe_headers.keys(): + removed = sorted(plugin_returned.keys() - safe_headers.keys()) + logger.warning( + "Plugin attempted to set headers blocked by security policy: %s (agent=%s, flag=%s)", + removed, + agent.name, + settings.enable_sensitive_header_passthrough, ) - prepared.headers.update(safe_headers) - - # Log security-blocked headers for forensic awareness - if plugin_returned.keys() - safe_headers.keys(): - removed = sorted(plugin_returned.keys() - safe_headers.keys()) - logger.warning( - "Plugin attempted to set headers blocked by security policy: %s (agent=%s, flag=%s)", - removed, - agent.name, - settings.enable_sensitive_header_passthrough, - ) except PluginViolationError as e: logger.error("Plugin RBAC violation for A2A agent %s: %s", agent_id, e) raise A2AAgentError(f"Plugin RBAC violation: {e}") from e diff --git a/mcpgateway/services/tool_service.py b/mcpgateway/services/tool_service.py index b24727b818..9f16bc0a66 100644 --- a/mcpgateway/services/tool_service.py +++ b/mcpgateway/services/tool_service.py @@ -35,7 +35,6 @@ import anyio from cpex.framework import ( GlobalContext, - HttpHeaderPayload, PluginContextTable, PluginError, PluginViolationError, @@ -72,7 +71,7 @@ from mcpgateway.db import Tool as DbTool from mcpgateway.db import ToolMetric, ToolMetricsHourly from mcpgateway.observability import create_child_span, create_span, inject_trace_context_headers, otel_context_active, set_span_attribute, set_span_error -from mcpgateway.plugins.utils import build_request_extensions, record_plugin_metrics +from mcpgateway.plugins.utils import build_hook_extensions, build_request_extensions, headers_from_modified_extensions, record_plugin_metrics from mcpgateway.schemas import AuthenticationValues, ToolCreate, ToolMetrics, ToolRead, ToolUpdate, TopPerformer from mcpgateway.services.a2a_protocol import prepare_a2a_invocation from mcpgateway.services.audit_trail_service import get_audit_trail_service @@ -283,7 +282,7 @@ def _header_payload_keys(value: Any) -> Optional[list[str]]: return _mapping_keys(root) -def _log_tool_pre_invoke_result(tool_name: str, original_args: Any, original_headers: Any, pre_result: Any) -> None: +def _log_tool_pre_invoke_result(tool_name: str, original_args: Any, original_headers: Any, pre_result: Any, hook_extensions: Any = None) -> None: """Log sanitized TOOL_PRE_INVOKE output shape for plugin diagnostics.""" try: if not logger.isEnabledFor(logging.DEBUG): @@ -292,9 +291,9 @@ def _log_tool_pre_invoke_result(tool_name: str, original_args: Any, original_hea sanitized_tool_name = sanitize_for_log(tool_name) modified_payload = getattr(pre_result, "modified_payload", None) before_arg_keys = _mapping_keys(original_args) - before_header_keys = _header_payload_keys(original_headers) + before_header_keys = _mapping_keys(original_headers) if isinstance(original_headers, dict) else _header_payload_keys(original_headers) - if modified_payload is None: + if modified_payload is None and getattr(pre_result, "modified_extensions", None) is None: logger.debug( "tool_pre_invoke completed for %s: modified_payload=None, arg_keys_before=%s, header_keys_before=%s", sanitized_tool_name, @@ -303,19 +302,24 @@ def _log_tool_pre_invoke_result(tool_name: str, original_args: Any, original_hea ) return - after_arg_keys = _mapping_keys(getattr(modified_payload, "args", None)) - after_header_keys = _header_payload_keys(getattr(modified_payload, "headers", None)) + after_arg_keys = _mapping_keys(getattr(modified_payload, "args", None)) if modified_payload is not None else before_arg_keys + after_headers = headers_from_modified_extensions(pre_result) + if after_headers is not None: + after_header_keys = _mapping_keys(after_headers) + else: + after_header_keys = _header_payload_keys(getattr(modified_payload, "headers", None)) if modified_payload is not None else before_header_keys before_arg_set = set(before_arg_keys or []) after_arg_set = set(after_arg_keys or []) before_header_set = set(before_header_keys or []) after_header_set = set(after_header_keys or []) - modified_name = sanitize_for_log(getattr(modified_payload, "name", None)) + modified_name = sanitize_for_log(getattr(modified_payload, "name", None)) if modified_payload is not None else None logger.debug( - "tool_pre_invoke completed for %s: modified_payload=True, modified_name=%s, " + "tool_pre_invoke completed for %s: modified_payload=%s, modified_name=%s, " "arg_keys_before=%s, arg_keys_after=%s, removed_arg_keys=%s, added_arg_keys=%s, " "header_keys_before=%s, header_keys_after=%s, removed_header_keys=%s, added_header_keys=%s", sanitized_tool_name, + modified_payload is not None, modified_name, before_arg_keys, after_arg_keys, @@ -4445,26 +4449,26 @@ async def prepare_rust_mcp_tool_execution( # inject credentials and clean arguments before the Rust direct call. modified_args = arguments if has_pre_invoke and arguments is not None: - pre_invoke_headers = HttpHeaderPayload(root=dict(runtime_headers)) + ext_in = build_hook_extensions(runtime_headers) pre_result, _ = await plugin_manager.invoke_hook( ToolHookType.TOOL_PRE_INVOKE, - payload=ToolPreInvokePayload(name=name, args=arguments, headers=pre_invoke_headers), + payload=ToolPreInvokePayload(name=name, args=arguments), global_context=hook_global_context, local_contexts=plugin_context_table, violations_as_exceptions=True, - extensions=build_request_extensions(), + extensions=ext_in, ) record_plugin_metrics(current_trace_id.get(), pre_result.metadata) - _log_tool_pre_invoke_result(name, arguments, pre_invoke_headers, pre_result) + _log_tool_pre_invoke_result(name, arguments, runtime_headers, pre_result, ext_in) if pre_result.modified_payload: modified_args = pre_result.modified_payload.args if pre_result.modified_payload.name and pre_result.modified_payload.name != name: tool_name_original = pre_result.modified_payload.name - if pre_result.modified_payload.headers is not None: - plugin_headers = pre_result.modified_payload.headers.root if hasattr(pre_result.modified_payload.headers, "root") else {} - for hk, hv in plugin_headers.items(): - if hk and hv: - runtime_headers[str(hk).lower()] = str(hv) + plugin_headers = headers_from_modified_extensions(pre_result) + if plugin_headers is not None: + for hk, hv in plugin_headers.items(): + if hk and hv: + runtime_headers[str(hk).lower()] = str(hv) # Defense in depth: strip X-Vault-Tokens (case-insensitive) from outbound # headers. The Vault plugin removes this header when it processes the token, @@ -5356,23 +5360,24 @@ async def invoke_tool( # Use pre-created Pydantic model from Phase 2 (no ORM access) if tool_metadata: global_context.metadata[TOOL_METADATA] = tool_metadata - pre_invoke_headers = HttpHeaderPayload(root=headers) + ext_in = build_hook_extensions(headers) pre_result, context_table = await plugin_manager.invoke_hook( ToolHookType.TOOL_PRE_INVOKE, - payload=ToolPreInvokePayload(name=name, args=arguments, headers=pre_invoke_headers), + payload=ToolPreInvokePayload(name=name, args=arguments), global_context=global_context, local_contexts=context_table, # Pass context from previous hooks violations_as_exceptions=True, - extensions=build_request_extensions(), + extensions=ext_in, ) record_plugin_metrics(current_trace_id.get(), pre_result.metadata) - _log_tool_pre_invoke_result(name, arguments, pre_invoke_headers, pre_result) + _log_tool_pre_invoke_result(name, arguments, headers, pre_result, ext_in) if pre_result.modified_payload: payload = pre_result.modified_payload name = payload.name arguments = payload.args - if payload.headers is not None: - headers = payload.headers.model_dump() + plugin_headers = headers_from_modified_extensions(pre_result) + if plugin_headers is not None: + headers = plugin_headers # Build the payload based on integration type payload = arguments.copy() @@ -6199,23 +6204,24 @@ async def connect_to_streamablehttp_server(server_url: str, headers: dict = head global_context.metadata[TOOL_METADATA] = tool_metadata if gateway_metadata: global_context.metadata[GATEWAY_METADATA] = gateway_metadata - pre_invoke_headers = HttpHeaderPayload(root=headers) + ext_in = build_hook_extensions(headers) pre_result, context_table = await plugin_manager.invoke_hook( ToolHookType.TOOL_PRE_INVOKE, - payload=ToolPreInvokePayload(name=name, args=arguments, headers=pre_invoke_headers), + payload=ToolPreInvokePayload(name=name, args=arguments), global_context=global_context, local_contexts=None, violations_as_exceptions=True, - extensions=build_request_extensions(), + extensions=ext_in, ) record_plugin_metrics(current_trace_id.get(), pre_result.metadata) - _log_tool_pre_invoke_result(name, arguments, pre_invoke_headers, pre_result) + _log_tool_pre_invoke_result(name, arguments, headers, pre_result, ext_in) if pre_result.modified_payload: payload = pre_result.modified_payload name = payload.name arguments = payload.args - if payload.headers is not None: - headers = payload.headers.model_dump() + plugin_headers = headers_from_modified_extensions(pre_result) + if plugin_headers is not None: + headers = plugin_headers # Defense in depth: strip X-Vault-Tokens (case-insensitive) from outbound # headers. The Vault plugin removes this header when it processes the token, @@ -6285,29 +6291,29 @@ async def connect_to_streamablehttp_server(server_url: str, headers: dict = head if plugin_manager and plugin_manager.has_hooks_for(ToolHookType.TOOL_PRE_INVOKE) and not skip_pre_invoke: if tool_metadata: global_context.metadata[TOOL_METADATA] = tool_metadata - pre_invoke_headers = HttpHeaderPayload(root=plugin_headers) + ext_in = build_hook_extensions(plugin_headers) pre_result, context_table = await plugin_manager.invoke_hook( ToolHookType.TOOL_PRE_INVOKE, - payload=ToolPreInvokePayload(name=name, args=arguments, headers=pre_invoke_headers), + payload=ToolPreInvokePayload(name=name, args=arguments), global_context=global_context, local_contexts=context_table, violations_as_exceptions=True, - extensions=build_request_extensions(), + extensions=ext_in, ) record_plugin_metrics(current_trace_id.get(), pre_result.metadata) - _log_tool_pre_invoke_result(name, arguments, pre_invoke_headers, pre_result) + _log_tool_pre_invoke_result(name, arguments, plugin_headers, pre_result, ext_in) if pre_result.modified_payload: payload = pre_result.modified_payload name = payload.name arguments = payload.args - if payload.headers is not None: - plugin_returned_headers = payload.headers.model_dump() - if a2a_allowlist: - allowlist_lower = {h.lower() for h in a2a_allowlist} - safe_headers = {k: v for k, v in plugin_returned_headers.items() if k.lower() in allowlist_lower} - if not settings.enable_sensitive_header_passthrough: - safe_headers = filter_sensitive_headers(safe_headers) - headers.update(safe_headers) + plugin_returned_headers = headers_from_modified_extensions(pre_result) + if plugin_returned_headers is not None: + if a2a_allowlist: + allowlist_lower = {h.lower() for h in a2a_allowlist} + safe_headers = {k: v for k, v in plugin_returned_headers.items() if k.lower() in allowlist_lower} + if not settings.enable_sensitive_header_passthrough: + safe_headers = filter_sensitive_headers(safe_headers) + headers.update(safe_headers) prepared = prepare_a2a_invocation( agent_type=a2a_agent_type, diff --git a/plugins/config.yaml b/plugins/config.yaml index 2aad937eb5..c9ef8d099c 100644 --- a/plugins/config.yaml +++ b/plugins/config.yaml @@ -1063,6 +1063,7 @@ plugins: tags: ["security", "vault", "OAUTH2"] mode: "disabled" # enforce | enforce_ignore_error | permissive | disabled priority: 10 + capabilities: ["write_headers"] conditions: - prompts: [] server_ids: [] @@ -1229,6 +1230,7 @@ plugins: tags: ["telemetry", "observability", "opentelemetry", "monitoring"] mode: "disabled" # enforce | enforce_ignore_error | permissive | disabled priority: 200 # Run late to capture all context + capabilities: ["read_headers"] conditions: [] # Apply to all tools config: export_full_payload: false # opt in only if you explicitly want result content in telemetry @@ -1268,6 +1270,7 @@ plugins: tags: ["security", "policy", "access-control", "pdp", "rbac", "mac"] mode: "disabled" # set to "sequential" to activate priority: 10 # Run early — access control before other plugins + capabilities: ["read_headers"] conditions: [] config: engines: @@ -1316,6 +1319,7 @@ plugins: tags: ["auth", "jwt", "claims", "rbac", "abac", "rfc9396"] mode: "disabled" # set to "transform" to activate priority: 10 # Run early — before other auth plugins + capabilities: ["read_headers"] conditions: [] config: context_key: jwt_claims diff --git a/plugins/examples/custom_auth_example/custom_auth.py b/plugins/examples/custom_auth_example/custom_auth.py index 87a039eb5f..27d446f3ef 100644 --- a/plugins/examples/custom_auth_example/custom_auth.py +++ b/plugins/examples/custom_auth_example/custom_auth.py @@ -44,6 +44,7 @@ PluginViolation, PluginViolationError, ) +from cpex.framework.extensions import Extensions, HttpExtension logger = logging.getLogger(__name__) @@ -96,6 +97,7 @@ async def http_pre_request( self, payload: HttpPreRequestPayload, context: PluginContext, + extensions: Extensions | None = None, ) -> PluginResult[HttpHeaderPayload]: """Transform custom authentication headers before authentication. @@ -110,6 +112,7 @@ async def http_pre_request( Args: payload: HTTP pre-request payload with headers. context: Plugin execution context. + extensions: Hook extensions (preferred source for headers). Returns: Result with modified headers if transformation applied. @@ -117,7 +120,10 @@ async def http_pre_request( if not self._cfg.transform_headers: return PluginResult(continue_processing=True) - headers = dict(payload.headers.root) + if extensions and extensions.http: + headers = dict(extensions.http.headers) + else: + headers = dict(payload.headers.root) # Check if custom API key header is present api_key_header = self._cfg.api_key_header.lower() @@ -128,10 +134,10 @@ async def http_pre_request( logger.info(f"Transforming {self._cfg.api_key_header} to Authorization header") headers["authorization"] = f"Bearer {api_key}" - # Return modified headers - modified_headers = HttpHeaderPayload(root=headers) + new_ext = (extensions or Extensions()).model_copy(update={"http": HttpExtension(headers=headers)}) return PluginResult( - modified_payload=modified_headers, + modified_payload=HttpHeaderPayload(root=headers), + modified_extensions=new_ext, metadata={"transformed": True, "original_header": self._cfg.api_key_header}, continue_processing=True, ) @@ -142,6 +148,7 @@ async def http_auth_resolve_user( self, payload: HttpAuthResolveUserPayload, context: PluginContext, + extensions: Extensions | None = None, ) -> PluginResult[dict]: """Resolve user identity using custom authentication mechanisms. @@ -158,12 +165,16 @@ async def http_auth_resolve_user( Args: payload: Auth resolution payload with credentials and headers. context: Plugin execution context. + extensions: Hook extensions (preferred source for headers). Returns: Result with authenticated user dict if successful, or continue_processing=True to fall back to standard JWT authentication. """ - headers = dict(payload.headers.root) + if extensions and extensions.http: + headers = dict(extensions.http.headers) + else: + headers = dict(payload.headers.root) # Example 1: API Key Authentication with Error Handling # Check if we have a bearer token that matches our API key mapping @@ -264,6 +275,7 @@ async def http_post_request( self, payload: HttpPostRequestPayload, context: PluginContext, + extensions: Extensions | None = None, ) -> PluginResult[HttpHeaderPayload]: """Add custom headers to response after request completion. @@ -279,14 +291,18 @@ async def http_post_request( Args: payload: HTTP post-request payload with response information. context: Plugin execution context. + extensions: Hook extensions (preferred source for request headers). Returns: Result with modified response headers if applicable. """ response_headers = dict(payload.response_headers.root) if payload.response_headers else {} - # Add correlation ID from request to response (if present) - request_headers = dict(payload.headers.root) + # Prefer extensions.http.headers; fall back to dual-written payload.headers + if extensions and extensions.http: + request_headers = dict(extensions.http.headers) + else: + request_headers = dict(payload.headers.root) if "x-correlation-id" in request_headers: response_headers["x-correlation-id"] = request_headers["x-correlation-id"] @@ -305,7 +321,9 @@ async def http_post_request( # Note: context.global_context.request_id is the same across all hooks for this request logger.info(f"[{context.global_context.request_id}] Auth request completed: path={payload.path} method={payload.method} status={payload.status_code} client={payload.client_host}") + new_ext = (extensions or Extensions()).model_copy(update={"http": HttpExtension(headers=response_headers)}) return PluginResult( modified_payload=HttpHeaderPayload(root=response_headers), + modified_extensions=new_ext, continue_processing=True, ) diff --git a/plugins/examples/simple_token_auth/simple_token_auth.py b/plugins/examples/simple_token_auth/simple_token_auth.py index e908cd3baa..499a7a367d 100644 --- a/plugins/examples/simple_token_auth/simple_token_auth.py +++ b/plugins/examples/simple_token_auth/simple_token_auth.py @@ -14,6 +14,7 @@ HttpAuthCheckPermissionPayload, HttpAuthCheckPermissionResultPayload, HttpAuthResolveUserPayload, + HttpHeaderPayload, HttpHookType, HttpPostRequestPayload, HttpPreRequestPayload, @@ -24,6 +25,7 @@ PluginViolation, PluginViolationError, ) +from cpex.framework.extensions import Extensions, HttpExtension from plugins.examples.simple_token_auth.token_storage import TokenStorage from pydantic import BaseModel @@ -88,12 +90,13 @@ def storage(self) -> TokenStorage: """Expose token storage for external access (e.g., login endpoints).""" return self._storage - async def http_pre_request(self, payload: HttpPreRequestPayload, context: PluginContext) -> PluginResult: + async def http_pre_request(self, payload: HttpPreRequestPayload, context: PluginContext, extensions: Extensions | None = None) -> PluginResult: """Transform X-Auth-Token to Authorization: Bearer if configured. Args: payload: HTTP pre-request payload context: Plugin context + extensions: Hook extensions (preferred source for headers) Returns: PluginResult with potentially modified headers @@ -101,7 +104,10 @@ async def http_pre_request(self, payload: HttpPreRequestPayload, context: Plugin if not self._cfg.transform_to_bearer: return PluginResult(continue_processing=True) - headers = dict(payload.headers.root) + if extensions and extensions.http: + headers = dict(extensions.http.headers) + else: + headers = dict(payload.headers.root) token_header = self._cfg.token_header.lower() logger.info(f"[SimpleTokenAuth] http_pre_request - Looking for header: {token_header}, headers: {list(headers.keys())}") @@ -121,10 +127,10 @@ async def http_pre_request(self, payload: HttpPreRequestPayload, context: Plugin logger.info(f"[SimpleTokenAuth] Transformed {token_header} to Authorization: Bearer {token[:20]}...") - from cpex.framework import HttpHeaderPayload - + new_ext = (extensions or Extensions()).model_copy(update={"http": HttpExtension(headers=headers)}) return PluginResult( modified_payload=HttpHeaderPayload(root=headers), + modified_extensions=new_ext, metadata={"transformed": True, "original_header": token_header}, continue_processing=True, ) @@ -238,22 +244,24 @@ async def http_auth_check_permission(self, payload: HttpAuthCheckPermissionPaylo continue_processing=True, # Permission granted, let middleware handle the response ) - async def http_post_request(self, payload: HttpPostRequestPayload, context: PluginContext) -> PluginResult: + async def http_post_request(self, payload: HttpPostRequestPayload, context: PluginContext, extensions: Extensions | None = None) -> PluginResult: """Add authentication status headers to responses. Args: payload: HTTP post-request payload context: Plugin context + extensions: Hook extensions (preferred source for request headers) Returns: PluginResult with modified response headers """ - from cpex.framework import HttpHeaderPayload - response_headers = dict(payload.response_headers.root) if payload.response_headers else {} - # Add correlation ID if present in request - request_headers = dict(payload.headers.root) + # Prefer extensions.http.headers; fall back to dual-written payload.headers + if extensions and extensions.http: + request_headers = dict(extensions.http.headers) + else: + request_headers = dict(payload.headers.root) if "x-correlation-id" in request_headers: response_headers["x-correlation-id"] = request_headers["x-correlation-id"] @@ -274,7 +282,12 @@ async def http_post_request(self, payload: HttpPostRequestPayload, context: Plug elif payload.status_code == 401: response_headers["x-auth-status"] = "failed" - return PluginResult(modified_payload=HttpHeaderPayload(root=response_headers), continue_processing=True) + new_ext = (extensions or Extensions()).model_copy(update={"http": HttpExtension(headers=response_headers)}) + return PluginResult( + modified_payload=HttpHeaderPayload(root=response_headers), + modified_extensions=new_ext, + continue_processing=True, + ) def get_supported_hooks(self) -> list[str]: """Return list of supported hook types.""" diff --git a/plugins/header_filter/header_filter_plugin.py b/plugins/header_filter/header_filter_plugin.py index bf68a0b68b..af7189a79f 100644 --- a/plugins/header_filter/header_filter_plugin.py +++ b/plugins/header_filter/header_filter_plugin.py @@ -19,13 +19,13 @@ from cpex.framework import ( AgentPreInvokePayload, AgentPreInvokeResult, - HttpHeaderPayload, Plugin, PluginConfig, PluginContext, ToolPreInvokePayload, ToolPreInvokeResult, ) +from cpex.framework.extensions import Extensions, HttpExtension from mcpgateway.services.logging_service import LoggingService # Initialize logging service @@ -63,7 +63,7 @@ class HeaderFilter(Plugin): This plugin prevents sensitive authentication and authorization headers from being leaked to MCP servers. It runs on pre-invoke hooks for tools and agents where HTTP - headers are available in the payload. + headers are available via ``extensions.http.headers``. Security considerations: - Headers are filtered case-insensitively @@ -123,55 +123,65 @@ def _filter_headers(self, headers: dict[str, str], context_name: str) -> tuple[d return filtered_headers, removed_headers - async def tool_pre_invoke(self, payload: ToolPreInvokePayload, context: PluginContext) -> ToolPreInvokeResult: # pylint: disable=unused-argument + @staticmethod + def _headers_from_extensions(extensions: Extensions | None) -> dict[str, str]: + """Read HTTP headers from hook extensions.""" + if extensions and extensions.http: + return dict(extensions.http.headers) + return {} + + def _result_with_headers(self, result_cls, extensions: Extensions | None, headers: dict[str, str]): + """Build a hook result that returns updated headers via modified_extensions.""" + new_ext = (extensions or Extensions()).model_copy(update={"http": HttpExtension(headers=headers)}) + return result_cls(modified_extensions=new_ext) + + async def tool_pre_invoke(self, payload: ToolPreInvokePayload, context: PluginContext, extensions: Extensions | None = None) -> ToolPreInvokeResult: # pylint: disable=unused-argument """Filter headers before tool invocation. Args: - payload: The tool payload containing headers. + payload: The tool payload. context: Plugin execution context. + extensions: Hook extensions (headers on ``extensions.http``). Returns: - Result with filtered headers. + Result with filtered headers in ``modified_extensions``. """ - if not payload.headers: + headers = self._headers_from_extensions(extensions) + if not headers: return ToolPreInvokeResult() - headers = payload.headers.model_dump() context_name = f"tool:{payload.name}" - filtered_headers, removed = self._filter_headers(headers, context_name) if removed: if self._sconfig.log_filtered_headers: logger.info(f"Filtered {len(removed)} header(s) from {context_name}: {', '.join(removed)}") - modified = payload.model_copy(update={"headers": HttpHeaderPayload(root=filtered_headers)}) - return ToolPreInvokeResult(modified_payload=modified) + return self._result_with_headers(ToolPreInvokeResult, extensions, filtered_headers) return ToolPreInvokeResult() - async def agent_pre_invoke(self, payload: AgentPreInvokePayload, context: PluginContext) -> AgentPreInvokeResult: # pylint: disable=unused-argument + async def agent_pre_invoke(self, payload: AgentPreInvokePayload, context: PluginContext, extensions: Extensions | None = None) -> AgentPreInvokeResult: # pylint: disable=unused-argument """Filter headers before agent invocation. Args: - payload: The agent payload containing headers. + payload: The agent payload. context: Plugin execution context. + extensions: Hook extensions (headers on ``extensions.http``). Returns: - Result with filtered headers. + Result with filtered headers in ``modified_extensions``. """ - if not payload.headers: + headers = self._headers_from_extensions(extensions) + if not headers: return AgentPreInvokeResult() - headers = payload.headers.model_dump() context_name = f"agent:{payload.agent_id}" - filtered_headers, removed = self._filter_headers(headers, context_name) if removed: if self._sconfig.log_filtered_headers: logger.info(f"Filtered {len(removed)} header(s) from {context_name}: {', '.join(removed)}") - modified = payload.model_copy(update={"headers": HttpHeaderPayload(root=filtered_headers)}) - return AgentPreInvokeResult(modified_payload=modified) + return self._result_with_headers(AgentPreInvokeResult, extensions, filtered_headers) return AgentPreInvokeResult() diff --git a/plugins/jwt_claims_extraction/jwt_claims_extraction.py b/plugins/jwt_claims_extraction/jwt_claims_extraction.py index 168abd672c..2b299d6b8f 100644 --- a/plugins/jwt_claims_extraction/jwt_claims_extraction.py +++ b/plugins/jwt_claims_extraction/jwt_claims_extraction.py @@ -45,6 +45,7 @@ PluginContext, PluginResult, ) +from cpex.framework.extensions import Extensions logger = logging.getLogger(__name__) @@ -85,6 +86,7 @@ async def http_auth_resolve_user( self, payload: HttpAuthResolveUserPayload, context: PluginContext, + extensions: Extensions | None = None, ) -> PluginResult[dict]: """Extract JWT claims and store in global context state. @@ -95,12 +97,13 @@ async def http_auth_resolve_user( Args: payload: Auth payload with credentials and headers. context: Plugin execution context with global_context. + extensions: Hook extensions (headers on ``extensions.http``). Returns: PluginResult with continue_processing=True (passthrough). """ try: - token = self._extract_token(payload) + token = self._extract_token(payload, extensions) if not token: logger.debug("No JWT token found in request, skipping claims extraction") @@ -130,11 +133,12 @@ async def http_auth_resolve_user( metadata={"jwt_claims_extracted": False}, ) - def _extract_token(self, payload: HttpAuthResolveUserPayload) -> Optional[str]: + def _extract_token(self, payload: HttpAuthResolveUserPayload, extensions: Extensions | None = None) -> Optional[str]: """Extract JWT token from request credentials or Authorization header. Args: payload: Auth payload with credentials and headers. + extensions: Hook extensions (preferred source for headers). Returns: JWT token string or None if not found. @@ -146,8 +150,11 @@ def _extract_token(self, payload: HttpAuthResolveUserPayload) -> Optional[str]: if token: return token - # Fallback to Authorization header - headers_dict = getattr(payload.headers, "root", {}) + # Prefer extensions.http.headers; fall back to dual-written payload.headers + if extensions and extensions.http: + headers_dict = extensions.http.headers + else: + headers_dict = getattr(payload.headers, "root", {}) if headers_dict: auth_header = headers_dict.get("authorization") or headers_dict.get("Authorization") if auth_header and auth_header.startswith("Bearer "): diff --git a/plugins/sparc_static_validator/sparc_static_validator.py b/plugins/sparc_static_validator/sparc_static_validator.py index 1c407a9b7f..03662a4640 100644 --- a/plugins/sparc_static_validator/sparc_static_validator.py +++ b/plugins/sparc_static_validator/sparc_static_validator.py @@ -415,7 +415,6 @@ async def tool_pre_invoke(self, payload: ToolPreInvokePayload, context: PluginCo modified_payload = ToolPreInvokePayload( name=payload.name, args=correction, - headers=payload.headers, ) return ToolPreInvokeResult( continue_processing=True, diff --git a/plugins/tools_telemetry_exporter/telemetry_exporter.py b/plugins/tools_telemetry_exporter/telemetry_exporter.py index 9e844bfc00..8f791b3362 100644 --- a/plugins/tools_telemetry_exporter/telemetry_exporter.py +++ b/plugins/tools_telemetry_exporter/telemetry_exporter.py @@ -17,6 +17,7 @@ # First-Party from cpex.framework import get_attr, Plugin, PluginConfig, PluginContext from cpex.framework.constants import GATEWAY_METADATA, TOOL_METADATA +from cpex.framework.extensions import Extensions from cpex.framework.hooks.tools import ToolPostInvokePayload, ToolPostInvokeResult, ToolPreInvokePayload, ToolPreInvokeResult from mcpgateway.services.logging_service import LoggingService @@ -133,23 +134,24 @@ def _is_sensitive_header_name(name: str) -> bool: return any(pattern.match(name) for pattern in _SENSITIVE_HEADER_PATTERNS) @classmethod - def _serialize_headers(cls, payload: ToolPreInvokePayload) -> str: + def _serialize_headers(cls, headers: dict[str, str] | None) -> str: """Serialize headers for telemetry, redacting sensitive values by default.""" - if not payload.headers: + if not headers: return "{}" sanitized_headers = {} - for key, value in payload.headers.root.items(): + for key, value in headers.items(): sanitized_headers[key] = _MASKED_HEADER_VALUE if cls._is_sensitive_header_name(key) else value return orjson.dumps(sanitized_headers, default=str).decode() - async def tool_pre_invoke(self, payload: ToolPreInvokePayload, context: PluginContext) -> ToolPreInvokeResult: + async def tool_pre_invoke(self, payload: ToolPreInvokePayload, context: PluginContext, extensions: Extensions | None = None) -> ToolPreInvokeResult: """Capture pre-invocation telemetry for tools. Args: payload: The tool payload containing arguments. context: Plugin execution context. + extensions: Hook extensions (headers on ``extensions.http``). Returns: Result with potentially modified tool arguments. @@ -157,6 +159,7 @@ async def tool_pre_invoke(self, payload: ToolPreInvokePayload, context: PluginCo logger.info("ToolsTelemetryExporter: Capturing pre-invocation tool telemetry.") context_attributes = self._get_pre_invoke_context_attributes(context) + header_map = dict(extensions.http.headers) if extensions and extensions.http else {} export_attributes = { "request_id": context_attributes["request_id"], "user": context_attributes["user"], @@ -169,7 +172,7 @@ async def tool_pre_invoke(self, payload: ToolPreInvokePayload, context: PluginCo "tool.target_tool_name": context_attributes["tool"]["target_tool_name"], "tool.description": context_attributes["tool"]["description"], "tool.invocation.args": orjson.dumps(payload.args, default=str).decode(), - "headers": self._serialize_headers(payload), + "headers": self._serialize_headers(header_map), } await self._export_telemetry(attributes=export_attributes, span_name="tool.pre_invoke") diff --git a/plugins/unified_pdp/unified_pdp.py b/plugins/unified_pdp/unified_pdp.py index f0f1240cfe..995d2d1ed1 100644 --- a/plugins/unified_pdp/unified_pdp.py +++ b/plugins/unified_pdp/unified_pdp.py @@ -44,6 +44,7 @@ PluginContext, PluginViolation, ) +from cpex.framework.extensions import Extensions from cpex.framework.hooks.tools import ( ToolPreInvokePayload, ToolPreInvokeResult, @@ -235,6 +236,7 @@ async def tool_pre_invoke( self, payload: ToolPreInvokePayload, context: PluginContext, + extensions: Extensions | None = None, ) -> ToolPreInvokeResult: """Called before every tool invocation. @@ -245,6 +247,7 @@ async def tool_pre_invoke( Args: payload: Contains the tool name and invocation arguments. context: Gateway-provided request context (user, tenant, etc.). + extensions: Hook extensions (headers on ``extensions.http``). Returns: A ToolPreInvokeResult — either pass-through or blocked with a @@ -253,7 +256,8 @@ async def tool_pre_invoke( subject = self._extract_subject(context) # Extract HTTP metadata from headers for IP and user_agent - http_meta = self._extract_http_metadata(payload.headers) + header_source = extensions.http.headers if extensions and extensions.http else payload.headers + http_meta = self._extract_http_metadata(header_source) # Extract classification_level from tool args if provided (for MAC engine) tool_args = payload.args or {} diff --git a/plugins/vault/vault_plugin.py b/plugins/vault/vault_plugin.py index d4e535b1ce..e691d02599 100644 --- a/plugins/vault/vault_plugin.py +++ b/plugins/vault/vault_plugin.py @@ -21,7 +21,6 @@ # First-Party from cpex.framework import ( - HttpHeaderPayload, Plugin, PluginConfig, PluginContext, @@ -29,6 +28,7 @@ ToolPreInvokeResult, get_attr, ) +from cpex.framework.extensions import Extensions, HttpExtension from mcpgateway.db import get_db from mcpgateway.services.gateway_service import GatewayService from mcpgateway.services.logging_service import LoggingService @@ -112,12 +112,13 @@ def _parse_vault_token_key(self, key: str) -> tuple[str, str | None, str | None, token_name = parts[3] if len(parts) > 3 else None return system, scope, token_type, token_name - async def tool_pre_invoke(self, payload: ToolPreInvokePayload, context: PluginContext) -> ToolPreInvokeResult: + async def tool_pre_invoke(self, payload: ToolPreInvokePayload, context: PluginContext, extensions: Extensions | None = None) -> ToolPreInvokeResult: """Generate bearer tokens from vault-saved tokens before tool invocation. Args: payload: The tool payload containing arguments. context: Plugin execution context. + extensions: Hook extensions (headers on ``extensions.http``). Returns: Result with potentially modified headers containing bearer token. @@ -171,19 +172,21 @@ async def tool_pre_invoke(self, payload: ToolPreInvokePayload, context: PluginCo finally: gen.close() + def _with_headers(hdrs: dict[str, str]) -> ToolPreInvokeResult: + new_ext = (extensions or Extensions()).model_copy(update={"http": HttpExtension(headers=hdrs)}) + return ToolPreInvokeResult(modified_extensions=new_ext) + + headers: dict[str, str] = {k.lower(): v for k, v in (extensions.http.headers.items() if extensions and extensions.http else [])} + if not system_key: logger.warning("System cannot be determined from gateway metadata.") # SECURITY: Strip vault header even when system cannot be determined - if payload.headers: - safe_headers = {k.lower(): v for k, v in payload.headers.root.items()} - if self._vault_header_key in safe_headers: - del safe_headers[self._vault_header_key] - payload = payload.model_copy(update={"headers": HttpHeaderPayload(root=safe_headers)}) - return ToolPreInvokeResult(modified_payload=payload) + if self._vault_header_key in headers: + del headers[self._vault_header_key] + return _with_headers(headers) return ToolPreInvokeResult() modified = False - headers: dict[str, str] = {k.lower(): v for k, v in payload.headers.root.items()} if payload.headers else {} # Check if vault header exists if self._vault_header_key not in headers: @@ -196,8 +199,7 @@ async def tool_pre_invoke(self, payload: ToolPreInvokePayload, context: PluginCo logger.error("Failed to parse vault tokens from header: %s", e) # SECURITY: Always remove vault header even on parse error del headers[self._vault_header_key] - payload = payload.model_copy(update={"headers": HttpHeaderPayload(root=headers)}) - return ToolPreInvokeResult(modified_payload=payload) + return _with_headers(headers) # SECURITY: Always remove vault header immediately after successful parsing # This header should NEVER be sent to the MCP server @@ -205,8 +207,7 @@ async def tool_pre_invoke(self, payload: ToolPreInvokePayload, context: PluginCo if not isinstance(vault_tokens, dict): logger.error("Vault tokens header is not a JSON object: %s", type(vault_tokens).__name__) - payload = payload.model_copy(update={"headers": HttpHeaderPayload(root=headers)}) - return ToolPreInvokeResult(modified_payload=payload) + return _with_headers(headers) logger.debug("Removed vault header '%s' from headers", self._vault_header_key) vault_handling = self._sconfig.vault_handling @@ -265,9 +266,8 @@ async def tool_pre_invoke(self, payload: ToolPreInvokePayload, context: PluginCo # Even if we didn't modify headers (no token match), we still removed the vault header logger.warning("Vault tokens provided but no match found for system '%s' - possible misconfiguration", system_key) - # Always return modified payload since the vault header was stripped - payload = payload.model_copy(update={"headers": HttpHeaderPayload(root=headers)}) - return ToolPreInvokeResult(modified_payload=payload) + # Always return modified extensions since the vault header was stripped + return _with_headers(headers) async def shutdown(self) -> None: """Shutdown the plugin gracefully. diff --git a/tests/unit/mcpgateway/plugins/fixtures/configs/tool_headers_metadata_plugin.yaml b/tests/unit/mcpgateway/plugins/fixtures/configs/tool_headers_metadata_plugin.yaml index 4b31c75e75..a58e605da1 100644 --- a/tests/unit/mcpgateway/plugins/fixtures/configs/tool_headers_metadata_plugin.yaml +++ b/tests/unit/mcpgateway/plugins/fixtures/configs/tool_headers_metadata_plugin.yaml @@ -9,6 +9,7 @@ plugins: tags: ["plugin", "headers"] mode: "sequential" # sequential | transform | disabled priority: 150 + capabilities: ["read_headers", "write_headers"] conditions: # Apply to specific tools/servers - prompts: [] diff --git a/tests/unit/mcpgateway/plugins/fixtures/configs/tool_headers_plugin.yaml b/tests/unit/mcpgateway/plugins/fixtures/configs/tool_headers_plugin.yaml index 48fcf2990e..60e1310aef 100644 --- a/tests/unit/mcpgateway/plugins/fixtures/configs/tool_headers_plugin.yaml +++ b/tests/unit/mcpgateway/plugins/fixtures/configs/tool_headers_plugin.yaml @@ -9,6 +9,7 @@ plugins: tags: ["plugin", "headers"] mode: "sequential" # sequential | transform | disabled priority: 150 + capabilities: ["read_headers", "write_headers"] conditions: # Apply to specific tools/servers - prompts: [] diff --git a/tests/unit/mcpgateway/plugins/fixtures/plugins/agent_plugins.py b/tests/unit/mcpgateway/plugins/fixtures/plugins/agent_plugins.py index 2876119b13..911ffdef23 100644 --- a/tests/unit/mcpgateway/plugins/fixtures/plugins/agent_plugins.py +++ b/tests/unit/mcpgateway/plugins/fixtures/plugins/agent_plugins.py @@ -85,7 +85,6 @@ async def agent_pre_invoke(self, payload: AgentPreInvokePayload, context: Plugin agent_id=payload.agent_id, messages=filtered_messages, tools=payload.tools, - headers=payload.headers, model=payload.model, system_prompt=payload.system_prompt, parameters=payload.parameters, diff --git a/tests/unit/mcpgateway/plugins/fixtures/plugins/headers.py b/tests/unit/mcpgateway/plugins/fixtures/plugins/headers.py index 72d3f0807b..c8c7a985cd 100644 --- a/tests/unit/mcpgateway/plugins/fixtures/plugins/headers.py +++ b/tests/unit/mcpgateway/plugins/fixtures/plugins/headers.py @@ -13,7 +13,6 @@ from cpex.framework import ( PluginContext, Plugin, - HttpHeaderPayload, PromptPosthookPayload, PromptPosthookResult, PromptPrehookPayload, @@ -27,6 +26,7 @@ ToolPreInvokePayload, ToolPreInvokeResult, ) +from cpex.framework.extensions import Extensions, HttpExtension logger = logging.getLogger("header_plugin") @@ -56,12 +56,13 @@ async def prompt_post_fetch(self, payload: PromptPosthookPayload, context: Plugi """ raise ValueError("Sadly! Prompt postfetch is broken!") - async def tool_pre_invoke(self, payload: ToolPreInvokePayload, context: PluginContext) -> ToolPreInvokeResult: + async def tool_pre_invoke(self, payload: ToolPreInvokePayload, context: PluginContext, extensions: Extensions | None = None) -> ToolPreInvokeResult: """Plugin hook run before a tool is invoked. Args: payload: The tool payload to be analyzed. context: Contextual information about the hook call. + extensions: Hook extensions (headers on ``extensions.http``). Returns: The result of the plugin's analysis, including whether the tool can proceed. @@ -71,11 +72,11 @@ async def tool_pre_invoke(self, payload: ToolPreInvokePayload, context: PluginCo assert tool_meta.original_name == "test_tool" assert tool_meta.url.host == "example.com" assert tool_meta.integration_type == "REST" or tool_meta.integration_type == "MCP" - headers = payload.headers.model_dump() if payload.headers else {} + headers = dict(extensions.http.headers) if extensions and extensions.http else {} if tool_meta.integration_type == "REST": - assert payload.headers - assert "Content-Type" in payload.headers - assert payload.headers["Content-Type"] == "application/json" + assert headers + assert "Content-Type" in headers + assert headers["Content-Type"] == "application/json" elif tool_meta.integration_type == "MCP": assert GATEWAY_METADATA in context.global_context.metadata gateway_meta = context.global_context.metadata[GATEWAY_METADATA] @@ -85,9 +86,11 @@ async def tool_pre_invoke(self, payload: ToolPreInvokePayload, context: PluginCo headers["User-Agent"] = "Mozilla/5.0" headers["Connection"] = "keep-alive" - modified_payload = payload.model_copy(update={"headers": HttpHeaderPayload(headers)}) - return ToolPreInvokeResult(continue_processing=True, modified_payload=modified_payload) + return ToolPreInvokeResult( + continue_processing=True, + modified_extensions=Extensions(http=HttpExtension(headers=headers)), + ) async def tool_post_invoke(self, payload: ToolPostInvokePayload, context: PluginContext) -> ToolPostInvokeResult: """Plugin hook run after a tool is invoked. @@ -162,24 +165,27 @@ async def prompt_post_fetch(self, payload: PromptPosthookPayload, context: Plugi """ raise ValueError("Sadly! Prompt postfetch is broken!") - async def tool_pre_invoke(self, payload: ToolPreInvokePayload, context: PluginContext) -> ToolPreInvokeResult: + async def tool_pre_invoke(self, payload: ToolPreInvokePayload, context: PluginContext, extensions: Extensions | None = None) -> ToolPreInvokeResult: """Plugin hook run before a tool is invoked. Args: payload: The tool payload to be analyzed. context: Contextual information about the hook call. + extensions: Hook extensions (headers on ``extensions.http``). Returns: The result of the plugin's analysis, including whether the tool can proceed. """ - headers = payload.headers.model_dump() if payload.headers else {} - if payload.headers: - assert "Content-Type" in payload.headers - assert payload.headers["Content-Type"] == "application/json" + headers = dict(extensions.http.headers) if extensions and extensions.http else {} + if headers: + assert "Content-Type" in headers + assert headers["Content-Type"] == "application/json" headers["User-Agent"] = "Mozilla/5.0" headers["Connection"] = "keep-alive" - modified_payload = payload.model_copy(update={"headers": HttpHeaderPayload(headers)}) - return ToolPreInvokeResult(continue_processing=True, modified_payload=modified_payload) + return ToolPreInvokeResult( + continue_processing=True, + modified_extensions=Extensions(http=HttpExtension(headers=headers)), + ) async def tool_post_invoke(self, payload: ToolPostInvokePayload, context: PluginContext) -> ToolPostInvokeResult: """Plugin hook run after a tool is invoked. diff --git a/tests/unit/mcpgateway/plugins/plugins/header_filter/test_header_filter_plugin.py b/tests/unit/mcpgateway/plugins/plugins/header_filter/test_header_filter_plugin.py index 51bd23af27..9bdaa81cc5 100644 --- a/tests/unit/mcpgateway/plugins/plugins/header_filter/test_header_filter_plugin.py +++ b/tests/unit/mcpgateway/plugins/plugins/header_filter/test_header_filter_plugin.py @@ -13,13 +13,13 @@ from cpex.framework import ( AgentPreInvokePayload, GlobalContext, - HttpHeaderPayload, PluginConfig, PluginContext, PluginMode, ToolHookType, ToolPreInvokePayload, ) +from cpex.framework.extensions import Extensions, HttpExtension # Import the Header Filter plugin from plugins.header_filter.header_filter_plugin import HeaderFilter, HeaderFilterConfig @@ -41,6 +41,7 @@ def plugin_config(self) -> PluginConfig: tags=["test", "header_filter"], mode=PluginMode.SEQUENTIAL, priority=20, + capabilities=["write_headers"], config={ "filter_headers": ["Authorization", "Cookie", "X-API-Key"], "log_filtered_headers": True, @@ -60,97 +61,97 @@ def plugin_context(self) -> PluginContext: async def test_no_headers_returns_empty_result(self, plugin_config, plugin_context): """Test that missing headers returns empty result.""" plugin = HeaderFilter(plugin_config) - payload = ToolPreInvokePayload(name="test_tool", args={}, headers=None) + payload = ToolPreInvokePayload(name="test_tool", args={}) - result = await plugin.tool_pre_invoke(payload, plugin_context) + result = await plugin.tool_pre_invoke(payload, plugin_context, None) - assert result.modified_payload is None + assert result.modified_extensions is None assert result.continue_processing @pytest.mark.asyncio async def test_authorization_header_is_filtered(self, plugin_config, plugin_context): """Test that Authorization header is filtered.""" plugin = HeaderFilter(plugin_config) - payload = ToolPreInvokePayload( - name="test_tool", - args={}, - headers=HttpHeaderPayload({"Content-Type": "application/json", "Authorization": "Bearer secret_token"}), - ) + payload = ToolPreInvokePayload(name="test_tool", args={}) + ext = Extensions(http=HttpExtension(headers={"Content-Type": "application/json", "Authorization": "Bearer secret_token"})) - result = await plugin.tool_pre_invoke(payload, plugin_context) + result = await plugin.tool_pre_invoke(payload, plugin_context, ext) - assert result.modified_payload is not None - assert "Authorization" not in result.modified_payload.headers.root - assert "Content-Type" in result.modified_payload.headers.root + assert result.modified_extensions is not None + assert result.modified_extensions.http is not None + hdrs = result.modified_extensions.http.headers + assert "Authorization" not in hdrs + assert "Content-Type" in hdrs assert result.continue_processing @pytest.mark.asyncio async def test_cookie_header_is_filtered(self, plugin_config, plugin_context): """Test that Cookie header is filtered.""" plugin = HeaderFilter(plugin_config) - payload = ToolPreInvokePayload( - name="test_tool", - args={}, - headers=HttpHeaderPayload({"Content-Type": "application/json", "Cookie": "session=abc123"}), - ) + payload = ToolPreInvokePayload(name="test_tool", args={}) + ext = Extensions(http=HttpExtension(headers={"Content-Type": "application/json", "Cookie": "session=abc123"})) - result = await plugin.tool_pre_invoke(payload, plugin_context) + result = await plugin.tool_pre_invoke(payload, plugin_context, ext) - assert result.modified_payload is not None - assert "Cookie" not in result.modified_payload.headers.root - assert "Content-Type" in result.modified_payload.headers.root + assert result.modified_extensions is not None + assert result.modified_extensions.http is not None + hdrs = result.modified_extensions.http.headers + assert "Cookie" not in hdrs + assert "Content-Type" in hdrs @pytest.mark.asyncio async def test_multiple_sensitive_headers_filtered(self, plugin_config, plugin_context): """Test that multiple sensitive headers are filtered.""" plugin = HeaderFilter(plugin_config) - payload = ToolPreInvokePayload( - name="test_tool", - args={}, - headers=HttpHeaderPayload( - { + payload = ToolPreInvokePayload(name="test_tool", args={}) + ext = Extensions( + http=HttpExtension( + headers={ "Content-Type": "application/json", "Authorization": "Bearer token", "Cookie": "session=xyz", "X-API-Key": "secret_key", # pragma: allowlist secret "User-Agent": "TestClient/1.0", } - ), + ) ) - result = await plugin.tool_pre_invoke(payload, plugin_context) + result = await plugin.tool_pre_invoke(payload, plugin_context, ext) - assert result.modified_payload is not None - assert "Authorization" not in result.modified_payload.headers.root - assert "Cookie" not in result.modified_payload.headers.root - assert "X-API-Key" not in result.modified_payload.headers.root - assert "Content-Type" in result.modified_payload.headers.root - assert "User-Agent" in result.modified_payload.headers.root + assert result.modified_extensions is not None + assert result.modified_extensions.http is not None + hdrs = result.modified_extensions.http.headers + assert "Authorization" not in hdrs + assert "Cookie" not in hdrs + assert "X-API-Key" not in hdrs + assert "Content-Type" in hdrs + assert "User-Agent" in hdrs @pytest.mark.asyncio async def test_case_insensitive_filtering(self, plugin_config, plugin_context): """Test that header filtering is case-insensitive.""" plugin = HeaderFilter(plugin_config) - payload = ToolPreInvokePayload( - name="test_tool", - args={}, - headers=HttpHeaderPayload( - { + payload = ToolPreInvokePayload(name="test_tool", args={}) + ext = Extensions( + http=HttpExtension( + headers={ "content-type": "application/json", "authorization": "Bearer token", "COOKIE": "session=xyz", "X-Api-Key": "secret", } - ), + ) ) - result = await plugin.tool_pre_invoke(payload, plugin_context) + result = await plugin.tool_pre_invoke(payload, plugin_context, ext) - assert result.modified_payload is not None - assert "authorization" not in result.modified_payload.headers.root - assert "COOKIE" not in result.modified_payload.headers.root - assert "X-Api-Key" not in result.modified_payload.headers.root - assert "content-type" in result.modified_payload.headers.root + assert result.modified_extensions is not None + assert result.modified_extensions.http is not None + hdrs = result.modified_extensions.http.headers + assert "authorization" not in hdrs + assert "COOKIE" not in hdrs + assert "X-Api-Key" not in hdrs + assert "content-type" in hdrs @pytest.mark.asyncio async def test_passthrough_headers_not_filtered(self, plugin_context): @@ -165,68 +166,64 @@ async def test_passthrough_headers_not_filtered(self, plugin_context): tags=["test"], mode=PluginMode.SEQUENTIAL, priority=20, + capabilities=["write_headers"], config={ "filter_headers": ["Authorization", "Cookie"], "allow_passthrough_headers": ["Authorization"], }, ) plugin = HeaderFilter(config) - payload = ToolPreInvokePayload( - name="test_tool", - args={}, - headers=HttpHeaderPayload({"Authorization": "Bearer token", "Cookie": "session=xyz"}), - ) + payload = ToolPreInvokePayload(name="test_tool", args={}) + ext = Extensions(http=HttpExtension(headers={"Authorization": "Bearer token", "Cookie": "session=xyz"})) - result = await plugin.tool_pre_invoke(payload, plugin_context) + result = await plugin.tool_pre_invoke(payload, plugin_context, ext) - assert result.modified_payload is not None - assert "Authorization" in result.modified_payload.headers.root - assert "Cookie" not in result.modified_payload.headers.root + assert result.modified_extensions is not None + assert result.modified_extensions.http is not None + hdrs = result.modified_extensions.http.headers + assert "Authorization" in hdrs + assert "Cookie" not in hdrs @pytest.mark.asyncio async def test_no_filtered_headers_returns_empty_result(self, plugin_config, plugin_context): """Test that when no headers are filtered, empty result is returned.""" plugin = HeaderFilter(plugin_config) - payload = ToolPreInvokePayload( - name="test_tool", - args={}, - headers=HttpHeaderPayload({"Content-Type": "application/json", "User-Agent": "TestClient/1.0"}), - ) + payload = ToolPreInvokePayload(name="test_tool", args={}) + ext = Extensions(http=HttpExtension(headers={"Content-Type": "application/json", "User-Agent": "TestClient/1.0"})) - result = await plugin.tool_pre_invoke(payload, plugin_context) + result = await plugin.tool_pre_invoke(payload, plugin_context, ext) - assert result.modified_payload is None + assert result.modified_extensions is None assert result.continue_processing @pytest.mark.asyncio async def test_empty_headers_dict_returns_empty_result(self, plugin_config, plugin_context): """Test that empty headers dict returns empty result.""" plugin = HeaderFilter(plugin_config) - payload = ToolPreInvokePayload(name="test_tool", args={}, headers=HttpHeaderPayload({})) + payload = ToolPreInvokePayload(name="test_tool", args={}) + ext = Extensions(http=HttpExtension(headers={})) - result = await plugin.tool_pre_invoke(payload, plugin_context) + result = await plugin.tool_pre_invoke(payload, plugin_context, ext) - assert result.modified_payload is None + assert result.modified_extensions is None assert result.continue_processing @pytest.mark.asyncio async def test_original_payload_not_mutated(self, plugin_config, plugin_context): - """Test that the original payload is not mutated (frozen model compliance).""" + """Test that the original extensions headers are not mutated (frozen model compliance).""" plugin = HeaderFilter(plugin_config) original_headers = {"Content-Type": "application/json", "Authorization": "Bearer token"} - payload = ToolPreInvokePayload( - name="test_tool", - args={}, - headers=HttpHeaderPayload(original_headers.copy()), - ) + payload = ToolPreInvokePayload(name="test_tool", args={}) + ext = Extensions(http=HttpExtension(headers=original_headers.copy())) - result = await plugin.tool_pre_invoke(payload, plugin_context) + result = await plugin.tool_pre_invoke(payload, plugin_context, ext) - # Original payload should still have Authorization - assert "Authorization" in payload.headers.root - # Modified payload should not - assert result.modified_payload is not None - assert "Authorization" not in result.modified_payload.headers.root + # Original extensions should still have Authorization + assert "Authorization" in ext.http.headers + # Modified extensions should not + assert result.modified_extensions is not None + assert result.modified_extensions.http is not None + assert "Authorization" not in result.modified_extensions.http.headers # ── agent_pre_invoke tests ──────────────────────────────────────── @@ -234,35 +231,36 @@ async def test_original_payload_not_mutated(self, plugin_config, plugin_context) async def test_agent_no_headers_returns_empty_result(self, plugin_config, plugin_context): """Test agent_pre_invoke with no headers.""" plugin = HeaderFilter(plugin_config) - payload = AgentPreInvokePayload(agent_id="test-agent", messages=[], headers=None) + payload = AgentPreInvokePayload(agent_id="test-agent", messages=[]) - result = await plugin.agent_pre_invoke(payload, plugin_context) + result = await plugin.agent_pre_invoke(payload, plugin_context, None) - assert result.modified_payload is None + assert result.modified_extensions is None assert result.continue_processing @pytest.mark.asyncio async def test_agent_headers_filtered(self, plugin_config, plugin_context): """Test agent_pre_invoke filters sensitive headers.""" plugin = HeaderFilter(plugin_config) - payload = AgentPreInvokePayload( - agent_id="test-agent", - messages=[], - headers=HttpHeaderPayload( - { + payload = AgentPreInvokePayload(agent_id="test-agent", messages=[]) + ext = Extensions( + http=HttpExtension( + headers={ "Content-Type": "application/json", "Authorization": "Bearer secret", "Cookie": "session=abc", } - ), + ) ) - result = await plugin.agent_pre_invoke(payload, plugin_context) + result = await plugin.agent_pre_invoke(payload, plugin_context, ext) - assert result.modified_payload is not None - assert "Authorization" not in result.modified_payload.headers.root - assert "Cookie" not in result.modified_payload.headers.root - assert "Content-Type" in result.modified_payload.headers.root + assert result.modified_extensions is not None + assert result.modified_extensions.http is not None + hdrs = result.modified_extensions.http.headers + assert "Authorization" not in hdrs + assert "Cookie" not in hdrs + assert "Content-Type" in hdrs @pytest.mark.asyncio async def test_agent_passthrough_headers(self, plugin_context): @@ -277,37 +275,34 @@ async def test_agent_passthrough_headers(self, plugin_context): tags=["test"], mode=PluginMode.SEQUENTIAL, priority=20, + capabilities=["write_headers"], config={ "filter_headers": ["Authorization", "Cookie"], "allow_passthrough_headers": ["Authorization"], }, ) plugin = HeaderFilter(config) - payload = AgentPreInvokePayload( - agent_id="test-agent", - messages=[], - headers=HttpHeaderPayload({"Authorization": "Bearer token", "Cookie": "session=xyz"}), - ) + payload = AgentPreInvokePayload(agent_id="test-agent", messages=[]) + ext = Extensions(http=HttpExtension(headers={"Authorization": "Bearer token", "Cookie": "session=xyz"})) - result = await plugin.agent_pre_invoke(payload, plugin_context) + result = await plugin.agent_pre_invoke(payload, plugin_context, ext) - assert result.modified_payload is not None - assert "Authorization" in result.modified_payload.headers.root - assert "Cookie" not in result.modified_payload.headers.root + assert result.modified_extensions is not None + assert result.modified_extensions.http is not None + hdrs = result.modified_extensions.http.headers + assert "Authorization" in hdrs + assert "Cookie" not in hdrs @pytest.mark.asyncio async def test_agent_no_filtered_headers_returns_empty_result(self, plugin_config, plugin_context): """Test agent_pre_invoke returns empty result when no headers filtered.""" plugin = HeaderFilter(plugin_config) - payload = AgentPreInvokePayload( - agent_id="test-agent", - messages=[], - headers=HttpHeaderPayload({"Content-Type": "application/json"}), - ) + payload = AgentPreInvokePayload(agent_id="test-agent", messages=[]) + ext = Extensions(http=HttpExtension(headers={"Content-Type": "application/json"})) - result = await plugin.agent_pre_invoke(payload, plugin_context) + result = await plugin.agent_pre_invoke(payload, plugin_context, ext) - assert result.modified_payload is None + assert result.modified_extensions is None assert result.continue_processing # ── Config and initialization tests ─────────────────────────────── @@ -325,21 +320,20 @@ async def test_default_config_when_config_is_none(self, plugin_context): tags=["test"], mode=PluginMode.SEQUENTIAL, priority=20, + capabilities=["write_headers"], config=None, ) plugin = HeaderFilter(config) - payload = ToolPreInvokePayload( - name="test_tool", - args={}, - headers=HttpHeaderPayload({"Authorization": "Bearer token", "Content-Type": "application/json"}), - ) + payload = ToolPreInvokePayload(name="test_tool", args={}) + ext = Extensions(http=HttpExtension(headers={"Authorization": "Bearer token", "Content-Type": "application/json"})) - result = await plugin.tool_pre_invoke(payload, plugin_context) + result = await plugin.tool_pre_invoke(payload, plugin_context, ext) # Default config should filter Authorization - assert result.modified_payload is not None - assert "Authorization" not in result.modified_payload.headers.root + assert result.modified_extensions is not None + assert result.modified_extensions.http is not None + assert "Authorization" not in result.modified_extensions.http.headers @pytest.mark.asyncio async def test_default_config_when_config_causes_validation_error(self, plugin_context): @@ -354,21 +348,20 @@ async def test_default_config_when_config_causes_validation_error(self, plugin_c tags=["test"], mode=PluginMode.SEQUENTIAL, priority=20, + capabilities=["write_headers"], config={"filter_headers": "not-a-list", "log_filtered_headers": "not-a-bool"}, ) plugin = HeaderFilter(config) - payload = ToolPreInvokePayload( - name="test_tool", - args={}, - headers=HttpHeaderPayload({"Authorization": "Bearer token", "Content-Type": "application/json"}), - ) + payload = ToolPreInvokePayload(name="test_tool", args={}) + ext = Extensions(http=HttpExtension(headers={"Authorization": "Bearer token", "Content-Type": "application/json"})) - result = await plugin.tool_pre_invoke(payload, plugin_context) + result = await plugin.tool_pre_invoke(payload, plugin_context, ext) # Default config fallback should filter Authorization - assert result.modified_payload is not None - assert "Authorization" not in result.modified_payload.headers.root + assert result.modified_extensions is not None + assert result.modified_extensions.http is not None + assert "Authorization" not in result.modified_extensions.http.headers def test_header_filter_config_defaults(self): """Test HeaderFilterConfig has sensible defaults.""" @@ -445,33 +438,35 @@ async def test_passthrough_for_vault_integration(self, plugin_context): tags=["test"], mode=PluginMode.SEQUENTIAL, priority=20, + capabilities=["write_headers"], config={ "filter_headers": ["Authorization", "Cookie", "X-API-Key"], "allow_passthrough_headers": ["Authorization"], }, ) plugin = HeaderFilter(config) - payload = ToolPreInvokePayload( - name="test_tool", - args={}, - headers=HttpHeaderPayload( - { + payload = ToolPreInvokePayload(name="test_tool", args={}) + ext = Extensions( + http=HttpExtension( + headers={ "Authorization": "Bearer vault_token", "Cookie": "session=abc", "X-API-Key": "secret_key", # pragma: allowlist secret "Content-Type": "application/json", } - ), + ) ) - result = await plugin.tool_pre_invoke(payload, plugin_context) + result = await plugin.tool_pre_invoke(payload, plugin_context, ext) - assert result.modified_payload is not None - assert "Authorization" in result.modified_payload.headers.root - assert result.modified_payload.headers.root["Authorization"] == "Bearer vault_token" - assert "Cookie" not in result.modified_payload.headers.root - assert "X-API-Key" not in result.modified_payload.headers.root - assert "Content-Type" in result.modified_payload.headers.root + assert result.modified_extensions is not None + assert result.modified_extensions.http is not None + hdrs = result.modified_extensions.http.headers + assert "Authorization" in hdrs + assert hdrs["Authorization"] == "Bearer vault_token" + assert "Cookie" not in hdrs + assert "X-API-Key" not in hdrs + assert "Content-Type" in hdrs @pytest.mark.asyncio async def test_multiple_passthrough_headers(self, plugin_context): @@ -486,32 +481,34 @@ async def test_multiple_passthrough_headers(self, plugin_context): tags=["test"], mode=PluginMode.SEQUENTIAL, priority=20, + capabilities=["write_headers"], config={ "filter_headers": ["Authorization", "Cookie", "X-API-Key", "X-Custom-Header"], "allow_passthrough_headers": ["Authorization", "X-Custom-Header"], }, ) plugin = HeaderFilter(config) - payload = ToolPreInvokePayload( - name="test_tool", - args={}, - headers=HttpHeaderPayload( - { + payload = ToolPreInvokePayload(name="test_tool", args={}) + ext = Extensions( + http=HttpExtension( + headers={ "Authorization": "Bearer token", "Cookie": "session=xyz", "X-API-Key": "api_key", "X-Custom-Header": "custom_value", } - ), + ) ) - result = await plugin.tool_pre_invoke(payload, plugin_context) + result = await plugin.tool_pre_invoke(payload, plugin_context, ext) - assert result.modified_payload is not None - assert "Authorization" in result.modified_payload.headers.root - assert "X-Custom-Header" in result.modified_payload.headers.root - assert "Cookie" not in result.modified_payload.headers.root - assert "X-API-Key" not in result.modified_payload.headers.root + assert result.modified_extensions is not None + assert result.modified_extensions.http is not None + hdrs = result.modified_extensions.http.headers + assert "Authorization" in hdrs + assert "X-Custom-Header" in hdrs + assert "Cookie" not in hdrs + assert "X-API-Key" not in hdrs @pytest.mark.asyncio async def test_passthrough_case_insensitive(self, plugin_context): @@ -526,28 +523,30 @@ async def test_passthrough_case_insensitive(self, plugin_context): tags=["test"], mode=PluginMode.SEQUENTIAL, priority=20, + capabilities=["write_headers"], config={ "filter_headers": ["Authorization", "Cookie"], "allow_passthrough_headers": ["authorization"], }, ) plugin = HeaderFilter(config) - payload = ToolPreInvokePayload( - name="test_tool", - args={}, - headers=HttpHeaderPayload( - { + payload = ToolPreInvokePayload(name="test_tool", args={}) + ext = Extensions( + http=HttpExtension( + headers={ "Authorization": "Bearer token", "COOKIE": "session=xyz", } - ), + ) ) - result = await plugin.tool_pre_invoke(payload, plugin_context) + result = await plugin.tool_pre_invoke(payload, plugin_context, ext) - assert result.modified_payload is not None - assert "Authorization" in result.modified_payload.headers.root - assert "COOKIE" not in result.modified_payload.headers.root + assert result.modified_extensions is not None + assert result.modified_extensions.http is not None + hdrs = result.modified_extensions.http.headers + assert "Authorization" in hdrs + assert "COOKIE" not in hdrs # ── Shutdown test ───────────────────────────────────────────────── @@ -575,24 +574,24 @@ async def test_log_filtered_headers_disabled(self, plugin_context): tags=["test"], mode=PluginMode.SEQUENTIAL, priority=20, + capabilities=["write_headers"], config={ "filter_headers": ["Authorization"], "log_filtered_headers": False, }, ) plugin = HeaderFilter(config) - payload = ToolPreInvokePayload( - name="test_tool", - args={}, - headers=HttpHeaderPayload({"Authorization": "Bearer token", "Content-Type": "application/json"}), - ) + payload = ToolPreInvokePayload(name="test_tool", args={}) + ext = Extensions(http=HttpExtension(headers={"Authorization": "Bearer token", "Content-Type": "application/json"})) - result = await plugin.tool_pre_invoke(payload, plugin_context) + result = await plugin.tool_pre_invoke(payload, plugin_context, ext) # Should still filter even with logging disabled - assert result.modified_payload is not None - assert "Authorization" not in result.modified_payload.headers.root - assert "Content-Type" in result.modified_payload.headers.root + assert result.modified_extensions is not None + assert result.modified_extensions.http is not None + hdrs = result.modified_extensions.http.headers + assert "Authorization" not in hdrs + assert "Content-Type" in hdrs if __name__ == "__main__": diff --git a/tests/unit/mcpgateway/plugins/plugins/tools_telemetry_exporter/test_tools_telemetry_exporter.py b/tests/unit/mcpgateway/plugins/plugins/tools_telemetry_exporter/test_tools_telemetry_exporter.py index 15d61811d9..31bdd963ac 100644 --- a/tests/unit/mcpgateway/plugins/plugins/tools_telemetry_exporter/test_tools_telemetry_exporter.py +++ b/tests/unit/mcpgateway/plugins/plugins/tools_telemetry_exporter/test_tools_telemetry_exporter.py @@ -14,7 +14,8 @@ import pytest # First-Party -from cpex.framework import GlobalContext, HttpHeaderPayload, PluginConfig, PluginContext, ToolHookType, ToolPostInvokePayload, ToolPreInvokePayload +from cpex.framework import GlobalContext, PluginConfig, PluginContext, ToolHookType, ToolPostInvokePayload, ToolPreInvokePayload +from cpex.framework.extensions import Extensions, HttpExtension from plugins.tools_telemetry_exporter.telemetry_exporter import ToolsTelemetryExporterPlugin @@ -55,8 +56,10 @@ async def test_pre_invoke_redacts_sensitive_headers(self): payload = ToolPreInvokePayload( name="test_tool", args={"input": "hello"}, - headers=HttpHeaderPayload( - { + ) + ext = Extensions( + http=HttpExtension( + headers={ "Authorization": "Bearer secret-token", "Cookie": "jwt_token=abc123; theme=dark", "X-API-Key": "top-secret", # pragma: allowlist secret @@ -64,10 +67,10 @@ async def test_pre_invoke_redacts_sensitive_headers(self): "Content-Type": "application/json", "X-Request-Id": "req-123", } - ), + ) ) - await plugin.tool_pre_invoke(payload, _create_context()) + await plugin.tool_pre_invoke(payload, _create_context(), ext) attrs = plugin._export_telemetry.await_args.kwargs["attributes"] exported_headers = json.loads(attrs["headers"]) @@ -86,16 +89,18 @@ async def test_pre_invoke_redacts_broad_token_header_patterns(self): payload = ToolPreInvokePayload( name="test_tool", args={}, - headers=HttpHeaderPayload( - { + ) + ext = Extensions( + http=HttpExtension( + headers={ "X-Delegation-Token": "delegation-secret", "Upstream-Authorization": "Bearer upstream-secret", "X-Session-Key": "session-secret", } - ), + ) ) - await plugin.tool_pre_invoke(payload, _create_context()) + await plugin.tool_pre_invoke(payload, _create_context(), ext) attrs = plugin._export_telemetry.await_args.kwargs["attributes"] exported_headers = json.loads(attrs["headers"]) diff --git a/tests/unit/mcpgateway/plugins/plugins/vault/test_vault_plugin.py b/tests/unit/mcpgateway/plugins/plugins/vault/test_vault_plugin.py index 5b71beaf4e..7bf8ac8d42 100644 --- a/tests/unit/mcpgateway/plugins/plugins/vault/test_vault_plugin.py +++ b/tests/unit/mcpgateway/plugins/plugins/vault/test_vault_plugin.py @@ -15,13 +15,13 @@ # First-Party from cpex.framework import ( GlobalContext, - HttpHeaderPayload, PluginConfig, PluginContext, PluginMode, ToolHookType, ToolPreInvokePayload, ) +from cpex.framework.extensions import Extensions, HttpExtension # Import the Vault plugin from plugins.vault.vault_plugin import Vault @@ -43,6 +43,7 @@ def plugin_config(self) -> PluginConfig: tags=["test", "vault"], mode=PluginMode.SEQUENTIAL, priority=10, + capabilities=["write_headers"], config={ "system_tag_prefix": "system", "vault_header_name": "X-Vault-Tokens", @@ -67,11 +68,12 @@ async def test_no_vault_header_returns_empty_result(self, plugin_config, plugin_ plugin = Vault(plugin_config) # Create payload without vault header - payload = ToolPreInvokePayload(name="test_tool", arguments={}, headers=HttpHeaderPayload(root={"Content-Type": "application/json"})) + payload = ToolPreInvokePayload(name="test_tool", args={}) + ext = Extensions(http=HttpExtension(headers={"Content-Type": "application/json"})) - result = await plugin.tool_pre_invoke(payload, plugin_context) + result = await plugin.tool_pre_invoke(payload, plugin_context, ext) - assert result.modified_payload is None + assert result.modified_extensions is None assert result.continue_processing @pytest.mark.asyncio @@ -83,14 +85,17 @@ async def test_vault_token_added_to_authorization_header(self, plugin_config, pl vault_tokens = {"github.com": "ghp_test123456789"} # Create payload with vault header (lowercase per ASGI spec) - payload = ToolPreInvokePayload(name="test_tool", arguments={}, headers=HttpHeaderPayload(root={"content-type": "application/json", "x-vault-tokens": json.dumps(vault_tokens)})) + payload = ToolPreInvokePayload(name="test_tool", args={}) + ext = Extensions(http=HttpExtension(headers={"content-type": "application/json", "x-vault-tokens": json.dumps(vault_tokens)})) - result = await plugin.tool_pre_invoke(payload, plugin_context) + result = await plugin.tool_pre_invoke(payload, plugin_context, ext) - assert result.modified_payload is not None - assert "authorization" in result.modified_payload.headers.root - assert result.modified_payload.headers.root["authorization"] == "Bearer ghp_test123456789" - assert "x-vault-tokens" not in result.modified_payload.headers.root + assert result.modified_extensions is not None + assert result.modified_extensions.http is not None + hdrs = result.modified_extensions.http.headers + assert "authorization" in hdrs + assert hdrs["authorization"] == "Bearer ghp_test123456789" + assert "x-vault-tokens" not in hdrs @pytest.mark.asyncio async def test_pat_token_uses_custom_header(self, plugin_config, plugin_context): @@ -101,14 +106,17 @@ async def test_pat_token_uses_custom_header(self, plugin_config, plugin_context) vault_tokens = {"github.com:USER:PAT:TOKEN": "ghp_pat_token123"} # Create payload with vault header (lowercase per ASGI spec) - payload = ToolPreInvokePayload(name="test_tool", arguments={}, headers=HttpHeaderPayload(root={"content-type": "application/json", "x-vault-tokens": json.dumps(vault_tokens)})) + payload = ToolPreInvokePayload(name="test_tool", args={}) + ext = Extensions(http=HttpExtension(headers={"content-type": "application/json", "x-vault-tokens": json.dumps(vault_tokens)})) - result = await plugin.tool_pre_invoke(payload, plugin_context) + result = await plugin.tool_pre_invoke(payload, plugin_context, ext) - assert result.modified_payload is not None - assert "x-github-token" in result.modified_payload.headers.root - assert result.modified_payload.headers.root["x-github-token"] == "ghp_pat_token123" - assert "x-vault-tokens" not in result.modified_payload.headers.root + assert result.modified_extensions is not None + assert result.modified_extensions.http is not None + hdrs = result.modified_extensions.http.headers + assert "x-github-token" in hdrs + assert hdrs["x-github-token"] == "ghp_pat_token123" + assert "x-vault-tokens" not in hdrs @pytest.mark.asyncio async def test_invalid_json_in_vault_header(self, plugin_config, plugin_context): @@ -116,13 +124,16 @@ async def test_invalid_json_in_vault_header(self, plugin_config, plugin_context) plugin = Vault(plugin_config) # Create payload with invalid JSON (lowercase per ASGI spec) - payload = ToolPreInvokePayload(name="test_tool", arguments={}, headers=HttpHeaderPayload(root={"content-type": "application/json", "x-vault-tokens": "invalid json"})) + payload = ToolPreInvokePayload(name="test_tool", args={}) + ext = Extensions(http=HttpExtension(headers={"content-type": "application/json", "x-vault-tokens": "invalid json"})) - result = await plugin.tool_pre_invoke(payload, plugin_context) + result = await plugin.tool_pre_invoke(payload, plugin_context, ext) # SECURITY: Vault header must be removed even on parse error - assert result.modified_payload is not None - assert "x-vault-tokens" not in result.modified_payload.headers.root + assert result.modified_extensions is not None + assert result.modified_extensions.http is not None + hdrs = result.modified_extensions.http.headers + assert "x-vault-tokens" not in hdrs assert result.continue_processing @pytest.mark.asyncio @@ -139,13 +150,16 @@ async def test_no_system_tag_strips_vault_header(self, plugin_config): vault_tokens = {"github.com": "token123"} - payload = ToolPreInvokePayload(name="test_tool", arguments={}, headers=HttpHeaderPayload(root={"x-vault-tokens": json.dumps(vault_tokens)})) + payload = ToolPreInvokePayload(name="test_tool", args={}) + ext = Extensions(http=HttpExtension(headers={"x-vault-tokens": json.dumps(vault_tokens)})) - result = await plugin.tool_pre_invoke(payload, context) + result = await plugin.tool_pre_invoke(payload, context, ext) # SECURITY: Vault header must be removed even when system tag is missing - assert result.modified_payload is not None - assert "x-vault-tokens" not in result.modified_payload.headers.root + assert result.modified_extensions is not None + assert result.modified_extensions.http is not None + hdrs = result.modified_extensions.http.headers + assert "x-vault-tokens" not in hdrs assert result.continue_processing @pytest.mark.asyncio @@ -157,11 +171,12 @@ async def test_no_system_tag_no_vault_header_returns_empty(self, plugin_config): global_context = GlobalContext(request_id="test-3", metadata={"gateway": gateway_metadata}) context = PluginContext(global_context=global_context) - payload = ToolPreInvokePayload(name="test_tool", arguments={}, headers=HttpHeaderPayload(root={"Content-Type": "application/json"})) + payload = ToolPreInvokePayload(name="test_tool", args={}) + ext = Extensions(http=HttpExtension(headers={"Content-Type": "application/json"})) - result = await plugin.tool_pre_invoke(payload, context) + result = await plugin.tool_pre_invoke(payload, context, ext) - assert result.modified_payload is None + assert result.modified_extensions is None assert result.continue_processing @pytest.mark.asyncio @@ -170,15 +185,18 @@ async def test_non_dict_json_vault_tokens_stripped(self, plugin_config, plugin_c plugin = Vault(plugin_config) # JSON array is valid JSON but not a dict — must not be treated as tokens (lowercase per ASGI spec) - payload = ToolPreInvokePayload(name="test_tool", arguments={}, headers=HttpHeaderPayload(root={"content-type": "application/json", "x-vault-tokens": '["not", "a", "dict"]'})) + payload = ToolPreInvokePayload(name="test_tool", args={}) + ext = Extensions(http=HttpExtension(headers={"content-type": "application/json", "x-vault-tokens": '["not", "a", "dict"]'})) - result = await plugin.tool_pre_invoke(payload, plugin_context) + result = await plugin.tool_pre_invoke(payload, plugin_context, ext) # SECURITY: Vault header must be removed - assert result.modified_payload is not None - assert "x-vault-tokens" not in result.modified_payload.headers.root + assert result.modified_extensions is not None + assert result.modified_extensions.http is not None + hdrs = result.modified_extensions.http.headers + assert "x-vault-tokens" not in hdrs # No Authorization header should be injected - assert "authorization" not in result.modified_payload.headers.root + assert "authorization" not in hdrs assert result.continue_processing @pytest.mark.asyncio @@ -189,13 +207,16 @@ async def test_complex_token_key_parsing(self, plugin_config, plugin_context): # Create vault tokens with complex key vault_tokens = {"github.com:USER:OAUTH2:ACCESS_TOKEN": "oauth_token_123"} - payload = ToolPreInvokePayload(name="test_tool", arguments={}, headers=HttpHeaderPayload(root={"x-vault-tokens": json.dumps(vault_tokens)})) + payload = ToolPreInvokePayload(name="test_tool", args={}) + ext = Extensions(http=HttpExtension(headers={"x-vault-tokens": json.dumps(vault_tokens)})) - result = await plugin.tool_pre_invoke(payload, plugin_context) + result = await plugin.tool_pre_invoke(payload, plugin_context, ext) - assert result.modified_payload is not None - assert "authorization" in result.modified_payload.headers.root - assert result.modified_payload.headers.root["authorization"] == "Bearer oauth_token_123" + assert result.modified_extensions is not None + assert result.modified_extensions.http is not None + hdrs = result.modified_extensions.http.headers + assert "authorization" in hdrs + assert hdrs["authorization"] == "Bearer oauth_token_123" def test_parse_vault_token_key(self, plugin_config): """Test the _parse_vault_token_key method.""" @@ -224,20 +245,27 @@ async def test_existing_bearer_token_is_replaced(self, plugin_config, plugin_con vault_tokens = {"github.com": "ghp_new_token_from_vault"} # Create payload with existing Authorization header (lowercase per ASGI spec) - payload = ToolPreInvokePayload( - name="test_tool", - arguments={}, - headers=HttpHeaderPayload(root={"content-type": "application/json", "authorization": "Bearer old_default_token", "x-vault-tokens": json.dumps(vault_tokens)}), + payload = ToolPreInvokePayload(name="test_tool", args={}) + ext = Extensions( + http=HttpExtension( + headers={ + "content-type": "application/json", + "authorization": "Bearer old_default_token", + "x-vault-tokens": json.dumps(vault_tokens), + } + ) ) - result = await plugin.tool_pre_invoke(payload, plugin_context) + result = await plugin.tool_pre_invoke(payload, plugin_context, ext) # Verify the old token was replaced with the new one from vault - assert result.modified_payload is not None - assert "authorization" in result.modified_payload.headers.root - assert result.modified_payload.headers.root["authorization"] == "Bearer ghp_new_token_from_vault" - assert result.modified_payload.headers.root["authorization"] != "Bearer old_default_token" - assert "x-vault-tokens" not in result.modified_payload.headers.root + assert result.modified_extensions is not None + assert result.modified_extensions.http is not None + hdrs = result.modified_extensions.http.headers + assert "authorization" in hdrs + assert hdrs["authorization"] == "Bearer ghp_new_token_from_vault" + assert hdrs["authorization"] != "Bearer old_default_token" + assert "x-vault-tokens" not in hdrs @pytest.mark.asyncio async def test_existing_custom_header_is_replaced_with_pat(self, plugin_config, plugin_context): @@ -248,18 +276,27 @@ async def test_existing_custom_header_is_replaced_with_pat(self, plugin_config, vault_tokens = {"github.com:USER:PAT:TOKEN": "ghp_new_pat_token"} # Create payload with existing custom header (lowercase per ASGI spec) - payload = ToolPreInvokePayload( - name="test_tool", arguments={}, headers=HttpHeaderPayload(root={"content-type": "application/json", "x-github-token": "old_github_token", "x-vault-tokens": json.dumps(vault_tokens)}) + payload = ToolPreInvokePayload(name="test_tool", args={}) + ext = Extensions( + http=HttpExtension( + headers={ + "content-type": "application/json", + "x-github-token": "old_github_token", + "x-vault-tokens": json.dumps(vault_tokens), + } + ) ) - result = await plugin.tool_pre_invoke(payload, plugin_context) + result = await plugin.tool_pre_invoke(payload, plugin_context, ext) # Verify the old custom header was replaced with the new PAT token - assert result.modified_payload is not None - assert "x-github-token" in result.modified_payload.headers.root - assert result.modified_payload.headers.root["x-github-token"] == "ghp_new_pat_token" - assert result.modified_payload.headers.root["x-github-token"] != "old_github_token" - assert "x-vault-tokens" not in result.modified_payload.headers.root + assert result.modified_extensions is not None + assert result.modified_extensions.http is not None + hdrs = result.modified_extensions.http.headers + assert "x-github-token" in hdrs + assert hdrs["x-github-token"] == "ghp_new_pat_token" + assert hdrs["x-github-token"] != "old_github_token" + assert "x-vault-tokens" not in hdrs @pytest.mark.asyncio async def test_vault_header_removed_when_no_token_match(self, plugin_config, plugin_context): @@ -270,15 +307,18 @@ async def test_vault_header_removed_when_no_token_match(self, plugin_config, plu vault_tokens = {"gitlab.com": "glpat_different_system_token"} # Create payload with vault header but no matching system (lowercase per ASGI spec) - payload = ToolPreInvokePayload(name="test_tool", arguments={}, headers=HttpHeaderPayload(root={"content-type": "application/json", "x-vault-tokens": json.dumps(vault_tokens)})) + payload = ToolPreInvokePayload(name="test_tool", args={}) + ext = Extensions(http=HttpExtension(headers={"content-type": "application/json", "x-vault-tokens": json.dumps(vault_tokens)})) - result = await plugin.tool_pre_invoke(payload, plugin_context) + result = await plugin.tool_pre_invoke(payload, plugin_context, ext) # SECURITY: Vault header must be removed even when no token match is found - assert result.modified_payload is not None - assert "x-vault-tokens" not in result.modified_payload.headers.root + assert result.modified_extensions is not None + assert result.modified_extensions.http is not None + hdrs = result.modified_extensions.http.headers + assert "x-vault-tokens" not in hdrs # No Authorization header should be added since there's no match - assert "authorization" not in result.modified_payload.headers.root + assert "authorization" not in hdrs assert result.continue_processing @pytest.mark.asyncio @@ -298,26 +338,27 @@ async def test_case_insensitive_vault_header_detection(self, plugin_config, plug ] for header_name in test_cases: - payload = ToolPreInvokePayload(name="test_tool", arguments={}, headers=HttpHeaderPayload(root={"Content-Type": "application/json", header_name: json.dumps(vault_tokens)})) + payload = ToolPreInvokePayload(name="test_tool", args={}) + ext = Extensions(http=HttpExtension(headers={"Content-Type": "application/json", header_name: json.dumps(vault_tokens)})) - result = await plugin.tool_pre_invoke(payload, plugin_context) + result = await plugin.tool_pre_invoke(payload, plugin_context, ext) # Should work regardless of case - assert result.modified_payload is not None, f"Failed for header case: {header_name}" - assert "authorization" in result.modified_payload.headers.root, f"Authorization not found for case: {header_name}" - assert result.modified_payload.headers.root["authorization"] == "Bearer ghp_test_case_insensitive" + assert result.modified_extensions is not None, f"Failed for header case: {header_name}" + assert result.modified_extensions.http is not None + hdrs = result.modified_extensions.http.headers + assert "authorization" in hdrs, f"Authorization not found for case: {header_name}" + assert hdrs["authorization"] == "Bearer ghp_test_case_insensitive" # Vault header should be removed (check all possible cases) - for key in result.modified_payload.headers.root.keys(): + for key in hdrs.keys(): assert key.lower() != "x-vault-tokens", f"Vault header not removed for case: {header_name}" @pytest.mark.asyncio async def test_copyonwritedict_root_attribute_access(self, plugin_config, plugin_context): - """Test that plugin correctly accesses headers via .root attribute. + """Test that plugin correctly processes headers via extensions.http.headers. - This test validates the fix for the CopyOnWriteDict.model_dump() bug where - model_dump() returned an empty dict in production (with real HTTP requests) - instead of the actual headers. While model_dump() works in unit tests, - the plugin now uses .root directly for consistency and reliability. + This test validates header access through the extensions path (replacing the + former CopyOnWriteDict.model_dump() / .root pattern on payload.headers). """ plugin = Vault(plugin_config) @@ -325,25 +366,29 @@ async def test_copyonwritedict_root_attribute_access(self, plugin_config, plugin vault_tokens = {"github.com": "ghp_root_access_test"} # Create payload with vault header - payload = ToolPreInvokePayload(name="test_tool", arguments={}, headers=HttpHeaderPayload(root={"Content-Type": "application/json", "X-Vault-Tokens": json.dumps(vault_tokens), "X-Custom-Header": "custom_value"})) + payload = ToolPreInvokePayload(name="test_tool", args={}) + headers_in = {"Content-Type": "application/json", "X-Vault-Tokens": json.dumps(vault_tokens), "X-Custom-Header": "custom_value"} + ext = Extensions(http=HttpExtension(headers=headers_in)) - # Verify .root contains actual headers (the correct access pattern) - assert len(payload.headers.root) == 3, "Headers.root should contain all headers" - assert "X-Vault-Tokens" in payload.headers.root - assert "X-Custom-Header" in payload.headers.root - assert "Content-Type" in payload.headers.root + # Verify extensions contain actual headers (the correct access pattern) + assert len(ext.http.headers) == 3, "extensions.http.headers should contain all headers" + assert "X-Vault-Tokens" in ext.http.headers + assert "X-Custom-Header" in ext.http.headers + assert "Content-Type" in ext.http.headers - # Plugin should work correctly by using .root - result = await plugin.tool_pre_invoke(payload, plugin_context) + # Plugin should work correctly by using extensions.http.headers + result = await plugin.tool_pre_invoke(payload, plugin_context, ext) # Verify plugin processed the headers correctly - assert result.modified_payload is not None - assert "authorization" in result.modified_payload.headers.root - assert result.modified_payload.headers.root["authorization"] == "Bearer ghp_root_access_test" - assert "x-vault-tokens" not in result.modified_payload.headers.root + assert result.modified_extensions is not None + assert result.modified_extensions.http is not None + hdrs = result.modified_extensions.http.headers + assert "authorization" in hdrs + assert hdrs["authorization"] == "Bearer ghp_root_access_test" + assert "x-vault-tokens" not in hdrs # Custom header should be preserved (normalized to lowercase) - assert "x-custom-header" in result.modified_payload.headers.root - assert result.modified_payload.headers.root["x-custom-header"] == "custom_value" + assert "x-custom-header" in hdrs + assert hdrs["x-custom-header"] == "custom_value" @pytest.mark.asyncio async def test_case_insensitive_header_with_config_variations(self, plugin_context): @@ -379,6 +424,7 @@ async def test_case_insensitive_header_with_config_variations(self, plugin_conte tags=["test", "vault"], mode=PluginMode.SEQUENTIAL, priority=10, + capabilities=["write_headers"], config={ "system_tag_prefix": "system", "vault_header_name": config_header, # Variable case @@ -391,21 +437,20 @@ async def test_case_insensitive_header_with_config_variations(self, plugin_conte plugin = Vault(plugin_config) # Create payload with specific request header case - payload = ToolPreInvokePayload( - name="test_tool", - arguments={}, - headers=HttpHeaderPayload(root={"Content-Type": "application/json", request_header: json.dumps(vault_tokens)}), # Variable case - ) + payload = ToolPreInvokePayload(name="test_tool", args={}) + ext = Extensions(http=HttpExtension(headers={"Content-Type": "application/json", request_header: json.dumps(vault_tokens)})) - result = await plugin.tool_pre_invoke(payload, plugin_context) + result = await plugin.tool_pre_invoke(payload, plugin_context, ext) # Should work regardless of case combination - assert result.modified_payload is not None, f"Failed for config='{config_header}', request='{request_header}'" - assert "authorization" in result.modified_payload.headers.root, f"Authorization not found for config='{config_header}', request='{request_header}'" - assert result.modified_payload.headers.root["authorization"] == "Bearer ghp_case_test" + assert result.modified_extensions is not None, f"Failed for config='{config_header}', request='{request_header}'" + assert result.modified_extensions.http is not None + hdrs = result.modified_extensions.http.headers + assert "authorization" in hdrs, f"Authorization not found for config='{config_header}', request='{request_header}'" + assert hdrs["authorization"] == "Bearer ghp_case_test" # Vault header should be removed (check all possible cases) - for key in result.modified_payload.headers.root.keys(): + for key in hdrs.keys(): assert key.lower() != config_header.lower(), f"Vault header not removed for config='{config_header}', request='{request_header}'" @pytest.mark.asyncio @@ -423,15 +468,18 @@ async def test_vault_header_lowercase_headers(self, plugin_config, plugin_contex vault_tokens = {"github.com": "ghp_production_token_456"} # Create payload with LOWERCASE vault header (as ASGI middleware provides) - payload = ToolPreInvokePayload(name="test_tool", arguments={}, headers=HttpHeaderPayload(root={"content-type": "application/json", "x-vault-tokens": json.dumps(vault_tokens)})) + payload = ToolPreInvokePayload(name="test_tool", args={}) + ext = Extensions(http=HttpExtension(headers={"content-type": "application/json", "x-vault-tokens": json.dumps(vault_tokens)})) - result = await plugin.tool_pre_invoke(payload, plugin_context) + result = await plugin.tool_pre_invoke(payload, plugin_context, ext) # Token should be found and processed despite case mismatch - assert result.modified_payload is not None - assert "authorization" in result.modified_payload.headers.root - assert result.modified_payload.headers.root["authorization"] == "Bearer ghp_production_token_456" - assert "x-vault-tokens" not in result.modified_payload.headers.root + assert result.modified_extensions is not None + assert result.modified_extensions.http is not None + hdrs = result.modified_extensions.http.headers + assert "authorization" in hdrs + assert hdrs["authorization"] == "Bearer ghp_production_token_456" + assert "x-vault-tokens" not in hdrs @pytest.mark.asyncio async def test_vault_header_lowercase_config(self, plugin_context): @@ -452,6 +500,7 @@ async def test_vault_header_lowercase_config(self, plugin_context): tags=["test", "vault"], mode=PluginMode.SEQUENTIAL, priority=10, + capabilities=["write_headers"], config={ "system_tag_prefix": "system", "vault_header_name": "x-vault-tokens", # lowercase config @@ -467,15 +516,18 @@ async def test_vault_header_lowercase_config(self, plugin_context): vault_tokens = {"github.com": "ghp_lowercase_config_token"} # Create payload with lowercase vault header - payload = ToolPreInvokePayload(name="test_tool", arguments={}, headers=HttpHeaderPayload(root={"x-vault-tokens": json.dumps(vault_tokens)})) + payload = ToolPreInvokePayload(name="test_tool", args={}) + ext = Extensions(http=HttpExtension(headers={"x-vault-tokens": json.dumps(vault_tokens)})) - result = await plugin.tool_pre_invoke(payload, plugin_context) + result = await plugin.tool_pre_invoke(payload, plugin_context, ext) # Token should be found and processed - assert result.modified_payload is not None - assert "authorization" in result.modified_payload.headers.root - assert result.modified_payload.headers.root["authorization"] == "Bearer ghp_lowercase_config_token" - assert "x-vault-tokens" not in result.modified_payload.headers.root + assert result.modified_extensions is not None + assert result.modified_extensions.http is not None + hdrs = result.modified_extensions.http.headers + assert "authorization" in hdrs + assert hdrs["authorization"] == "Bearer ghp_lowercase_config_token" + assert "x-vault-tokens" not in hdrs @pytest.mark.asyncio async def test_pat_token_custom_header_lowercase(self, plugin_config, plugin_context): @@ -491,15 +543,18 @@ async def test_pat_token_custom_header_lowercase(self, plugin_config, plugin_con vault_tokens = {"github.com:USER:PAT:TOKEN": "ghp_pat_lowercase_header"} # Create payload with lowercase headers (production ASGI flow) - payload = ToolPreInvokePayload(name="test_tool", arguments={}, headers=HttpHeaderPayload(root={"content-type": "application/json", "x-vault-tokens": json.dumps(vault_tokens)})) + payload = ToolPreInvokePayload(name="test_tool", args={}) + ext = Extensions(http=HttpExtension(headers={"content-type": "application/json", "x-vault-tokens": json.dumps(vault_tokens)})) - result = await plugin.tool_pre_invoke(payload, plugin_context) + result = await plugin.tool_pre_invoke(payload, plugin_context, ext) # PAT token should be written with lowercase header name - assert result.modified_payload is not None - assert "x-github-token" in result.modified_payload.headers.root - assert result.modified_payload.headers.root["x-github-token"] == "ghp_pat_lowercase_header" - assert "x-vault-tokens" not in result.modified_payload.headers.root + assert result.modified_extensions is not None + assert result.modified_extensions.http is not None + hdrs = result.modified_extensions.http.headers + assert "x-github-token" in hdrs + assert hdrs["x-github-token"] == "ghp_pat_lowercase_header" + assert "x-vault-tokens" not in hdrs if __name__ == "__main__": diff --git a/tests/unit/mcpgateway/services/test_a2a_agent_invoke_hooks.py b/tests/unit/mcpgateway/services/test_a2a_agent_invoke_hooks.py index ff4e8009eb..af3266d9bf 100644 --- a/tests/unit/mcpgateway/services/test_a2a_agent_invoke_hooks.py +++ b/tests/unit/mcpgateway/services/test_a2a_agent_invoke_hooks.py @@ -12,7 +12,7 @@ additional mocking of `_get_plugin_manager` to control plugin hook behaviour. Tests cover: - - PRE_INVOKE fires with correct AgentPreInvokePayload (headers=request_headers) + - PRE_INVOKE fires with correct AgentPreInvokePayload (headers via extensions.http) - PRE_INVOKE applies modified headers to prepared request - PRE_INVOKE applies modified parameters - PRE_INVOKE PluginViolationError is raised as A2AAgentError @@ -91,7 +91,7 @@ def _has_hooks(hook_type): return False pm.has_hooks_for = MagicMock(side_effect=_has_hooks) - pm.invoke_hook = AsyncMock(return_value=(SimpleNamespace(modified_payload=None, retry_delay_ms=0, metadata=None), {})) + pm.invoke_hook = AsyncMock(return_value=(SimpleNamespace(modified_payload=None, modified_extensions=None, retry_delay_ms=0, metadata=None), {})) return pm @@ -147,12 +147,14 @@ async def test_request_headers_filtered_by_passthrough_whitelist( ) payload = pm.invoke_hook.await_args_list[0].kwargs["payload"] - filtered = payload.headers.root + extensions = pm.invoke_hook.await_args_list[0].kwargs.get("extensions") + filtered = extensions.http.headers if extensions and extensions.http else {} assert "x-tenant-id" in filtered assert "x-request-id" in filtered assert "user-agent" not in filtered assert "referer" not in filtered assert "authorization" not in filtered + assert payload.headers is None @patch("mcpgateway.services.metrics_buffer_service.get_metrics_buffer_service") @patch("mcpgateway.services.a2a_service.fresh_db_session") @@ -196,7 +198,9 @@ async def test_no_passthrough_headers_strips_all( ) payload = pm.invoke_hook.await_args_list[0].kwargs["payload"] - assert payload.headers.root == {} + extensions = pm.invoke_hook.await_args_list[0].kwargs.get("extensions") + assert extensions is None or extensions.http is None or extensions.http.headers == {} + assert payload.headers is None @patch("mcpgateway.services.metrics_buffer_service.get_metrics_buffer_service") @patch("mcpgateway.services.a2a_service.fresh_db_session") @@ -212,7 +216,7 @@ async def test_pre_invoke_hook_receives_request_headers( mock_db, mock_agent, ): - """PRE_INVOKE receives the inbound request_headers as AgentPreInvokePayload.headers.""" + """PRE_INVOKE receives the inbound request_headers via extensions.http.headers.""" # Third-Party from cpex.framework import AgentHookType, AgentPreInvokePayload @@ -243,10 +247,11 @@ async def test_pre_invoke_hook_receives_request_headers( first_call = pm.invoke_hook.await_args_list[0] assert first_call.args[0] == AgentHookType.AGENT_PRE_INVOKE payload = first_call.kwargs["payload"] + extensions = first_call.kwargs.get("extensions") assert isinstance(payload, AgentPreInvokePayload) assert payload.agent_id == mock_agent.id - assert payload.headers is not None - assert payload.headers.root == inbound_headers + assert extensions is not None and extensions.http is not None + assert extensions.http.headers == inbound_headers @patch("mcpgateway.services.metrics_buffer_service.get_metrics_buffer_service") @patch("mcpgateway.services.a2a_service.fresh_db_session") @@ -287,7 +292,9 @@ async def test_pre_invoke_default_headers_when_no_request_headers( ) payload = pm.invoke_hook.await_args_list[0].kwargs["payload"] - assert payload.headers.root == {} + extensions = pm.invoke_hook.await_args_list[0].kwargs.get("extensions") + assert extensions is None or extensions.http is None or extensions.http.headers == {} + assert payload.headers is None @patch("mcpgateway.services.metrics_buffer_service.get_metrics_buffer_service") @patch("mcpgateway.services.a2a_service.fresh_db_session") @@ -305,7 +312,7 @@ async def test_pre_invoke_modified_headers_applied_to_prepared( ): """Headers modified by the plugin are applied to the outbound prepared request headers.""" # Third-Party - from cpex.framework import HttpHeaderPayload + from cpex.framework.extensions import Extensions, HttpExtension # Add headers to passthrough whitelist so they pass Layer 1 filtering mock_agent.passthrough_headers = ["x-custom", "x-request-id"] @@ -314,8 +321,9 @@ async def test_pre_invoke_modified_headers_applied_to_prepared( modified = SimpleNamespace( modified_payload=SimpleNamespace( parameters=None, - headers=HttpHeaderPayload(root={"X-Custom": "value", "X-Request-ID": "plugin-req-123"}), + headers=None, ), + modified_extensions=Extensions(http=HttpExtension(headers={"X-Custom": "value", "X-Request-ID": "plugin-req-123"})), retry_delay_ms=0, metadata=None, ) @@ -501,7 +509,7 @@ async def _invoke_hook_side_effect(*args, **kwargs): nonlocal call_count call_count += 1 if call_count == 1: - return (SimpleNamespace(modified_payload=None, metadata=None), {}) + return (SimpleNamespace(modified_payload=None, modified_extensions=None, metadata=None), {}) raise RuntimeError("post crash") pm.invoke_hook = AsyncMock(side_effect=_invoke_hook_side_effect) @@ -698,7 +706,7 @@ async def _capture_plugin_manager(agent_context_id): async def _invoke_hook(hook_type, payload, global_context=None, local_contexts=None, violations_as_exceptions=True, extensions=None): captured_context["global"] = global_context - return (SimpleNamespace(modified_payload=None, retry_delay_ms=0, metadata=None), {}) + return (SimpleNamespace(modified_payload=None, modified_extensions=None, retry_delay_ms=0, metadata=None), {}) pm.invoke_hook = _invoke_hook return pm @@ -754,7 +762,7 @@ async def test_global_context_metadata_contains_agent( async def _invoke_hook(hook_type, payload, global_context=None, local_contexts=None, violations_as_exceptions=True, extensions=None): captured_context["global"] = global_context - return (SimpleNamespace(modified_payload=None, retry_delay_ms=0, metadata=None), {}) + return (SimpleNamespace(modified_payload=None, modified_extensions=None, retry_delay_ms=0, metadata=None), {}) pm.invoke_hook = _invoke_hook @@ -809,7 +817,7 @@ async def test_content_type_in_agent_metadata( async def _invoke_hook(hook_type, payload, global_context=None, local_contexts=None, violations_as_exceptions=True, extensions=None): captured_context["global"] = global_context - return (SimpleNamespace(modified_payload=None, retry_delay_ms=0, metadata=None), {}) + return (SimpleNamespace(modified_payload=None, modified_extensions=None, retry_delay_ms=0, metadata=None), {}) pm.invoke_hook = _invoke_hook diff --git a/tests/unit/mcpgateway/services/test_tool_pre_invoke_logging.py b/tests/unit/mcpgateway/services/test_tool_pre_invoke_logging.py index f2adb54e88..fd674726b7 100644 --- a/tests/unit/mcpgateway/services/test_tool_pre_invoke_logging.py +++ b/tests/unit/mcpgateway/services/test_tool_pre_invoke_logging.py @@ -13,7 +13,7 @@ from types import SimpleNamespace # Third-Party -from cpex.framework import HttpHeaderPayload +from cpex.framework.extensions import Extensions, HttpExtension from pydantic import BaseModel # First-Party @@ -24,8 +24,8 @@ def test_tool_pre_invoke_logging_without_modified_payload_logs_only_keys(caplog): """No modified payload should be logged without argument or header values.""" original_args = {"normal\nkey": "visible-value", "wxo_auth": "secret-token"} # pragma: allowlist secret - original_headers = HttpHeaderPayload(root={"Authorization": "Bearer secret", "x-wxo-access-token": "secret"}) - pre_result = SimpleNamespace(modified_payload=None) + original_headers = {"Authorization": "Bearer secret", "x-wxo-access-token": "secret"} + pre_result = SimpleNamespace(modified_payload=None, modified_extensions=None) with caplog.at_level(logging.DEBUG, logger="mcpgateway.services.tool_service"): _log_tool_pre_invoke_result("rishiserver-list-all-secrets", original_args, original_headers, pre_result) @@ -47,13 +47,16 @@ def test_tool_pre_invoke_logging_with_modified_payload_logs_key_diffs(caplog): "wxo_connection_id": "", "wxo_environment_id": "draft", } - original_headers = HttpHeaderPayload(root={"Authorization": "Bearer secret", "x-old": "old-value"}) + original_headers = {"Authorization": "Bearer secret", "x-old": "old-value"} modified_payload = SimpleNamespace( name="renamed-tool", args={"real_arg": "changed-value"}, - headers=HttpHeaderPayload(root={"x-connection": "connection-secret"}), + headers=None, + ) + pre_result = SimpleNamespace( + modified_payload=modified_payload, + modified_extensions=Extensions(http=HttpExtension(headers={"x-connection": "connection-secret"})), ) - pre_result = SimpleNamespace(modified_payload=modified_payload) with caplog.at_level(logging.DEBUG, logger="mcpgateway.services.tool_service"): _log_tool_pre_invoke_result("rishiserver-list-all-secrets", original_args, original_headers, pre_result) @@ -72,7 +75,7 @@ def test_tool_pre_invoke_logging_with_modified_payload_logs_key_diffs(caplog): def test_tool_pre_invoke_logging_handles_missing_mappings(caplog): """Diagnostics should tolerate absent headers and non-mapping args.""" - pre_result = SimpleNamespace(modified_payload=SimpleNamespace(name="tool", args=None, headers=None)) + pre_result = SimpleNamespace(modified_payload=SimpleNamespace(name="tool", args=None, headers=None), modified_extensions=None) with caplog.at_level(logging.DEBUG, logger="mcpgateway.services.tool_service"): _log_tool_pre_invoke_result("tool", None, None, pre_result) @@ -87,7 +90,7 @@ def test_tool_pre_invoke_logging_handles_missing_mappings(caplog): def test_tool_pre_invoke_logging_sanitizes_tool_and_modified_names(caplog): """Caller and plugin-controlled tool names should be sanitized before logging.""" modified_payload = SimpleNamespace(name="renamed\nCRITICAL\x1b[31m", args={}, headers=None) - pre_result = SimpleNamespace(modified_payload=modified_payload) + pre_result = SimpleNamespace(modified_payload=modified_payload, modified_extensions=None) with caplog.at_level(logging.DEBUG, logger="mcpgateway.services.tool_service"): _log_tool_pre_invoke_result("tool\r\nERROR\x1b[31m", {}, None, pre_result) @@ -111,7 +114,7 @@ def fail_if_called(_value): old_level = logger.level try: logger.setLevel(logging.INFO) - _log_tool_pre_invoke_result("tool", {"key": "value"}, None, SimpleNamespace(modified_payload=None)) + _log_tool_pre_invoke_result("tool", {"key": "value"}, None, SimpleNamespace(modified_payload=None, modified_extensions=None)) finally: logger.setLevel(old_level) @@ -125,7 +128,7 @@ def broken_sanitizer(_value): monkeypatch.setattr(tool_service_module, "sanitize_for_log", broken_sanitizer) with caplog.at_level(logging.DEBUG, logger="mcpgateway.services.tool_service"): - _log_tool_pre_invoke_result("tool", {"key": "value"}, None, SimpleNamespace(modified_payload=None)) + _log_tool_pre_invoke_result("tool", {"key": "value"}, None, SimpleNamespace(modified_payload=None, modified_extensions=None)) assert "tool_pre_invoke diagnostic logging failed" in caplog.text @@ -142,7 +145,7 @@ def debug(self, *_args, **_kwargs): monkeypatch.setattr(tool_service_module, "logger", RaisingLogger()) - _log_tool_pre_invoke_result("tool", {}, None, SimpleNamespace(modified_payload=None)) + _log_tool_pre_invoke_result("tool", {}, None, SimpleNamespace(modified_payload=None, modified_extensions=None)) def test_tool_pre_invoke_logging_handles_pydantic_model_args(caplog): @@ -153,7 +156,7 @@ class ArgsModel(BaseModel): wxo_auth: str args = ArgsModel(visible="keep-me", wxo_auth="secret-token") # pragma: allowlist secret - pre_result = SimpleNamespace(modified_payload=None) + pre_result = SimpleNamespace(modified_payload=None, modified_extensions=None) with caplog.at_level(logging.DEBUG, logger="mcpgateway.services.tool_service"): _log_tool_pre_invoke_result("tool", args, None, pre_result) diff --git a/tests/unit/mcpgateway/services/test_tool_service.py b/tests/unit/mcpgateway/services/test_tool_service.py index a758cd0130..6ae62050f2 100644 --- a/tests/unit/mcpgateway/services/test_tool_service.py +++ b/tests/unit/mcpgateway/services/test_tool_service.py @@ -10292,7 +10292,8 @@ async def test_prepare_rust_mcp_tool_execution_oauth_authorization_code_plugin_i post-hook check sees Authorization is present and lets the plan through. """ # Third-Party - from cpex.framework import HttpHeaderPayload, PluginResult, ToolPreInvokePayload + from cpex.framework import PluginResult, ToolPreInvokePayload + from cpex.framework.extensions import Extensions, HttpExtension cache = self._cache_mock( self._cache_payload( @@ -10317,9 +10318,15 @@ async def mock_invoke_hook(hook_type, payload, global_context, local_contexts=No modified = ToolPreInvokePayload( name=payload.name, args=payload.args, - headers=HttpHeaderPayload({"Authorization": "Bearer plugin-injected-token"}), ) - return PluginResult(modified_payload=modified, continue_processing=True), {} + return ( + PluginResult( + modified_payload=modified, + modified_extensions=Extensions(http=HttpExtension(headers={"Authorization": "Bearer plugin-injected-token"})), + continue_processing=True, + ), + {}, + ) mock_pm.invoke_hook = mock_invoke_hook @@ -10478,7 +10485,8 @@ def _inject(headers): async def test_prepare_rust_mcp_pre_invoke_only_returns_eligible_plan_with_hooks(self, tool_service): """Pre-invoke hooks only (no post-invoke) should produce eligible plan with hook results.""" # Third-Party - from cpex.framework import HttpHeaderPayload, ToolPreInvokePayload + from cpex.framework import ToolPreInvokePayload + from cpex.framework.extensions import Extensions, HttpExtension from cpex.framework.models import PluginResult cache = self._cache_mock(self._cache_payload(timeout_ms=2500)) @@ -10487,14 +10495,20 @@ async def test_prepare_rust_mcp_pre_invoke_only_returns_eligible_plan_with_hooks mock_pm = MagicMock() mock_pm.has_hooks_for = MagicMock(side_effect=lambda hook_type: hook_type == ToolHookType.TOOL_PRE_INVOKE) - # Mock invoke_hook to return modified args and headers + # Mock invoke_hook to return modified args and headers via extensions async def mock_invoke_hook(hook_type, payload, global_context, local_contexts=None, violations_as_exceptions=False, **_kwargs): # noqa: ARG001 modified = ToolPreInvokePayload( name=payload.name, args={"cleaned_arg": "value"}, - headers=HttpHeaderPayload({"x-injected-cred": "secret123"}), # pragma: allowlist secret ) - return PluginResult(modified_payload=modified, continue_processing=True), {} + return ( + PluginResult( + modified_payload=modified, + modified_extensions=Extensions(http=HttpExtension(headers={"x-injected-cred": "secret123"})), # pragma: allowlist secret + continue_processing=True, + ), + {}, + ) mock_pm.invoke_hook = mock_invoke_hook @@ -10595,8 +10609,9 @@ async def test_prepare_rust_mcp_pre_invoke_passes_runtime_headers_not_request_he received_headers = {} - async def mock_invoke_hook(hook_type, payload, global_context, local_contexts=None, violations_as_exceptions=False, **_kwargs): # noqa: ARG001 - received_headers.update(payload.headers.root) + async def mock_invoke_hook(hook_type, payload, global_context, local_contexts=None, violations_as_exceptions=False, extensions=None, **_kwargs): # noqa: ARG001 + if extensions and extensions.http: + received_headers.update(extensions.http.headers) return PluginResult(continue_processing=True), {} mock_pm.invoke_hook = mock_invoke_hook diff --git a/tests/unit/mcpgateway/services/test_tool_service_coverage.py b/tests/unit/mcpgateway/services/test_tool_service_coverage.py index f7435027b9..3f6e09ef82 100644 --- a/tests/unit/mcpgateway/services/test_tool_service_coverage.py +++ b/tests/unit/mcpgateway/services/test_tool_service_coverage.py @@ -6314,8 +6314,8 @@ async def test_rest_timeout_triggers_cb_and_post_hook_and_metrics_counter_failur plugin_manager.has_hooks_for = MagicMock(return_value=True) plugin_manager.invoke_hook = AsyncMock( side_effect=[ - (SimpleNamespace(modified_payload=None, metadata=None), context_table), # pre-invoke - (SimpleNamespace(modified_payload=None, retry_delay_ms=0, metadata=None), context_table), # post-invoke (timeout handler) + (SimpleNamespace(modified_payload=None, modified_extensions=None, metadata=None), context_table), # pre-invoke + (SimpleNamespace(modified_payload=None, modified_extensions=None, retry_delay_ms=0, metadata=None), context_table), # post-invoke (timeout handler) ] ) @@ -6474,7 +6474,7 @@ def _has_hooks_for(hook_type): plugin_manager.has_hooks_for = MagicMock(side_effect=_has_hooks_for) modified_payload = SimpleNamespace(name="test_tool", args={"k": "v"}, headers=None) - plugin_manager.invoke_hook = AsyncMock(return_value=(SimpleNamespace(modified_payload=modified_payload, metadata=None), {})) + plugin_manager.invoke_hook = AsyncMock(return_value=(SimpleNamespace(modified_payload=modified_payload, modified_extensions=None, metadata=None), {})) mock_response = MagicMock() mock_response.status_code = 200 @@ -7552,7 +7552,7 @@ async def fake_get(*a, **kw): plugin_manager.has_hooks_for = MagicMock(return_value=True) plugin_manager.invoke_hook = AsyncMock( side_effect=[ - (SimpleNamespace(modified_payload=None, retry_delay_ms=0, metadata=None), {}), # pre-invoke + (SimpleNamespace(modified_payload=None, modified_extensions=None, retry_delay_ms=0, metadata=None), {}), # pre-invoke (SimpleNamespace(modified_payload=SimpleNamespace(result={"status": "transformed", "valid": False}), retry_delay_ms=0, metadata=None), {}), # post-invoke ] ) @@ -7605,7 +7605,7 @@ async def fake_get(*a, **kw): plugin_manager.has_hooks_for = MagicMock(return_value=True) plugin_manager.invoke_hook = AsyncMock( side_effect=[ - (SimpleNamespace(modified_payload=None, retry_delay_ms=0, metadata=None), {}), # pre-invoke + (SimpleNamespace(modified_payload=None, modified_extensions=None, retry_delay_ms=0, metadata=None), {}), # pre-invoke (SimpleNamespace(modified_payload=SimpleNamespace(result={"unserializable", "set", "values"}), retry_delay_ms=0, metadata=None), {}), # post-invoke ] ) @@ -8073,12 +8073,24 @@ def _has_hooks_for(hook_type): return hook_type == ToolHookType.TOOL_PRE_INVOKE plugin_manager.has_hooks_for = MagicMock(side_effect=_has_hooks_for) + # Third-Party + from cpex.framework.extensions import Extensions, HttpExtension + modified_payload = SimpleNamespace( name="test_tool", args={"interaction_type": "query", "foo": "bar"}, - headers=SimpleNamespace(model_dump=lambda: {"Content-Type": "application/json", "X-Test": "1"}), + headers=None, + ) + plugin_manager.invoke_hook = AsyncMock( + return_value=( + SimpleNamespace( + modified_payload=modified_payload, + modified_extensions=Extensions(http=HttpExtension(headers={"Content-Type": "application/json", "X-Test": "1"})), + metadata=None, + ), + {}, + ) ) - plugin_manager.invoke_hook = AsyncMock(return_value=(SimpleNamespace(modified_payload=modified_payload, metadata=None), {})) captured = {} mock_http_response = MagicMock() @@ -8182,14 +8194,15 @@ async def fake_post(url, json=None, headers=None): ) assert result is not None - payload = plugin_manager.invoke_hook.await_args.kwargs["payload"] - assert payload.headers.root == { + extensions = plugin_manager.invoke_hook.await_args.kwargs.get("extensions") + hook_headers = extensions.http.headers if extensions and extensions.http else {} + assert hook_headers == { "Content-Type": "application/json", "X-Tenant-Id": "tenant-123", } - assert "Authorization" not in payload.headers.root - assert "X-Global-Id" not in payload.headers.root - assert "X-Blocked" not in payload.headers.root + assert "Authorization" not in hook_headers + assert "X-Global-Id" not in hook_headers + assert "X-Blocked" not in hook_headers assert captured_headers["Authorization"] == "Bearer client-token" assert captured_headers["X-Tenant-Id"] == "tenant-123" @@ -8255,12 +8268,13 @@ async def fake_post(url, json=None, headers=None): ) assert result is not None - payload = plugin_manager.invoke_hook.await_args.kwargs["payload"] - assert payload.headers.root == { + extensions = plugin_manager.invoke_hook.await_args.kwargs.get("extensions") + hook_headers = extensions.http.headers if extensions and extensions.http else {} + assert hook_headers == { "Content-Type": "application/json", "X-Tenant-Id": "tenant-123", } - assert "Authorization" not in payload.headers.root + assert "Authorization" not in hook_headers assert "Authorization" not in captured_headers assert captured_headers["X-Tenant-Id"] == "tenant-123" @@ -8282,18 +8296,32 @@ async def test_a2a_pre_invoke_modified_headers_are_refiltered(self, tool_service plugin_manager = MagicMock() plugin_manager.has_hooks_for = MagicMock(side_effect=lambda hook_type: hook_type == ToolHookType.TOOL_PRE_INVOKE) + # Third-Party + from cpex.framework.extensions import Extensions, HttpExtension + modified_payload = SimpleNamespace( name="test_tool", args={"interaction_type": "query"}, - headers=SimpleNamespace( - model_dump=lambda: { - "Authorization": "Bearer plugin-token", - "X-Tenant-Id": "tenant-from-plugin", - "X-Other": "drop-me", - } - ), + headers=None, + ) + plugin_manager.invoke_hook = AsyncMock( + return_value=( + SimpleNamespace( + modified_payload=modified_payload, + modified_extensions=Extensions( + http=HttpExtension( + headers={ + "Authorization": "Bearer plugin-token", + "X-Tenant-Id": "tenant-from-plugin", + "X-Other": "drop-me", + } + ) + ), + metadata=None, + ), + {}, + ) ) - plugin_manager.invoke_hook = AsyncMock(return_value=(PluginResult(modified_payload=modified_payload), {})) captured_headers = {} mock_http_response = MagicMock() @@ -8342,6 +8370,7 @@ async def test_a2a_pre_invoke_blocks_headers_when_agent_allowlist_unset(self, to """A2A tool invocation matches the direct A2A default-deny behavior when the agent allowlist is unset.""" # Third-Party from cpex.framework import PluginResult, ToolHookType + from cpex.framework.extensions import Extensions, HttpExtension tp = _make_tool_payload( integration_type="A2A", @@ -8358,15 +8387,26 @@ async def test_a2a_pre_invoke_blocks_headers_when_agent_allowlist_unset(self, to modified_payload = SimpleNamespace( name="test_tool", args={"interaction_type": "query"}, - headers=SimpleNamespace( - model_dump=lambda: { - "Authorization": "Bearer plugin-token", - "X-Tenant-Id": "tenant-from-plugin", - "X-Other": "drop-me", - } - ), + headers=None, + ) + plugin_manager.invoke_hook = AsyncMock( + return_value=( + SimpleNamespace( + modified_payload=modified_payload, + modified_extensions=Extensions( + http=HttpExtension( + headers={ + "Authorization": "Bearer plugin-token", + "X-Tenant-Id": "tenant-from-plugin", + "X-Other": "drop-me", + } + ) + ), + metadata=None, + ), + {}, + ) ) - plugin_manager.invoke_hook = AsyncMock(return_value=(PluginResult(modified_payload=modified_payload), {})) captured_headers = {} mock_http_response = MagicMock() @@ -8406,8 +8446,9 @@ async def fake_post(url, json=None, headers=None): request_headers={"X-Tenant-Id": "tenant-from-request"}, ) - payload = plugin_manager.invoke_hook.await_args.kwargs["payload"] - assert payload.headers.root == {"Content-Type": "application/json"} + extensions = plugin_manager.invoke_hook.await_args.kwargs.get("extensions") + hook_headers = extensions.http.headers if extensions and extensions.http else {} + assert hook_headers == {"Content-Type": "application/json"} assert "X-Tenant-Id" not in captured_headers assert "Authorization" not in captured_headers assert "X-Other" not in captured_headers @@ -8844,7 +8885,7 @@ def _has_hooks_for(hook_type): return hook_type == ToolHookType.TOOL_POST_INVOKE plugin_manager.has_hooks_for = MagicMock(side_effect=_has_hooks_for) - plugin_manager.invoke_hook = AsyncMock(return_value=(SimpleNamespace(modified_payload=None, retry_delay_ms=0, metadata=None), context_table)) + plugin_manager.invoke_hook = AsyncMock(return_value=(SimpleNamespace(modified_payload=None, modified_extensions=None, retry_delay_ms=0, metadata=None), context_table)) with ( _setup_cache_for_invoke(tp), @@ -9075,7 +9116,8 @@ async def test_mcp_gateway_oauth_authorization_code_plugin_injects_auth(self, to plugin-injected value. """ # Third-Party - from cpex.framework import HttpHeaderPayload, PluginResult, ToolPreInvokePayload + from cpex.framework import PluginResult, ToolPreInvokePayload + from cpex.framework.extensions import Extensions, HttpExtension tp = _make_tool_payload(integration_type="MCP", request_type="SSE", gateway_id="gw-uuid-1", jsonpath_filter="") gp = _make_gateway_payload(auth_type="oauth", oauth_config={"grant_type": "authorization_code"}) @@ -9114,10 +9156,17 @@ async def __aexit__(self, *exc): async def mock_invoke_hook(hook_type, payload, global_context, local_contexts=None, violations_as_exceptions=False, extensions=None): # noqa: ARG001 if hook_type != ToolHookType.TOOL_PRE_INVOKE: return PluginResult(modified_payload=None, continue_processing=True), {} - new_headers = dict(payload.headers.root) if payload.headers else {} + new_headers = dict(extensions.http.headers) if extensions and extensions.http else {} new_headers["Authorization"] = "Bearer plugin-injected-token" - modified = ToolPreInvokePayload(name=payload.name, args=payload.args, headers=HttpHeaderPayload(new_headers)) - return PluginResult(modified_payload=modified, continue_processing=True), {} + modified = ToolPreInvokePayload(name=payload.name, args=payload.args) + return ( + PluginResult( + modified_payload=modified, + modified_extensions=Extensions(http=HttpExtension(headers=new_headers)), + continue_processing=True, + ), + {}, + ) mock_pm.invoke_hook = mock_invoke_hook @@ -9747,7 +9796,7 @@ def _has_hooks_for(hook_type): return hook_type == ToolHookType.TOOL_POST_INVOKE plugin_manager.has_hooks_for = MagicMock(side_effect=_has_hooks_for) - plugin_manager.invoke_hook = AsyncMock(return_value=(SimpleNamespace(modified_payload=None, retry_delay_ms=0, metadata=None), context_table)) + plugin_manager.invoke_hook = AsyncMock(return_value=(SimpleNamespace(modified_payload=None, modified_extensions=None, retry_delay_ms=0, metadata=None), context_table)) def fake_sse_client(*, url=None, headers=None, httpx_client_factory=None, **_kw): class _CM: @@ -9873,7 +9922,7 @@ def _has_hooks_for(hook_type): return hook_type == ToolHookType.TOOL_PRE_INVOKE plugin_manager.has_hooks_for = MagicMock(side_effect=_has_hooks_for) - plugin_manager.invoke_hook = AsyncMock(return_value=(SimpleNamespace(modified_payload=None, metadata=None), {})) + plugin_manager.invoke_hook = AsyncMock(return_value=(SimpleNamespace(modified_payload=None, modified_extensions=None, metadata=None), {})) def fake_streamablehttp_client(*, url=None, headers=None, httpx_client_factory=None, **_kw): class _CM: @@ -9945,7 +9994,7 @@ def _has_hooks_for(hook_type): plugin_manager.has_hooks_for = MagicMock(side_effect=_has_hooks_for) modified_payload = SimpleNamespace(name="test_tool", args={}, headers=None) - plugin_manager.invoke_hook = AsyncMock(return_value=(SimpleNamespace(modified_payload=modified_payload, metadata=None), {})) + plugin_manager.invoke_hook = AsyncMock(return_value=(SimpleNamespace(modified_payload=modified_payload, modified_extensions=None, metadata=None), {})) upstream_session = AsyncMock() upstream_session.call_tool = AsyncMock(return_value=ToolResult(content=[TextContent(type="text", text="ok")], is_error=False)) @@ -10005,7 +10054,7 @@ def _has_hooks_for(hook_type): return hook_type == ToolHookType.TOOL_POST_INVOKE plugin_manager.has_hooks_for = MagicMock(side_effect=_has_hooks_for) - plugin_manager.invoke_hook = AsyncMock(return_value=(SimpleNamespace(modified_payload=None, retry_delay_ms=0, metadata=None), None)) + plugin_manager.invoke_hook = AsyncMock(return_value=(SimpleNamespace(modified_payload=None, modified_extensions=None, retry_delay_ms=0, metadata=None), None)) def fake_streamablehttp_client(*, url=None, headers=None, httpx_client_factory=None, **_kw): class _CM: diff --git a/tests/unit/mcpgateway/test_a2a_passthrough_headers.py b/tests/unit/mcpgateway/test_a2a_passthrough_headers.py index 3823cd3944..77e968d934 100644 --- a/tests/unit/mcpgateway/test_a2a_passthrough_headers.py +++ b/tests/unit/mcpgateway/test_a2a_passthrough_headers.py @@ -25,6 +25,7 @@ # Third-Party import pytest +from cpex.framework.extensions import Extensions, HttpExtension from fastapi.testclient import TestClient from sqlalchemy import create_engine from sqlalchemy.orm import Session, sessionmaker @@ -1163,16 +1164,17 @@ async def test_invoke_agent_with_plugin_header_security_warning(self, mock_setti mock_plugin_manager = MagicMock() mock_plugin_manager.has_hooks_for = MagicMock(return_value=True) - # Mock the plugin hook result + # Mock the plugin hook result — headers via modified_extensions mock_pre_result = MagicMock() mock_pre_result.modified_payload.parameters = None - mock_pre_result.modified_payload.headers = MagicMock() - mock_pre_result.modified_payload.headers.model_dump = MagicMock( - return_value={ - "x-custom": "allowed", - "authorization": "Bearer plugin", # Will be filtered # pragma: allowlist secret - "x-api-key": "secret", # Will be filtered # pragma: allowlist secret - } + mock_pre_result.modified_extensions = Extensions( + http=HttpExtension( + headers={ + "x-custom": "allowed", + "authorization": "Bearer plugin", # Will be filtered # pragma: allowlist secret + "x-api-key": "secret", # Will be filtered # pragma: allowlist secret + } + ) ) mock_plugin_manager.invoke_hook = AsyncMock(return_value=(mock_pre_result, {})) diff --git a/tests/unit/mcpgateway/test_a2a_plugin_header_security.py b/tests/unit/mcpgateway/test_a2a_plugin_header_security.py index 91c019f0d8..ba7bd48699 100644 --- a/tests/unit/mcpgateway/test_a2a_plugin_header_security.py +++ b/tests/unit/mcpgateway/test_a2a_plugin_header_security.py @@ -5,7 +5,7 @@ Tests for plugin header modification security (PR #5183 review fix). -Validates that plugin-returned headers in modified_payload.headers are +Validates that plugin-returned headers in modified_extensions.http.headers are subject to the same filtering and whitelist enforcement as inbound headers. This prevents malicious or compromised plugins from injecting sensitive diff --git a/tests/unit/plugins/test_jwt_claims_extraction.py b/tests/unit/plugins/test_jwt_claims_extraction.py index b4942690d6..a7028b9a12 100644 --- a/tests/unit/plugins/test_jwt_claims_extraction.py +++ b/tests/unit/plugins/test_jwt_claims_extraction.py @@ -19,6 +19,7 @@ PluginConfig, PluginContext, ) +from cpex.framework.extensions import Extensions, HttpExtension from cpex.framework.hooks.http import ( HttpAuthResolveUserPayload, HttpHeaderPayload, @@ -65,6 +66,11 @@ def _make_context(request_id: str = "test-123") -> PluginContext: return PluginContext(global_context=GlobalContext(request_id=request_id)) +def _empty_ext() -> Extensions: + """Empty HTTP extensions for dual-write schema calls.""" + return Extensions(http=HttpExtension(headers={})) + + class TestJwtClaimsExtractionPlugin: """Test JWT claims extraction plugin.""" @@ -77,7 +83,7 @@ async def test_extract_claims_from_credentials(self, plugin: JwtClaimsExtraction ) ctx = _make_context("test-creds") - result = await plugin.http_auth_resolve_user(payload, ctx) + result = await plugin.http_auth_resolve_user(payload, ctx, _empty_ext()) assert result.continue_processing is True claims = ctx.global_context.state["jwt_claims"] @@ -91,13 +97,16 @@ async def test_extract_claims_from_credentials(self, plugin: JwtClaimsExtraction @pytest.mark.asyncio async def test_extract_claims_from_authorization_header(self, plugin: JwtClaimsExtractionPlugin, sample_jwt_token: str) -> None: """Test extracting claims from Authorization header fallback.""" + auth_headers = {"Authorization": f"Bearer {sample_jwt_token}"} + # Dual-write: schema still requires payload.headers; real data lives on extensions. payload = HttpAuthResolveUserPayload( credentials=None, - headers=HttpHeaderPayload(root={"Authorization": f"Bearer {sample_jwt_token}"}), + headers=HttpHeaderPayload(root={}), ) ctx = _make_context("test-header") + ext = Extensions(http=HttpExtension(headers=auth_headers)) - result = await plugin.http_auth_resolve_user(payload, ctx) + result = await plugin.http_auth_resolve_user(payload, ctx, ext) assert result.continue_processing is True assert ctx.global_context.state["jwt_claims"]["sub"] == "user123" @@ -111,7 +120,7 @@ async def test_no_token_present(self, plugin: JwtClaimsExtractionPlugin) -> None ) ctx = _make_context("test-empty") - result = await plugin.http_auth_resolve_user(payload, ctx) + result = await plugin.http_auth_resolve_user(payload, ctx, _empty_ext()) assert result.continue_processing is True assert "jwt_claims" not in ctx.global_context.state @@ -139,7 +148,7 @@ async def test_extract_rfc9396_authorization_details(self, plugin: JwtClaimsExtr ) ctx = _make_context("test-rfc9396") - await plugin.http_auth_resolve_user(payload, ctx) + await plugin.http_auth_resolve_user(payload, ctx, _empty_ext()) claims = ctx.global_context.state["jwt_claims"] assert "authorization_details" in claims @@ -154,7 +163,7 @@ async def test_malformed_token_error_handling(self, plugin: JwtClaimsExtractionP ) ctx = _make_context("test-error") - result = await plugin.http_auth_resolve_user(payload, ctx) + result = await plugin.http_auth_resolve_user(payload, ctx, _empty_ext()) assert result.continue_processing is True assert result.metadata["jwt_claims_extracted"] is False @@ -169,7 +178,7 @@ async def test_ignores_non_bearer_scheme(self, plugin: JwtClaimsExtractionPlugin ) ctx = _make_context("test-basic") - result = await plugin.http_auth_resolve_user(payload, ctx) + result = await plugin.http_auth_resolve_user(payload, ctx, _empty_ext()) assert result.continue_processing is True assert "jwt_claims" not in ctx.global_context.state @@ -194,7 +203,7 @@ async def test_custom_context_key(self, sample_jwt_token: str) -> None: ) ctx = _make_context("test-custom-key") - await custom_plugin.http_auth_resolve_user(payload, ctx) + await custom_plugin.http_auth_resolve_user(payload, ctx, _empty_ext()) assert "custom_claims" in ctx.global_context.state assert ctx.global_context.state["custom_claims"]["sub"] == "user123" diff --git a/tests/unit/plugins/test_unified_pdp_plugin.py b/tests/unit/plugins/test_unified_pdp_plugin.py index af769b1958..0bdd156bf1 100644 --- a/tests/unit/plugins/test_unified_pdp_plugin.py +++ b/tests/unit/plugins/test_unified_pdp_plugin.py @@ -81,7 +81,7 @@ async def test_tool_pre_invoke_allow(self): plugin = self._plugin() plugin._pdp.check_access = AsyncMock(return_value=_allow_decision()) - payload = ToolPreInvokePayload(name="db-query", args={"sql": "SELECT 1"}, headers={}) + payload = ToolPreInvokePayload(name="db-query", args={"sql": "SELECT 1"}) result = await plugin.tool_pre_invoke(payload, _make_context()) assert result.continue_processing is not False @@ -95,7 +95,7 @@ async def test_tool_pre_invoke_deny(self): plugin = self._plugin() plugin._pdp.check_access = AsyncMock(return_value=_deny_decision()) - payload = ToolPreInvokePayload(name="db-query", args={}, headers={}) + payload = ToolPreInvokePayload(name="db-query", args={}) result = await plugin.tool_pre_invoke(payload, _make_context()) assert result.continue_processing is False @@ -110,7 +110,7 @@ async def test_tool_pre_invoke_passes_correct_action(self): plugin = self._plugin() plugin._pdp.check_access = AsyncMock(return_value=_allow_decision()) - payload = ToolPreInvokePayload(name="my-tool", args={}, headers={}) + payload = ToolPreInvokePayload(name="my-tool", args={}) await plugin.tool_pre_invoke(payload, _make_context()) call_args = plugin._pdp.check_access.call_args @@ -167,7 +167,7 @@ async def test_subject_extracted_from_dict_user(self): plugin._pdp.check_access = AsyncMock(return_value=_allow_decision()) user = {"email": "bob@x.com", "roles": ["admin"], "team_id": "ops", "mfa_verified": True} - payload = ToolPreInvokePayload(name="t", args={}, headers={}) + payload = ToolPreInvokePayload(name="t", args={}) await plugin.tool_pre_invoke(payload, _make_context(user=user)) subject = plugin._pdp.check_access.call_args[0][0] @@ -182,7 +182,7 @@ async def test_subject_extracted_from_string_user(self): plugin = self._plugin() plugin._pdp.check_access = AsyncMock(return_value=_allow_decision()) - payload = ToolPreInvokePayload(name="t", args={}, headers={}) + payload = ToolPreInvokePayload(name="t", args={}) await plugin.tool_pre_invoke(payload, _make_context(user="simple-user-id")) subject = plugin._pdp.check_access.call_args[0][0] @@ -201,7 +201,7 @@ async def test_subject_anonymous_when_user_is_none(self): ctx.global_context.server_id = "test-server" ctx.global_context.request_id = "req-123" ctx.global_context.tenant_id = "tenant-1" - payload = ToolPreInvokePayload(name="t", args={}, headers={}) + payload = ToolPreInvokePayload(name="t", args={}) await plugin.tool_pre_invoke(payload, ctx) subject = plugin._pdp.check_access.call_args[0][0] @@ -214,7 +214,7 @@ async def test_resource_type_is_tool_on_tool_hook(self): plugin = self._plugin() plugin._pdp.check_access = AsyncMock(return_value=_allow_decision()) - payload = ToolPreInvokePayload(name="my-tool", args={}, headers={}) + payload = ToolPreInvokePayload(name="my-tool", args={}) await plugin.tool_pre_invoke(payload, _make_context()) resource = plugin._pdp.check_access.call_args[0][2]