Skip to content

Commit 4fc496d

Browse files
committed
fix(bulk): address selector review feedback
1 parent 414aa49 commit 4fc496d

9 files changed

Lines changed: 304 additions & 70 deletions

File tree

src/ha_mcp/policy/middleware.py

Lines changed: 12 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -64,6 +64,11 @@ def _passes_ungated(name: str, args: dict[str, Any]) -> bool:
6464
return name in PROXY_META_TOOLS or _is_approval_management(name, args)
6565

6666

67+
def _has_dynamic_selector_targets(name: str, args: dict[str, Any]) -> bool:
68+
"""Return whether identical arguments can resolve to different targets later."""
69+
return name == "ha_bulk_control" and args.get("selector") is not None
70+
71+
6772
class PolicyMiddleware(Middleware):
6873
"""Gate tool calls against a Policy, blocking with progress heartbeats."""
6974

@@ -120,15 +125,19 @@ async def on_call_tool(
120125

121126
rule = find_matching_rule(name, args, policy)
122127
args_hash = compute_args_hash(args)
128+
dynamic_targets = _has_dynamic_selector_targets(name, args)
129+
remember_minutes = (
130+
0 if dynamic_targets else rule.remember_minutes if rule else 0
131+
)
123132

124-
if self._queue.is_remembered(name, args_hash):
133+
if not dynamic_targets and self._queue.is_remembered(name, args_hash):
125134
return await call_next(context)
126135

127136
existing = self._queue.find(name, args_hash)
128137
if existing and existing.decision == "approved":
129138
self._queue.consume_and_maybe_remember(
130139
existing,
131-
remember_minutes=rule.remember_minutes if rule else 0,
140+
remember_minutes=remember_minutes,
132141
)
133142
return await call_next(context)
134143
if existing and existing.decision == "denied":
@@ -155,7 +164,7 @@ async def on_call_tool(
155164
if pending.decision == "approved":
156165
self._queue.consume_and_maybe_remember(
157166
pending,
158-
remember_minutes=rule.remember_minutes if rule else 0,
167+
remember_minutes=remember_minutes,
159168
)
160169
return await call_next(context)
161170
if pending.decision == "denied":

src/ha_mcp/tools/bulk_selector.py

Lines changed: 34 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -96,9 +96,20 @@ def _expand_entity(
9696
*,
9797
path: tuple[str, ...],
9898
expanded_groups: set[str],
99+
cache: dict[str, tuple[frozenset[str], frozenset[str]]],
99100
parameter: str,
100101
) -> set[str]:
101102
"""Expand generic membership recursively and reject incomplete graph walks."""
103+
if entity_id in path:
104+
cycle_start = path.index(entity_id)
105+
cycle = " -> ".join((*path[cycle_start:], entity_id))
106+
raise BulkSelectorValidationError(
107+
f"Aggregate membership cycle detected: {cycle}"
108+
)
109+
if cached := cache.get(entity_id):
110+
cached_leaves, cached_groups = cached
111+
expanded_groups.update(cached_groups)
112+
return set(cached_leaves)
102113
state = states.get(entity_id)
103114
if state is None:
104115
raise BulkSelectorValidationError(
@@ -107,27 +118,29 @@ def _expand_entity(
107118
)
108119
members = normalize_member_entity_ids(state.get("attributes"))
109120
if members is None:
121+
cache[entity_id] = (frozenset({entity_id}), frozenset())
110122
return {entity_id}
111-
if entity_id in path:
112-
cycle_start = path.index(entity_id)
113-
cycle = " -> ".join((*path[cycle_start:], entity_id))
114-
raise BulkSelectorValidationError(
115-
f"Aggregate membership cycle detected: {cycle}"
116-
)
117123
expanded_groups.add(entity_id)
118124
if not members:
125+
cache[entity_id] = (frozenset(), frozenset({entity_id}))
119126
return set()
120127
leaves: set[str] = set()
128+
nested_groups: set[str] = {entity_id}
121129
for member_id in members:
130+
member_groups: set[str] = set()
122131
leaves.update(
123132
_expand_entity(
124133
member_id,
125134
states,
126135
path=(*path, entity_id),
127-
expanded_groups=expanded_groups,
136+
expanded_groups=member_groups,
137+
cache=cache,
128138
parameter=parameter,
129139
)
130140
)
141+
nested_groups.update(member_groups)
142+
expanded_groups.update(nested_groups)
143+
cache[entity_id] = (frozenset(leaves), frozenset(nested_groups))
131144
return leaves
132145

133146

@@ -149,7 +162,7 @@ def _validate_selector(
149162
parameter="selector.domain",
150163
)
151164
domain = domain.strip().lower()
152-
normalized_action = action.strip() if isinstance(action, str) else ""
165+
normalized_action = action.strip().lower() if isinstance(action, str) else ""
153166
valid_actions = get_domain_handler(domain).get(
154167
"valid_actions", ["on", "off", "toggle"]
155168
)
@@ -214,6 +227,7 @@ def _expand_roots(
214227
roots: list[str],
215228
states: Mapping[str, Mapping[str, Any]],
216229
expanded_groups: set[str],
230+
cache: dict[str, tuple[frozenset[str], frozenset[str]]],
217231
parameter: str,
218232
) -> set[str]:
219233
"""Expand a list of aggregate or leaf roots into a deduplicated leaf set."""
@@ -225,6 +239,7 @@ def _expand_roots(
225239
states,
226240
path=(),
227241
expanded_groups=expanded_groups,
242+
cache=cache,
228243
parameter=parameter,
229244
)
230245
)
@@ -315,21 +330,24 @@ async def resolve_bulk_selector(
315330
hidden = await _load_hidden_entities(
316331
client, entity_result, states_result, device_result, entity_registry
317332
)
318-
candidate_roots = sorted(
333+
matching_roots = {
319334
entity_id
320335
for entity_id in states
321-
if entity_id not in hidden
322-
and entity_id.startswith(f"{domain}.")
336+
if entity_id.startswith(f"{domain}.")
323337
and _entity_area_id(entity_id, entity_registry, device_areas) in selected_areas
324-
)
338+
}
339+
directly_hidden = matching_roots & hidden
340+
candidate_roots = sorted(matching_roots - hidden)
325341
expanded_groups: set[str] = set()
342+
expansion_cache: dict[str, tuple[frozenset[str], frozenset[str]]] = {}
326343
selected_leaves = _expand_roots(
327-
candidate_roots, states, expanded_groups, "selector"
344+
candidate_roots, states, expanded_groups, expansion_cache, "selector"
328345
)
329346
excluded_leaves = _expand_roots(
330347
excluded_roots,
331348
states,
332349
expanded_groups,
350+
expansion_cache,
333351
"selector.exclude_entity_ids",
334352
)
335353
selected_leaves = {
@@ -367,5 +385,7 @@ async def resolve_bulk_selector(
367385
excluded_entity_ids=sorted(effective_excluded - hidden),
368386
selected_area_ids=sorted(selected_areas),
369387
expanded_group_ids=sorted(expanded_groups - hidden),
370-
hidden_entity_count=len(hidden_selected | (effective_excluded & hidden)),
388+
hidden_entity_count=len(
389+
directly_hidden | hidden_selected | (effective_excluded & hidden)
390+
),
371391
)

src/ha_mcp/tools/tools_service.py

Lines changed: 11 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1657,6 +1657,10 @@ async def ha_bulk_control(
16571657
)
16581658
if not isinstance(parsed_selector, dict):
16591659
raise BulkSelectorValidationError("selector must be a JSON object")
1660+
except (ValueError, BulkSelectorValidationError) as exc:
1661+
parameter = getattr(exc, "parameter", "selector")
1662+
raise_tool_error(create_validation_error(str(exc), parameter=parameter))
1663+
try:
16601664
resolution = await resolve_bulk_selector(
16611665
self._client,
16621666
parsed_selector,
@@ -1665,9 +1669,15 @@ async def ha_bulk_control(
16651669
timeout_seconds=timeout_seconds,
16661670
validate_first=validate_first,
16671671
)
1668-
except (ValueError, BulkSelectorValidationError) as exc:
1672+
except BulkSelectorValidationError as exc:
16691673
parameter = getattr(exc, "parameter", "selector")
16701674
raise_tool_error(create_validation_error(str(exc), parameter=parameter))
1675+
except Exception as exc:
1676+
exception_to_structured_error(
1677+
exc,
1678+
context={"operation": "resolve bulk selector"},
1679+
)
1680+
raise # unreachable: exception_to_structured_error always raises
16711681
if dry_run:
16721682
return {
16731683
"success": True,

tests/src/e2e/workflows/core/test_bulk.py

Lines changed: 74 additions & 46 deletions
Original file line numberDiff line numberDiff line change
@@ -13,7 +13,13 @@
1313

1414
import pytest
1515

16-
from ...utilities.assertions import assert_mcp_success, parse_mcp_result, safe_call_tool
16+
from ...utilities.assertions import (
17+
MCPAssertions,
18+
assert_mcp_success,
19+
parse_mcp_result,
20+
safe_call_tool,
21+
)
22+
from ...utilities.wait_helpers import wait_for_entity_state
1723

1824
logger = logging.getLogger(__name__)
1925

@@ -46,56 +52,78 @@ def _extract_bulk_boolean_entity_id(data: dict) -> str | None:
4652
class TestBulkControl:
4753
"""Test ha_bulk_control tool functionality."""
4854

49-
async def test_selector_dry_run_resolves_area_and_exclusion(
50-
self, mcp_client, cleanup_tracker
51-
):
55+
async def test_selector_dry_run_resolves_area_and_exclusion(self, mcp_client):
5256
"""Preview exact area leaves after applying an entity exclusion."""
5357
suffix = uuid4().hex[:8]
54-
area_result = await mcp_client.call_tool(
55-
"ha_set_area_or_floor",
56-
{"kind": "area", "name": f"Bulk selector {suffix}"},
57-
)
58-
area_data = assert_mcp_success(area_result, "Create selector area")
59-
area_id = area_data["area_id"]
60-
cleanup_tracker.track("area", area_id)
61-
62-
entity_ids = []
63-
for label in ("included", "excluded"):
64-
create_result = await mcp_client.call_tool(
65-
"ha_config_set_helper",
66-
{
67-
"helper_type": "input_boolean",
68-
"name": f"Bulk selector {label} {suffix}",
69-
},
70-
)
71-
create_data = assert_mcp_success(create_result, "Create selector helper")
72-
entity_id = _extract_bulk_boolean_entity_id(create_data)
73-
assert entity_id, f"Missing helper entity_id: {create_data}"
74-
cleanup_tracker.track("input_boolean", entity_id)
75-
entity_ids.append(entity_id)
76-
assign_result = await mcp_client.call_tool(
77-
"ha_set_entity", {"entity_id": entity_id, "area_id": area_id}
58+
area_id: str | None = None
59+
entity_ids: list[str] = []
60+
try:
61+
area_result = await mcp_client.call_tool(
62+
"ha_set_area_or_floor",
63+
{"kind": "area", "name": f"Bulk selector {suffix}"},
7864
)
79-
assert_mcp_success(assign_result, "Assign selector helper to area")
65+
area_data = assert_mcp_success(area_result, "Create selector area")
66+
area_id = area_data["area_id"]
67+
68+
for label in ("included", "excluded"):
69+
create_result = await mcp_client.call_tool(
70+
"ha_config_set_helper",
71+
{
72+
"helper_type": "input_boolean",
73+
"name": f"Bulk selector {label} {suffix}",
74+
},
75+
)
76+
create_data = assert_mcp_success(
77+
create_result, "Create selector helper"
78+
)
79+
entity_id = _extract_bulk_boolean_entity_id(create_data)
80+
assert entity_id, f"Missing helper entity_id: {create_data}"
81+
entity_ids.append(entity_id)
82+
assign_result = await mcp_client.call_tool(
83+
"ha_set_entity", {"entity_id": entity_id, "area_id": area_id}
84+
)
85+
assert_mcp_success(assign_result, "Assign selector helper to area")
8086

81-
result = await mcp_client.call_tool(
82-
"ha_bulk_control",
83-
{
84-
"selector": {
85-
"domain": "input_boolean",
86-
"area_ids": [area_id],
87-
"exclude_entity_ids": [entity_ids[1]],
88-
},
89-
"action": "off",
90-
"dry_run": True,
91-
},
92-
)
93-
data = assert_mcp_success(result, "Preview structural bulk selection")
87+
for entity_id in entity_ids:
88+
assert await wait_for_entity_state(mcp_client, entity_id, "off"), (
89+
f"Selector helper {entity_id} was not registered in time"
90+
)
9491

95-
assert data["dry_run"] is True
96-
assert data["dispatched"] is False
97-
assert data["resolution"]["resolved_entity_ids"] == [entity_ids[0]]
98-
assert data["resolution"]["excluded_entity_ids"] == [entity_ids[1]]
92+
async with MCPAssertions(mcp_client) as mcp:
93+
data = await mcp.call_tool_success(
94+
"ha_bulk_control",
95+
{
96+
"selector": {
97+
"domain": "input_boolean",
98+
"area_ids": [area_id],
99+
"exclude_entity_ids": [entity_ids[1]],
100+
},
101+
"action": "off",
102+
"dry_run": True,
103+
},
104+
)
105+
106+
assert data["dry_run"] is True
107+
assert data["dispatched"] is False
108+
assert data["resolution"]["resolved_entity_ids"] == [entity_ids[0]]
109+
assert data["resolution"]["excluded_entity_ids"] == [entity_ids[1]]
110+
finally:
111+
for entity_id in entity_ids:
112+
await safe_call_tool(
113+
mcp_client,
114+
"ha_remove_helpers_integrations",
115+
{
116+
"helper_type": "input_boolean",
117+
"target": entity_id,
118+
"confirm": True,
119+
},
120+
)
121+
if area_id is not None:
122+
await safe_call_tool(
123+
mcp_client,
124+
"ha_remove_area_or_floor",
125+
{"kind": "area", "id": area_id},
126+
)
99127

100128
async def test_bulk_turn_on_single_light(self, mcp_client, test_light_entity):
101129
"""Test bulk_control with a single light entity."""

tests/src/e2e/workflows/groups/test_lifecycle.py

Lines changed: 8 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -95,10 +95,14 @@ async def test_group_basic_lifecycle(self, mcp_client, cleanup_tracker):
9595
},
9696
)
9797
search_record = next(
98-
entity
99-
for entity in search_data["entities"]
100-
if entity["entity_id"] == f"group.{object_id}"
101-
)
98+
(
99+
entity
100+
for entity in search_data["entities"]
101+
if entity["entity_id"] == f"group.{object_id}"
102+
),
103+
None,
104+
)
105+
assert search_record is not None, search_data
102106
assert search_record["is_group"] is True
103107
assert search_record["member_entity_ids"] == [
104108
"light.bed_light",

tests/src/unit/policy/test_evaluator.py

Lines changed: 10 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -416,7 +416,7 @@ def test_operations_rule_force_gates_selector(self):
416416
Predicate(
417417
path="args.operations.*.entity_id",
418418
op="regex",
419-
value=r"^lock\\.",
419+
value=r"^lock\.",
420420
)
421421
],
422422
)
@@ -456,7 +456,7 @@ def test_operations_calls_keep_normal_predicate_semantics(self):
456456
Predicate(
457457
path="args.operations.*.entity_id",
458458
op="regex",
459-
value=r"^lock\\.",
459+
value=r"^lock\.",
460460
)
461461
],
462462
)
@@ -471,3 +471,11 @@ def test_operations_calls_keep_normal_predicate_semantics(self):
471471
)
472472
== Verdict.ALLOW
473473
)
474+
assert (
475+
evaluate(
476+
"ha_bulk_control",
477+
{"operations": [{"entity_id": "lock.front", "action": "lock"}]},
478+
policy,
479+
)
480+
== Verdict.REQUIRE_APPROVAL
481+
)

0 commit comments

Comments
 (0)