feat(muse): harden choice watcher concurrency, add rules dictionary, resume pool, and session bind

This commit is contained in:
operator
2026-10-07 01:50:18 +00:00
parent 90f4ef661a
commit 094bd7d691
10 changed files with 2166 additions and 40 deletions
+353 -38
View File
@@ -28,6 +28,7 @@ Usage:
"""
import argparse
import fcntl
import hashlib
import json
import os
@@ -53,6 +54,7 @@ POLL_INTERVAL = 0.5
STABILITY_POLLS = 2
MAX_ANSWERS_PER_HOUR = 20
ANSWERED_TTL_SECONDS = 600 # identical prompt back after 10m => stuck, allow one recovery answer
CLAIM_TTL_SECONDS = 30 # concurrent-claim window: bounds wedge if winner dies pre-send
RULES_FILE = os.path.join(REPO_ROOT, "muse-choices-rules.json")
HOLD_WINDOW_SECONDS = 120 # D3: short hold window, then expire to approve
NEGATIVE_KEYS = {"muse-approval": "2", "yn": "n"} # D2 deny keys
@@ -163,6 +165,13 @@ def logfile_for(socket_path, pane_id):
)
def answered_file_for(socket_path, pane_id):
"""On-disk answered-sig store: restarts must not re-answer prompts."""
return os.path.join(
STATE_DIR, "%s-%s-%s.answered.json" % (FILE_PREFIX, slug_socket(socket_path), clean_pane(pane_id))
)
class WatcherLog:
"""Never-raising JSON-lines logger with best-effort rotation."""
@@ -624,21 +633,49 @@ def _child_pids(pid):
def muse_approval_flags(cmd_argv):
"""Approval posture of a muse argv: {"auto_approve", "flags"}."""
"""Approval posture of a muse argv.
Returns {"auto_approve", "flags", "profile", "mode", "bypass"}.
profile is the --permission-profile id (built-ins :read-only,
:standard, :unrestricted) or None; mode is "yolo" for --yolo, the
profile id when one was passed, else "default"; bypass is True
when tool-approval prompts are disabled at launch (yolo,
--disable-approval, or --approval-mode=never), in which case the
pane shows no approval dialogs (questions may still render, so the
watcher stays active).
"""
flags = []
argv = cmd_argv or []
if "--yolo" in argv:
flags.append("yolo")
if "--disable-approval" in argv:
flags.append("disable-approval")
if "--disable-sandbox" in argv:
flags.append("disable-sandbox")
if "--trust-workspace" in argv:
flags.append("trust-workspace")
profile = None
for i, arg in enumerate(argv):
if arg == "--approval-mode" and i + 1 < len(argv):
flags.append("approval-mode=%s" % argv[i + 1])
elif arg.startswith("--approval-mode="):
flags.append("approval-mode=%s" % arg.split("=", 1)[1])
elif arg == "--permission-profile" and i + 1 < len(argv):
profile = argv[i + 1]
flags.append("permission-profile=%s" % profile)
elif arg.startswith("--permission-profile="):
profile = arg.split("=", 1)[1]
flags.append("permission-profile=%s" % profile)
auto = ("yolo" in flags or "disable-approval" in flags
or "approval-mode=never" in flags)
return {"auto_approve": auto, "flags": flags}
if "yolo" in flags:
mode = "yolo"
elif profile is not None:
mode = profile
else:
mode = "default"
return {"auto_approve": auto, "flags": flags, "profile": profile,
"mode": mode, "bypass": auto}
def launch_opt_out(cmd_argv):
@@ -688,9 +725,11 @@ def pane_muse_argv(socket_path, pane_id):
# Minimum pane geometry for reliable approval rendering. Empirically
# derived: a 35x7 tile drops approval text the matcher needs, while
# 35x35/36x35/71x27 panes answer cleanly. Below either bound the pane
# is flagged squeezed (see `box runtime layout` / `spread`).
MIN_APPROVAL_WIDTH = 40
# 35x35/36x35/71x27 panes answer cleanly (width 35 works when tall
# enough; capture uses -J so wrapping is width-independent). Below
# either bound the pane is flagged squeezed (see `box runtime layout` /
# `spread`).
MIN_APPROVAL_WIDTH = 35
MIN_APPROVAL_HEIGHT = 12
# Session naming convention (see NODES.md): <node>--<role>--<id>
@@ -749,7 +788,8 @@ def runtime_rows(socket_path):
or pane_height < MIN_APPROVAL_HEIGHT))
node = node_from_session(session)
is_muse = "muse-bin" in cmd or "muse-code" in cmd
posture = {"auto_approve": None, "flags": []}
posture = {"auto_approve": None, "flags": [], "profile": None,
"mode": None, "bypass": None}
if is_muse and pane_pid:
candidates = [pane_pid] + _child_pids(pane_pid)
for cand in candidates:
@@ -763,7 +803,7 @@ def runtime_rows(socket_path):
continue
st = runtime_state(text)
match = st["match"] or {}
watcher_pid = is_running(socket_path, pane_id)
watcher_pid = watcher_alive(socket_path, pane_id)
rows.append({
"socket": socket_path, "session": session,
"window": window, "pane": pane_id, "cmd": cmd,
@@ -773,6 +813,9 @@ def runtime_rows(socket_path):
"squeezed": squeezed,
"auto_approve": posture["auto_approve"],
"approval_flags": posture["flags"],
"permission_mode": posture["mode"],
"permission_profile": posture["profile"],
"permission_bypass": posture["bypass"],
"state": st["state"],
"prompt_kind": match.get("kind"),
"prompt_key": match.get("key"),
@@ -813,14 +856,124 @@ def pane_state(socket_path, pane_id):
class WatcherState:
"""Tracks prompt stability and once-per-prompt answering."""
"""Tracks prompt stability and once-per-prompt answering.
def __init__(self):
When persist_path is set, answered sigs survive restarts (atomic
JSON store): a daemon that restarts while an answered prompt is
still visible must not answer it again (observed live: same sig
re-answered minutes later, stray "1" landing in the input box).
"""
def __init__(self, persist_path=None):
self.pending_sig = None
self.stable_count = 0
self.answered_sigs = {}
self.answer_times = []
self.last_capped_sig = None
self._persist_path = persist_path
if persist_path:
self._load_answered()
def _load_answered(self):
try:
with open(self._persist_path) as f:
data = json.load(f)
except Exception:
return
if not isinstance(data, dict):
return
now = time.time()
for sig, ts in data.items():
if (isinstance(sig, str) and isinstance(ts, (int, float))
and now - ts < ANSWERED_TTL_SECONDS):
self.answered_sigs[sig] = ts
def _save_answered(self):
if not self._persist_path:
return
try:
tmp = "%s.tmp.%d" % (self._persist_path, os.getpid())
with open(tmp, "w") as f:
json.dump(self.answered_sigs, f)
os.replace(tmp, self._persist_path)
except Exception:
pass
def refresh_answered(self):
"""Merge on-disk answered sigs into memory. Never raises.
Cross-process once-per-prompt: a peer watcher that answered
after our startup persisted its sig; re-reading before we type
suppresses the duplicate ('11' in the input box).
"""
if not self._persist_path:
return
try:
with open(self._persist_path) as f:
data = json.load(f)
except Exception:
return
if not isinstance(data, dict):
return
now = time.time()
for sig, ts in data.items():
if (isinstance(sig, str) and isinstance(ts, (int, float))
and now - ts < ANSWERED_TTL_SECONDS):
if sig not in self.answered_sigs:
self.answered_sigs[sig] = ts
def try_claim(self, sig, now=None):
"""Atomically claim a sig for this process. True iff we won.
O_CREAT|O_EXCL makes the first claimer win even when two
watchers reach the same stable prompt in the same poll window;
the loser suppresses instead of double-typing. Claims expire
after CLAIM_TTL_SECONDS (concurrent window only; stuck-dialog
recovery is governed by the answered store's longer TTL), so a
winner that dies between claim and send wedges at most briefly.
Memory-only states (no persist path) always win. Never raises.
"""
if not self._persist_path:
return True
now = time.time() if now is None else now
base = os.path.basename(self._persist_path)
directory = os.path.dirname(self._persist_path) or STATE_DIR
# Prune expired claims for this pane (best-effort).
try:
for name in os.listdir(directory):
if not name.startswith(base + ".") or not name.endswith(".claim"):
continue
p = os.path.join(directory, name)
try:
with open(p) as f:
rec = json.load(f)
ts = float(rec.get("ts", 0))
except Exception:
ts = 0
try:
if now - ts >= CLAIM_TTL_SECONDS:
os.remove(p)
except OSError:
pass
except Exception:
pass
path = "%s.%s.claim" % (self._persist_path, sig)
try:
fd = os.open(path, os.O_CREAT | os.O_EXCL | os.O_WRONLY)
except FileExistsError:
return False
except OSError:
return True # claim store unavailable: fail open, send once
try:
os.write(fd, json.dumps({"pid": os.getpid(),
"ts": now}).encode())
except Exception:
pass
try:
os.close(fd)
except Exception:
pass
return True
def _prune_times(self, now):
cutoff = now - 3600
@@ -869,6 +1022,7 @@ class WatcherState:
self.answer_times.append(now)
self.pending_sig = None
self.stable_count = 0
self._save_answered()
def _tmux(socket_path, *args, timeout=5):
@@ -887,9 +1041,15 @@ def pane_exists(socket_path, pane_id):
def capture_pane(socket_path, pane_id, history=80):
"""Capture pane text with wrapped rows joined (-J).
-J makes matching width-independent: narrow panes wrap the same
dialog onto more physical rows, which otherwise pushes cue/option
spans apart and breaks the matcher.
"""
try:
r = _tmux(socket_path, "capture-pane", "-p", "-t", pane_id, "-S", "-%d" % history,
timeout=5)
r = _tmux(socket_path, "capture-pane", "-p", "-J", "-t", pane_id,
"-S", "-%d" % history, timeout=5)
if r.returncode != 0:
return None
return r.stdout
@@ -1065,10 +1225,53 @@ def list_holds():
return out
def _peer_suppressed(state, sig, now, log):
"""True when a peer watcher already owns sig; suppress our send.
Checks the shared on-disk answered store first (peer answered
earlier and persisted), then attempts an atomic claim (peer racing
us in the same window). On suppression the pending prompt is
reset so later polls re-observe cleanly; answered-store hits are
already merged into memory, claim-loss is not recorded (peer's
send owns it, and its claim expires quickly if it dies). Never
raises; memory-only states never suppress.
"""
try:
state.refresh_answered()
except Exception:
pass
try:
ts = state.answered_sigs.get(sig)
if ts is not None and now - ts < ANSWERED_TTL_SECONDS:
try:
log.log("info", "duplicate suppressed (peer answered)",
sig=sig)
except Exception:
pass
state.pending_sig = None
state.stable_count = 0
return True
except Exception:
pass
try:
if not state.try_claim(sig, now):
try:
log.log("info", "duplicate suppressed (peer claimed)",
sig=sig)
except Exception:
pass
state.pending_sig = None
state.stable_count = 0
return True
except Exception:
pass
return False
def _check_hold(socket_path, pane_id, state, match, now, log):
"""Rule-hold gate. Returns None (fresh: evaluate rules), "released"
(hold over: approve WITHOUT re-evaluating, else expiry would
re-hold forever), "held", or "denied".
re-hold forever), "held", "denied", or "duplicate-suppressed".
The holdfile is the single source of truth (crash-safe,
box-visible): present + fresh => suppress; directive deny => deny
@@ -1093,6 +1296,8 @@ def _check_hold(socket_path, pane_id, state, match, now, log):
log.log("warn", "resolve-deny refused: D2 never denies questions",
sig=sig, kind=match["kind"])
return "held"
if _peer_suppressed(state, sig, now, log):
return "duplicate-suppressed"
ok = send_answer(socket_path, pane_id, neg, enter=True)
clear_hold(socket_path, pane_id)
state.record_answer(sig, now)
@@ -1132,8 +1337,8 @@ def _poll_once(socket_path, pane_id, state, log, dry_run=False):
now = time.time()
gate = _check_hold(socket_path, pane_id, state, match, now, log) \
if match is not None else None
if gate in ("held", "denied"):
if gate == "denied":
if gate in ("held", "denied", "duplicate-suppressed"):
if gate in ("denied", "duplicate-suppressed"):
state.pending_sig = None
state.stable_count = 0
return gate
@@ -1175,6 +1380,8 @@ def _poll_once(socket_path, pane_id, state, log, dry_run=False):
rule=rule["id"] if rule else None)
state.record_answer(match["sig"], now)
return "dry-denied"
if _peer_suppressed(state, match["sig"], now, log):
return "duplicate-suppressed"
ok = send_answer(socket_path, pane_id, neg, enter=True)
state.record_answer(match["sig"], now)
log.log("info" if ok else "error", "denied %s" % neg,
@@ -1216,6 +1423,8 @@ def _poll_once(socket_path, pane_id, state, log, dry_run=False):
options=match["options"], cue=match["cue"])
state.record_answer(match["sig"], now)
return "dry-answered"
if _peer_suppressed(state, match["sig"], now, log):
return "duplicate-suppressed"
ok = send_answer(socket_path, pane_id, key,
enter=match.get("enter", True))
state.record_answer(match["sig"], now)
@@ -1242,12 +1451,45 @@ def _poll_once(socket_path, pane_id, state, log, dry_run=False):
return "waiting"
def _log_posture(socket_path, pane_id, log):
"""Log the pane's permission posture once at watcher start.
The watcher answers with per-choice logging in every mode (that
trail is the default path and informs policy); a bypass posture
(yolo / approval disabled) additionally gets a box audit record,
since the session then makes choices outside the trail. Never
raises.
"""
try:
posture = muse_approval_flags(pane_muse_argv(socket_path, pane_id))
except Exception:
return
try:
log.log("info", "pane posture", mode=posture["mode"],
bypass=posture["bypass"], flags=posture["flags"])
except Exception:
pass
try:
if posture["bypass"] or posture["mode"] not in ("default", None):
audit("muse-choice-posture",
name="%s:%s" % (os.path.basename(socket_path), pane_id),
extra={"socket": socket_path, "pane": pane_id,
"mode": posture["mode"],
"profile": posture["profile"],
"bypass": posture["bypass"],
"flags": posture["flags"]})
except Exception:
pass
def watch_loop(socket_path, pane_id, dry_run=False):
"""Main daemon loop. Returns only when the pane is gone or signalled."""
log = WatcherLog(logfile_for(socket_path, pane_id))
state = WatcherState()
state = WatcherState(
persist_path=answered_file_for(socket_path, pane_id))
log.log("info", "watcher started", socket=socket_path, pane=pane_id,
dry_run=dry_run, pid=os.getpid())
_log_posture(socket_path, pane_id, log)
polls = 0
answers = 0
last_heartbeat = time.time()
@@ -1346,6 +1588,30 @@ def is_running(socket_path, pane_id):
return None
def watcher_alive(socket_path, pane_id):
"""Pid of the live watcher for a pane, pidfile or orphan.
is_running covers the normal pidfile case; _watch_procs catches an
identical watcher alive with a lost pidfile (/tmp cleaned under
it). Starters must consult this (not is_running alone) or they
spawn a second daemon onto the same pane (double answers). Never
raises; None means no live watcher.
"""
try:
pid = is_running(socket_path, pane_id)
if pid:
return pid
except Exception:
pass
try:
for proc in _watch_procs():
if proc["socket"] == socket_path and proc["pane"] == pane_id:
return proc["pid"]
except Exception:
pass
return None
def _daemonize():
"""Double-fork away from the controlling terminal (survives shell exit)."""
if os.fork() != 0:
@@ -1361,27 +1627,85 @@ def _daemonize():
os.close(devnull)
_PIDFILE_LOCK_FH = None
def _claim_pidfile(pidfile):
"""Claim a pidfile for this process. Returns False if another live
watcher owns it (concurrent box + timer starts must not pile up)."""
"""Claim a pidfile for this process. Returns False if another
starter holds it. Kernel-enforced: the winner holds an exclusive
flock for its whole lifetime, so concurrent box + timer starts can
never pile two daemons onto one pane (double answers, '11' in the
input box). Call _release_pidfile() on exit."""
global _PIDFILE_LOCK_FH
me = os.getpid()
try:
with open(pidfile) as f:
other = int(f.read().strip())
fh = open(pidfile, "a+")
except OSError:
return False
try:
fcntl.flock(fh, fcntl.LOCK_EX | fcntl.LOCK_NB)
except (OSError, IOError):
fh.close()
return False
try:
fh.seek(0)
content = fh.read().strip()
other = int(content) if content else None
except Exception:
other = None
if other and other != me and _pid_alive(other) and _pid_is_watcher(other):
# Live foreign owner (e.g. a daemon from before locking existed,
# or a pid recycled into a watcher): yield to it.
try:
fcntl.flock(fh, fcntl.LOCK_UN)
except Exception:
pass
fh.close()
return False
try:
with open(pidfile, "w") as f:
f.write(str(me))
fh.seek(0)
fh.truncate()
fh.write(str(me))
fh.flush()
except Exception:
pass
if _PIDFILE_LOCK_FH is not None:
try:
_PIDFILE_LOCK_FH.close()
except Exception:
pass
_PIDFILE_LOCK_FH = fh
return True
def _release_pidfile(pidfile):
"""Drop the claim taken by _claim_pidfile: remove the pidfile only
if we still own it, then release the lock. Never raises."""
global _PIDFILE_LOCK_FH
try:
with open(pidfile) as f:
owner = f.read().strip()
except Exception:
owner = ""
if owner == str(os.getpid()):
try:
os.remove(pidfile)
except OSError:
pass
if _PIDFILE_LOCK_FH is not None:
try:
fcntl.flock(_PIDFILE_LOCK_FH, fcntl.LOCK_UN)
except Exception:
pass
try:
_PIDFILE_LOCK_FH.close()
except Exception:
pass
_PIDFILE_LOCK_FH = None
def start_watcher(socket_path, pane_id, dry_run=False):
pid = is_running(socket_path, pane_id)
pid = watcher_alive(socket_path, pane_id)
if pid:
return {"ok": False, "status": "already_running", "pid": pid}
if not pane_exists(socket_path, pane_id):
@@ -1400,7 +1724,7 @@ def start_watcher(socket_path, pane_id, dry_run=False):
def stop_watcher(socket_path, pane_id, timeout=5):
pid = is_running(socket_path, pane_id)
pid = watcher_alive(socket_path, pane_id)
if not pid:
return {"ok": True, "status": "not_running"}
try:
@@ -1448,7 +1772,7 @@ def _start_detached(socket_path, pane_id, dry_run=False):
os._exit(0)
os.waitpid(pid, 0)
time.sleep(0.2)
return is_running(socket_path, pane_id) is not None
return watcher_alive(socket_path, pane_id) is not None
def start_all(dry_run=False, sockets=None):
@@ -1461,7 +1785,7 @@ def start_all(dry_run=False, sockets=None):
if not panes:
results.append({"socket": sock, "status": "no_muse_panes"})
for pane in panes:
if is_running(sock, pane):
if watcher_alive(sock, pane):
results.append({"socket": sock, "pane": pane,
"status": "already_running"})
continue
@@ -1524,7 +1848,7 @@ def reconcile(sockets=None):
if not os.path.exists(sock):
continue
for pane in muse_panes(sock):
if is_running(sock, pane):
if watcher_alive(sock, pane):
already.append("%s:%s" % (sock, pane))
continue
if _start_detached(sock, pane, dry_run=desired["dry_run"]):
@@ -1660,7 +1984,7 @@ def main(argv=None):
args = ap.parse_args(argv)
if args.cmd == "start":
existing = is_running(args.socket, args.pane)
existing = watcher_alive(args.socket, args.pane)
if existing:
print(json.dumps({"ok": True, "status": "already_running",
"pid": existing,
@@ -1674,7 +1998,7 @@ def main(argv=None):
os._exit(0)
_, status = os.waitpid(pid, 0)
time.sleep(0.3)
running = is_running(args.socket, args.pane)
running = watcher_alive(args.socket, args.pane)
print(json.dumps({"ok": running is not None, "pid": running,
"log": logfile_for(args.socket, args.pane)}))
return 0 if running else 1
@@ -1716,16 +2040,7 @@ def main(argv=None):
try:
return watch_loop(args.socket, args.pane, dry_run=args.dry_run)
finally:
try:
with open(pidfile) as f:
owner = f.read().strip()
except Exception:
owner = ""
if owner == str(os.getpid()):
try:
os.remove(pidfile)
except OSError:
pass
_release_pidfile(pidfile)
if args.cmd == "match":
text = sys.stdin.read()
print(json.dumps(find_choice_prompt(text, tail_window=args.tail_window), indent=1))