diff --git a/docs/data-sources/remote-mcp-servers.md b/docs/data-sources/remote-mcp-servers.md index 82c58dbed3..c0bf297fbe 100644 --- a/docs/data-sources/remote-mcp-servers.md +++ b/docs/data-sources/remote-mcp-servers.md @@ -46,6 +46,57 @@ mcp_servers: llm_instructions: "This server provides general data access capabilities. Use it when you need to retrieve external information or perform remote operations that aren't covered by other toolsets." ``` +#### Dynamic Headers with Request Context + +MCP servers can use dynamic headers that are populated from the incoming HTTP request context. This is useful for passing authentication tokens or other request-specific headers to your MCP server. + +Use the `extra_headers` field (instead of `headers`) with template variables to reference headers from the incoming request: + +```yaml +mcp_servers: + my_server: + description: "My MCP server with dynamic authentication" + config: + url: "http://example.com:8000/mcp/messages" + mode: streamable-http + extra_headers: + X-Auth-Token: "{{ request_context.headers['X-Auth-Token'] }}" + X-User-Id: "{{ request_context.headers['X-User-Id'] }}" + llm_instructions: "Use this server to access resources with per-request authentication." +``` + +**How it works:** + +- When a request comes to HolmesGPT (via the server API), headers from that request are available in `request_context.headers` +- Header lookups are case-insensitive (e.g., `X-Auth-Token`, `x-auth-token`, and `X-AUTH-TOKEN` all work) +- The template is rendered when calling the MCP server, passing the header value through +- You can also use environment variables: `"{{ env.MY_VAR }}"` or combine them: `"Bearer {{ request_context.headers['token'] }}"` + +**Example use case:** + +This is particularly useful when your MCP server needs to authenticate with external services using tokens that are specific to each request/user. + +```yaml +mcp_servers: + remote_api_server: + description: "Remote API MCP Server" + config: + url: "http://mcp-server:8000/mcp" + mode: streamable-http + extra_headers: + X-Auth-Token: "{{ request_context.headers['X-Auth-Token'] }}" + llm_instructions: "Use this server to interact with remote APIs." +``` + +When making requests to HolmesGPT, include the required header: + +```bash +curl -X POST http://holmes-server/api/investigate \ + -H "X-Auth-Token: your-auth-token-here" \ + -H "Content-Type: application/json" \ + -d '{"question": "Check system status"}' +``` + ### URL Format The URL should point to the MCP server endpoint. The exact path depends on your server configuration: diff --git a/holmes/core/investigation.py b/holmes/core/investigation.py index f1f4fac2c5..b59294e072 100644 --- a/holmes/core/investigation.py +++ b/holmes/core/investigation.py @@ -1,5 +1,5 @@ import logging -from typing import Optional +from typing import Any, Dict, Optional from holmes.config import Config from holmes.core.investigation_structured_output import ( @@ -26,6 +26,7 @@ def investigate_issues( model: Optional[str] = None, trace_span=DummySpan(), runbooks: Optional[RunbookCatalog] = None, + request_context: Optional[Dict[str, Any]] = None, ) -> InvestigationResult: context = dal.get_issue_data(investigate_request.context.get("robusta_issue_id")) @@ -57,6 +58,7 @@ def investigate_issues( sections=investigate_request.sections, trace_span=trace_span, runbooks=runbooks, + request_context=request_context, ) (text_response, sections) = process_response_into_sections(investigation.result) diff --git a/holmes/core/tool_calling_llm.py b/holmes/core/tool_calling_llm.py index 486eed8738..f50617c47b 100644 --- a/holmes/core/tool_calling_llm.py +++ b/holmes/core/tool_calling_llm.py @@ -258,7 +258,10 @@ def reset_interaction_state(self) -> None: self._runbook_in_use = False def process_tool_decisions( - self, messages: List[Dict[str, Any]], tool_decisions: List[ToolApprovalDecision] + self, + messages: List[Dict[str, Any]], + tool_decisions: List[ToolApprovalDecision], + request_context: Optional[Dict[str, Any]] = None, ) -> tuple[List[Dict[str, Any]], list[StreamMessage]]: """ Process tool approval decisions and execute approved tools. @@ -319,6 +322,7 @@ def process_tool_decisions( tool_number=None, user_approved=True, session_approved_prefixes=session_prefixes, + request_context=request_context, ) else: # Tool was rejected or no decision found, add rejection message @@ -369,6 +373,7 @@ def prompt_call( response_format: Optional[Union[dict, Type[BaseModel]]] = None, sections: Optional[InputSectionsDataType] = None, trace_span=DummySpan(), + request_context: Optional[Dict[str, Any]] = None, ) -> LLMResult: messages = [ {"role": "system", "content": system_prompt}, @@ -380,6 +385,7 @@ def prompt_call( user_prompt=user_prompt, sections=sections, trace_span=trace_span, + request_context=request_context, ) def messages_call( @@ -387,9 +393,13 @@ def messages_call( messages: List[Dict[str, str]], response_format: Optional[Union[dict, Type[BaseModel]]] = None, trace_span=DummySpan(), + request_context: Optional[Dict[str, Any]] = None, ) -> LLMResult: return self.call( - messages, response_format=response_format, trace_span=trace_span + messages, + response_format=response_format, + trace_span=trace_span, + request_context=request_context, ) def _should_include_restricted_tools(self) -> bool: @@ -412,6 +422,7 @@ def call( # type: ignore sections: Optional[InputSectionsDataType] = None, trace_span=DummySpan(), tool_number_offset: int = 0, + request_context: Optional[Dict[str, Any]] = None, ) -> LLMResult: tool_calls: list[ dict @@ -546,6 +557,7 @@ def call( # type: ignore previous_tool_calls=tool_calls, trace_span=trace_span, tool_number=tool_number, + request_context=request_context, ) futures_tool_numbers[future] = tool_number futures.append(future) @@ -567,6 +579,7 @@ def call( # type: ignore tool_call_result=tool_call_result, tool_number=tool_number, trace_span=trace_span, + request_context=request_context, ) tool_result_response_dict = ( @@ -603,6 +616,7 @@ def _directly_invoke_tool_call( tool_call_id: str, tool_number: Optional[int] = None, session_approved_prefixes: Optional[List[str]] = None, + request_context: Optional[Dict[str, Any]] = None, ) -> StructuredToolResult: tool = self.tool_executor.get_tool_by_name(tool_name) if not tool: @@ -624,6 +638,7 @@ def _directly_invoke_tool_call( tool_name=tool_name, tool_call_id=tool_call_id, session_approved_prefixes=session_approved_prefixes or [], + request_context=request_context, ) tool_response = tool.invoke(tool_params, context=invoke_context) @@ -656,6 +671,7 @@ def _get_tool_call_result( previous_tool_calls: list[dict], tool_number: Optional[int] = None, session_approved_prefixes: Optional[List[str]] = None, + request_context: Optional[Dict[str, Any]] = None, ) -> ToolCallResult: tool_params = {} try: @@ -681,6 +697,7 @@ def _get_tool_call_result( tool_number=tool_number, tool_call_id=tool_call_id, session_approved_prefixes=session_approved_prefixes, + request_context=request_context, ) if not isinstance(tool_response, StructuredToolResult): @@ -750,6 +767,7 @@ def _invoke_llm_tool_call( tool_number=None, user_approved: bool = False, session_approved_prefixes: Optional[List[str]] = None, + request_context: Optional[Dict[str, Any]] = None, ) -> ToolCallResult: if trace_span is None: trace_span = DummySpan() @@ -784,6 +802,7 @@ def _invoke_llm_tool_call( tool_number=tool_number, user_approved=user_approved, session_approved_prefixes=session_approved_prefixes, + request_context=request_context, ) original_token_count = prevent_overly_big_tool_response( @@ -818,6 +837,7 @@ def _handle_tool_call_approval( tool_call_result: ToolCallResult, tool_number: Optional[int], trace_span: Any, + request_context: Optional[Dict[str, Any]] = None, ) -> ToolCallResult: """ Handle approval for a single tool call if required. @@ -869,6 +889,7 @@ def _handle_tool_call_approval( user_approved=True, tool_number=tool_number, tool_call_id=tool_call_result.tool_call_id, + request_context=request_context, ) tool_call_result.result = new_response else: @@ -891,6 +912,7 @@ def call_stream( msgs: Optional[list[dict]] = None, enable_tool_approval: bool = False, tool_decisions: List[ToolApprovalDecision] | None = None, + request_context: Optional[Dict[str, Any]] = None, ): """ This function DOES NOT call llm.completion(stream=true). @@ -899,7 +921,9 @@ def call_stream( # Process tool decisions if provided if msgs and tool_decisions: logging.info(f"Processing {len(tool_decisions)} tool decisions") - msgs, events = self.process_tool_decisions(msgs, tool_decisions) + msgs, events = self.process_tool_decisions( + msgs, tool_decisions, request_context + ) yield from events messages: list[dict] = [] @@ -1040,6 +1064,7 @@ def call_stream( trace_span=DummySpan(), # Streaming mode doesn't support tracing yet tool_number=tool_number, session_approved_prefixes=session_prefixes, + request_context=request_context, ) futures.append(future) yield StreamMessage( @@ -1179,6 +1204,7 @@ def investigate( sections: Optional[InputSectionsDataType] = None, trace_span=DummySpan(), runbooks: Optional[RunbookCatalog] = None, + request_context: Optional[Dict[str, Any]] = None, ) -> LLMResult: issue_runbooks = self.runbook_manager.get_instructions_for_issue(issue) @@ -1252,6 +1278,7 @@ def investigate( response_format=response_format, sections=sections, trace_span=trace_span, + request_context=request_context, ) res.instructions = issue_runbooks return res diff --git a/holmes/core/tools.py b/holmes/core/tools.py index 8b751d62e2..89d9398c78 100644 --- a/holmes/core/tools.py +++ b/holmes/core/tools.py @@ -170,6 +170,22 @@ class ToolInvokeContext(BaseModel): session_approved_prefixes: List[ str ] = [] # Bash prefixes approved during this session + request_context: Optional[Dict[str, Any]] = None + + def model_dump(self, **kwargs): + """Override to exclude sensitive context from serialization""" + data = super().model_dump(**kwargs) + if data.get("request_context"): + # Sanitize: show keys but not values + data["request_context"] = { + k: "***REDACTED***" for k in data["request_context"].keys() + } + return data + + def __str__(self): + """Override to prevent accidental context leakage in logs""" + context_keys = list((self.request_context or {}).keys()) + return f"ToolInvokeContext(tool_number={self.tool_number}, user_approved={self.user_approved}, context_keys={context_keys})" class Tool(ABC, BaseModel): diff --git a/holmes/plugins/toolsets/mcp/toolset_mcp.py b/holmes/plugins/toolsets/mcp/toolset_mcp.py index 1bd0db38e5..7d5874f2ac 100644 --- a/holmes/plugins/toolsets/mcp/toolset_mcp.py +++ b/holmes/plugins/toolsets/mcp/toolset_mcp.py @@ -1,12 +1,14 @@ import asyncio import json import logging +import os import threading from contextlib import asynccontextmanager from enum import Enum from typing import Any, ClassVar, Dict, List, Optional, Tuple, Type, Union import httpx +from jinja2 import Template from mcp.client.session import ClientSession from mcp.client.sse import sse_client from mcp.client.stdio import StdioServerParameters, stdio_client @@ -30,6 +32,17 @@ _locks_lock = threading.Lock() +class CaseInsensitiveDict(dict): + """Dictionary with case-insensitive key lookup for HTTP headers.""" + + def __getitem__(self, key): + if isinstance(key, str): + for k, v in self.items(): + if k.lower() == key.lower(): + return v + raise KeyError(key) + + def create_mcp_http_client_factory(verify_ssl: bool = True): """Create a factory function for httpx clients with configurable SSL verification.""" @@ -89,6 +102,16 @@ class MCPConfig(BaseModel): description="Whether to verify SSL certificates (set to false for local/dev servers without valid SSL).", examples=[False], ) + extra_headers: Optional[Dict[str, str]] = Field( + default=None, + description="Template headers that will be rendered with request context and environment variables.", + examples=[ + { + "X-Custom-Header": "{{ request_context.headers['X-Custom-Header'] }}", + "X-Api-Key": "{{ env.API_KEY }}", + } + ], + ) def get_lock_string(self) -> str: return str(self.url) @@ -120,7 +143,9 @@ def get_lock_string(self) -> str: @asynccontextmanager -async def get_initialized_mcp_session(toolset: "RemoteMCPToolset"): +async def get_initialized_mcp_session( + toolset: "RemoteMCPToolset", request_context: Optional[Dict[str, Any]] = None +): if toolset._mcp_config is None: raise ValueError("MCP config is not initialized") @@ -140,9 +165,10 @@ async def get_initialized_mcp_session(toolset: "RemoteMCPToolset"): elif toolset._mcp_config.mode == MCPMode.SSE: url = str(toolset._mcp_config.url) httpx_factory = create_mcp_http_client_factory(toolset._mcp_config.verify_ssl) + rendered_headers = toolset._render_headers(request_context) async with sse_client( url, - toolset._mcp_config.headers, + rendered_headers, sse_read_timeout=SSE_READ_TIMEOUT, httpx_client_factory=httpx_factory, ) as ( @@ -155,9 +181,10 @@ async def get_initialized_mcp_session(toolset: "RemoteMCPToolset"): else: url = str(toolset._mcp_config.url) httpx_factory = create_mcp_http_client_factory(toolset._mcp_config.verify_ssl) + rendered_headers = toolset._render_headers(request_context) async with streamablehttp_client( url, - headers=toolset._mcp_config.headers, + headers=rendered_headers, sse_read_timeout=SSE_READ_TIMEOUT, httpx_client_factory=httpx_factory, ) as ( @@ -182,7 +209,7 @@ def _invoke(self, params: dict, context: ToolInvokeContext) -> StructuredToolRes lock = get_server_lock(str(self.toolset._mcp_config.get_lock_string())) with lock: - return asyncio.run(self._invoke_async(params)) + return asyncio.run(self._invoke_async(params, context.request_context)) except Exception as e: return StructuredToolResult( status=StructuredToolResultStatus.ERROR, @@ -200,8 +227,12 @@ def _is_content_error(content: str) -> bool: except Exception: return False - async def _invoke_async(self, params: Dict) -> StructuredToolResult: - async with get_initialized_mcp_session(self.toolset) as session: + async def _invoke_async( + self, params: Dict, request_context: Optional[Dict[str, Any]] + ) -> StructuredToolResult: + async with get_initialized_mcp_session( + self.toolset, request_context + ) as session: tool_result = await session.call_tool(self.name, params) merged_text = " ".join(c.text for c in tool_result.content if c.type == "text") @@ -274,6 +305,97 @@ class RemoteMCPToolset(Toolset): icon_url: str = "https://registry.npmmirror.com/@lobehub/icons-static-png/1.46.0/files/light/mcp.png" _mcp_config: Optional[Union[MCPConfig, StdioMCPConfig]] = None + def _render_headers( + self, request_context: Optional[Dict[str, Any]] = None + ) -> Optional[Dict[str, str]]: + """ + Merge and render headers for MCP connection. + + Process: + 1. Start with 'headers' field (backward compatibility, passed as-is) + 2. Render 'extra_headers' templates with request_context and env vars + 3. Merge them (extra_headers takes precedence) + + Template sources for extra_headers: + - {{ request_context.headers['foo'] }}: Pass-through from client request + - {{ env.CORALOGIX_API_KEY }}: From environment variables + - "hardcoded value": Static hardcoded values + + Returns: + Merged headers dictionary or None + + Example of mcp_config: + mcp_servers: + my_mcp_server: + config: + ... + headers: + Header-Name: "hardcoded value" + extra_headers: + Header-Name-1: "hardcoded value" + Header-Name-2: "{{ request_context.headers['foo'] }}" + Header-Name-3: "{{ env.CORALOGIX_API_KEY }}" + """ + if not isinstance(self._mcp_config, MCPConfig): + return None + + # Start with direct headers (no rendering, backward compatibility) + final_headers = {} + if self._mcp_config.headers: + final_headers.update(self._mcp_config.headers) + + # Render and merge extra_headers + if self._mcp_config.extra_headers: + for header_name, header_template in self._mcp_config.extra_headers.items(): + try: + rendered_value = self._render_template( + header_template, request_context + ) + final_headers[header_name] = rendered_value + except Exception as e: # noqa: BLE001 + logging.warning( + f"MCP toolset '{self.name}': Failed to render header template " + f"'{header_name}': {e}" + ) + + return final_headers if final_headers else None + + def _render_template( + self, template_str: str, request_context: Optional[Dict[str, Any]] = None + ) -> str: + """ + Render a single template string using Jinja2. + + Supports: + - {{ request_context.headers['foo'] }} - case-insensitive header lookup + - {{ env.API_KEY }} - environment variables + - Plain strings (no template syntax) + """ + # Build context for Jinja2 template rendering + context: Dict[str, Any] = { + "env": os.environ, + } + + if request_context: + # Wrap headers in CaseInsensitiveDict for case-insensitive lookup + request_context_copy = request_context.copy() + if "headers" in request_context_copy: + request_context_copy["headers"] = CaseInsensitiveDict( + request_context_copy["headers"] + ) + context["request_context"] = request_context_copy + else: + context["request_context"] = {"headers": CaseInsensitiveDict()} + + try: + template = Template(template_str) + return template.render(context) + except Exception as e: + logging.warning( + f"MCP toolset '{self.name}': Failed to render template '{template_str}': {e}" + ) + return template_str + def model_post_init(self, __context: Any) -> None: self.prerequisites = [ CallablePrerequisite(callable=self.prerequisites_callable) @@ -353,5 +475,5 @@ def prerequisites_callable(self, config) -> Tuple[bool, str]: ) async def _get_server_tools(self): - async with get_initialized_mcp_session(self) as session: + async with get_initialized_mcp_session(self, None) as session: return await session.list_tools() diff --git a/server.py b/server.py index e42813bfdb..588e71ba40 100644 --- a/server.py +++ b/server.py @@ -10,58 +10,59 @@ # DO NOT ADD ANY IMPORTS OR CODE ABOVE THIS LINE # IMPORTING ABOVE MIGHT INITIALIZE AN HTTPS CLIENT THAT DOESN'T TRUST THE CUSTOM CERTIFICATE import json -from typing import List, Optional -from holmes.utils.global_instructions import generate_runbooks_args -from holmes.core.prompt import generate_user_prompt -import litellm -import sentry_sdk -from holmes import get_version, is_official_release - -from holmes.core import investigation -from holmes.utils.connection_utils import patch_socket_create_connection -from holmes.utils.holmes_status import update_holmes_status_in_db import logging -import uvicorn -import colorlog import threading import time +from typing import List, Optional -from litellm.exceptions import AuthenticationError +import colorlog +import litellm +import sentry_sdk +import uvicorn from fastapi import FastAPI, HTTPException, Request from fastapi.responses import StreamingResponse -from holmes.utils.stream import stream_investigate_formatter, stream_chat_formatter +from litellm.exceptions import AuthenticationError + +from holmes import get_version, is_official_release from holmes.common.env_vars import ( + DEVELOPMENT_MODE, ENABLE_CONNECTION_KEEPALIVE, + ENABLE_TELEMETRY, HOLMES_HOST, HOLMES_PORT, LOG_PERFORMANCE, SENTRY_DSN, - ENABLE_TELEMETRY, - DEVELOPMENT_MODE, SENTRY_TRACES_SAMPLE_RATE, TOOLSET_STATUS_REFRESH_INTERVAL_SECONDS, ) from holmes.config import Config +from holmes.core import investigation from holmes.core.conversations import ( build_chat_messages, build_issue_chat_messages, build_workload_health_chat_messages, ) +from holmes.core.investigation_structured_output import clear_json_markdown from holmes.core.models import ( - FollowUpAction, - InvestigationResult, - InvestigateRequest, - WorkloadHealthRequest, ChatRequest, ChatResponse, + FollowUpAction, + InvestigateRequest, + InvestigationResult, IssueChatRequest, WorkloadHealthChatRequest, + WorkloadHealthRequest, workload_health_structured_output, ) -from holmes.core.investigation_structured_output import clear_json_markdown +from holmes.core.prompt import generate_user_prompt from holmes.plugins.prompts import load_and_render_prompt +from holmes.utils.connection_utils import patch_socket_create_connection +from holmes.utils.global_instructions import generate_runbooks_args +from holmes.utils.holmes_status import update_holmes_status_in_db from holmes.utils.holmes_sync_toolsets import holmes_sync_toolsets_status from holmes.utils.log import EndpointFilter +from holmes.utils.stream import stream_chat_formatter, stream_investigate_formatter + # removed: add_runbooks_to_user_prompt @@ -194,15 +195,17 @@ async def log_requests(request: Request, call_next): @app.post("/api/investigate") -def investigate_issues(investigate_request: InvestigateRequest): +def investigate_issues(investigate_request: InvestigateRequest, http_request: Request): try: runbooks = config.get_runbook_catalog() + request_context = extract_passthrough_headers(http_request) result = investigation.investigate_issues( investigate_request=investigate_request, dal=dal, config=config, model=investigate_request.model, runbooks=runbooks, + request_context=request_context, ) return result @@ -216,11 +219,12 @@ def investigate_issues(investigate_request: InvestigateRequest): @app.post("/api/stream/investigate") -def stream_investigate_issues(req: InvestigateRequest): +def stream_investigate_issues(req: InvestigateRequest, http_request: Request): try: ai, system_prompt, user_prompt, response_format, sections, runbooks = ( investigation.get_investigation_context(req, dal, config) ) + request_context = extract_passthrough_headers(http_request) return StreamingResponse( stream_investigate_formatter( @@ -229,6 +233,7 @@ def stream_investigate_issues(req: InvestigateRequest): user_prompt=user_prompt, response_format=response_format, sections=sections, + request_context=request_context, ), runbooks, ), @@ -243,7 +248,7 @@ def stream_investigate_issues(req: InvestigateRequest): @app.post("/api/workload_health_check") -def workload_health_check(request: WorkloadHealthRequest): +def workload_health_check(request: WorkloadHealthRequest, http_request: Request): try: runbooks = config.get_runbook_catalog() resource = request.resource @@ -285,10 +290,12 @@ def workload_health_check(request: WorkloadHealthRequest): }, ) + request_context = extract_passthrough_headers(http_request) ai_call = ai.prompt_call( system_prompt, request.ask, workload_health_structured_output, + request_context=request_context, ) ai_call.result = clear_json_markdown(ai_call.result) @@ -312,6 +319,7 @@ def workload_health_check(request: WorkloadHealthRequest): @app.post("/api/workload_health_chat") def workload_health_conversation( request: WorkloadHealthChatRequest, + http_request: Request, ): try: ai = config.create_toolcalling_llm(dal=dal, model=request.model) @@ -323,7 +331,8 @@ def workload_health_conversation( config=config, global_instructions=global_instructions, ) - llm_call = ai.messages_call(messages=messages) + request_context = extract_passthrough_headers(http_request) + llm_call = ai.messages_call(messages=messages, request_context=request_context) return ChatResponse( analysis=llm_call.result, @@ -341,7 +350,7 @@ def workload_health_conversation( @app.post("/api/issue_chat") -def issue_conversation(issue_chat_request: IssueChatRequest): +def issue_conversation(issue_chat_request: IssueChatRequest, http_request: Request): try: runbooks = config.get_runbook_catalog() ai = config.create_toolcalling_llm(dal=dal, model=issue_chat_request.model) @@ -354,7 +363,8 @@ def issue_conversation(issue_chat_request: IssueChatRequest): global_instructions=global_instructions, runbooks=runbooks, ) - llm_call = ai.messages_call(messages=messages) + request_context = extract_passthrough_headers(http_request) + llm_call = ai.messages_call(messages=messages, request_context=request_context) return ChatResponse( analysis=llm_call.result, @@ -381,8 +391,36 @@ def already_answered(conversation_history: Optional[List[dict]]) -> bool: return False +def extract_passthrough_headers(request: Request) -> dict: + """ + Extract pass-through headers from the request, excluding sensitive auth headers. + These headers are forwarded to MCP servers for authentication and context. + + The blocked headers can be configured via the HOLMES_PASSTHROUGH_BLOCKED_HEADERS + environment variable (comma-separated list). Defaults to "authorization,cookie,set-cookie". + + Returns: + dict: {"headers": {"X-Foo-Bar": "...", "ABC": "...", ...}} + """ + # Get blocked headers from environment variable or use defaults + blocked_headers_str = os.environ.get( + "HOLMES_PASSTHROUGH_BLOCKED_HEADERS", "authorization,cookie,set-cookie" + ) + blocked_headers = { + h.strip().lower() for h in blocked_headers_str.split(",") if h.strip() + } + + passthrough_headers = {} + for header_name, header_value in request.headers.items(): + if header_name.lower() not in blocked_headers: + # Preserve original case from request (no normalization) + passthrough_headers[header_name] = header_value + + return {"headers": passthrough_headers} if passthrough_headers else {} + + @app.post("/api/chat") -def chat(chat_request: ChatRequest): +def chat(chat_request: ChatRequest, http_request: Request): try: # Log incoming request details has_images = bool(chat_request.images) @@ -406,6 +444,7 @@ def chat(chat_request: ChatRequest): runbooks=runbooks, images=chat_request.images, ) + request_context = extract_passthrough_headers(http_request) follow_up_actions = [] if not already_answered(chat_request.conversation_history): @@ -438,6 +477,7 @@ def chat(chat_request: ChatRequest): enable_tool_approval=chat_request.enable_tool_approval or False, tool_decisions=chat_request.tool_decisions, response_format=chat_request.response_format, + request_context=request_context, ), [f.model_dump() for f in follow_up_actions], ), @@ -447,6 +487,7 @@ def chat(chat_request: ChatRequest): llm_call = ai.messages_call( messages=messages, response_format=chat_request.response_format, + request_context=request_context, ) # For non-streaming, we need to handle approvals differently diff --git a/tests/test_approval_workflow.py b/tests/test_approval_workflow.py index ffb9a22dc2..67c3146cca 100644 --- a/tests/test_approval_workflow.py +++ b/tests/test_approval_workflow.py @@ -1,5 +1,5 @@ import json -from typing import List, Optional +from typing import Any, Dict, List, Optional from unittest.mock import MagicMock, patch import pytest @@ -139,6 +139,7 @@ def mock_invoke_tool( tool_call_id: str, tool_number: Optional[int] = None, session_approved_prefixes: Optional[List[str]] = None, + request_context: Optional[Dict[str, Any]] = None, ) -> StructuredToolResult: return StructuredToolResult( status=StructuredToolResultStatus.APPROVAL_REQUIRED, @@ -266,7 +267,7 @@ def test_streaming_chat_approval_workflow_approve_and_execute( # Mock process_tool_decisions to simulate approval and execution ai.process_tool_decisions = MagicMock( - side_effect=lambda messages, tool_decisions: ( + side_effect=lambda messages, tool_decisions, request_context=None: ( messages + [ { @@ -413,7 +414,7 @@ def test_streaming_chat_approval_workflow_reject_command( # Mock process_tool_decisions to simulate rejection ai.process_tool_decisions = MagicMock( - side_effect=lambda messages, tool_decisions: ( + side_effect=lambda messages, tool_decisions, request_context=None: ( messages + [ { diff --git a/tests/test_mcp_toolset.py b/tests/test_mcp_toolset.py index 8a3903a5f6..678cb4e4fe 100644 --- a/tests/test_mcp_toolset.py +++ b/tests/test_mcp_toolset.py @@ -1,5 +1,6 @@ import asyncio import logging +import os import shutil import subprocess from unittest.mock import AsyncMock, patch @@ -9,6 +10,7 @@ from holmes.core.tools import ( StructuredToolResultStatus, + ToolInvokeContext, ToolParameter, ) from holmes.plugins.toolsets.mcp.toolset_mcp import ( @@ -374,7 +376,7 @@ async def mock_get_server_tools(): ) with client_patch, session_patch: - result = asyncio.run(mcp_tool._invoke_async(params)) + result = asyncio.run(mcp_tool._invoke_async(params, None)) assert result.status == StructuredToolResultStatus.SUCCESS assert response_text in result.data @@ -535,7 +537,7 @@ async def mock_get_server_tools(): ) with client_patch, session_patch: - result = asyncio.run(mcp_tool._invoke_async(params)) + result = asyncio.run(mcp_tool._invoke_async(params, None)) assert result.status == StructuredToolResultStatus.SUCCESS assert response_text in result.data @@ -898,7 +900,7 @@ async def mock_get_server_tools(): ) with client_patch, session_patch: - result = asyncio.run(mcp_tool._invoke_async(params)) + result = asyncio.run(mcp_tool._invoke_async(params, None)) assert result.status == StructuredToolResultStatus.SUCCESS assert response_text in result.data @@ -1081,9 +1083,18 @@ def test_everything_stdio_tool_invocation(self, suppress_migration_warnings): if greet_tool is None: pytest.skip("greet tool not found in MCP server") - # Actually invoke the tool on the real server with timeout + context = ToolInvokeContext.model_construct( + tool_number=1, + user_approved=True, + llm=None, + max_token_count=1000, + tool_call_id="test-id", + tool_name="greet", + request_context=None, + ) + try: - invoke_result = greet_tool._invoke({"name": "Alice"}, None) + invoke_result = greet_tool._invoke({"name": "Alice"}, context) except Exception as e: pytest.fail(f"Tool invocation failed: {e}") @@ -1129,3 +1140,372 @@ async def run_test(): # Verify the tools loaded in the toolset match what we got from list_tools assert len(toolset.tools) == len(list_result.tools) + + +class TestHeaderRendering: + def test_render_headers_with_static_headers_only(self): + mcp_toolset = RemoteMCPToolset( + name="test_mcp", + description="Test toolset", + config={ + "url": "http://localhost:1234", + "headers": {"Authorization": "Bearer token123"}, + }, + ) + + mcp_toolset.prerequisites_callable(config=mcp_toolset.config) + rendered = mcp_toolset._render_headers(None) + + assert rendered is not None + assert rendered["Authorization"] == "Bearer token123" + + def test_render_headers_with_extra_headers_static(self): + mcp_toolset = RemoteMCPToolset( + name="test_mcp", + description="Test toolset", + config={ + "url": "http://localhost:1234", + "extra_headers": {"X-Custom": "static-value"}, + }, + ) + + mcp_toolset.prerequisites_callable(config=mcp_toolset.config) + rendered = mcp_toolset._render_headers(None) + + assert rendered is not None + assert rendered["X-Custom"] == "static-value" + + def test_render_headers_with_request_context_template(self): + mcp_toolset = RemoteMCPToolset( + name="test_mcp", + description="Test toolset", + config={ + "url": "http://localhost:1234", + "extra_headers": { + "X-Tenant-Id": "{{ request_context.headers['X-Tenant-Id'] }}" + }, + }, + ) + + mcp_toolset.prerequisites_callable(config=mcp_toolset.config) + request_context = {"headers": {"X-Tenant-Id": "tenant-123"}} + rendered = mcp_toolset._render_headers(request_context) + + assert rendered is not None + assert rendered["X-Tenant-Id"] == "tenant-123" + + def test_render_headers_with_request_context_case_insensitive(self): + mcp_toolset = RemoteMCPToolset( + name="test_mcp", + description="Test toolset", + config={ + "url": "http://localhost:1234", + "extra_headers": { + "X-Tenant-Id": "{{ request_context.headers['x-tenant-id'] }}" + }, + }, + ) + + mcp_toolset.prerequisites_callable(config=mcp_toolset.config) + request_context = {"headers": {"X-Tenant-Id": "tenant-456"}} + rendered = mcp_toolset._render_headers(request_context) + + assert rendered is not None + assert rendered["X-Tenant-Id"] == "tenant-456" + + def test_render_headers_with_env_var_template(self, monkeypatch): + monkeypatch.setenv("TEST_API_KEY", "secret-key-789") + + mcp_toolset = RemoteMCPToolset( + name="test_mcp", + description="Test toolset", + config={ + "url": "http://localhost:1234", + "extra_headers": {"X-Api-Key": "{{ env.TEST_API_KEY }}"}, + }, + ) + + mcp_toolset.prerequisites_callable(config=mcp_toolset.config) + rendered = mcp_toolset._render_headers(None) + + assert rendered is not None + assert rendered["X-Api-Key"] == "secret-key-789" + + def test_render_headers_merge_static_and_extra(self): + mcp_toolset = RemoteMCPToolset( + name="test_mcp", + description="Test toolset", + config={ + "url": "http://localhost:1234", + "headers": {"Authorization": "Bearer static"}, + "extra_headers": {"X-Custom": "dynamic"}, + }, + ) + + mcp_toolset.prerequisites_callable(config=mcp_toolset.config) + rendered = mcp_toolset._render_headers(None) + + assert rendered is not None + assert rendered["Authorization"] == "Bearer static" + assert rendered["X-Custom"] == "dynamic" + + def test_render_headers_extra_overrides_static(self): + mcp_toolset = RemoteMCPToolset( + name="test_mcp", + description="Test toolset", + config={ + "url": "http://localhost:1234", + "headers": {"X-Header": "old-value"}, + "extra_headers": {"X-Header": "new-value"}, + }, + ) + + mcp_toolset.prerequisites_callable(config=mcp_toolset.config) + rendered = mcp_toolset._render_headers(None) + + assert rendered is not None + assert rendered["X-Header"] == "new-value" + + def test_render_headers_with_missing_request_context_header(self): + mcp_toolset = RemoteMCPToolset( + name="test_mcp", + description="Test toolset", + config={ + "url": "http://localhost:1234", + "extra_headers": { + "X-Missing": "{{ request_context.headers['X-Missing'] }}" + }, + }, + ) + + mcp_toolset.prerequisites_callable(config=mcp_toolset.config) + request_context = {"headers": {"X-Other": "value"}} + rendered = mcp_toolset._render_headers(request_context) + + assert rendered is not None + # Jinja2 default behavior is to render undefined variables as empty strings + assert rendered["X-Missing"] == "" + + def test_render_headers_with_missing_env_var(self): + mcp_toolset = RemoteMCPToolset( + name="test_mcp", + description="Test toolset", + config={ + "url": "http://localhost:1234", + "extra_headers": {"X-Key": "{{ env.NONEXISTENT_VAR }}"}, + }, + ) + + mcp_toolset.prerequisites_callable(config=mcp_toolset.config) + rendered = mcp_toolset._render_headers(None) + + assert rendered is not None + # Jinja2 default behavior is to render undefined variables as empty strings + assert rendered["X-Key"] == "" + + def test_render_headers_mixed_templates(self, monkeypatch): + monkeypatch.setenv("API_KEY", "env-secret") + + mcp_toolset = RemoteMCPToolset( + name="test_mcp", + description="Test toolset", + config={ + "url": "http://localhost:1234", + "extra_headers": { + "Authorization": "Bearer {{ env.API_KEY }}", + "X-Tenant": "{{ request_context.headers['X-Tenant'] }}", + "X-Static": "static-value", + }, + }, + ) + + mcp_toolset.prerequisites_callable(config=mcp_toolset.config) + request_context = {"headers": {"X-Tenant": "tenant-999"}} + rendered = mcp_toolset._render_headers(request_context) + + assert rendered is not None + assert rendered["Authorization"] == "Bearer env-secret" + assert rendered["X-Tenant"] == "tenant-999" + assert rendered["X-Static"] == "static-value" + + def test_render_headers_stdio_config_returns_none(self): + mcp_toolset = RemoteMCPToolset( + name="test_mcp", + description="Test toolset", + config={ + "mode": "stdio", + "command": "python", + "args": ["server.py"], + }, + ) + + mcp_toolset._mcp_config = StdioMCPConfig( + mode=MCPMode.STDIO, command="python", args=["server.py"] + ) + rendered = mcp_toolset._render_headers(None) + + assert rendered is None + + +class TestRequestContextPassthrough: + def test_tool_invoke_context_sanitizes_request_context(self): + context = ToolInvokeContext.model_construct( + tool_number=1, + user_approved=True, + llm=None, + max_token_count=1000, + tool_call_id="test-id", + tool_name="test-tool", + request_context={"headers": {"Authorization": "Bearer secret"}}, + ) + + dumped = context.model_dump() + assert "request_context" in dumped + assert dumped["request_context"]["headers"] == "***REDACTED***" + + def test_tool_invoke_context_str_hides_values(self): + context = ToolInvokeContext.model_construct( + tool_number=1, + user_approved=True, + llm=None, + max_token_count=1000, + tool_call_id="test-id", + tool_name="test-tool", + request_context={"headers": {"X-Tenant": "secret-tenant"}}, + ) + + str_repr = str(context) + assert "secret-tenant" not in str_repr + assert "context_keys=['headers']" in str_repr + + def test_get_initialized_mcp_session_passes_request_context( + self, monkeypatch + ): + mcp_toolset = RemoteMCPToolset( + name="test_mcp", + description="Test toolset", + config={ + "url": "http://localhost:1234", + "extra_headers": { + "X-Tenant": "{{ request_context.headers['X-Tenant'] }}" + }, + }, + ) + + async def mock_get_server_tools(): + return ListToolsResult(tools=[]) + + monkeypatch.setattr(mcp_toolset, "_get_server_tools", mock_get_server_tools) + mcp_toolset.prerequisites_callable(config=mcp_toolset.config) + + mock_read_stream = AsyncMock() + mock_write_stream = AsyncMock() + mock_session = AsyncMock() + mock_session.initialize = AsyncMock(return_value=None) + + mock_client_context = AsyncMock() + mock_client_context.__aenter__ = AsyncMock( + return_value=(mock_read_stream, mock_write_stream) + ) + mock_client_context.__aexit__ = AsyncMock(return_value=None) + + mock_session_context = AsyncMock() + mock_session_context.__aenter__ = AsyncMock(return_value=mock_session) + mock_session_context.__aexit__ = AsyncMock(return_value=None) + + captured_headers = None + + def capture_sse_client_call(_url, headers, *, sse_read_timeout, httpx_client_factory=None): + nonlocal captured_headers + captured_headers = headers + return mock_client_context + + with patch( + "holmes.plugins.toolsets.mcp.toolset_mcp.sse_client", + side_effect=capture_sse_client_call, + ): + with patch( + "holmes.plugins.toolsets.mcp.toolset_mcp.ClientSession", + return_value=mock_session_context, + ): + request_context = {"headers": {"X-Tenant": "tenant-abc"}} + + async def run_test(): + async with get_initialized_mcp_session( + mcp_toolset, request_context + ) as _: + pass + + asyncio.run(run_test()) + + assert captured_headers is not None + assert captured_headers["X-Tenant"] == "tenant-abc" + + def test_tool_invoke_async_passes_request_context(self, monkeypatch): + tool = Tool( + name="test_tool", + inputSchema={"type": "object", "properties": {}, "required": []}, + description="Test tool", + ) + + mock_toolset = RemoteMCPToolset( + name="test_toolset", + description="Test toolset", + config={ + "url": "http://localhost:1234", + "extra_headers": { + "X-Context": "{{ request_context.headers['X-Context'] }}" + }, + }, + ) + + async def mock_get_server_tools(): + return ListToolsResult(tools=[]) + + monkeypatch.setattr(mock_toolset, "_get_server_tools", mock_get_server_tools) + mock_toolset.prerequisites_callable(config=mock_toolset.config) + + mcp_tool = RemoteMCPTool.create(tool, mock_toolset) + + mock_session = AsyncMock() + mock_session.initialize = AsyncMock(return_value=None) + mock_session.call_tool = AsyncMock( + return_value=CallToolResult( + content=[TextContent(type="text", text="success")], isError=False + ) + ) + + mock_read_stream = AsyncMock() + mock_write_stream = AsyncMock() + + mock_client_context = AsyncMock() + mock_client_context.__aenter__ = AsyncMock( + return_value=(mock_read_stream, mock_write_stream) + ) + mock_client_context.__aexit__ = AsyncMock(return_value=None) + + mock_session_context = AsyncMock() + mock_session_context.__aenter__ = AsyncMock(return_value=mock_session) + mock_session_context.__aexit__ = AsyncMock(return_value=None) + + captured_headers = None + + def capture_sse_client_call(_url, headers, *, sse_read_timeout, httpx_client_factory=None): + nonlocal captured_headers + captured_headers = headers + return mock_client_context + + with patch( + "holmes.plugins.toolsets.mcp.toolset_mcp.sse_client", + side_effect=capture_sse_client_call, + ): + with patch( + "holmes.plugins.toolsets.mcp.toolset_mcp.ClientSession", + return_value=mock_session_context, + ): + request_context = {"headers": {"X-Context": "ctx-value"}} + result = asyncio.run(mcp_tool._invoke_async({}, request_context)) + + assert result.status == StructuredToolResultStatus.SUCCESS + assert captured_headers is not None + assert captured_headers["X-Context"] == "ctx-value" diff --git a/tests/test_server_endpoints.py b/tests/test_server_endpoints.py index 1c1ba7f0c1..cc7018ecb3 100644 --- a/tests/test_server_endpoints.py +++ b/tests/test_server_endpoints.py @@ -1,9 +1,10 @@ from unittest.mock import MagicMock, patch import pytest +from fastapi import Request from fastapi.testclient import TestClient -from server import app +from server import app, extract_passthrough_headers @pytest.fixture @@ -404,3 +405,156 @@ def test_api_workload_health_check( assert "tool_name" in tool_call assert "description" in tool_call assert "result" in tool_call + + +class TestExtractPassthroughHeaders: + def test_extract_normal_headers(self): + scope = { + "type": "http", + "headers": [ + (b"x-tenant-id", b"tenant-123"), + (b"x-custom-header", b"custom-value"), + (b"content-type", b"application/json"), + ], + } + request = Request(scope) + result = extract_passthrough_headers(request) + + assert result == { + "headers": { + "x-tenant-id": "tenant-123", + "x-custom-header": "custom-value", + "content-type": "application/json", + } + } + + def test_blocks_authorization_header(self): + scope = { + "type": "http", + "headers": [ + (b"authorization", b"Bearer secret-token"), + (b"x-tenant-id", b"tenant-123"), + ], + } + request = Request(scope) + result = extract_passthrough_headers(request) + + assert result == {"headers": {"x-tenant-id": "tenant-123"}} + assert "authorization" not in result["headers"] + + def test_blocks_cookie_headers(self): + scope = { + "type": "http", + "headers": [ + (b"cookie", b"session=abc123"), + (b"set-cookie", b"session=abc123; Path=/"), + (b"x-tenant-id", b"tenant-123"), + ], + } + request = Request(scope) + result = extract_passthrough_headers(request) + + assert result == {"headers": {"x-tenant-id": "tenant-123"}} + assert "cookie" not in result["headers"] + assert "set-cookie" not in result["headers"] + + def test_case_insensitive_blocking(self): + scope = { + "type": "http", + "headers": [ + (b"Authorization", b"Bearer secret"), + (b"COOKIE", b"session=abc"), + (b"Set-Cookie", b"session=abc; Path=/"), + (b"x-tenant-id", b"tenant-123"), + ], + } + request = Request(scope) + result = extract_passthrough_headers(request) + + assert result == {"headers": {"x-tenant-id": "tenant-123"}} + assert "Authorization" not in result["headers"] + assert "COOKIE" not in result["headers"] + assert "Set-Cookie" not in result["headers"] + + def test_empty_headers(self): + scope = {"type": "http", "headers": []} + request = Request(scope) + result = extract_passthrough_headers(request) + + assert result == {} + + def test_all_blocked_headers(self): + scope = { + "type": "http", + "headers": [ + (b"authorization", b"Bearer secret"), + (b"cookie", b"session=abc"), + (b"set-cookie", b"session=abc; Path=/"), + ], + } + request = Request(scope) + result = extract_passthrough_headers(request) + + assert result == {} + + def test_preserves_header_case(self): + scope = { + "type": "http", + "headers": [ + (b"X-Tenant-ID", b"tenant-123"), + (b"X-Custom-Header", b"value"), + ], + } + request = Request(scope) + result = extract_passthrough_headers(request) + + assert "X-Tenant-ID" in result["headers"] + assert "X-Custom-Header" in result["headers"] + assert result["headers"]["X-Tenant-ID"] == "tenant-123" + assert result["headers"]["X-Custom-Header"] == "value" + + def test_custom_blocked_headers_via_env(self, monkeypatch): + """Test that HOLMES_PASSTHROUGH_BLOCKED_HEADERS env var works""" + # Set custom blocked headers via environment variable + monkeypatch.setenv("HOLMES_PASSTHROUGH_BLOCKED_HEADERS", "x-internal-token,x-secret") + + scope = { + "type": "http", + "headers": [ + (b"x-internal-token", b"secret-value"), + (b"x-secret", b"another-secret"), + (b"authorization", b"Bearer token"), # Not in custom list, should pass + (b"x-tenant-id", b"tenant-123"), + ], + } + request = Request(scope) + result = extract_passthrough_headers(request) + + # Custom blocked headers should be filtered + assert "x-internal-token" not in result["headers"] + assert "x-secret" not in result["headers"] + # Authorization is not in custom list, so it should pass through + assert "authorization" in result["headers"] + assert result["headers"]["authorization"] == "Bearer token" + # Regular headers should pass + assert result["headers"]["x-tenant-id"] == "tenant-123" + + def test_empty_blocked_headers_env(self, monkeypatch): + """Test that empty HOLMES_PASSTHROUGH_BLOCKED_HEADERS allows all headers""" + monkeypatch.setenv("HOLMES_PASSTHROUGH_BLOCKED_HEADERS", "") + + scope = { + "type": "http", + "headers": [ + (b"authorization", b"Bearer token"), + (b"cookie", b"session=abc"), + (b"x-tenant-id", b"tenant-123"), + ], + } + request = Request(scope) + result = extract_passthrough_headers(request) + + # With empty blocklist, all headers should pass through + assert "authorization" in result["headers"] + assert "cookie" in result["headers"] + assert "x-tenant-id" in result["headers"]