diff --git a/.sampo/changesets/mcp-lowlevel-tool-ownership.md b/.sampo/changesets/mcp-lowlevel-tool-ownership.md new file mode 100644 index 000000000..ec0a54a2c --- /dev/null +++ b/.sampo/changesets/mcp-lowlevel-tool-ownership.md @@ -0,0 +1,7 @@ +--- +pypi/posthog: minor +--- + +Add `resolve_original_tool` for low-level MCP servers. Fresh server instances can now remove PostHog-owned arguments before strict tool validation. + +Raw low-level servers also remove PostHog-owned arguments after `tools/list`. A tool-owned `context` remains tool data. diff --git a/posthog/mcp/_argument_ownership.py b/posthog/mcp/_argument_ownership.py new file mode 100644 index 000000000..3ade16fec --- /dev/null +++ b/posthog/mcp/_argument_ownership.py @@ -0,0 +1,98 @@ +"""Resolve ownership of arguments that PostHog adds to MCP tool schemas.""" + +from __future__ import annotations + +import inspect +from collections.abc import Mapping +from typing import Any, Dict, FrozenSet, Optional, Tuple + +from ._context_parameters import is_context_enabled, schema_has_param +from ._model_parameters import is_capture_model_enabled +from .logger import log + +_COMPLEX_SCHEMA_KEYS = ("$ref", "oneOf", "allOf", "anyOf") + + +def analytics_owned_parameters( + options: Any, + input_schema: Any, +) -> FrozenSet[str]: + """Return the enabled arguments that PostHog can add to this schema.""" + enabled = set() + if is_context_enabled(options.context): + enabled.add("context") + if options.enable_conversation_id: + enabled.add("conversation_id") + if is_capture_model_enabled(options.capture_model): + enabled.add("llm_model") + + if isinstance(input_schema, dict) and any( + input_schema.get(key) for key in _COMPLEX_SCHEMA_KEYS + ): + return frozenset() + return frozenset( + name for name in enabled if not schema_has_param(input_schema, name) + ) + + +def cache_listed_tool_ownership( + data: Any, tool: Any, *, schema_attribute: str +) -> FrozenSet[str]: + """Cache ownership from the host schema before PostHog changes it.""" + name = getattr(tool, "name", None) + if not isinstance(name, str): + return frozenset() + schema = getattr(tool, schema_attribute, None) + ownership = analytics_owned_parameters(data.options, schema) + data.tool_analytics_parameter_ownership[name] = ownership + if isinstance(schema, dict): + data.tool_input_schemas[name] = schema + else: + data.tool_input_schemas.pop(name, None) + return ownership + + +async def resolve_lowlevel_tool_ownership( + data: Any, name: str +) -> Tuple[Optional[FrozenSet[str]], Optional[Dict[str, Any]]]: + """Resolve one raw tool. A served listing on this instance has priority.""" + if name in data.tool_analytics_parameter_ownership: + return ( + data.tool_analytics_parameter_ownership[name], + data.tool_input_schemas.get(name), + ) + + resolver = data.options.resolve_original_tool + if resolver is None: + return None, None + + try: + descriptor = resolver(name) + if inspect.isawaitable(descriptor): + descriptor = await descriptor + if descriptor is None: + return None, None + schema = _descriptor_input_schema(descriptor) + if not isinstance(schema, dict): + log( + f"Warning: resolve_original_tool failed for tool {name!r}: " + "resolver returned no usable input schema" + ) + return None, None + return analytics_owned_parameters(data.options, schema), schema + except Exception as error: # noqa: BLE001 - analytics must not break dispatch + log(f"Warning: resolve_original_tool failed for tool {name!r}: {error}") + return None, None + + +def _descriptor_input_schema(descriptor: Any) -> Any: + """Read MCP 1.x, MCP 2.x, and dictionary tool descriptors.""" + if isinstance(descriptor, Mapping): + if "inputSchema" in descriptor: + return descriptor["inputSchema"] + return descriptor.get("input_schema") + + schema = getattr(descriptor, "input_schema", None) + if schema is not None: + return schema + return getattr(descriptor, "inputSchema", None) diff --git a/posthog/mcp/_context_parameters.py b/posthog/mcp/_context_parameters.py index 015663d6b..943d13764 100644 --- a/posthog/mcp/_context_parameters.py +++ b/posthog/mcp/_context_parameters.py @@ -48,10 +48,11 @@ def add_context_parameter_to_schema( """Return a new JSON Schema dict with a ``context`` string property added. Returns the input unchanged (logging a warning) for schemas that already - define ``context`` or use ``oneOf``/``allOf``/``anyOf``. ``required`` controls - whether ``context`` is added to the schema's ``required`` list — pass ``False`` - where the advertised schema is also used to validate inbound calls (the - low-level server), so a call omitting ``context`` is not rejected.""" + define ``context`` or use ``$ref``/``oneOf``/``allOf``/``anyOf``. ``required`` + controls whether ``context`` is added to the schema's ``required`` list — + pass ``False`` where the advertised schema is also used to validate inbound + calls (the low-level server), so a call omitting ``context`` is not rejected. + """ schema = input_schema if ( @@ -64,9 +65,10 @@ def add_context_parameter_to_schema( ) return schema - if schema and (schema.get("oneOf") or schema.get("allOf") or schema.get("anyOf")): + if schema and any(schema.get(key) for key in ("$ref", "oneOf", "allOf", "anyOf")): log( - f'WARN: Tool "{tool_name}" has complex schema (oneOf/allOf/anyOf). Skipping context injection.' + f'WARN: Tool "{tool_name}" has complex schema ' + "($ref/oneOf/allOf/anyOf). Skipping context injection." ) return schema diff --git a/posthog/mcp/_conversation_id.py b/posthog/mcp/_conversation_id.py index 506f878ce..d3c774516 100644 --- a/posthog/mcp/_conversation_id.py +++ b/posthog/mcp/_conversation_id.py @@ -35,7 +35,9 @@ def add_conversation_id_to_schema( input_schema: Optional[Dict[str, Any]], tool_name: str = "unknown" ) -> Optional[Dict[str, Any]]: """Return a new JSON Schema with an optional ``conversation_id`` string property. - Skips schemas that already define it or use ``oneOf``/``allOf``/``anyOf``.""" + Skips schemas that already define it or use + ``$ref``/``oneOf``/``allOf``/``anyOf``. + """ schema = input_schema if ( schema @@ -48,7 +50,7 @@ def add_conversation_id_to_schema( f"WARN: Tool \"{tool_name}\" already has '{CONVERSATION_ID_PARAM_NAME}'. Skipping injection." ) return schema - if schema and (schema.get("oneOf") or schema.get("allOf") or schema.get("anyOf")): + if schema and any(schema.get(key) for key in ("$ref", "oneOf", "allOf", "anyOf")): log( f'WARN: Tool "{tool_name}" has complex schema. Skipping conversation_id injection.' ) diff --git a/posthog/mcp/_instrument_lowlevel.py b/posthog/mcp/_instrument_lowlevel.py index f8e476283..00bd4ef4d 100644 --- a/posthog/mcp/_instrument_lowlevel.py +++ b/posthog/mcp/_instrument_lowlevel.py @@ -22,6 +22,10 @@ import mcp.types as mcp_types +from ._argument_ownership import ( + cache_listed_tool_ownership, + resolve_lowlevel_tool_ownership, +) from ._context_parameters import is_context_enabled, schema_has_param from ._conversation_id import build_prompt_back from ._event_types import MCPAnalyticsEventType @@ -29,6 +33,7 @@ advertised_tool_names, apply_virtual_tool_injection, collect_listed_tools, + copy_tools_list_result, extract_tools, is_first_listing_page, mutate_tool_schema, @@ -54,9 +59,11 @@ def instrument_low_level(server: Any, data: MCPAnalyticsData) -> None: - """Instrument a raw ``mcp.server.Server``. ``context`` is injected as an - optional schema property and NOT stripped — that schema is also the call's - validation schema, and a typical ``(name, arguments)`` handler ignores extra keys.""" + """Instrument a raw ``mcp.server.Server``. + + The adapter removes arguments that a listing or resolver proves PostHog + owns. Unknown arguments pass through unchanged. + """ data.server_name = getattr(server, "name", None) data.server_version = getattr(server, "version", None) _wrap_call_tool(server, data, strip_injected=False) @@ -186,11 +193,21 @@ def _wrap_call_tool( async def handler(req: Any) -> Any: name = req.params.name arguments = dict(req.params.arguments or {}) - strip, model_ours, input_schema = ( - await _standalone_ownership(data, high_level, name, req.params.meta) - if strip_injected - else (set(), data.tool_model_parameter_injected.get(name), None) - ) + parameter_ownership = None + if strip_injected: + strip, model_ours, input_schema = await _standalone_ownership( + data, high_level, name, req.params.meta + ) + else: + parameter_ownership, input_schema = await resolve_lowlevel_tool_ownership( + data, name + ) + strip = set(parameter_ownership or ()) + model_ours = ( + "llm_model" in parameter_ownership + if parameter_ownership is not None + else data.tool_model_parameter_injected.get(name) + ) client_name, client_version = _client_info(server) protocol_version = _protocol_version(server) mcp_session_id = _mcp_session_id(server) @@ -213,6 +230,7 @@ async def handler(req: Any) -> Any: protocol_version=protocol_version, extra={"session_id": mcp_session_id, "ctx": _request_context(server)}, input_schema=input_schema, + analytics_owned_parameters=parameter_ownership, ) if lifecycle.is_missing_capability and ( @@ -327,6 +345,7 @@ async def handler(req: Any) -> Any: def _inject_tool_schemas( + server: Any, data: MCPAnalyticsData, tools: list, *, @@ -345,6 +364,7 @@ def _inject_tool_schemas( verdicts: Dict[str, bool] = {} for tool in tools: schema = getattr(tool, "inputSchema", None) + cache_listed_tool_ownership(data, tool, schema_attribute="inputSchema") mutate_tool_schema( data, tool, @@ -363,6 +383,11 @@ def _inject_tool_schemas( # Which one dispatches is unknown, so the strip fails closed. data.tool_model_parameter_injected[tool.name] = False + cache = getattr(server, "_tool_cache", None) + if isinstance(cache, dict): + for tool in tools: + cache[tool.name] = tool + def _wrap_list_tools( server: Any, @@ -404,16 +429,18 @@ async def probe_raw_tool_names(_ctx: Any = None) -> Optional[Set[str]]: "registering your handlers." ) return None - result = await original(mcp_types.ListToolsRequest(method="tools/list")) + result = copy_tools_list_result( + await original(mcp_types.ListToolsRequest(method="tools/list")) + ) tools = extract_tools(result) - # `original` is usually the SDK's own list_tools decorator, which rebuilds - # `Server._tool_cache` from these un-injected schemas every time it runs. - # That cache is what the SDK validates real tool arguments against, so - # without re-injecting here the next real call is rejected for sending the - # `context` we advertised. Same reason the `req is None` branch below - # injects. + # Resolve ownership from the host listing. Apply injection only to the + # copy so a shared host descriptor remains original. _inject_tool_schemas( - data, tools, context_required=context_required, high_level=high_level + server, + data, + tools, + context_required=context_required, + high_level=high_level, ) return advertised_tool_names(tools) @@ -421,17 +448,17 @@ async def probe_raw_tool_names(_ctx: Any = None) -> Optional[Set[str]]: async def handler(req: Any) -> Any: # The server calls the handler with None to populate its tool cache. - # Skip analytics there — but still inject, because that cache is the - # schema the SDK validates calls against. This adapter advertises - # `context`/`conversation_id` without stripping them, so a cache built - # from un-injected schemas rejects the very arguments we told the agent - # to send ("Additional properties are not allowed") on any tool with - # `additionalProperties: false`. + # Skip analytics there, but return the same injected schema as a client + # listing. if req is None: - result = await original(req) + result = copy_tools_list_result(await original(req)) tools = extract_tools(result) _inject_tool_schemas( - data, tools, context_required=context_required, high_level=high_level + server, + data, + tools, + context_required=context_required, + high_level=high_level, ) return result @@ -464,7 +491,7 @@ async def handler(req: Any) -> Any: start = time.monotonic() try: - result = await original(req) + result = copy_tools_list_result(await original(req)) except Exception as error: await lifecycle.record_error(error, (time.monotonic() - start) * 1000) raise @@ -481,7 +508,11 @@ async def handler(req: Any) -> Any: ) _inject_tool_schemas( - data, tools, context_required=context_required, high_level=high_level + server, + data, + tools, + context_required=context_required, + high_level=high_level, ) result = apply_virtual_tool_injection( diff --git a/posthog/mcp/_instrument_v2.py b/posthog/mcp/_instrument_v2.py index 8191eedc8..843f2d032 100644 --- a/posthog/mcp/_instrument_v2.py +++ b/posthog/mcp/_instrument_v2.py @@ -37,6 +37,10 @@ import mcp.types as mcp_types +from ._argument_ownership import ( + cache_listed_tool_ownership, + resolve_lowlevel_tool_ownership, +) from ._context_parameters import is_context_enabled, schema_has_param from ._conversation_id import build_prompt_back from ._event_types import MCPAnalyticsEventType @@ -44,6 +48,7 @@ advertised_tool_names, apply_virtual_tool_injection, collect_listed_tools, + copy_tools_list_result, is_first_listing_page, mutate_tool_schema, params_to_request_dict, @@ -503,21 +508,30 @@ async def handler(ctx: Any, params: Any) -> Any: # reads the self-reported model anyway; only a listing that proved the # application owns `llm_model` stops it (posthog-js ADR-0011). analytics_owns_model = data.tool_model_parameter_injected.get(name) is not False - input_schema = None standalone = data.standalone_fastmcp() if data.standalone_fastmcp else None + parameter_ownership = None + input_schema = None if standalone is not None: version = _requested_tool_version(ctx) injected, input_schema = await _standalone_injected_parameters( standalone, data, name, version ) if injected is not None: + parameter_ownership = injected analytics_owns_model = "llm_model" in injected - call_arguments = { - key: value - for key, value in arguments.items() - if key not in injected - } - params = params.model_copy(update={"arguments": call_arguments}) + else: + parameter_ownership, input_schema = await resolve_lowlevel_tool_ownership( + data, name + ) + if parameter_ownership is not None: + analytics_owns_model = "llm_model" in parameter_ownership + if parameter_ownership is not None: + call_arguments = { + key: value + for key, value in arguments.items() + if key not in parameter_ownership + } + params = params.model_copy(update={"arguments": call_arguments}) token, client_name, client_version, protocol_version, mcp_session_id = ( _resolve_ctx(ctx) ) @@ -534,6 +548,7 @@ async def handler(ctx: Any, params: Any) -> Any: protocol_version=protocol_version, extra={"session_id": mcp_session_id, "ctx": ctx}, input_schema=input_schema, + analytics_owned_parameters=parameter_ownership, ) # No tool registry on a raw low-level server, so ownership is settled @@ -700,7 +715,7 @@ async def handler(ctx: Any, params: Any) -> Any: start = time.monotonic() try: - result = await original(ctx, params) + result = copy_tools_list_result(await original(ctx, params)) except Exception as error: await lifecycle.record_error(error, (time.monotonic() - start) * 1000) raise @@ -715,6 +730,7 @@ async def handler(ctx: Any, params: Any) -> Any: for tool in tools: schema = getattr(tool, "input_schema", None) + cache_listed_tool_ownership(data, tool, schema_attribute="input_schema") owns_context = ( _tool_owns_param_v2(high_level, tool.name, "context") if high_level is not None @@ -728,7 +744,6 @@ async def handler(ctx: Any, params: Any) -> Any: context_required=context_required, is_sdk_virtual_tool=False, ) - result = apply_virtual_tool_injection( result, injection, names, data, schema_field="input_schema" ) diff --git a/posthog/mcp/_instrumentation.py b/posthog/mcp/_instrumentation.py index 5a82fa035..167969069 100644 --- a/posthog/mcp/_instrumentation.py +++ b/posthog/mcp/_instrumentation.py @@ -10,11 +10,12 @@ import asyncio import concurrent.futures +import copy import os import threading from dataclasses import dataclass from datetime import datetime, timezone -from typing import Any, Dict, List, Literal, Optional, Set +from typing import Any, Dict, FrozenSet, List, Literal, Optional, Set from ._capture import capture_event from ._context_parameters import ( @@ -434,6 +435,7 @@ class ToolCallLifecycle: arguments: Optional[Dict[str, Any]] request_meta: Optional[Dict[str, Any]] allow_self_reported_model: bool + analytics_owned_parameters: Optional[FrozenSet[str]] request: Dict[str, Any] extra: Dict[str, Any] mcp_session_id: Optional[str] @@ -541,6 +543,7 @@ async def record_error(self, error: Any, duration_ms: float) -> None: arguments=self.arguments, request_meta=self.request_meta, allow_self_reported_model=self.allow_self_reported_model, + analytics_owned_parameters=self.analytics_owned_parameters, error=error, duration_ms=duration_ms, client_name=self.client_name, @@ -563,6 +566,7 @@ async def record_result( arguments=self.arguments, request_meta=self.request_meta, allow_self_reported_model=self.allow_self_reported_model, + analytics_owned_parameters=self.analytics_owned_parameters, result=result, duration_ms=duration_ms, client_name=self.client_name, @@ -588,6 +592,7 @@ def start_tool_call_lifecycle( protocol_version: Optional[str], extra: Dict[str, Any], input_schema: Any = None, + analytics_owned_parameters: Optional[FrozenSet[str]] = None, ) -> ToolCallLifecycle: """Resolve adapter-independent policy for a tool call without dispatching it.""" enabled = enabled_virtual_tool_names(data) @@ -597,7 +602,12 @@ def start_tool_call_lifecycle( # running the host's `on_feedback` handler read the configured options. feedback_options = resolve_collect_feedback_options(data.options.collect_feedback) conversation_id, minted = resolve_conversation_id( - data.options.enable_conversation_id, arguments + data.options.enable_conversation_id + and ( + analytics_owned_parameters is None + or "conversation_id" in analytics_owned_parameters + ), + arguments, ) # A carried session stays stable until the agent supplies its own handle. has_carried_session = token is not None or bool(mcp_session_id) @@ -609,6 +619,7 @@ def start_tool_call_lifecycle( arguments=arguments, request_meta=request_meta, allow_self_reported_model=allow_self_reported_model, + analytics_owned_parameters=analytics_owned_parameters, request=build_tool_call_request(name, arguments), extra=extra, mcp_session_id=mcp_session_id, @@ -642,6 +653,7 @@ async def record_tool_call( conversation_id: Optional[str] = None, extra: Optional[Dict[str, Any]] = None, input_schema: Any = None, + analytics_owned_parameters: Optional[FrozenSet[str]] = None, ) -> None: # Analytics must never change what the tool returns or raises: any failure # building/publishing the event is logged and swallowed here. @@ -654,7 +666,13 @@ async def record_tool_call( "tool_description": data.tool_descriptions.get(name), "tool_category": data.tool_categories.get(name), "parameters": build_captured_mcp_parameters( - request, strip_llm_model=allow_self_reported_model + request, + strip_llm_model=allow_self_reported_model, + strip_argument_names=( + set(analytics_owned_parameters) + if analytics_owned_parameters is not None + else None + ), ), "duration": duration_ms, "client_name": client_name, @@ -663,7 +681,18 @@ async def record_tool_call( "conversation_id": conversation_id, "is_error": False, } - set_event_intent(event, await resolve_tool_call_intent(data, request, extra)) + set_event_intent( + event, + await resolve_tool_call_intent( + data, + request, + extra, + allow_context_argument=( + analytics_owned_parameters is None + or "context" in analytics_owned_parameters + ), + ), + ) if is_capture_model_enabled(data.options.capture_model): model, source = resolve_model( request_meta, @@ -714,6 +743,17 @@ def extract_tools(result: Any) -> list: return list(getattr(root, "tools", []) or []) +def copy_tools_list_result(result: Any) -> Any: + """Copy a tool listing before schema injection changes its descriptors.""" + try: + return result.model_copy(deep=True) + except Exception: # noqa: BLE001 - analytics must not break a listing + try: + return copy.deepcopy(result) + except Exception: # noqa: BLE001 + return result + + def tools_list_envelope(result: Any) -> Optional[Dict[str, Any]]: """The ``tools/list`` result minus its tools, or None when nothing else is set. diff --git a/posthog/mcp/_intent.py b/posthog/mcp/_intent.py index 076a3e1d8..e8a3b3165 100644 --- a/posthog/mcp/_intent.py +++ b/posthog/mcp/_intent.py @@ -55,6 +55,8 @@ async def resolve_tool_call_intent( data: MCPAnalyticsData, request: Dict[str, Any], extra: Optional[Dict[str, Any]] = None, + *, + allow_context_argument: bool = True, ) -> Optional[ResolvedIntent]: from ._instrumentation import ( VIRTUAL_TOOL_MISSING_CAPABILITY, @@ -71,6 +73,7 @@ async def resolve_tool_call_intent( missing_name = enabled_virtual_tool_names(data).get(VIRTUAL_TOOL_MISSING_CAPABILITY) if ( is_context_enabled(data.options.context) + and allow_context_argument and (missing_name is None or name != missing_name) and context_argument ): diff --git a/posthog/mcp/_internal.py b/posthog/mcp/_internal.py index e1cdecef0..12bc83d28 100644 --- a/posthog/mcp/_internal.py +++ b/posthog/mcp/_internal.py @@ -17,7 +17,7 @@ from collections import OrderedDict from dataclasses import dataclass, field from datetime import datetime, timezone -from typing import Any, Awaitable, Callable, Dict, Optional, Set, Tuple +from typing import Any, Awaitable, Callable, Dict, FrozenSet, Optional, Set, Tuple from .logger import log from ._sink import McpEventSink @@ -91,6 +91,14 @@ class MCPAnalyticsData: # True only when PostHog added llm_model to this tool's advertised schema. # Missing/False fails closed so an application-owned field is never read or stripped. tool_model_parameter_injected: Dict[str, bool] = field(default_factory=dict) + # Ownership learned from tools/list on this instance has priority over the + # low-level resolver callback. + tool_analytics_parameter_ownership: Dict[str, FrozenSet[str]] = field( + default_factory=dict + ) + # Original schemas used to resolve ownership. They also identify safe input + # names without recording argument values. + tool_input_schemas: Dict[str, Dict[str, Any]] = field(default_factory=dict) # Which tools got `_mcp_instructions` declared on their advertised output # schema at tools/list. Only those may be mirrored into on a call — writing # an undeclared key fails the customer's whole result under diff --git a/posthog/mcp/_sanitization.py b/posthog/mcp/_sanitization.py index 0600ab69a..b14f4796d 100644 --- a/posthog/mcp/_sanitization.py +++ b/posthog/mcp/_sanitization.py @@ -11,7 +11,7 @@ from __future__ import annotations import re -from typing import Any, Dict, List, Tuple +from typing import Any, Dict, List, Optional, Set, Tuple from urllib.parse import SplitResult, parse_qsl, urlencode, urlsplit, urlunsplit # SDK-injected arguments stripped from captured $mcp_parameters (they surface as @@ -740,7 +740,10 @@ def _sanitize_resource_block(block: Dict[str, Any]) -> Any: def build_captured_mcp_parameters( - request: Any, *, strip_llm_model: bool = False + request: Any, + *, + strip_llm_model: bool = False, + strip_argument_names: Optional[Set[str]] = None, ) -> Dict[str, Any]: """Build the sanitized ``$mcp_parameters`` payload from a request, stripping the injected ``context`` argument before logging.""" @@ -752,35 +755,45 @@ def build_captured_mcp_parameters( if key in request: captured_request[key] = sanitize_captured_value(request[key]) + names_to_strip = strip_argument_names + if names_to_strip is None: + names_to_strip = set(_INJECTED_ARGUMENT_NAMES) + if strip_llm_model: + names_to_strip.add("llm_model") + if "params" in request: captured_request["params"] = _build_captured_mcp_params( - request["params"], strip_llm_model=strip_llm_model + request["params"], strip_argument_names=names_to_strip ) return {"request": captured_request} -def _build_captured_mcp_params(params: Any, *, strip_llm_model: bool) -> Any: +def _build_captured_mcp_params(params: Any, *, strip_argument_names: Set[str]) -> Any: if not _is_record(params): return sanitize_captured_value(params) captured: Dict[str, Any] = {} for key, value in params.items(): captured[key] = ( - _build_captured_mcp_arguments(value, strip_llm_model=strip_llm_model) + _build_captured_mcp_arguments( + value, strip_argument_names=strip_argument_names + ) if key == "arguments" else sanitize_captured_value(value) ) return captured -def _build_captured_mcp_arguments(arguments: Any, *, strip_llm_model: bool) -> Any: +def _build_captured_mcp_arguments( + arguments: Any, *, strip_argument_names: Set[str] +) -> Any: if not _is_record(arguments): return sanitize_captured_value(arguments) captured: Dict[str, Any] = {} for key, value in arguments.items(): - if key in _INJECTED_ARGUMENT_NAMES or (strip_llm_model and key == "llm_model"): + if key in strip_argument_names: continue captured[key] = sanitize_captured_value(value) return captured diff --git a/posthog/mcp/types.py b/posthog/mcp/types.py index f65ae0d26..7db240edb 100644 --- a/posthog/mcp/types.py +++ b/posthog/mcp/types.py @@ -242,6 +242,10 @@ class MCPAnalyticsOptions: # Return the alternative names that one tool accepts. The SDK records alias # use but does not change tool arguments. resolve_input_aliases: Optional[ResolveInputAliasesFn] = None + # Return the original tool descriptor for a low-level server. The SDK uses + # its input schema to remove only PostHog-owned arguments on a fresh server + # instance that did not serve tools/list. + resolve_original_tool: Optional[Callable[[str], Any]] = None @dataclass diff --git a/posthog/test/mcp/test_features_m4.py b/posthog/test/mcp/test_features_m4.py index f1f7bc37f..561162cad 100644 --- a/posthog/test/mcp/test_features_m4.py +++ b/posthog/test/mcp/test_features_m4.py @@ -193,7 +193,7 @@ async def test_fastmcp_keeps_input_names_after_conversation_anchoring(): assert properties["$mcp_input_keys"] == ["a", "b"] -async def test_lowlevel_does_not_reuse_a_listed_schema_for_input_names(): +async def test_lowlevel_reuses_a_listed_schema_for_input_names(): server = make_lowlevel() client = FakeClient() instrument(server, client) @@ -205,7 +205,7 @@ async def test_lowlevel_does_not_reuse_a_listed_schema_for_input_names(): await _flush() properties = _events(client, "$mcp_tool_call")[0]["properties"] - assert properties["$mcp_input_keys"] == ["[redacted]"] + assert properties["$mcp_input_keys"] == ["msg"] async def test_lowlevel_conversation_id_captured_and_prompt_back(): diff --git a/posthog/test/mcp/test_lowlevel.py b/posthog/test/mcp/test_lowlevel.py index ac84541c8..b51a5a10a 100644 --- a/posthog/test/mcp/test_lowlevel.py +++ b/posthog/test/mcp/test_lowlevel.py @@ -386,3 +386,259 @@ async def test_initialize_emitted_once(): assert len(_events(client, "$mcp_initialize")) == 1 assert len(_events(client, "$mcp_tool_call")) == 2 + + +def _make_strict_ownership_server(input_schema): + server = Server("strict-lowlevel") + seen = [] + + @server.list_tools() + async def list_tools(): + return [ + mcp_types.Tool( + name="search_docs", + inputSchema=input_schema, + ) + ] + + @server.call_tool() + async def call_tool(name, arguments): + seen.append(dict(arguments or {})) + allowed = set(input_schema.get("properties", {})) + unexpected = set(arguments or {}) - allowed + if unexpected: + raise ValueError(f"Unexpected arguments: {sorted(unexpected)}") + return [mcp_types.TextContent(type="text", text="ok")] + + return server, seen + + +async def test_listing_does_not_mutate_tool_used_by_fresh_resolver(): + schema = { + "type": "object", + "properties": {"query": {"type": "string"}}, + "additionalProperties": False, + } + tool = mcp_types.Tool(name="search_docs", inputSchema=schema) + + def make_server(): + server = Server("reused-tool") + seen = [] + + @server.list_tools() + async def list_tools(): + return [tool] + + @server.call_tool() + async def call_tool(name, arguments): + seen.append(dict(arguments or {})) + return [mcp_types.TextContent(type="text", text="ok")] + + return server, seen + + first, _ = make_server() + instrument(first, FakeClient(), MCPAnalyticsOptions(enable_conversation_id=False)) + await first.request_handlers[mcp_types.ListToolsRequest]( + mcp_types.ListToolsRequest(method="tools/list") + ) + assert tool.inputSchema == schema + + second, seen = make_server() + instrument( + second, + FakeClient(), + MCPAnalyticsOptions( + enable_conversation_id=False, + resolve_original_tool=lambda _name: tool, + ), + ) + await second.request_handlers[mcp_types.CallToolRequest]( + _call_request( + "search_docs", + {"query": "flags", "context": "find docs", "llm_model": "model-a"}, + ) + ) + + assert seen == [{"query": "flags"}] + + +async def test_lowlevel_keeps_top_level_reference_schema_unchanged(): + schema = { + "$ref": "#/$defs/Input", + "$defs": { + "Input": { + "type": "object", + "properties": { + "context": {"type": "string"}, + "conversation_id": {"type": "string"}, + "llm_model": {"type": "string"}, + }, + } + }, + } + server, _ = _make_strict_ownership_server(schema) + instrument(server, FakeClient()) + + result = await server.request_handlers[mcp_types.ListToolsRequest]( + mcp_types.ListToolsRequest(method="tools/list") + ) + + assert result.root.tools[0].inputSchema == schema + + +async def test_fresh_lowlevel_resolver_strips_posthog_arguments(): + schema = { + "type": "object", + "properties": {"query": {"type": "string"}}, + "additionalProperties": False, + } + server, seen = _make_strict_ownership_server(schema) + client = FakeClient() + instrument( + server, + client, + MCPAnalyticsOptions( + enable_conversation_id=False, + resolve_original_tool=lambda _name: {"inputSchema": schema}, + ), + ) + + result = await server.request_handlers[mcp_types.CallToolRequest]( + _call_request( + "search_docs", + {"query": "flags", "context": "find docs", "llm_model": "model-a"}, + ) + ) + await _flush() + + assert result.root.isError is False + assert seen == [{"query": "flags"}] + props = _events(client, "$mcp_tool_call")[0]["properties"] + assert props["$mcp_intent"] == "find docs" + assert props["$mcp_llm_model"] == "model-a" + assert props["$mcp_input_keys"] == ["query"] + + +async def test_fresh_lowlevel_resolver_preserves_tool_owned_context(): + schema = { + "type": "object", + "properties": { + "query": {"type": "string"}, + "context": {"type": "string"}, + }, + "additionalProperties": False, + } + server, seen = _make_strict_ownership_server(schema) + client = FakeClient() + instrument( + server, + client, + MCPAnalyticsOptions( + capture_model=False, + enable_conversation_id=False, + resolve_original_tool=lambda _name: {"input_schema": schema}, + ), + ) + + await server.request_handlers[mcp_types.CallToolRequest]( + _call_request("search_docs", {"query": "flags", "context": "tool context"}) + ) + await _flush() + + assert seen == [{"query": "flags", "context": "tool context"}] + props = _events(client, "$mcp_tool_call")[0]["properties"] + assert "$mcp_intent" not in props + captured = props["$mcp_parameters"]["request"]["params"]["arguments"] + assert captured["context"] == "tool context" + + +@pytest.mark.parametrize( + "failure", ["none", "raise", "missing_schema", "schema_property_raises"] +) +async def test_fresh_lowlevel_resolver_failure_keeps_arguments(failure): + schema = { + "type": "object", + "properties": {"query": {"type": "string"}}, + "additionalProperties": False, + } + server, seen = _make_strict_ownership_server(schema) + client = FakeClient() + messages = [] + + class BrokenDescriptor: + @property + def inputSchema(self): + raise RuntimeError("schema unavailable") + + def resolver(_name): + if failure == "raise": + raise RuntimeError("registry unavailable") + if failure == "missing_schema": + return {"name": "search_docs"} + if failure == "schema_property_raises": + return BrokenDescriptor() + return None + + instrument( + server, + client, + MCPAnalyticsOptions( + enable_conversation_id=False, + logger=messages.append, + resolve_original_tool=resolver, + ), + ) + + await server.request_handlers[mcp_types.CallToolRequest]( + _call_request( + "search_docs", + {"query": "flags", "context": "find docs", "llm_model": "model-a"}, + ) + ) + + assert seen == [{"query": "flags", "context": "find docs", "llm_model": "model-a"}] + warnings = [ + message for message in messages if "resolve_original_tool failed" in message + ] + assert len(warnings) == int(failure != "none") + + +async def test_lowlevel_served_listing_has_priority_over_resolver(): + schema = { + "type": "object", + "properties": { + "query": {"type": "string"}, + "context": {"type": "string"}, + }, + } + resolver_calls = [] + server, seen = _make_strict_ownership_server(schema) + + def resolver(name): + resolver_calls.append(name) + return { + "inputSchema": { + "type": "object", + "properties": {"query": {"type": "string"}}, + } + } + + client = FakeClient() + instrument( + server, + client, + MCPAnalyticsOptions( + capture_model=False, + enable_conversation_id=False, + resolve_original_tool=resolver, + ), + ) + await server.request_handlers[mcp_types.ListToolsRequest]( + mcp_types.ListToolsRequest(method="tools/list") + ) + await server.request_handlers[mcp_types.CallToolRequest]( + _call_request("search_docs", {"query": "flags", "context": "tool context"}) + ) + + assert resolver_calls == [] + assert seen == [{"query": "flags", "context": "tool context"}] diff --git a/posthog/test/mcp/test_types.py b/posthog/test/mcp/test_types.py index e8a3f5424..98e437aac 100644 --- a/posthog/test/mcp/test_types.py +++ b/posthog/test/mcp/test_types.py @@ -1,8 +1,23 @@ from unittest.mock import Mock +from types import SimpleNamespace import pytest from posthog.mcp.types import MCPAnalyticsOptions, PreparedToolCall, UserIdentity +from posthog.mcp._argument_ownership import _descriptor_input_schema + + +@pytest.mark.parametrize( + "descriptor", + [ + {"inputSchema": {"type": "object"}}, + {"input_schema": {"type": "object"}}, + SimpleNamespace(inputSchema={"type": "object"}), + SimpleNamespace(input_schema={"type": "object"}), + ], +) +def test_original_tool_descriptor_schema_shapes(descriptor): + assert _descriptor_input_schema(descriptor) == {"type": "object"} @pytest.mark.parametrize("capture_model", [False, True]) diff --git a/posthog/test/mcp/test_v2_lowlevel.py b/posthog/test/mcp/test_v2_lowlevel.py index d923be04d..905292b1d 100644 --- a/posthog/test/mcp/test_v2_lowlevel.py +++ b/posthog/test/mcp/test_v2_lowlevel.py @@ -761,3 +761,269 @@ async def on_list_tools(ctx, params): assert result.content[0].text == get_more_tools_result_text() assert _events(client, "$mcp_missing_capability") assert not [m for m in messages if "Cannot inject PostHog's" in m] + + +def _make_strict_ownership_server_v2(input_schema): + seen = [] + + async def on_call_tool(ctx, params): + arguments = dict(params.arguments or {}) + seen.append(arguments) + allowed = set(input_schema.get("properties", {})) + unexpected = set(arguments) - allowed + if unexpected: + raise ValueError(f"Unexpected arguments: {sorted(unexpected)}") + return mcp_types.CallToolResult( + content=[mcp_types.TextContent(type="text", text="ok")] + ) + + async def on_list_tools(ctx, params): + return mcp_types.ListToolsResult( + tools=[ + mcp_types.Tool( + name="search_docs", + input_schema=input_schema, + ) + ] + ) + + return ( + Server( + "strict-lowlevel-v2", + on_call_tool=on_call_tool, + on_list_tools=on_list_tools, + ), + seen, + ) + + +async def test_v2_fresh_lowlevel_resolver_strips_posthog_arguments(): + schema = { + "type": "object", + "properties": {"query": {"type": "string"}}, + "additionalProperties": False, + } + server, seen = _make_strict_ownership_server_v2(schema) + client = FakeClient() + instrument( + server, + client, + MCPAnalyticsOptions( + enable_conversation_id=False, + resolve_original_tool=lambda _name: {"input_schema": schema}, + ), + ) + + result = await _call_tool( + server, + "search_docs", + {"query": "flags", "context": "find docs", "llm_model": "model-a"}, + ) + await _flush() + + assert result.is_error is False + assert seen == [{"query": "flags"}] + props = _events(client, "$mcp_tool_call")[0]["properties"] + assert props["$mcp_intent"] == "find docs" + assert props["$mcp_llm_model"] == "model-a" + assert props["$mcp_input_keys"] == ["query"] + + +async def test_v2_listing_does_not_mutate_tool_used_by_fresh_resolver(): + schema = { + "type": "object", + "properties": {"query": {"type": "string"}}, + "additionalProperties": False, + } + tool = mcp_types.Tool(name="search_docs", input_schema=schema) + + def make_server(): + seen = [] + + async def on_call_tool(ctx, params): + seen.append(dict(params.arguments or {})) + return mcp_types.CallToolResult( + content=[mcp_types.TextContent(type="text", text="ok")] + ) + + async def on_list_tools(ctx, params): + return mcp_types.ListToolsResult(tools=[tool]) + + return ( + Server( + "reused-tool-v2", + on_call_tool=on_call_tool, + on_list_tools=on_list_tools, + ), + seen, + ) + + first, _ = make_server() + instrument(first, FakeClient(), MCPAnalyticsOptions(enable_conversation_id=False)) + await _list_tools(first) + assert tool.input_schema == schema + + second, seen = make_server() + instrument( + second, + FakeClient(), + MCPAnalyticsOptions( + enable_conversation_id=False, + resolve_original_tool=lambda _name: tool, + ), + ) + await _call_tool( + second, + "search_docs", + {"query": "flags", "context": "find docs", "llm_model": "model-a"}, + ) + + assert seen == [{"query": "flags"}] + + +async def test_v2_lowlevel_keeps_top_level_reference_schema_unchanged(): + schema = { + "$ref": "#/$defs/Input", + "$defs": { + "Input": { + "type": "object", + "properties": { + "context": {"type": "string"}, + "conversation_id": {"type": "string"}, + "llm_model": {"type": "string"}, + }, + } + }, + } + server, _ = _make_strict_ownership_server_v2(schema) + instrument(server, FakeClient()) + + result = await _list_tools(server) + + assert result.tools[0].input_schema == schema + + +async def test_v2_fresh_lowlevel_resolver_preserves_tool_owned_context(): + schema = { + "type": "object", + "properties": { + "query": {"type": "string"}, + "context": {"type": "string"}, + }, + "additionalProperties": False, + } + server, seen = _make_strict_ownership_server_v2(schema) + client = FakeClient() + instrument( + server, + client, + MCPAnalyticsOptions( + capture_model=False, + enable_conversation_id=False, + resolve_original_tool=lambda _name: {"inputSchema": schema}, + ), + ) + + await _call_tool( + server, + "search_docs", + {"query": "flags", "context": "tool context"}, + ) + await _flush() + + assert seen == [{"query": "flags", "context": "tool context"}] + props = _events(client, "$mcp_tool_call")[0]["properties"] + assert "$mcp_intent" not in props + captured = props["$mcp_parameters"]["request"]["params"]["arguments"] + assert captured["context"] == "tool context" + + +@pytest.mark.parametrize( + "failure", ["none", "raise", "missing_schema", "schema_property_raises"] +) +async def test_v2_fresh_lowlevel_resolver_failure_keeps_arguments(failure): + schema = { + "type": "object", + "properties": {"query": {"type": "string"}}, + "additionalProperties": False, + } + server, seen = _make_strict_ownership_server_v2(schema) + client = FakeClient() + messages = [] + + class BrokenDescriptor: + @property + def inputSchema(self): + raise RuntimeError("schema unavailable") + + def resolver(_name): + if failure == "raise": + raise RuntimeError("registry unavailable") + if failure == "missing_schema": + return {"name": "search_docs"} + if failure == "schema_property_raises": + return BrokenDescriptor() + return None + + instrument( + server, + client, + MCPAnalyticsOptions( + enable_conversation_id=False, + logger=messages.append, + resolve_original_tool=resolver, + ), + ) + + with pytest.raises(ValueError, match="Unexpected arguments"): + await _call_tool( + server, + "search_docs", + {"query": "flags", "context": "find docs", "llm_model": "model-a"}, + ) + + assert seen == [{"query": "flags", "context": "find docs", "llm_model": "model-a"}] + warnings = [ + message for message in messages if "resolve_original_tool failed" in message + ] + assert len(warnings) == int(failure != "none") + + +async def test_v2_lowlevel_served_listing_has_priority_over_resolver(): + schema = { + "type": "object", + "properties": { + "query": {"type": "string"}, + "context": {"type": "string"}, + }, + } + resolver_calls = [] + server, seen = _make_strict_ownership_server_v2(schema) + + def resolver(name): + resolver_calls.append(name) + return { + "input_schema": { + "type": "object", + "properties": {"query": {"type": "string"}}, + } + } + + instrument( + server, + FakeClient(), + MCPAnalyticsOptions( + capture_model=False, + enable_conversation_id=False, + resolve_original_tool=resolver, + ), + ) + await _list_tools(server) + await _call_tool( + server, + "search_docs", + {"query": "flags", "context": "tool context"}, + ) + + assert resolver_calls == [] + assert seen == [{"query": "flags", "context": "tool context"}] diff --git a/posthog/test/mcp/test_virtual_tools.py b/posthog/test/mcp/test_virtual_tools.py index 9c5bea29b..c02ecc7bf 100644 --- a/posthog/test/mcp/test_virtual_tools.py +++ b/posthog/test/mcp/test_virtual_tools.py @@ -290,10 +290,9 @@ async def test_renamed_tool_is_intercepted_and_the_default_name_is_not(): assert len(_events(client, "$mcp_missing_capability")) == 1 -async def test_real_tool_named_get_more_tools_keeps_normal_injection(): - # With report_missing off the SDK advertises no such tool, so one by that - # name is an ordinary application tool: it gets `context` injected and its - # value captured as $mcp_intent, like any other tool's. +async def test_real_tool_named_get_more_tools_keeps_tool_owned_context(): + # With report_missing off the SDK advertises no such virtual tool. A real + # tool with this name keeps its declared `context` argument as tool data. server = make_paged_lowlevel([[_REAL_GET_MORE_TOOLS]]) client = FakeClient() instrument(server, client, MCPAnalyticsOptions(report_missing=False, context=True)) @@ -305,7 +304,11 @@ async def test_real_tool_named_get_more_tools_keeps_normal_injection(): assert out.root.content[0].text == "real tool ran" calls = _events(client, "$mcp_tool_call") assert calls - assert calls[0]["properties"]["$mcp_intent"] == "delete a cohort" + properties = calls[0]["properties"] + assert "$mcp_intent" not in properties + assert properties["$mcp_parameters"]["request"]["params"]["arguments"] == { + "context": "delete a cohort" + } # --- a host that reuses one result object -------------------------------------- diff --git a/references/public_api_snapshot.txt b/references/public_api_snapshot.txt index 08ce2a36a..54e8c6b5e 100644 --- a/references/public_api_snapshot.txt +++ b/references/public_api_snapshot.txt @@ -853,6 +853,7 @@ attribute posthog.mcp.types.MCPAnalyticsOptions.logger: Optional[LoggerFn] = Non attribute posthog.mcp.types.MCPAnalyticsOptions.missing_capability_tool_name: Optional[str] = None attribute posthog.mcp.types.MCPAnalyticsOptions.report_missing: bool = False attribute posthog.mcp.types.MCPAnalyticsOptions.resolve_input_aliases: Optional[ResolveInputAliasesFn] = None +attribute posthog.mcp.types.MCPAnalyticsOptions.resolve_original_tool: Optional[Callable[[str], Any]] = None attribute posthog.mcp.types.MCPAnalyticsOptions.server_build: Optional[str] = None attribute posthog.mcp.types.MCPAnalyticsOptions.should_record_input_key: Optional[ShouldRecordInputKeyFn] = None attribute posthog.mcp.types.PreparedToolCall.args: Optional[JsonRecord] = None @@ -1067,7 +1068,7 @@ class posthog.mcp.types.CollectFeedbackOptions(tool_name: Optional[str] = None, class posthog.mcp.types.FeedbackReport(feedback_type: str = 'other', summary: str = '', sentiment: Optional[str] = None, friction_points: Optional[str] = None, suggested_improvement: Optional[str] = None, details: Optional[str] = None, tool_name: Optional[str] = None, task_completed: Optional[bool] = None, extras: JsonRecord = dict(), raw: JsonRecord = dict()) class posthog.mcp.types.MCPAnalyticsContextOptions(description: Optional[str] = None) class posthog.mcp.types.MCPAnalyticsModelOptions(description: Optional[str] = None) -class posthog.mcp.types.MCPAnalyticsOptions(logger: Optional[LoggerFn] = None, report_missing: bool = False, missing_capability_tool_name: Optional[str] = None, enable_conversation_id: bool = True, enable_exception_autocapture: bool = True, context: Union[bool, MCPAnalyticsContextOptions] = True, identify: Optional[Union[IdentifyFn, UserIdentity]] = None, intent_fallback: Optional[IntentFallbackFn] = None, before_send: Optional[BeforeSendFn] = None, event_properties: Optional[EventPropertiesFn] = None, capture_model: Union[bool, MCPAnalyticsModelOptions] = True, collect_feedback: Union[bool, CollectFeedbackOptions] = False, server_build: Optional[str] = None, should_record_input_key: Optional[ShouldRecordInputKeyFn] = None, resolve_input_aliases: Optional[ResolveInputAliasesFn] = None) +class posthog.mcp.types.MCPAnalyticsOptions(logger: Optional[LoggerFn] = None, report_missing: bool = False, missing_capability_tool_name: Optional[str] = None, enable_conversation_id: bool = True, enable_exception_autocapture: bool = True, context: Union[bool, MCPAnalyticsContextOptions] = True, identify: Optional[Union[IdentifyFn, UserIdentity]] = None, intent_fallback: Optional[IntentFallbackFn] = None, before_send: Optional[BeforeSendFn] = None, event_properties: Optional[EventPropertiesFn] = None, capture_model: Union[bool, MCPAnalyticsModelOptions] = True, collect_feedback: Union[bool, CollectFeedbackOptions] = False, server_build: Optional[str] = None, should_record_input_key: Optional[ShouldRecordInputKeyFn] = None, resolve_input_aliases: Optional[ResolveInputAliasesFn] = None, resolve_original_tool: Optional[Callable[[str], Any]] = None) class posthog.mcp.types.PreparedToolCall(args: Optional[JsonRecord] = None, intent: Optional[str] = None, intent_source: Optional[str] = None, is_missing_capability: bool = False, llm_model: Optional[str] = None, llm_model_source: Optional[MCPAnalyticsModelSource] = None, is_feedback: bool = False, feedback_report: Optional[FeedbackReport] = None, session_id: Optional[str] = None, conversation_id: Optional[str] = None, _conversation_state: Optional[PreparedConversationState] = None) class posthog.mcp.types.ToolInputOptions(should_record_input_key: Optional[ShouldRecordInputKeyFn] = None, input_aliases: Optional[InputAliasMap] = None) class posthog.mcp.types.UserIdentity(distinct_id: str, properties: Optional[JsonRecord] = None, groups: Optional[Dict[str, str]] = None)