diff --git a/tests/test_nudge_queue.py b/tests/test_nudge_queue.py index 3fb3e724..09d8d94a 100644 --- a/tests/test_nudge_queue.py +++ b/tests/test_nudge_queue.py @@ -2,6 +2,7 @@ from __future__ import annotations +import logging import threading import pytest @@ -401,8 +402,10 @@ class TestValidation: class TestValidUntil: - """``valid_until`` predicate: drain re-checks freshness; falsy / - raising predicates drop the entry without delivery. + """``valid_until`` predicate: drain re-checks freshness. Falsy + predicates drop the entry without delivery and log at ``info`` + (normal lifecycle outcome); raising predicates drop the entry and + log at ``warning`` with ``exc_info`` (misbehaving predicate). """ def test_valid_until_true_delivers(self): @@ -411,26 +414,52 @@ class TestValidUntil: out = q.drain({"any"}) assert out == [("a", "1", None)] - def test_valid_until_false_drops_silently(self): + def test_valid_until_false_drops_with_info_log(self, caplog: pytest.LogCaptureFixture): q = NudgeQueue() q.enqueue("a", "1", "any", valid_until=lambda: False) - out = q.drain({"any"}) + with caplog.at_level(logging.INFO, logger="turnstone.core.nudge_queue"): + out = q.drain({"any"}) assert out == [] # Already removed from queue (drain partition removes BEFORE # predicate check — falsy doesn't return to queue). assert len(q) == 0 + # The drop emits a structured info record so a wiring + # regression (a predicate that always returns False) is still + # observable, without spamming ``warning`` for the routine + # lifecycle case where ``valid_until`` is doing its job. + # structlog renders the event name + extras into ``msg`` as a + # single rendered string, so substring-match like the + # ``watch_dispatch.queue_full`` assertion in + # tests/test_watch_dispatch.py. + drops = [r for r in caplog.records if "nudge_queue.predicate_dropped" in r.getMessage()] + assert len(drops) == 1 + assert drops[0].levelno == logging.INFO + assert "predicate_false" in drops[0].getMessage() + assert "'nudge_type': 'a'" in drops[0].getMessage() + assert "'channel': 'any'" in drops[0].getMessage() + assert "'text_len': 1" in drops[0].getMessage() - def test_valid_until_exception_drops_silently(self): + def test_valid_until_exception_drops_with_warning(self, caplog: pytest.LogCaptureFixture): q = NudgeQueue() def boom() -> bool: raise RuntimeError("predicate crash") q.enqueue("a", "1", "any", valid_until=boom) - out = q.drain({"any"}) + with caplog.at_level(logging.WARNING, logger="turnstone.core.nudge_queue"): + out = q.drain({"any"}) assert out == [] # Crash-on-predicate is treated as "no longer valid" — drop, not propagate. assert len(q) == 0 + # Stays at ``warning`` (with ``exc_info``) because a raising + # predicate is a bug, not a normal lifecycle outcome. + drops = [r for r in caplog.records if "nudge_queue.predicate_dropped" in r.getMessage()] + assert len(drops) == 1 + assert drops[0].levelno == logging.WARNING + rendered = drops[0].getMessage() + assert "predicate_raised" in rendered + assert "RuntimeError" in rendered + assert "predicate crash" in rendered def test_valid_until_evaluated_outside_lock(self): """The predicate may do non-trivial work (e.g. storage I/O) diff --git a/tests/test_watch_dispatch.py b/tests/test_watch_dispatch.py index 4b850965..0c8d8471 100644 --- a/tests/test_watch_dispatch.py +++ b/tests/test_watch_dispatch.py @@ -232,59 +232,52 @@ class TestSoftCap: # --------------------------------------------------------------------------- -# valid_until predicate +# Predicate independence # --------------------------------------------------------------------------- -class TestValidUntil: - """The ``valid_until`` predicate captured at dispatch time re-checks - the watch's ``active`` flag at drain time, so a cancelled watch's - last splat doesn't ride out a future wake. +class TestPredicateIndependence: + """The watch closure does NOT wire a ``valid_until`` predicate. + + Earlier the closure wired ``_still_active`` (re-reading + ``is_watch_active`` at drain time). That predicate raced + ``WatchRunner._poll_watch``'s commit of ``active=False`` and silently + dropped every terminal fire. The closure now enqueues without a + predicate; entries survive drain regardless of the row's ``active`` + column state. """ - def test_valid_until_drops_when_watch_inactive(self, tmp_db, monkeypatch): + def test_drain_delivers_even_when_storage_reports_inactive(self, tmp_db, monkeypatch): session = _make_session_for_dispatch() _runner, dispatch = _register_runner(session) - # Storage stub returns False at drain time. - is_active_calls = patch_session_storage(monkeypatch, active=False) - - dispatch(_reminder("body"), "watch-1") - # Drain fires the predicate; entry should NOT be delivered. - out = session._nudge_queue.drain({"any"}) - assert out == [] - # Predicate ran once with the dispatched watch_id. - assert is_active_calls == ["watch-1"] - - def test_valid_until_drops_when_storage_raises(self, tmp_db, monkeypatch): - """The closure's broad-except in the predicate translates a - storage-layer exception to ``False`` so the drain doesn't - propagate; the predicate captured ``watch_id`` correctly - (otherwise storage wouldn't even be touched). - """ - session = _make_session_for_dispatch() - _runner, dispatch = _register_runner(session) - - patch_session_storage(monkeypatch, raise_on_is_active=True) - - dispatch(_reminder("body"), "watch-bound-id") - out = session._nudge_queue.drain({"any"}) - assert out == [] - - def test_valid_until_delivers_when_watch_active(self, tmp_db, monkeypatch): - """Happy-path counter-test for the predicate above: the entry - DOES drain when the watch is still active. - """ - session = _make_session_for_dispatch() - _runner, dispatch = _register_runner(session) - - patch_session_storage(monkeypatch, active=True) + # Even if storage reports active=False, the entry should still + # drain — no predicate to drop it. + patch_session_storage(monkeypatch, active=False) dispatch(_reminder("body"), "watch-1") out = session._nudge_queue.drain({"any"}) assert len(out) == 1 assert out[0][0] == "watch_triggered" + def test_dispatch_never_calls_is_watch_active(self, tmp_db, monkeypatch): + """Pin the invariant directly: the closure must NOT consult + ``storage.is_watch_active`` anywhere along the enqueue + drain + path. Without this assertion, a future change that re-wires + an ``is_watch_active`` predicate would silently bring back the + bug that motivates this whole module. + """ + session = _make_session_for_dispatch() + _runner, dispatch = _register_runner(session) + + is_active_calls = patch_session_storage(monkeypatch, active=True) + + dispatch(_reminder("body"), "watch-bound-id") + session._nudge_queue.drain({"any"}) + assert is_active_calls == [], ( + f"watch closure must not call is_watch_active; got {is_active_calls!r}" + ) + # --------------------------------------------------------------------------- # Concurrency diff --git a/tests/test_watch_integration.py b/tests/test_watch_integration.py index ecf86f00..0c532c5f 100644 --- a/tests/test_watch_integration.py +++ b/tests/test_watch_integration.py @@ -24,11 +24,15 @@ the structural integration gate for the watch switchover. from __future__ import annotations +import contextlib from typing import Any from unittest.mock import MagicMock, patch +import pytest + from tests._helpers import patch_session_storage from turnstone.core.session import ChatSession +from turnstone.core.storage import get_storage from turnstone.core.watch import WatchRunner @@ -272,3 +276,283 @@ def test_watch_dispatch_through_restore_fn_lands_on_rehydrated_session(tmp_db, m # Original session's queue stays empty — the dispatch did NOT # accidentally route back to it. assert len(original._nudge_queue) == 0 + + +@pytest.mark.parametrize( + ("stop_on", "max_polls", "label"), + [ + ('"HIT" in output', 100, "stop_on_fired"), + (None, 1, "max_polls_reached"), + ], +) +def test_poll_watch_terminal_fire_survives_drain( + tmp_db: str, + monkeypatch: pytest.MonkeyPatch, + stop_on: str | None, + max_polls: int, + label: str, +) -> None: + """Regression for the dispatch-ordering bug. + + With the broken ordering (``update_watch(active=False)`` before + ``_dispatch_result``) plus the ``_still_active`` ``valid_until`` + predicate that re-reads ``is_watch_active`` at drain time, every + terminal watch fire was silently dropped — the closure enqueued + the entry but the predicate immediately invalidated it because + the row's ``active`` flag had already been flipped to ``0`` in + the same poll. The model never saw the fire. + + This test drives a REAL ``WatchRunner._poll_watch`` against a real + ``tmp_db`` watch row (no ``patch_session_storage(active=True)`` + stub — that stub is exactly what masked the bug in earlier tests). + Covers both terminal paths: ``stop_on`` condition matched and + ``poll_count >= max_polls`` reached. + """ + session = _make_session() + storage = get_storage() + + runner = WatchRunner(storage=storage, node_id="test-node") + session.set_watch_runner(runner) + + storage.create_watch( + watch_id=f"w-regression-{label}", + ws_id=session._ws_id, + node_id="test-node", + name=f"regression-{label}", + command="echo HIT", + interval_secs=10.0, + stop_on=stop_on, + max_polls=max_polls, + created_by="model", + next_poll="1970-01-01T00:00:00", + ) + + # Spy ``enqueue`` so the assertion can distinguish "dispatch never + # called" (a different bug class) from "dispatch enqueued but the + # predicate dropped it at drain" (this bug). + enqueue_calls: list[tuple[str, str, str]] = [] + real_enqueue = session._nudge_queue.enqueue + + def _spy_enqueue(*args: Any, **kwargs: Any) -> None: + enqueue_calls.append((args[0], args[1][:40], args[2])) + return real_enqueue(*args, **kwargs) + + monkeypatch.setattr(session._nudge_queue, "enqueue", _spy_enqueue) + + # For the max_polls=1 case the first poll has prev_output=None and + # would not normally fire on output change; the max_polls branch + # at watch.py:412-414 still marks is_final=True so dispatch runs. + due = storage.list_due_watches("2099-01-01T00:00:00") + matching = [r for r in due if r["watch_id"] == f"w-regression-{label}"] + assert len(matching) == 1, f"watch row not picked up by list_due_watches: {due!r}" + runner._poll_watch(matching[0]) + + assert len(enqueue_calls) == 1, ( + f"_poll_watch did not enqueue exactly one fire (got {enqueue_calls!r}); " + "this is a different bug from the predicate-drop regression" + ) + assert enqueue_calls[0][0] == "watch_triggered" + + assert storage.is_watch_active(f"w-regression-{label}") is False, ( + "terminal fire should have committed active=False on the row" + ) + + # The key assertion: drain delivers the entry. Pre-fix this + # returned ``[]`` because the ``_still_active`` predicate re-read + # ``active=0``. Post-fix the watch closure no longer wires a + # predicate and the entry survives. + out = session._nudge_queue.drain({"any"}) + assert len(out) == 1, ( + "Watch fire was enqueued but never reached drain — dispatch-ordering " + "regression. Check that WatchRunner._poll_watch dispatches BEFORE " + "committing active=False, and that the watch closure in " + "ChatSession.set_watch_runner does not wire an is_watch_active " + "predicate." + ) + nt, text, _meta = out[0] + assert nt == "watch_triggered" + assert "HIT" in text + + +def test_cancel_reports_already_completed_for_auto_cancelled_watch(tmp_db: str) -> None: + """After a watch fires and auto-cancels, the cancel-by-name path + should report 'already completed' rather than 'not found'. + + Pre-fix, ``_exec_watch`` cancel looked the watch up via + ``list_watches_for_ws`` which filters ``active==1``, so a recently- + auto-cancelled row was invisible and the model got the same + 'not found' message it would for a typo'd name. Post-fix the + cancel path uses ``find_watch_by_name`` (no active filter) and + branches on ``row["active"]``. + """ + session = _make_session() + storage = get_storage() + + storage.create_watch( + watch_id="w-completed-1", + ws_id=session._ws_id, + node_id="test-node", + name="completed-watch", + command="echo x", + interval_secs=10.0, + stop_on=None, + max_polls=100, + created_by="model", + next_poll="", + ) + # Simulate the post-fire state. + storage.update_watch("w-completed-1", active=False, next_poll="") + + _call_id, msg = session._exec_watch( + {"call_id": "c1", "action": "cancel", "watch_name": "completed-watch"} + ) + + assert "not found" not in msg.lower() + assert "completed" in msg.lower() + + +def test_cancel_reports_not_found_for_unknown_watch(tmp_db: str) -> None: + """The 'not found' message still applies when the watch genuinely + does not exist — make sure the new ``find_watch_by_name`` path + didn't accidentally turn every cancel into 'already completed'. + """ + session = _make_session() + + _call_id, msg = session._exec_watch( + {"call_id": "c1", "action": "cancel", "watch_name": "ghost-watch"} + ) + + assert "not found" in msg.lower() + + +def test_poll_watch_retry_deactivate_after_update_watch_failure( + tmp_db: str, monkeypatch: pytest.MonkeyPatch +) -> None: + """``_terminal_dispatched`` lifecycle: if ``update_watch`` raises + AFTER ``_dispatch_result`` shipped the reminder for a terminal + fire, the next ``_poll_watch`` tick MUST retry the row write + (so the row stops appearing in ``list_due_watches``) and MUST NOT + re-dispatch the reminder the model already saw. + + This is the keystone path that prevents duplicate-fire under + transient storage failure. Pre-this-test, the entire branch was + unexercised. + """ + session = _make_session() + storage = get_storage() + runner = WatchRunner(storage=storage, node_id="test-node") + session.set_watch_runner(runner) + + watch_id = "w-retry-1" + storage.create_watch( + watch_id=watch_id, + ws_id=session._ws_id, + node_id="test-node", + name="retry-watch", + command="echo HIT", + interval_secs=10.0, + stop_on='"HIT" in output', + max_polls=100, + created_by="model", + next_poll="1970-01-01T00:00:00", + ) + + enqueue_calls: list[tuple[str, str]] = [] + real_enqueue = session._nudge_queue.enqueue + + def _spy_enqueue(*args: Any, **kwargs: Any) -> None: + enqueue_calls.append((args[0], args[1][:32])) + return real_enqueue(*args, **kwargs) + + monkeypatch.setattr(session._nudge_queue, "enqueue", _spy_enqueue) + + # Stage 1 — first poll. ``update_watch`` raises AFTER dispatch. + real_update = storage.update_watch + update_raise = {"armed": True} + + def _failing_update(wid: str, **fields: Any) -> bool: + if update_raise["armed"]: + raise RuntimeError("simulated transient storage failure") + return real_update(wid, **fields) + + monkeypatch.setattr(storage, "update_watch", _failing_update) + + due = storage.list_due_watches("2099-01-01T00:00:00") + matching = [r for r in due if r["watch_id"] == watch_id] + assert len(matching) == 1 + # ``_poll_watch`` doesn't catch the storage error; the outer + # ``_tick`` would log it. Suppress here so the test owns the + # boundary and continues to its assertions. + with contextlib.suppress(RuntimeError): + runner._poll_watch(matching[0]) + + # Dispatch ran exactly once and the watch_id sits in the + # terminal-dispatched set awaiting retry. + assert len(enqueue_calls) == 1 + assert enqueue_calls[0][0] == "watch_triggered" + assert watch_id in runner._terminal_dispatched + + # The row is still active=1 because update_watch raised. It + # would re-appear in list_due_watches on the next tick. + assert storage.is_watch_active(watch_id) is True + + # Stage 2 — second poll. Storage now succeeds; retry-deactivate + # branch must commit active=False WITHOUT re-dispatching. + update_raise["armed"] = False + + due = storage.list_due_watches("2099-01-01T00:00:00") + matching = [r for r in due if r["watch_id"] == watch_id] + assert len(matching) == 1 + runner._poll_watch(matching[0]) + + # Exactly one dispatch in total — the retry path took the + # short-circuit return at the top of _poll_watch. + assert len(enqueue_calls) == 1, f"retry-deactivate must not re-dispatch; got {enqueue_calls!r}" + # Row is now inactive (the retry path's update_watch landed). + assert storage.is_watch_active(watch_id) is False + # Set is cleared so future watches with the same id (unlikely) / + # process memory doesn't accumulate. + assert watch_id not in runner._terminal_dispatched + + +def test_cancel_clears_pending_terminal_dispatched_entry( + tmp_db: str, monkeypatch: pytest.MonkeyPatch +) -> None: + """If ``update_watch`` raised after dispatch, leaving a pending + entry in ``_terminal_dispatched``, and the user then cancels the + watch out-of-band, the retry-deactivate branch never gets to run + (the cancel sets ``next_poll=""`` which removes the row from + ``list_due_watches``). The cancel path itself must discard the + pending entry; otherwise the runner leaks ``watch_id``s for the + process lifetime. + """ + session = _make_session() + storage = get_storage() + runner = WatchRunner(storage=storage, node_id="test-node") + session.set_watch_runner(runner) + + watch_id = "w-leak-1" + storage.create_watch( + watch_id=watch_id, + ws_id=session._ws_id, + node_id="test-node", + name="leak-watch", + command="echo x", + interval_secs=10.0, + stop_on=None, + max_polls=100, + created_by="model", + next_poll="", + ) + # Simulate: dispatch shipped, update_watch raised, watch_id sits + # in the runner's pending set. + with runner._terminal_dispatched_lock: + runner._terminal_dispatched.add(watch_id) + + # User cancels. Because the cancel writes active=False, next_poll="", + # the row leaves list_due_watches and the runner's retry-deactivate + # branch never executes for it. The cancel must discard the entry. + storage.update_watch(watch_id, active=False, next_poll="") + session._exec_watch({"call_id": "c1", "action": "cancel", "watch_name": "leak-watch"}) + + assert watch_id not in runner._terminal_dispatched diff --git a/tests/test_watch_storage.py b/tests/test_watch_storage.py index 84e106df..4b1d6a4a 100644 --- a/tests/test_watch_storage.py +++ b/tests/test_watch_storage.py @@ -2,6 +2,10 @@ from __future__ import annotations +import sqlalchemy as sa + +from turnstone.core.storage._schema import watches as watches_table + def _make_watch_kwargs(**overrides): """Build default kwargs for create_watch.""" @@ -101,6 +105,91 @@ class TestWatchListQueries: db.update_watch("w1", active=False) assert db.list_watches_for_ws("ws-1") == [] + def test_find_by_name_returns_inactive(self, db): + """``find_watch_by_name`` ignores the active filter — that is + what lets the cancel-by-name UX distinguish 'already completed' + from 'no such watch.' + """ + db.create_watch(**_make_watch_kwargs(watch_id="w1", ws_id="ws-1", name="completed")) + db.update_watch("w1", active=False) + + row = db.find_watch_by_name("ws-1", "completed") + assert row is not None + assert row["watch_id"] == "w1" + assert not row["active"] + + def test_find_by_name_matches_watch_id_prefix(self, db): + db.create_watch(**_make_watch_kwargs(watch_id="abcdef123", ws_id="ws-1", name="x")) + row = db.find_watch_by_name("ws-1", "abc") + assert row is not None + assert row["watch_id"] == "abcdef123" + + def test_find_by_name_scoped_to_ws(self, db): + db.create_watch(**_make_watch_kwargs(watch_id="w1", ws_id="ws-1", name="shared")) + db.create_watch(**_make_watch_kwargs(watch_id="w2", ws_id="ws-2", name="shared")) + + row = db.find_watch_by_name("ws-1", "shared") + assert row is not None + assert row["watch_id"] == "w1" + + def test_find_by_name_returns_none_when_missing(self, db): + assert db.find_watch_by_name("ws-1", "ghost") is None + + def test_find_by_name_empty_input_returns_none(self, db): + db.create_watch(**_make_watch_kwargs(watch_id="w1", ws_id="ws-1", name="x")) + assert db.find_watch_by_name("ws-1", "") is None + + def test_find_by_name_treats_percent_as_literal(self, db): + """A model-supplied '%' must NOT match arbitrary watch_ids. + + Pre-escape, ``watch_id.like(f"{name_or_prefix}%")`` would + interpret '%' as 'match anything' and pick up the first row in + the workstream regardless of name. + """ + db.create_watch(**_make_watch_kwargs(watch_id="w1", ws_id="ws-1", name="real-watch")) + assert db.find_watch_by_name("ws-1", "%") is None + + def test_find_by_name_treats_underscore_as_literal(self, db): + """Same as the '%' case for the single-char LIKE wildcard.""" + db.create_watch(**_make_watch_kwargs(watch_id="abcd", ws_id="ws-1", name="real-watch")) + # '_' would otherwise match any single char, picking up + # watch_ids beginning with 'a', 'b', etc. + assert db.find_watch_by_name("ws-1", "_") is None + + def test_find_by_name_prefers_active_over_newer_inactive(self, db): + """If a same-name pair exists where the inactive row is NEWER + than the active row, find_watch_by_name must still return the + active row. Pre-fix the query was ``ORDER BY created DESC + LIMIT 1`` — which would return the newer inactive row and + cause the cancel UX to report 'already completed' for a name + whose live row is still polling. + + Reachable in practice because storage allows out-of-band + writes (e.g. ``delete_watches_for_ws`` cleanup followed by + re-create, an admin manually flipping ``active``, or test + scaffolding) that bypass the create-time duplicate-name + guard. + """ + # Older active watch. + db.create_watch(**_make_watch_kwargs(watch_id="w-active", ws_id="ws-1", name="recurring")) + # Newer inactive watch with the same name. ``create_watch`` + # stamps ``created`` to ``now`` at second resolution, so we + # bypass the API to give the inactive row a deterministically + # later timestamp. + db.create_watch(**_make_watch_kwargs(watch_id="w-inactive", ws_id="ws-1", name="recurring")) + with db._conn() as conn: + conn.execute( + sa.update(watches_table) + .where(watches_table.c.watch_id == "w-inactive") + .values(active=0, next_poll="", created="2099-01-01T00:00:00") + ) + conn.commit() + + row = db.find_watch_by_name("ws-1", "recurring") + assert row is not None + assert row["watch_id"] == "w-active" + assert row["active"] + def test_list_for_node(self, db): db.create_watch(**_make_watch_kwargs(watch_id="w1", node_id="n1")) db.create_watch(**_make_watch_kwargs(watch_id="w2", node_id="n1")) diff --git a/turnstone/core/nudge_queue.py b/turnstone/core/nudge_queue.py index 6e217769..997ddac2 100644 --- a/turnstone/core/nudge_queue.py +++ b/turnstone/core/nudge_queue.py @@ -19,11 +19,11 @@ Channels: Drain preserves FIFO order; non-matching entries stay queued. Each entry can carry an optional ``valid_until`` predicate that drain evaluates outside the queue lock; entries whose predicate returns -``False`` (or raises) are silently dropped without delivery — used by -producers whose payload becomes stale if the underlying state changes -between enqueue and drain (e.g. ``idle_children`` re-checks the active -child set, dropping the nudge if every child finished while the queue -sat). Operations are atomic under an internal :class:`threading.Lock`. +``False`` are dropped (logged at ``info`` — normal lifecycle outcome, +e.g. ``idle_children`` after every child closed) and entries whose +predicate raises are dropped (logged at ``warning`` with ``exc_info`` +— a misbehaving predicate). Operations are atomic under an internal +:class:`threading.Lock`. """ from __future__ import annotations @@ -32,9 +32,13 @@ import threading from collections import deque from typing import TYPE_CHECKING, Any, Literal, NamedTuple +from turnstone.core.log import get_logger + if TYPE_CHECKING: from collections.abc import Callable +log = get_logger(__name__) + Channel = Literal["user", "tool", "any"] _VALID_CHANNELS: frozenset[str] = frozenset({"user", "tool", "any"}) @@ -139,7 +143,12 @@ class NudgeQueue: self._items = kept # Predicates evaluate outside the lock — they may do storage # I/O or other work that shouldn't block other producers / - # the drain consumer's other queues. + # the drain consumer's other queues. Drop-level distinction: + # a ``False`` return is a normal lifecycle outcome (the + # producer's snapshot is stale — e.g. ``idle_children`` after + # every child closed) and logs at ``info``; a raised exception + # is a wiring bug (predicate is misbehaving) and stays at + # ``warning`` with ``exc_info`` so the traceback surfaces. out: list[tuple[str, str, dict[str, Any] | None]] = [] for entry in candidates: if entry.valid_until is None: @@ -148,11 +157,27 @@ class NudgeQueue: try: if entry.valid_until(): out.append((entry.nudge_type, entry.text, entry.metadata)) + continue + log.info( + "nudge_queue.predicate_dropped", + extra={ + "nudge_type": entry.nudge_type, + "channel": entry.channel, + "reason": "predicate_false", + "text_len": len(entry.text), + }, + ) except Exception: - # Predicate raising is treated as "no longer valid" — - # drop silently rather than letting one bad predicate - # poison the whole drain batch. - pass + log.warning( + "nudge_queue.predicate_dropped", + extra={ + "nudge_type": entry.nudge_type, + "channel": entry.channel, + "reason": "predicate_raised", + "text_len": len(entry.text), + }, + exc_info=True, + ) return out def __len__(self) -> int: diff --git a/turnstone/core/session.py b/turnstone/core/session.py index cae05088..79429087 100644 --- a/turnstone/core/session.py +++ b/turnstone/core/session.py @@ -1670,13 +1670,17 @@ class ChatSession: The closure carries: - a soft cap on per-session ``"watch_triggered"`` depth via :data:`_WATCH_QUEUE_SOFT_CAP` + drop-oldest-on-saturation. - - a ``valid_until`` predicate that re-checks - ``storage.is_watch_active(watch_id)`` at drain time so a - cancelled watch's last splat doesn't ride out a future wake. - producer-side :func:`sanitize_payload` over the whole formatted message so steering-vector / control-char payloads sourced from arbitrary shell output can't tamper with the envelope at interpolation time. + + No ``valid_until`` predicate is wired: ``WatchRunner._poll_watch`` + commits ``active=False`` for terminal fires right after dispatch + returns, and an ``is_watch_active`` predicate would race that + write at drain time and drop the fire the model was meant to see. + A user-cancelled watch's last splat is informative (the reminder + carries ``is_final=True``), not stale-noise to suppress. """ self._watch_runner = runner nudge_queue = self._nudge_queue @@ -1707,18 +1711,6 @@ class ChatSession: _WATCH_QUEUE_SOFT_CAP, ) - def _still_active() -> bool: - # Re-checked at drain time outside the queue lock — if - # the watch was cancelled between fire and drain, the - # entry gets dropped silently rather than splicing a - # stale result onto the user's next turn. Single-column - # ``is_watch_active`` avoids the full-row marshal of - # ``get_watch`` on this hot path. - try: - return get_storage().is_watch_active(watch_id) - except Exception: - return False - def _maybe_sanitize(v: Any) -> Any: return sanitize_payload(v) if isinstance(v, str) else v @@ -1731,7 +1723,6 @@ class ChatSession: "watch_triggered", sanitized, "any", - valid_until=_still_active, metadata=metadata or None, ) @@ -10473,16 +10464,23 @@ class ChatSession: msg = "Error: storage unavailable" self._report_tool_result(call_id, "watch", msg, is_error=True) return call_id, msg - watches = storage.list_watches_for_ws(self._ws_id) - target = None - for w in watches: - if w["name"] == name or w["watch_id"].startswith(name): - target = w - break + target = storage.find_watch_by_name(self._ws_id, name) if target is None: msg = f'Watch "{name}" not found.' self._report_tool_result(call_id, "watch", msg, is_error=True) return call_id, msg + # In either branch below the row leaves ``list_due_watches`` + # view (already-inactive or just-cancelled with empty + # next_poll), so the runner's retry-deactivate branch will + # never reclaim a pending ``_terminal_dispatched`` entry. + # Clear it here to bound the lifetime of any leftover from + # a previous dispatch-then-failed-row-write. + if self._watch_runner is not None: + self._watch_runner.forget_terminal_dispatched(target["watch_id"]) + if not target["active"]: + msg = f'Watch "{target["name"]}" already completed (auto-cancelled).' + self._report_tool_result(call_id, "watch", msg) + return call_id, msg storage.update_watch(target["watch_id"], active=False, next_poll="") msg = f'Watch "{target["name"]}" cancelled.' self._report_tool_result(call_id, "watch", msg) diff --git a/turnstone/core/storage/_postgresql.py b/turnstone/core/storage/_postgresql.py index af9e992a..0ca64bf8 100644 --- a/turnstone/core/storage/_postgresql.py +++ b/turnstone/core/storage/_postgresql.py @@ -72,6 +72,9 @@ from turnstone.core.storage._schema import ( from turnstone.core.storage._utils import ( HEURISTIC_RULE_MUTABLE as _HEURISTIC_RULE_MUTABLE, ) +from turnstone.core.storage._utils import ( + LIKE_ESCAPE as _LIKE_ESCAPE, +) from turnstone.core.storage._utils import ( MCP_SERVER_MUTABLE as _MCP_SERVER_MUTABLE, ) @@ -102,6 +105,9 @@ from turnstone.core.storage._utils import ( from turnstone.core.storage._utils import ( VERDICT_MUTABLE as _VERDICT_MUTABLE, ) +from turnstone.core.storage._utils import ( + escape_like as _escape_like, +) from turnstone.core.storage._utils import ( normalize_search_terms as _normalize_search_terms, ) @@ -120,11 +126,6 @@ from turnstone.core.workstream import BULK_CLOSE_STATE_VALUES, WorkstreamKind log = get_logger(__name__) -def _escape_ilike(s: str) -> str: - """Escape ILIKE metacharacters for use with ESCAPE '\\\\'.""" - return s.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_") - - def _resolve_pg_listen_url(override: str, sqlalchemy_url: str) -> str: """Resolve the URL used by the dedicated LISTEN connection. @@ -1835,6 +1836,33 @@ class PostgreSQLBackend: ).fetchall() return [dict(r._mapping) for r in rows] + def find_watch_by_name(self, ws_id: str, name_or_prefix: str) -> dict[str, Any] | None: + + if not name_or_prefix: + return None + like_pattern = _escape_like(name_or_prefix) + "%" + with self._conn() as conn: + row = conn.execute( + sa.select(watches) + .where( + (watches.c.ws_id == ws_id) + & ( + (watches.c.name == name_or_prefix) + | watches.c.watch_id.like(like_pattern, escape=_LIKE_ESCAPE) + ) + ) + # Active rows win over inactive ones with the same name. + # _prepare_watch's duplicate-name guard filters active=1, + # so a model can recreate a name after the previous one + # auto-cancelled; a cancel-by-name request on the live + # row must not be shadowed by the older completed row. + .order_by(watches.c.active.desc(), watches.c.created.desc()) + .limit(1) + ).fetchone() + if row is None: + return None + return dict(row._mapping) + def list_watches_for_node(self, node_id: str) -> list[dict[str, Any]]: with self._conn() as conn: @@ -3671,7 +3699,7 @@ class PostgreSQLBackend: clauses = [] params: dict[str, str] = {} for i, t in enumerate(terms): - escaped = _escape_ilike(t) + escaped = _escape_like(t) clauses.append( f"(name ILIKE :n{i} ESCAPE '\\' " f"OR description ILIKE :d{i} ESCAPE '\\' " @@ -3750,7 +3778,7 @@ class PostgreSQLBackend: scope_clauses, params = self._build_scope_or_clause(scopes) term_clauses = [] for i, t in enumerate(terms): - escaped = _escape_ilike(t) + escaped = _escape_like(t) term_clauses.append( f"(name ILIKE :n{i} ESCAPE '\\' " f"OR description ILIKE :d{i} ESCAPE '\\' " diff --git a/turnstone/core/storage/_protocol.py b/turnstone/core/storage/_protocol.py index 3e9ed598..d916b4c9 100644 --- a/turnstone/core/storage/_protocol.py +++ b/turnstone/core/storage/_protocol.py @@ -1003,6 +1003,23 @@ class StorageBackend(Protocol): """Return active watches for a workstream, ordered by created DESC.""" ... + def find_watch_by_name(self, ws_id: str, name_or_prefix: str) -> dict[str, Any] | None: + """Return a watch in ``ws_id`` whose ``name`` matches + ``name_or_prefix`` exactly, or whose ``watch_id`` starts with it. + + Unlike :meth:`list_watches_for_ws` this DOES NOT filter on the + ``active`` flag — callers can inspect ``row["active"]`` to + distinguish a still-running watch from one that fired and + auto-cancelled. Returns ``None`` if no match. + + When multiple rows match, prefers active rows over inactive + ones, then most-recently-created. Without the active + preference, a recreated-after-completion name would let the + older inactive row shadow the new active one in the cancel + path. + """ + ... + def list_watches_for_node(self, node_id: str) -> list[dict[str, Any]]: """Return all active watches on a node, ordered by created DESC.""" ... diff --git a/turnstone/core/storage/_sqlite.py b/turnstone/core/storage/_sqlite.py index a392d3c2..a66c469a 100644 --- a/turnstone/core/storage/_sqlite.py +++ b/turnstone/core/storage/_sqlite.py @@ -72,6 +72,9 @@ from turnstone.core.storage._schema import ( from turnstone.core.storage._utils import ( HEURISTIC_RULE_MUTABLE as _HEURISTIC_RULE_MUTABLE, ) +from turnstone.core.storage._utils import ( + LIKE_ESCAPE as _LIKE_ESCAPE, +) from turnstone.core.storage._utils import ( MCP_SERVER_MUTABLE as _MCP_SERVER_MUTABLE, ) @@ -102,6 +105,9 @@ from turnstone.core.storage._utils import ( from turnstone.core.storage._utils import ( VERDICT_MUTABLE as _VERDICT_MUTABLE, ) +from turnstone.core.storage._utils import ( + escape_like as _escape_like, +) from turnstone.core.storage._utils import ( normalize_search_terms as _normalize_search_terms, ) @@ -120,11 +126,6 @@ from turnstone.core.workstream import BULK_CLOSE_STATE_VALUES, WorkstreamKind log = get_logger(__name__) -def _escape_like(s: str) -> str: - """Escape LIKE metacharacters for use with ESCAPE '\\\\'.""" - return s.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_") - - def _fts5_query(query: str) -> str: """Convert a plain search string into a safe FTS5 query.""" terms = query.split() @@ -1976,6 +1977,33 @@ class SQLiteBackend: ).fetchall() return [dict(r._mapping) for r in rows] + def find_watch_by_name(self, ws_id: str, name_or_prefix: str) -> dict[str, Any] | None: + + if not name_or_prefix: + return None + like_pattern = _escape_like(name_or_prefix) + "%" + with self._conn() as conn: + row = conn.execute( + sa.select(watches) + .where( + (watches.c.ws_id == ws_id) + & ( + (watches.c.name == name_or_prefix) + | watches.c.watch_id.like(like_pattern, escape=_LIKE_ESCAPE) + ) + ) + # Active rows win over inactive ones with the same name. + # _prepare_watch's duplicate-name guard filters active=1, + # so a model can recreate a name after the previous one + # auto-cancelled; a cancel-by-name request on the live + # row must not be shadowed by the older completed row. + .order_by(watches.c.active.desc(), watches.c.created.desc()) + .limit(1) + ).fetchone() + if row is None: + return None + return dict(row._mapping) + def list_watches_for_node(self, node_id: str) -> list[dict[str, Any]]: with self._conn() as conn: diff --git a/turnstone/core/storage/_utils.py b/turnstone/core/storage/_utils.py index b9cf6da2..cc4e9e90 100644 --- a/turnstone/core/storage/_utils.py +++ b/turnstone/core/storage/_utils.py @@ -98,6 +98,33 @@ def sanitize_text(value: str | None) -> str | None: return value +# --------------------------------------------------------------------------- +# SQL LIKE escaping +# --------------------------------------------------------------------------- + +# Use a non-default escape character so callers passing the result to +# SQLAlchemy's ``.like(pattern, escape=LIKE_ESCAPE)`` get the same +# semantics on SQLite and PostgreSQL. ``\`` is the SQL standard. +LIKE_ESCAPE = "\\" + + +def escape_like(value: str) -> str: + """Escape ``%`` and ``_`` (and the escape character itself) so the + string can be safely embedded in a SQL ``LIKE`` pattern. + + Pair with ``column.like(escape_like(prefix) + "%", escape=LIKE_ESCAPE)`` + to do a true prefix match against caller-supplied input. Without + this, untrusted text containing ``%`` or ``_`` is interpreted as a + wildcard — e.g. a model-supplied watch name of ``"%"`` would match + every row in the queried partition. + """ + return ( + value.replace(LIKE_ESCAPE, LIKE_ESCAPE * 2) + .replace("%", LIKE_ESCAPE + "%") + .replace("_", LIKE_ESCAPE + "_") + ) + + # --------------------------------------------------------------------------- # Row helper # --------------------------------------------------------------------------- diff --git a/turnstone/core/watch.py b/turnstone/core/watch.py index a8fb6356..f115de26 100644 --- a/turnstone/core/watch.py +++ b/turnstone/core/watch.py @@ -289,6 +289,17 @@ class WatchRunner: self._dispatch_fns: dict[str, Callable[[dict[str, Any], str], None]] = {} self._dispatch_lock = threading.Lock() + # Watch ids whose terminal reminder has already been dispatched + # but whose row write has not yet been confirmed. Populated + # between ``_dispatch_result`` and ``update_watch`` in + # :meth:`_poll_watch`; on a subsequent tick the same row will + # still appear in ``list_due_watches`` (active=1, next_poll + # unchanged) — the guard at the top of ``_poll_watch`` retries + # the row write WITHOUT re-dispatching. Bounded by transient + # storage failure depth (~MAX_WATCHES_PER_WS × num_ws). + self._terminal_dispatched: set[str] = set() + self._terminal_dispatched_lock = threading.Lock() + self._stop_event = threading.Event() self._thread: threading.Thread | None = None @@ -314,14 +325,18 @@ class WatchRunner: def set_dispatch_fn(self, ws_id: str, fn: Callable[[dict[str, Any], str], None]) -> None: """Register a per-workstream dispatch fn. - The fn signature is ``(reminder, watch_id)`` — the runner passes - the originating ``watch_id`` so dispatch closures can capture - per-watch metadata (e.g. a ``valid_until`` predicate that - re-checks ``storage.is_watch_active(watch_id)`` before - delivering a stale entry). ``reminder`` is the structured dict - returned by :func:`build_watch_reminder` — ``text`` carries the - formatted body, the remaining fields ride as queue-entry - metadata so the frontend can render a ``.msg.watch-result`` card. + The fn signature is ``(reminder, watch_id)``. ``reminder`` is + the structured dict returned by :func:`build_watch_reminder` — + ``text`` carries the formatted body, the remaining fields ride + as queue-entry metadata so the frontend can render a + ``.msg.watch-result`` card. ``watch_id`` is passed for + closures that need per-watch metadata in their queue plumbing + (e.g. correlating a fire back to the originating row in logs); + do NOT use it to gate delivery against + ``storage.is_watch_active(watch_id)`` — see + :meth:`ChatSession.set_watch_runner` for why that pattern + races :meth:`_poll_watch`'s commit of ``active=False`` and + drops fires the model was meant to see. """ with self._dispatch_lock: self._dispatch_fns[ws_id] = fn @@ -339,6 +354,22 @@ class WatchRunner: with self._dispatch_lock: return self._dispatch_fns.get(ws_id) + def forget_terminal_dispatched(self, watch_id: str) -> None: + """Discard ``watch_id`` from the pending-terminal-dispatched + set if present. Called by paths that take a watch out of + :meth:`StorageBackend.list_due_watches` view independent of + the runner's own poll (most importantly the user-cancel path + in :meth:`ChatSession._exec_watch`). Without this, a + ``_poll_watch`` whose row write failed AFTER dispatch would + leak ``watch_id`` in ``_terminal_dispatched`` indefinitely — + the user-cancel writes ``next_poll=''`` which excludes the + row from ``list_due_watches``, so the retry-deactivate branch + at the top of :meth:`_poll_watch` never fires to clear the + entry. + """ + with self._terminal_dispatched_lock: + self._terminal_dispatched.discard(watch_id) + # -- Main loop ----------------------------------------------------------- def _run(self) -> None: @@ -383,6 +414,21 @@ class WatchRunner: prev_output = watch_row.get("last_output") created = watch_row.get("created", "") + # Re-poll of a row whose terminal reminder already shipped but + # whose ``active=False`` write didn't land — retry just the row + # write so the row stops appearing in ``list_due_watches``; do + # NOT re-dispatch the reminder, which the model already saw. + with self._terminal_dispatched_lock: + already_dispatched = watch_id in self._terminal_dispatched + if already_dispatched: + try: + self._storage.update_watch(watch_id, active=False, next_poll="") + with self._terminal_dispatched_lock: + self._terminal_dispatched.discard(watch_id) + except Exception: + log.exception("watch_runner.retry_deactivate_failed", extra={"watch_id": watch_id}) + return + # Safety check blocked = is_command_blocked(command) if blocked: @@ -416,22 +462,17 @@ class WatchRunner: now = datetime.now(UTC) now_str = now.strftime("%Y-%m-%dT%H:%M:%S") - # Update DB - update_fields: dict[str, Any] = { - "poll_count": poll_count, - "last_output": output, - "last_exit_code": exit_code, - "last_poll": now_str, - } - if is_final: - update_fields["active"] = False - update_fields["next_poll"] = "" - else: - next_poll = now + timedelta(seconds=watch_row["interval_secs"]) - update_fields["next_poll"] = next_poll.strftime("%Y-%m-%dT%H:%M:%S") - self._storage.update_watch(watch_id, **update_fields) - - # Dispatch result if condition fired or final + # Dispatch before committing the row update. Belt-and-braces + # given the rest of the fix (closure no longer wires a + # ``valid_until`` predicate, cancel-by-name uses + # :meth:`find_watch_by_name` which ignores the ``active`` + # filter): either order would deliver the reminder today, but + # this ordering preserves the invariant against re-wiring an + # ``is_watch_active`` predicate or adding a new + # ``active``-filtered read on this hot path. Combined with the + # ``_terminal_dispatched`` guard above it also bounds the + # duplicate-fire blast radius if the row write fails after the + # reminder shipped. if fired or is_final: # Compute elapsed from created time elapsed_secs = 0.0 @@ -454,6 +495,33 @@ class WatchRunner: reason=reason, ) self._dispatch_result(ws_id, reminder, watch_id) + if is_final: + # Mark BEFORE the row write so a raise below routes the + # next tick into the retry-deactivate branch instead of + # re-firing the reminder. + with self._terminal_dispatched_lock: + self._terminal_dispatched.add(watch_id) + + # Update DB + update_fields: dict[str, Any] = { + "poll_count": poll_count, + "last_output": output, + "last_exit_code": exit_code, + "last_poll": now_str, + } + if is_final: + update_fields["active"] = False + update_fields["next_poll"] = "" + else: + next_poll = now + timedelta(seconds=watch_row["interval_secs"]) + update_fields["next_poll"] = next_poll.strftime("%Y-%m-%dT%H:%M:%S") + self._storage.update_watch(watch_id, **update_fields) + + if is_final: + # Row write committed; the retry-deactivate branch will + # never be reached for this watch_id. + with self._terminal_dispatched_lock: + self._terminal_dispatched.discard(watch_id) log.debug( "watch_runner.polled",