refactor(api): make tool_handlers require a DeviceManager

Eliminates the silent fallback to DEFAULT_MANAGER that produced the
DeviceNotFoundError incident. All existing callers already pass
manager explicitly.
This commit is contained in:
2026-07-21 13:52:01 +08:00
parent d1b0fffabb
commit 29b9a8c39a
2 changed files with 26 additions and 15 deletions
+17 -15
View File
@@ -3,7 +3,7 @@ from collections.abc import Callable
from typing import Any from typing import Any
from api.errors import call_with_semantic_errors from api.errors import call_with_semantic_errors
from device.manager import DEFAULT_MANAGER, DeviceManager from device.manager import DeviceManager
from tools.describe_screen import describe_screen from tools.describe_screen import describe_screen
from tools.find_icon import find_icon_on_screen from tools.find_icon import find_icon_on_screen
from tools.find_text import find_text_on_screen from tools.find_text import find_text_on_screen
@@ -17,12 +17,11 @@ from tools.ui_tree import get_ui_tree
def tool_handlers( def tool_handlers(
*, *,
manager: DeviceManager | None = None, manager: DeviceManager,
) -> dict[str, Callable[..., Any]]: ) -> dict[str, Callable[..., Any]]:
device_manager = manager or DEFAULT_MANAGER
def _screenshot(device_id: str | None = None) -> dict[str, Any]: def _screenshot(device_id: str | None = None) -> dict[str, Any]:
image = take_screenshot(device_id, manager=device_manager) image = take_screenshot(device_id, manager=manager)
return { return {
"ok": True, "ok": True,
"image_base64": base64.b64encode(image).decode("ascii"), "image_base64": base64.b64encode(image).decode("ascii"),
@@ -39,7 +38,7 @@ def tool_handlers(
x, x,
y, y,
device_id=device_id, device_id=device_id,
manager=device_manager, manager=manager,
), ),
"swipe": lambda start_x, start_y, end_x, end_y, duration_ms=500, device_id=None: call_with_semantic_errors( "swipe": lambda start_x, start_y, end_x, end_y, duration_ms=500, device_id=None: call_with_semantic_errors(
swipe, swipe,
@@ -49,55 +48,55 @@ def tool_handlers(
end_y, end_y,
duration_ms=duration_ms, duration_ms=duration_ms,
device_id=device_id, device_id=device_id,
manager=device_manager, manager=manager,
), ),
"input_text": lambda text, device_id=None: call_with_semantic_errors( "input_text": lambda text, device_id=None: call_with_semantic_errors(
input_text, input_text,
text, text,
device_id=device_id, device_id=device_id,
manager=device_manager, manager=manager,
), ),
"launch_app": lambda app_id, device_id=None: call_with_semantic_errors( "launch_app": lambda app_id, device_id=None: call_with_semantic_errors(
launch_app, launch_app,
app_id, app_id,
device_id=device_id, device_id=device_id,
manager=device_manager, manager=manager,
), ),
"find_text": lambda query, device_id=None: call_with_semantic_errors( "find_text": lambda query, device_id=None: call_with_semantic_errors(
find_text_on_screen, find_text_on_screen,
query, query,
device_id=device_id, device_id=device_id,
manager=device_manager, manager=manager,
), ),
"find_icon": lambda name, device_id=None: call_with_semantic_errors( "find_icon": lambda name, device_id=None: call_with_semantic_errors(
find_icon_on_screen, find_icon_on_screen,
name, name,
device_id=device_id, device_id=device_id,
manager=device_manager, manager=manager,
), ),
"get_ui_tree": lambda device_id=None, include_app_info=False: ( "get_ui_tree": lambda device_id=None, include_app_info=False: (
call_with_semantic_errors( call_with_semantic_errors(
get_ui_tree, get_ui_tree,
device_id, device_id,
manager=device_manager, manager=manager,
include_app_info=include_app_info, include_app_info=include_app_info,
) )
), ),
"describe_screen": lambda device_id=None: call_with_semantic_errors( "describe_screen": lambda device_id=None: call_with_semantic_errors(
lambda: describe_screen(device_id, manager=device_manager).to_dict() lambda: describe_screen(device_id, manager=manager).to_dict()
), ),
"list_devices": lambda: [ "list_devices": lambda: [
device.to_dict() for device in device_manager.list_devices() device.to_dict() for device in manager.list_devices()
], ],
"device_status": lambda device_id: call_with_semantic_errors( "device_status": lambda device_id: call_with_semantic_errors(
lambda: {"device_id": device_id, "status": device_manager.status(device_id)} lambda: {"device_id": device_id, "status": manager.status(device_id)}
), ),
} }
def create_mcp_server( def create_mcp_server(
*, *,
manager: DeviceManager | None = None, manager: DeviceManager,
skill_catalog_store: Any | None = None, skill_catalog_store: Any | None = None,
skill_active_subscriptions: set[str] | None = None, skill_active_subscriptions: set[str] | None = None,
skill_local_store: Any | None = None, skill_local_store: Any | None = None,
@@ -107,6 +106,9 @@ def create_mcp_server(
except ImportError as exc: except ImportError as exc:
raise RuntimeError("mcp SDK is not installed") from exc raise RuntimeError("mcp SDK is not installed") from exc
if manager is None:
raise ValueError("create_mcp_server requires a non-None manager")
handlers = tool_handlers(manager=manager) handlers = tool_handlers(manager=manager)
server = FastMCP("apex-agent") server = FastMCP("apex-agent")
+9
View File
@@ -64,3 +64,12 @@ def test_mcp_ui_tree_can_include_active_app_info() -> None:
"activity": ".MainActivity", "activity": ".MainActivity",
} }
assert isinstance(response["nodes"], list) assert isinstance(response["nodes"], list)
def test_tool_handlers_requires_manager() -> None:
"""D12: tool_handlers must not silently fall back to DEFAULT_MANAGER."""
from api.mcp import tool_handlers
import pytest
with pytest.raises(TypeError):
tool_handlers() # type: ignore[call-arg]