aminghadersohi commented on code in PR #44682:
URL: https://github.com/apache/superset/pull/44682#discussion_r4139206776
##########
tests/unit_tests/mcp_service/test_tool_search_transform.py:
##########
@@ -1139,3 +1142,259 @@ async def run() -> list[MagicMock]:
rendered_with = asyncio.run(run())
assert rendered_with == all_tools
+
+
[email protected]
+def bm25_transform() -> BM25SearchTransform:
+ """Build the production transform with the default search result budget."""
+ server = MagicMock()
+ _apply_tool_search_transform(
+ server,
+ {
+ "strategy": "bm25",
+ "max_results": 5,
+ "always_visible": ["health_check"],
+ },
+ )
+ return server.add_transform.call_args[0][0]
+
+
[email protected]
+def crowded_catalog() -> list[Tool]:
+ """Build more competing name mentions than fit in one search response."""
+ return [
+ Tool.from_function(
+ lambda: None,
+ name=f"peer_{index}",
+ description="Use generate_chart to create charts.",
+ )
+ for index in range(8)
+ ] + [
+ Tool.from_function(
+ lambda: None, name="generate_chart", description="Create charts."
+ )
+ ]
+
+
[email protected]("description_size", [1, 10000])
[email protected](
+ "query", ["generate_chart", " GENERATE_chart ", "generate chart"]
+)
+def test_bm25_promotes_exact_name(
+ bm25_transform: BM25SearchTransform,
+ crowded_catalog: list[Tool],
+ description_size: int,
+ query: str,
+) -> None:
+ """Exact names survive crowded descriptions and document-length
penalties."""
+ exact = crowded_catalog[-1]
+ exact.description = "Create charts with configurable visualization
options. " * (
+ description_size
+ )
+ baseline = asyncio.run(
+ BM25SearchTransform(max_results=5)._search(crowded_catalog,
"generate_chart")
+ )
+ assert exact not in baseline
+
+ results = asyncio.run(bm25_transform._search(crowded_catalog, query))
+
+ assert results[0] is exact
+ assert len(results) == 5
+ assert len({tool.name for tool in results}) == 5
+
+
[email protected]("query", ["charts", "create charts", "generate",
"no_match"])
+def test_bm25_non_exact_ranking_unchanged(
+ bm25_transform: BM25SearchTransform,
+ crowded_catalog: list[Tool],
+ query: str,
+) -> None:
+ """Non-exact searches preserve upstream BM25 ranking and result count."""
+ expected = asyncio.run(
+ BM25SearchTransform(max_results=5)._search(crowded_catalog, query)
+ )
+ assert asyncio.run(bm25_transform._search(crowded_catalog, query)) ==
expected
+
+
[email protected]("exclusion", ["permission", "catalog", "pinned"])
+def test_bm25_exact_name_respects_visibility(
+ bm25_transform: BM25SearchTransform,
+ crowded_catalog: list[Tool],
+ exclusion: str,
+) -> None:
+ """Promotion cannot recover RBAC-denied, catalog-hidden, or pinned
tools."""
+ exact = crowded_catalog[-1]
+ # Warm the same transform for a caller who could previously see this tool.
+ assert asyncio.run(bm25_transform._search(crowded_catalog, exact.name))[0]
is exact
+ catalog = crowded_catalog
+ if exclusion == "catalog":
+ catalog = crowded_catalog[:-1]
+ elif exclusion == "pinned":
+ bm25_transform._always_visible.add(exact.name)
+ else:
+ setattr(exact.fn, CLASS_PERMISSION_ATTR, "Chart")
+ setattr(exact.fn, METHOD_PERMISSION_ATTR, "write")
+
+ app = Flask(__name__)
+ app.config["MCP_RBAC_ENABLED"] = True
+ with (
+ app.app_context(),
+ patch.object(
+ bm25_transform, "get_tool_catalog", AsyncMock(return_value=catalog)
+ ),
+ patch(
+ "superset.mcp_service.auth.security_manager",
new_callable=MagicMock
+ ) as security_manager,
+ ):
+ g.user = SimpleNamespace(username="viewer")
+ security_manager.can_access.return_value = False
+ search_tool = bm25_transform._make_search_tool()
+ results = asyncio.run(search_tool.fn(query=exact.name, ctx=None))
+
+ assert exact.name not in [tool["name"] for tool in results]
+ assert len(results) == 5
+ if exclusion == "permission":
+ security_manager.can_access.assert_called_with("can_write", "Chart")
+
+
+def test_bm25_always_visible_tools_stay_pinned(
+ bm25_transform: BM25SearchTransform,
+ crowded_catalog: list[Tool],
+) -> None:
+ """Pinned tools remain listed and do not consume the search result
budget."""
+ health = Tool.from_function(lambda: None, name="health_check")
+ catalog = [health, *crowded_catalog]
+ listed = asyncio.run(bm25_transform.transform_tools(catalog))
+ assert {tool.name for tool in listed} == {
+ "health_check",
+ "search_tools",
+ "call_tool",
+ }
+ with patch.object(
+ bm25_transform, "get_tool_catalog", AsyncMock(return_value=catalog)
+ ):
+ results = asyncio.run(
+ bm25_transform._make_search_tool().fn(query="generate_chart",
ctx=None)
+ )
+ assert results[0]["name"] == "generate_chart"
+ assert len(results) == 5
+ assert "health_check" not in [tool["name"] for tool in results]
+
+
[email protected]
+def production_bm25_transform() -> BM25SearchTransform:
+ """Build the transform from the shipped search configuration."""
+ server = MagicMock()
+ _apply_tool_search_transform(server, dict(MCP_TOOL_SEARCH_CONFIG))
+ return server.add_transform.call_args[0][0]
+
+
[email protected]
+def registered_catalog() -> list[Tool]:
+ """Return every registered tool as the search catalog reads it."""
+ from superset.mcp_service.app import mcp
+
+ # search_tools reads the catalog through list_tools middleware.
+ catalog = list(asyncio.run(mcp.list_tools()))
+ assert {"generate_chart", "health_check"} <= {tool.name for tool in
catalog}
+ return catalog
+
+
+def _exact_name_search(
+ transform: BM25SearchTransform,
+ catalog: list[Tool],
+ *,
+ can_access: bool | Callable[[str, str], bool],
+ can_view_metadata: bool,
+) -> tuple[dict[str, list[str]], set[str]]:
+ """Search each tool's exact name through search_tools as one caller.
+
+ Returns the ranked names for every query and the caller's visible names.
+ """
+ app = Flask(__name__)
+ app.config["MCP_RBAC_ENABLED"] = True
+ with (
+ app.app_context(),
+ patch.object(transform, "get_tool_catalog",
AsyncMock(return_value=catalog)),
+ patch(
+ "superset.mcp_service.auth.security_manager",
new_callable=MagicMock
+ ) as security_manager,
+ patch(
+ "superset.mcp_service.privacy.user_can_view_data_model_metadata",
+ return_value=can_view_metadata,
+ ),
+ ):
+ g.user = SimpleNamespace(username="viewer")
+ if callable(can_access):
+ security_manager.can_access.side_effect = can_access
+ else:
+ security_manager.can_access.return_value = can_access
+ visible = {
+ tool.name for tool in
asyncio.run(transform._get_visible_tools(None))
+ }
+ search = transform._make_search_tool().fn
+ results = {
+ tool.name: [
+ result["name"] for result in
asyncio.run(search(query=tool.name))
+ ]
+ for tool in catalog
+ }
+ return results, visible
+
+
+def test_bm25_exact_name_finds_every_registered_tool(
+ production_bm25_transform: BM25SearchTransform,
+ registered_catalog: list[Tool],
+) -> None:
+ """Every searchable registered tool is the first result for its own name.
+
+ Long definitions such as generate_chart previously ranked below the
+ result limit for their own names.
+ """
+ results, visible = _exact_name_search(
+ production_bm25_transform,
+ registered_catalog,
+ can_access=True,
+ can_view_metadata=True,
+ )
+ pinned = set(MCP_TOOL_SEARCH_CONFIG["always_visible"])
+ assert visible == {tool.name for tool in registered_catalog} - pinned
+
+ not_first = {
Review Comment:
Addressed in f9e5e6431c691814c9f8130a3d0fe4aed6a8b6d5: added explicit types
for pinned, denied, leaked, and both not_first mappings. The tests.* mypy
override permits untyped definitions and neighboring tests rely on inferred
locals, so these annotations are not required by the checked-in configuration,
but they are a small clarity improvement. Validation: 70 tests passed;
file-scoped pre-commit hooks and a separate mypy run passed.
##########
tests/unit_tests/mcp_service/test_tool_search_transform.py:
##########
@@ -1139,3 +1142,259 @@ async def run() -> list[MagicMock]:
rendered_with = asyncio.run(run())
assert rendered_with == all_tools
+
+
[email protected]
+def bm25_transform() -> BM25SearchTransform:
+ """Build the production transform with the default search result budget."""
+ server = MagicMock()
+ _apply_tool_search_transform(
+ server,
+ {
+ "strategy": "bm25",
+ "max_results": 5,
+ "always_visible": ["health_check"],
+ },
+ )
+ return server.add_transform.call_args[0][0]
+
+
[email protected]
+def crowded_catalog() -> list[Tool]:
+ """Build more competing name mentions than fit in one search response."""
+ return [
+ Tool.from_function(
+ lambda: None,
+ name=f"peer_{index}",
+ description="Use generate_chart to create charts.",
+ )
+ for index in range(8)
+ ] + [
+ Tool.from_function(
+ lambda: None, name="generate_chart", description="Create charts."
+ )
+ ]
+
+
[email protected]("description_size", [1, 10000])
[email protected](
+ "query", ["generate_chart", " GENERATE_chart ", "generate chart"]
+)
+def test_bm25_promotes_exact_name(
+ bm25_transform: BM25SearchTransform,
+ crowded_catalog: list[Tool],
+ description_size: int,
+ query: str,
+) -> None:
+ """Exact names survive crowded descriptions and document-length
penalties."""
+ exact = crowded_catalog[-1]
+ exact.description = "Create charts with configurable visualization
options. " * (
+ description_size
+ )
+ baseline = asyncio.run(
+ BM25SearchTransform(max_results=5)._search(crowded_catalog,
"generate_chart")
+ )
+ assert exact not in baseline
+
+ results = asyncio.run(bm25_transform._search(crowded_catalog, query))
+
+ assert results[0] is exact
+ assert len(results) == 5
+ assert len({tool.name for tool in results}) == 5
+
+
[email protected]("query", ["charts", "create charts", "generate",
"no_match"])
+def test_bm25_non_exact_ranking_unchanged(
+ bm25_transform: BM25SearchTransform,
+ crowded_catalog: list[Tool],
+ query: str,
+) -> None:
+ """Non-exact searches preserve upstream BM25 ranking and result count."""
+ expected = asyncio.run(
+ BM25SearchTransform(max_results=5)._search(crowded_catalog, query)
+ )
+ assert asyncio.run(bm25_transform._search(crowded_catalog, query)) ==
expected
+
+
[email protected]("exclusion", ["permission", "catalog", "pinned"])
+def test_bm25_exact_name_respects_visibility(
+ bm25_transform: BM25SearchTransform,
+ crowded_catalog: list[Tool],
+ exclusion: str,
+) -> None:
+ """Promotion cannot recover RBAC-denied, catalog-hidden, or pinned
tools."""
+ exact = crowded_catalog[-1]
+ # Warm the same transform for a caller who could previously see this tool.
+ assert asyncio.run(bm25_transform._search(crowded_catalog, exact.name))[0]
is exact
+ catalog = crowded_catalog
+ if exclusion == "catalog":
+ catalog = crowded_catalog[:-1]
+ elif exclusion == "pinned":
+ bm25_transform._always_visible.add(exact.name)
+ else:
+ setattr(exact.fn, CLASS_PERMISSION_ATTR, "Chart")
+ setattr(exact.fn, METHOD_PERMISSION_ATTR, "write")
+
+ app = Flask(__name__)
+ app.config["MCP_RBAC_ENABLED"] = True
+ with (
+ app.app_context(),
+ patch.object(
+ bm25_transform, "get_tool_catalog", AsyncMock(return_value=catalog)
+ ),
+ patch(
+ "superset.mcp_service.auth.security_manager",
new_callable=MagicMock
+ ) as security_manager,
+ ):
+ g.user = SimpleNamespace(username="viewer")
+ security_manager.can_access.return_value = False
+ search_tool = bm25_transform._make_search_tool()
+ results = asyncio.run(search_tool.fn(query=exact.name, ctx=None))
+
+ assert exact.name not in [tool["name"] for tool in results]
+ assert len(results) == 5
+ if exclusion == "permission":
+ security_manager.can_access.assert_called_with("can_write", "Chart")
+
+
+def test_bm25_always_visible_tools_stay_pinned(
+ bm25_transform: BM25SearchTransform,
+ crowded_catalog: list[Tool],
+) -> None:
+ """Pinned tools remain listed and do not consume the search result
budget."""
+ health = Tool.from_function(lambda: None, name="health_check")
+ catalog = [health, *crowded_catalog]
+ listed = asyncio.run(bm25_transform.transform_tools(catalog))
+ assert {tool.name for tool in listed} == {
+ "health_check",
+ "search_tools",
+ "call_tool",
+ }
+ with patch.object(
+ bm25_transform, "get_tool_catalog", AsyncMock(return_value=catalog)
+ ):
+ results = asyncio.run(
+ bm25_transform._make_search_tool().fn(query="generate_chart",
ctx=None)
+ )
+ assert results[0]["name"] == "generate_chart"
+ assert len(results) == 5
+ assert "health_check" not in [tool["name"] for tool in results]
+
+
[email protected]
+def production_bm25_transform() -> BM25SearchTransform:
+ """Build the transform from the shipped search configuration."""
+ server = MagicMock()
+ _apply_tool_search_transform(server, dict(MCP_TOOL_SEARCH_CONFIG))
+ return server.add_transform.call_args[0][0]
+
+
[email protected]
+def registered_catalog() -> list[Tool]:
+ """Return every registered tool as the search catalog reads it."""
+ from superset.mcp_service.app import mcp
+
+ # search_tools reads the catalog through list_tools middleware.
+ catalog = list(asyncio.run(mcp.list_tools()))
+ assert {"generate_chart", "health_check"} <= {tool.name for tool in
catalog}
+ return catalog
+
+
+def _exact_name_search(
+ transform: BM25SearchTransform,
+ catalog: list[Tool],
+ *,
+ can_access: bool | Callable[[str, str], bool],
+ can_view_metadata: bool,
+) -> tuple[dict[str, list[str]], set[str]]:
+ """Search each tool's exact name through search_tools as one caller.
+
+ Returns the ranked names for every query and the caller's visible names.
+ """
+ app = Flask(__name__)
+ app.config["MCP_RBAC_ENABLED"] = True
+ with (
+ app.app_context(),
+ patch.object(transform, "get_tool_catalog",
AsyncMock(return_value=catalog)),
+ patch(
+ "superset.mcp_service.auth.security_manager",
new_callable=MagicMock
+ ) as security_manager,
+ patch(
+ "superset.mcp_service.privacy.user_can_view_data_model_metadata",
+ return_value=can_view_metadata,
+ ),
+ ):
+ g.user = SimpleNamespace(username="viewer")
+ if callable(can_access):
+ security_manager.can_access.side_effect = can_access
+ else:
+ security_manager.can_access.return_value = can_access
+ visible = {
+ tool.name for tool in
asyncio.run(transform._get_visible_tools(None))
+ }
+ search = transform._make_search_tool().fn
+ results = {
+ tool.name: [
+ result["name"] for result in
asyncio.run(search(query=tool.name))
+ ]
+ for tool in catalog
+ }
+ return results, visible
+
+
+def test_bm25_exact_name_finds_every_registered_tool(
+ production_bm25_transform: BM25SearchTransform,
+ registered_catalog: list[Tool],
+) -> None:
+ """Every searchable registered tool is the first result for its own name.
+
+ Long definitions such as generate_chart previously ranked below the
+ result limit for their own names.
+ """
+ results, visible = _exact_name_search(
+ production_bm25_transform,
+ registered_catalog,
+ can_access=True,
+ can_view_metadata=True,
+ )
+ pinned = set(MCP_TOOL_SEARCH_CONFIG["always_visible"])
+ assert visible == {tool.name for tool in registered_catalog} - pinned
+
+ not_first = {
+ name: ranked[:1]
+ for name, ranked in results.items()
+ if name in visible and ranked[:1] != [name]
+ }
+ assert not_first == {}
+ for name in pinned:
+ assert name not in results[name]
+
+
+def test_bm25_exact_name_never_surfaces_unauthorized_registered_tools(
+ production_bm25_transform: BM25SearchTransform,
+ registered_catalog: list[Tool],
+) -> None:
+ """Exact-name searches only return tools the caller is authorized to
see."""
+ results, visible = _exact_name_search(
+ production_bm25_transform,
+ registered_catalog,
+ # A read-only caller without data-model metadata access.
+ can_access=lambda permission, _view: permission in {"can_read",
"can_get"},
Review Comment:
Addressed in f9e5e6431c691814c9f8130a3d0fe4aed6a8b6d5: replaced the lambda
with _read_only_can_access(permission: str, _view: str) -> bool and a
docstring, preserving the permission predicate. Neighboring tests use lambdas
and the tests.* mypy override allows untyped definitions, but a named typed
helper is inexpensive and makes the callable contract explicit. Validation: 70
tests passed; file-scoped pre-commit hooks and a separate mypy run passed.
--
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.
To unsubscribe, e-mail: [email protected]
For queries about this service, please contact Infrastructure at:
[email protected]
---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]