Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 4 additions & 4 deletions mcpgateway/auth.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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(
Expand All @@ -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)

Expand Down
18 changes: 12 additions & 6 deletions mcpgateway/middleware/http_auth_middleware.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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(
Expand All @@ -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
Expand Down Expand Up @@ -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)

Expand Down
48 changes: 47 additions & 1 deletion mcpgateway/plugins/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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__)

Expand Down Expand Up @@ -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.

Expand Down
47 changes: 23 additions & 24 deletions mcpgateway/services/a2a_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -2250,7 +2250,6 @@ async def invoke_agent(
AgentHookType,
AgentPreInvokePayload,
GlobalContext,
HttpHeaderPayload,
PluginViolationError,
)

Expand Down Expand Up @@ -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
Expand Down
Loading