gabotorresruiz commented on code in PR #44581: URL: https://github.com/apache/superset/pull/44581#discussion_r4150385770
########## superset/mcp_service/worker.py: ########## @@ -0,0 +1,700 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +"""Bounded, deadline-aware execution of MCP tools outside the transport loop. + +The tool's coroutine runs on a worker-owned loop because its database APIs are +synchronous. Only transport notifications are marshalled back to the server +loop. Flask and SQLAlchemy lifetimes belong to the worker, not the waiting +request: a timed-out DBAPI call cannot outlive and reuse a torn-down session. +""" + +from __future__ import annotations + +import asyncio +import functools +import logging +import threading +import time +import uuid +from concurrent.futures import Future, ThreadPoolExecutor +from contextlib import contextmanager +from contextvars import ContextVar, copy_context +from typing import Any, Callable, Coroutine, Iterator, ParamSpec, TYPE_CHECKING, TypeVar +from weakref import WeakKeyDictionary + +from fastmcp.exceptions import ToolError +from flask import current_app, g, has_app_context, has_request_context +from sqlalchemy import inspect as sa_inspect +from sqlalchemy.orm.state import InstanceState +from sqlalchemy.pool import QueuePool + +from superset.mcp_service.session_scope import _mcp_session_token + +if TYPE_CHECKING: + from flask import Flask + + from superset.models.core import Database + +logger = logging.getLogger(__name__) +_active_call: ContextVar[WorkerCall | None] = ContextVar( + "mcp_worker_call", default=None +) +_metadata_context_owned: ContextVar[bool] = ContextVar( + "mcp_metadata_context_owned", default=False +) +_P = ParamSpec("_P") +_T = TypeVar("_T") +_pools_lock = threading.Lock() +_pools: WeakKeyDictionary[Flask, WorkerPool] = WeakKeyDictionary() + + +class WorkerDeadlineExceeded(BaseException): + """Stop abandoned work without being swallowed by tool error handlers.""" + + +BUSY_MESSAGE = "MCP server busy: all tool workers are occupied. Retry later." + +# Tools that only read Superset's metadata database. They are admitted under a +# separate, larger bound so they keep answering while slow warehouse queries +# hold every warehouse slot. Anything that can reach a warehouse, a semantic +# layer or a screenshot service stays in the warehouse bound. A listed tool +# that does reach a warehouse must still take a warehouse slot first (see +# WorkerCall.admit_warehouse), so a wrong entry cannot overdraw the pool. +METADATA_ONLY_TOOLS = frozenset( + { + "find_users", + "get_annotation_layer_info", + "get_chart_info", + "get_chart_type_schema", + "get_dashboard_info", + "get_database_info", + "get_dataset_info", + "get_instance_info", + "get_layer_annotation_info", + "get_query_info", + "get_report_info", + "get_rls_filter_info", + "get_role_info", + "get_saved_query_info", + "get_schema", + "get_tag_info", + "get_task_info", + "get_theme_info", + "get_user_info", + "health_check", + "list_annotation_layers", + "list_charts", + "list_dashboards", + "list_databases", + "list_datasets", + "list_layer_annotations", + "list_queries", + "list_reports", + "list_rls_filters", + "list_roles", + "list_saved_queries", + "list_tags", + "list_tasks", + "list_themes", + "list_users", + } +) + + +class WorkerPool: + """Bound submissions, including abandoned work; never queue behind a query.""" + + def __init__(self, size: int, metadata_size: int = 0) -> None: + if size < 1: + raise ValueError("MCP_TOOL_WORKERS must be positive") + if metadata_size < 0: + raise ValueError("MCP_METADATA_TOOL_WORKERS must not be negative") + self.slots = threading.BoundedSemaphore(size) + self.metadata_slots = ( + threading.BoundedSemaphore(metadata_size) if metadata_size else None + ) + self.executor = ThreadPoolExecutor( + size + metadata_size, thread_name_prefix="mcp-tool" + ) + # Cancellation must not wait behind the warehouse work it is cancelling. + # Only calls holding a warehouse slot register cancellation. + self.cancel_slots = threading.BoundedSemaphore(size) + self.cancellations = ThreadPoolExecutor(size, thread_name_prefix="mcp-cancel") + + def admission(self, metadata_only: bool) -> threading.BoundedSemaphore: + """Choose the bound a call is admitted under.""" + if metadata_only and self.metadata_slots is not None: + return self.metadata_slots + return self.slots + + def submit( + self, + fn: Callable[[], Any], + finished: Callable[[], None], + slots: threading.BoundedSemaphore | None = None, + ) -> Future[Any]: + """Admit immediately or report overload without retaining a queued call.""" + slots = self.slots if slots is None else slots + if not slots.acquire(blocking=False): + raise ToolError(BUSY_MESSAGE) + try: + future = self.executor.submit(fn) + except BaseException: + slots.release() + raise + future.add_done_callback(lambda _: finished()) + return future + + def cancel(self, fn: Callable[[], None], finished: Callable[[], None]) -> bool: + """Bound cancellation I/O separately so it cannot delay the caller.""" + if not self.cancel_slots.acquire(blocking=False): + return False + try: + future = self.cancellations.submit(fn) + except RuntimeError: + self.cancel_slots.release() + return False + + def completed(_: Future[None]) -> None: + """Release cancellation capacity before admitting another tool call.""" + self.cancel_slots.release() + finished() + + future.add_done_callback(completed) + return True + + +DEFAULT_TOOL_WORKERS = 16 +DEFAULT_METADATA_TOOL_WORKERS = 16 + + +def _metadata_pool_capacity(app: Flask) -> int | None: + """Return how many metadata connections can be checked out at once. + + ``None`` means a checkout never waits for another holder to return one + (e.g. ``NullPool``, per-thread pools, or unlimited overflow). + """ + from superset import db + + with app.app_context(): + pool = db.engine.pool + if not isinstance(pool, QueuePool): + return None + # SQLAlchemy has no public accessor for the configured overflow limit. + max_overflow = pool._max_overflow # pylint: disable=protected-access + return None if max_overflow < 0 else pool.size() + max_overflow + + +def tool_worker_count(app: Flask) -> int: + """Admit only as many calls as the metadata pool can always serve. + + An admitted call can hold one metadata connection for the whole of its + warehouse I/O, and its cancellation needs another. ``2 * workers + 1`` + connections therefore always leave one that is only held by short metadata + lookups, so cancellation and transport-side lookups (tools/list filtering, + audit logging) never wait for a warehouse query to end on its own. + """ + configured = app.config.get("MCP_TOOL_WORKERS") + capacity = _metadata_pool_capacity(app) + if capacity is None: + return DEFAULT_TOOL_WORKERS if configured is None else configured + limit = (capacity - 1) // 2 + if limit < 1: + raise ValueError( + f"The metadata database pool allows {capacity} connections; " + "MCP tool execution needs at least 3" + ) + if configured is None: + return min(DEFAULT_TOOL_WORKERS, limit) + if configured > limit: + logger.warning( + "MCP_TOOL_WORKERS=%s needs %s metadata database connections, but the " + "pool allows %s; admitting %s concurrent tool calls. Raise the pool's " + "pool_size/max_overflow in SQLALCHEMY_ENGINE_OPTIONS to admit more.", + configured, + 2 * configured + 1, + capacity, + limit, + ) + return limit + return configured + + +def metadata_tool_worker_count(app: Flask) -> int: + """Admit metadata-only calls independently of warehouse capacity. + + They hold a metadata connection only for their own short metadata queries, + never across warehouse I/O, so they cannot keep the connections reserved + for warehouse calls and their cancellation from cycling. They only wait + briefly for a connection, on a worker thread rather than the event loop. + """ + configured = app.config.get("MCP_METADATA_TOOL_WORKERS") + count = DEFAULT_METADATA_TOOL_WORKERS if configured is None else configured Review Comment: Not a blocker, and the docstring matches what I measured for the case this is designed for: with the real 5 plus 10 metadata pool and 7 warehouse calls each holding a connection across blocked warehouse I/O, 16 concurrent metadata-only calls finished in 0.07s. The one thing I would raise is that this bound is never clamped against the pool, unlike its sibling: the default 16 is larger than the default capacity of 15. When all 15 are checked out (the 14 the admission math allows, plus one held by a transport-side audit write) all 16 metadata-only calls block inside `QueuePool` for the full `pool_timeout` and come back as `MCP tool timed out after 30 seconds`, each holding a worker thread for that long. Typed and graceful rather than a correctness problem, but `briefly` becomes the whole deadline in that window. Is clamping the default against capacity worth it here, or is the narrowness of the window the reason to leave it alone? ########## superset/sql/execution/cancellation.py: ########## @@ -0,0 +1,170 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +"""Optional execution-owner hooks around a live warehouse cursor. + +Ordinary web and Celery execution are unchanged. An in-process execution owner +can register cancellation while keeping its async runtime out of model code. +""" + +from __future__ import annotations + +from contextlib import AbstractContextManager, contextmanager +from contextvars import ContextVar +from typing import Any, Callable, Iterator, TYPE_CHECKING + +from sqlalchemy import event +from sqlalchemy.engine import Engine + +if TYPE_CHECKING: + from superset.models.core import Database + +cursor_scope: ContextVar[ + Callable[[Database, Any, str | None, str | None], AbstractContextManager[None]] + | None +] = ContextVar("warehouse_cursor_scope", default=None) + + +@contextmanager +def cancellable_cursor( + database: Database, + cursor: Any, + catalog: str | None = None, + schema: str | None = None, +) -> Iterator[None]: + """Let the execution owner capture cancellation before blocking DBAPI I/O.""" + scope = cursor_scope.get() + if scope is None: + yield + else: + with scope(database, cursor, catalog, schema): + yield + + +before_warehouse_access: ContextVar[Callable[[], None] | None] = ContextVar( + "warehouse_before_access", default=None +) +check_deadline: ContextVar[Callable[[], None] | None] = ContextVar( + "warehouse_check_deadline", default=None +) +after_execute: ContextVar[Callable[[], None] | None] = ContextVar( + "warehouse_after_execute", default=None +) + + +def check_query_deadline() -> None: + """Stop before metadata access or another statement after warehouse I/O.""" + if check := check_deadline.get(): + check() + + +def query_executed() -> None: + """Refresh cancellation handles exposed only after driver execution.""" + if refresh := after_execute.get(): + refresh() + + +_engine_scope: ContextVar[tuple[Database, Engine, str | None, str | None] | None] = ( + ContextVar("warehouse_engine_scope", default=None) +) + + +@contextmanager +def without_execution_hooks() -> Iterator[None]: + """Keep cancellation I/O independent of the query it is cancelling. + + Preserve unrelated context, such as authentication and tenant routing, while + preventing prequeries from checking the abandoned deadline or registering + cancellation recursively. + """ + cursor_token = cursor_scope.set(None) + access_token = before_warehouse_access.set(None) + deadline_token = check_deadline.set(None) + execute_token = after_execute.set(None) + engine_token = _engine_scope.set(None) + try: + yield + finally: + _engine_scope.reset(engine_token) + after_execute.reset(execute_token) + check_deadline.reset(deadline_token) + before_warehouse_access.reset(access_token) + cursor_scope.reset(cursor_token) + + +@contextmanager +def cancellable_engine( + database: Database, engine: Engine, catalog: str | None, schema: str | None +) -> Iterator[None]: + """Cover SQLAlchemy warehouse statements, including metadata discovery. + + Listeners are installed once, never added/removed on shared cached engines. + Matching the engine excludes Superset metadata queries in the same context. + """ + if cursor_scope.get() is None: + yield + return + if admit := before_warehouse_access.get(): + # Before any connection is opened, including prequeries on connect. + admit() Review Comment: This line is what has `unit-tests (current)` red. That job runs `pytest --cov=superset/sql/ ./tests/unit_tests/sql/ --cov-fail-under=100`, and `admit()` plus its branch are the only things missing: ``` superset/sql/execution/cancellation.py 69 1 16 1 98% 122 TOTAL 2151 1 710 1 99% FAIL Required test coverage of 100% not reached. Total coverage: 99.93% ``` I reproduced the same command locally at this head and got the same numbers, and at the merge base `e90b9fb751` it is `100.00%`. The hook is covered, just from a suite the gate does not measure: `test_metadata_only_call_takes_warehouse_slot_before_warehouse_io` exercises it from `tests/unit_tests/mcp_service`. `test_without_execution_hooks_restores_context` does enter `cancellable_engine` with `cursor_scope` set, but it never sets `before_warehouse_access`, so only the false branch is taken. This closed it for me, back to `100.00%` with 2,012 passed: ```python def test_engine_admits_before_opening_a_connection() -> None: """An execution owner is asked to admit warehouse access before connecting.""" admit = Mock() engine = create_engine("sqlite://") with ExitStack() as stack: stack.callback( cancellation.cursor_scope.reset, cancellation.cursor_scope.set(Mock()) ) stack.callback( cancellation.before_warehouse_access.reset, cancellation.before_warehouse_access.set(admit), ) stack.callback(engine.dispose) with cancellation.cancellable_engine(Mock(), engine, None, None): admit.assert_called_once_with() ``` -- 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]
