diff --git a/strix/core/agents.py b/strix/core/agents.py index eb1b76bc..efd7cb60 100644 --- a/strix/core/agents.py +++ b/strix/core/agents.py @@ -291,13 +291,29 @@ class AgentCoordinator: async def mark_running(self, agent_id: str) -> None: async with self._lock: if agent_id in self.statuses: - self.statuses[agent_id] = "running" - self.errors.pop(agent_id, None) - self.wait_kinds.pop(agent_id, None) - self.runtimes.setdefault(agent_id, AgentRuntime()).user_wake_required = False - self._parent_notified.discard(agent_id) + self._set_running_locked(agent_id) await self._maybe_snapshot() + async def resume_silent_user_wait(self, agent_id: str) -> bool: + """Undo a park on the user, unless the agent's state has moved on since. + + Returns False without touching anything when the agent is no longer + waiting on the user (a stop or a delivered message got there first). + """ + async with self._lock: + if self.statuses.get(agent_id) != "waiting" or self.wait_kinds.get(agent_id) != "user": + return False + self._set_running_locked(agent_id) + await self._maybe_snapshot() + return True + + def _set_running_locked(self, agent_id: str) -> None: + self.statuses[agent_id] = "running" + self.errors.pop(agent_id, None) + self.wait_kinds.pop(agent_id, None) + self.runtimes.setdefault(agent_id, AgentRuntime()).user_wake_required = False + self._parent_notified.discard(agent_id) + async def park_waiting(self, agent_id: str, *, wait_kind: WaitKind) -> None: """Park an agent, recording what it is waiting on so the driver can time it.""" async with self._lock: diff --git a/strix/core/execution.py b/strix/core/execution.py index ac44c53f..91d94498 100644 --- a/strix/core/execution.py +++ b/strix/core/execution.py @@ -574,10 +574,12 @@ async def _run_until_lifecycle( continue said_to_user = said_to_user or _said_to_user(result) - status = await _agent_status(coordinator, agent_id) + # Atomic: only an agent still parked on the user is put back to work, so + # a stop that lands in between is never overwritten. silent_yield = ( - interactive and not said_to_user and await _parked_for_user(coordinator, agent_id) + interactive and not said_to_user and await coordinator.resume_silent_user_wait(agent_id) ) + status = await _agent_status(coordinator, agent_id) if status != "running" and not silent_yield: await coordinator.reset_recovery(agent_id) return result @@ -591,9 +593,6 @@ async def _run_until_lifecycle( interactive=interactive, silent_yield=silent_yield, ) - if silent_yield: - await coordinator.mark_running(agent_id) - if recoveries >= recovery_limit: return await _exhausted_recovery(coordinator, agent_id, result, interactive=interactive) @@ -895,14 +894,6 @@ async def _agent_status(coordinator: AgentCoordinator, agent_id: str) -> Status return coordinator.statuses.get(agent_id) -async def _parked_for_user(coordinator: AgentCoordinator, agent_id: str) -> bool: - async with coordinator._lock: - return ( - coordinator.statuses.get(agent_id) == "waiting" - and coordinator.wait_kinds.get(agent_id) == "user" - ) - - def _log_recovery( agent_id: str, result: RunResultBase | None, diff --git a/tests/test_execution.py b/tests/test_execution.py index 62c5814f..83d3829b 100644 --- a/tests/test_execution.py +++ b/tests/test_execution.py @@ -1436,6 +1436,48 @@ async def test_whitespace_only_text_does_not_count_as_a_reply( assert len(calls) == 2 +@pytest.mark.asyncio +async def test_a_stop_during_a_silent_yield_is_kept( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """An operator stop that lands while the agent is parked must survive the bounce.""" + coordinator = AgentCoordinator() + await coordinator.register("root", "strix", parent_id=None) + calls: list[Any] = [] + + async def _park_then_get_stopped(*_args: Any, **kwargs: Any) -> Any: + calls.append(kwargs.get("input_data")) + await coordinator.park_waiting("root", wait_kind="user") + await coordinator.request_stop("root") + return MagicMock(final_output="", new_items=[]) + + monkeypatch.setattr(execution, "_run_cycle_parked", _park_then_get_stopped) + + await _drive(coordinator, "root", interactive=True) + + assert len(calls) == 1 + assert coordinator.statuses["root"] == "stopped" + + +@pytest.mark.asyncio +async def test_resume_silent_user_wait_only_touches_a_user_park() -> None: + coordinator = AgentCoordinator() + await coordinator.register("root", "strix", parent_id=None) + + await coordinator.park_waiting("root", wait_kind="user") + assert await coordinator.resume_silent_user_wait("root") is True + assert coordinator.statuses["root"] == "running" + assert "root" not in coordinator.wait_kinds + + await coordinator.park_waiting("root", wait_kind="agents") + assert await coordinator.resume_silent_user_wait("root") is False + assert coordinator.statuses.get("root") == "waiting" + + await coordinator.request_stop("root") + assert await coordinator.resume_silent_user_wait("root") is False + assert coordinator.statuses.get("root") == "stopped" + + @pytest.mark.asyncio async def test_a_wait_on_agents_is_never_a_silent_yield( monkeypatch: pytest.MonkeyPatch,