#!/usr/bin/env python3 """muse_resume_pool: per-repo, profile-aware Muse Code resume pool. Why this exists --------------- ``muse resume`` scopes its picker by workspace but is blind to muse-auth profiles, and the TUI ``/resume`` reads session logs itself and lumps every workspace into one heap. A session resumed under a different credential than the one that created it fails server-side: the continuation is cryptographically bound to the creating account, so the server rejects the resume. This tool lists only the sessions that can actually resume here and now, and guards ``resume`` calls before they fail. Data sources (read-only, no secrets) ------------------------------------ - ``~/.local/share/muse/session-index.db`` ``sessions`` table (``mode=ro``). Fallback when the index is missing: scan ``sessions/*/*/*/*/session.jsonl`` for ``runtime.session.metadata`` (workspace_root) and ``session.name.changed`` (session_name) records, mirroring ``/resume``. - ``~/.config/muse/active_profile``, ``session_profiles.json``, ``switch_history.jsonl``: profile *names* only. ``auth.json`` token bytes are never read, logged, or compared. Profile resolution mirrors ``muse-auth``: cached ``session_profiles.json`` mapping wins; otherwise the latest switch at or before session start; else the earliest switch; else the active profile as fallback. When the auth dir is unreadable (e.g. inside the Muse sandbox, which masks it), the profile is ``unknown`` and only proven mismatches are hidden/blocked. """ import argparse import glob import json import os import sqlite3 import subprocess import sys INDEX_COLUMNS = ( "session_id", "session_name", "workspace_root", "workspace_key", "provider_id", "model_id", "git_branch", "title", "first_user_prompt", "created_at_us", "updated_at_us", "prompt_count", "status", ) def default_paths(): home = os.path.expanduser("~") data_home = os.environ.get("XDG_DATA_HOME", os.path.join(home, ".local", "share")) return { "index_db": os.path.join(data_home, "muse", "session-index.db"), "sessions_dir": os.path.join(data_home, "muse", "sessions"), "config_dir": os.path.expanduser("~/.config/muse"), } def canonical_workspace(cwd=None): """Repo root for the pool: git top-level, else real cwd.""" cwd = cwd or os.getcwd() try: out = subprocess.run( ["git", "-C", cwd, "rev-parse", "--show-toplevel"], stdout=subprocess.PIPE, stderr=subprocess.DEVNULL, text=True, timeout=10, ) if out.returncode == 0 and out.stdout.strip(): return os.path.realpath(out.stdout.strip()) except Exception: pass return os.path.realpath(cwd) def load_index_rows(index_db): """Read session rows from the index (read-only). None if unavailable.""" if not os.path.exists(index_db): return None cols = ", ".join(INDEX_COLUMNS) try: uri = "file:{}?mode=ro".format(index_db.replace("?", "%3F")) conn = sqlite3.connect(uri, uri=True, timeout=5) try: conn.row_factory = sqlite3.Row cur = conn.execute( "SELECT {} FROM sessions ORDER BY " "updated_at_us DESC, created_at_us DESC, session_id ASC".format(cols) ) return [dict(r) for r in cur.fetchall()] finally: conn.close() except sqlite3.Error: return None def _scan_log_for_session(session_log): """Extract workspace/name/title from one session.jsonl (bounded read).""" workspace = None name = None title = None first_prompt = None created_at_us = None updated_at_us = None prompt_count = 0 try: with open(session_log, "r", errors="replace") as fh: for line in fh: line = line.strip() if not line: continue try: rec = json.loads(line) except ValueError: continue if "children" in rec: # retained permission frame wrapper continue rec_at = rec.get("recorded_at") if isinstance(rec_at, int): if created_at_us is None: created_at_us = rec_at updated_at_us = rec_at ptype = rec.get("payload_type", "") payload = rec.get("payload", {}) if isinstance(rec.get("payload"), dict) else {} if ptype == "runtime.session.metadata": record = payload.get("record", {}) workspace = workspace or record.get("workspace_root") elif ptype == "session.name.changed": if payload.get("new_name"): name = payload["new_name"] elif ptype == "runtime.session": event = payload.get("event", {}) if event.get("kind") == "started" and not first_prompt: prompt = event.get("prompt") or "" first_prompt = prompt[:200] title = title or prompt[:80] prompt_count += 1 except OSError: return None if workspace is None and name is None and created_at_us is None: return None session_id = os.path.basename(os.path.dirname(session_log)) return { "session_id": session_id, "session_name": name, "workspace_root": workspace, "workspace_key": workspace, "provider_id": None, "model_id": None, "git_branch": None, "title": title or "New session", "first_user_prompt": first_prompt, "created_at_us": created_at_us, "updated_at_us": updated_at_us or created_at_us, "prompt_count": prompt_count, "status": "valid", } def scan_session_logs(sessions_dir): """Fallback pool source: scan top-level session.jsonl files directly.""" pattern = os.path.join(sessions_dir, "*", "*", "*", "*", "session.jsonl") rows = [] for path in glob.glob(pattern): row = _scan_log_for_session(path) if row: rows.append(row) rows.sort( key=lambda r: ( r.get("updated_at_us") or 0, r.get("created_at_us") or 0, r.get("session_id") or "", ), reverse=True, ) return rows def load_auth_state(config_dir): """Load profile names only. Never touches auth.json token bytes.""" state = { "readable": False, "active": None, "session_profiles": {}, "switch_history": [], } if not os.path.isdir(config_dir): return state if not (os.access(config_dir, os.R_OK) and os.access(config_dir, os.X_OK)): return state state["readable"] = True try: with open(os.path.join(config_dir, "active_profile"), "r") as fh: state["active"] = fh.read().strip() or None except OSError: pass try: with open(os.path.join(config_dir, "session_profiles.json"), "r") as fh: data = json.load(fh) if isinstance(data, dict): state["session_profiles"] = { str(k): str(v) for k, v in data.items() } except (OSError, ValueError): pass try: history = [] with open(os.path.join(config_dir, "switch_history.jsonl"), "r") as fh: for line in fh: line = line.strip() if not line: continue try: entry = json.loads(line) except ValueError: continue if entry.get("profile"): history.append(entry) history.sort(key=lambda e: e.get("epoch", 0)) state["switch_history"] = history except OSError: pass return state def resolve_profile(session_id, created_at_us, auth_state): """Return (profile_or_None, source). Mirrors muse-auth resolution order.""" if not auth_state.get("readable"): return None, "unknown" cached = auth_state.get("session_profiles", {}) if session_id in cached: return cached[session_id], "cached" history = auth_state.get("switch_history", []) start_epoch = (created_at_us / 1e6) if created_at_us else None if history and start_epoch: for switch in reversed(history): if switch.get("epoch", 0) <= start_epoch: return switch.get("profile"), "history" return history[0].get("profile"), "history" if auth_state.get("active"): return auth_state["active"], "fallback" return None, "unknown" def annotate_rows(rows, auth_state): """Attach profile + source to each row (mutates and returns rows).""" for row in rows: profile, source = resolve_profile( row.get("session_id"), row.get("created_at_us"), auth_state ) row["auth_profile"] = profile row["auth_source"] = source return rows def pool_for_workspace(rows, workspace): """Exact workspace_key match (workspace_root fallback), index order kept.""" pool = [] for row in rows: key = row.get("workspace_key") or row.get("workspace_root") if key and os.path.realpath(key) == workspace: pool.append(row) return pool def split_resumable(pool, active_profile): """(resumable, blocked): only proven profile mismatches are blocked. Unknown profiles (sandboxed auth dir, no mapping/history) stay resumable but flagged, since blocking them would hide possibly valid sessions. """ resumable, blocked = [], [] for row in pool: profile = row.get("auth_profile") if active_profile and profile and profile != active_profile: blocked.append(row) else: resumable.append(row) return resumable, blocked def resolve_ref(rows, ref): """Resolve UUID / UUID prefix / session name. Returns (matches, kind).""" ref = (ref or "").strip() if not ref: return [], "empty" exact = [r for r in rows if r.get("session_id") == ref] if exact: return exact, "uuid" named = [r for r in rows if (r.get("session_name") or "") == ref] if named: return named, "name" if len(ref) >= 8: prefixed = [ r for r in rows if (r.get("session_id") or "").startswith(ref) ] if prefixed: return prefixed, "prefix" return [], "none" def check_resume(rows, ref, workspace, auth_state): """Guard decision for resuming ``ref`` from ``workspace``. Returns dict(ok=bool, reason=str, detail=str, fix=str, row=row|None). """ matches, kind = resolve_ref(rows, ref) if kind == "empty" or not matches: return { "ok": False, "reason": "unknown-session", "detail": "No session matches '{}'.".format(ref), "fix": "List this repo's pool: muse_resume_pool.py pool", "row": None, } if len(matches) > 1: ids = ", ".join(m["session_id"][:12] for m in matches[:5]) return { "ok": False, "reason": "ambiguous", "detail": "'{}' matches {} sessions: {}".format(ref, len(matches), ids), "fix": "Use a longer UUID prefix or the full session id.", "row": None, } row = matches[0] key = row.get("workspace_key") or row.get("workspace_root") if not key or os.path.realpath(key) != workspace: return { "ok": False, "reason": "wrong-workspace", "detail": "Session '{}' belongs to workspace '{}', not '{}'.".format( row.get("session_name") or row["session_id"][:12], key, workspace ), "fix": "cd '{}' first, or pick a session from this repo's pool.".format(key or "?"), "row": row, } if row.get("status") and row["status"] != "valid": return { "ok": False, "reason": "bad-status", "detail": "Session '{}' has status '{}'.".format( row.get("session_name") or row["session_id"][:12], row["status"] ), "fix": "Pick a session with status 'valid' from this repo's pool.", "row": row, } active = auth_state.get("active") profile = row.get("auth_profile") if active and profile and profile != active: return { "ok": False, "reason": "wrong-profile", "detail": "Session '{}' was created under muse-auth profile '{}' " "but the active profile is '{}'; the server would reject the " "resume (continuation is bound to the creating credential).".format( row.get("session_name") or row["session_id"][:12], profile, active ), "fix": "Run `muse-auth use {}` (outside the sandbox), then resume.".format(profile), "row": row, } detail = "Session '{}' is resumable here.".format( row.get("session_name") or row["session_id"][:12] ) if not auth_state.get("readable"): detail += " (Profile unverified: auth dir unreadable from this shell.)" elif not profile: detail += " (Profile unknown: no mapping or switch history.)" return { "ok": True, "reason": "ok", "detail": detail, "fix": "", "row": row, } def load_rows(paths): """Index first, log-scan fallback. Returns (rows, source).""" rows = load_index_rows(paths["index_db"]) if rows is not None: return rows, "index" return scan_session_logs(paths["sessions_dir"]), "log-scan" def format_pool_table(resumable, blocked, show_all): lines = [] header = "{:<18} {:<12} {:<10} {:<5} {}".format( "NAME", "SESSION", "PROFILE", "MSGS", "TITLE" ) lines.append(header) for row in resumable: flag = "?" if not row.get("auth_profile") else " " lines.append( "{:<18} {:<12} {:<10} {:<5} {}{}".format( (row.get("session_name") or "-")[:18], (row.get("session_id") or "")[:12], (row.get("auth_profile") or "?")[:10], row.get("prompt_count", 0), flag, (row.get("title") or "")[:60], ) ) if show_all: for row in blocked: lines.append( "{:<18} {:<12} {:<10} {:<5} {} [BLOCKED: profile mismatch]".format( (row.get("session_name") or "-")[:18], (row.get("session_id") or "")[:12], (row.get("auth_profile") or "?")[:10], row.get("prompt_count", 0), (row.get("title") or "")[:60], ) ) elif blocked: lines.append( "({} session(s) hidden: wrong muse-auth profile; use --all to show)".format( len(blocked) ) ) return "\n".join(lines) def cmd_pool(args, paths): workspace = os.path.realpath(args.workspace or canonical_workspace()) rows, source = load_rows(paths) auth_state = load_auth_state(paths["config_dir"]) annotate_rows(rows, auth_state) pool = pool_for_workspace(rows, workspace) resumable, blocked = split_resumable(pool, auth_state.get("active")) if args.json: print(json.dumps({ "workspace": workspace, "source": source, "active_profile": auth_state.get("active"), "auth_readable": auth_state.get("readable"), "resumable": resumable, "blocked": blocked if args.all else [], "blocked_count": len(blocked), }, indent=2, default=str)) return 0 print("workspace: {} (source: {})".format(workspace, source)) if auth_state.get("readable"): print("active profile: {}".format(auth_state.get("active") or "(none)")) else: print("active profile: ? (auth dir unreadable from this shell)") if not pool: print("No sessions for this workspace.") return 0 print(format_pool_table(resumable, blocked, args.all)) return 0 def cmd_check(args, paths): workspace = os.path.realpath(args.workspace or canonical_workspace()) rows, _ = load_rows(paths) auth_state = load_auth_state(paths["config_dir"]) annotate_rows(rows, auth_state) decision = check_resume(rows, args.ref, workspace, auth_state) if args.json: row = dict(decision["row"]) if decision["row"] else None print(json.dumps({ "ok": decision["ok"], "reason": decision["reason"], "detail": decision["detail"], "fix": decision["fix"], "row": row, }, indent=2, default=str)) else: status = "OK" if decision["ok"] else "BLOCKED ({})".format(decision["reason"]) print("{}: {}".format(status, decision["detail"])) if decision["fix"]: print("fix: {}".format(decision["fix"])) return 0 if decision["ok"] else 1 def cmd_resume(args, paths): workspace = os.path.realpath(args.workspace or canonical_workspace()) rows, _ = load_rows(paths) auth_state = load_auth_state(paths["config_dir"]) annotate_rows(rows, auth_state) if args.ref == "--last": pool = pool_for_workspace(rows, workspace) resumable, _ = split_resumable(pool, auth_state.get("active")) if not resumable: print("No resumable sessions for this workspace.", file=sys.stderr) return 1 session_id = resumable[0]["session_id"] else: decision = check_resume(rows, args.ref, workspace, auth_state) if not decision["ok"]: print("refusing to resume: {}".format(decision["detail"]), file=sys.stderr) if decision["fix"]: print("fix: {}".format(decision["fix"]), file=sys.stderr) return 1 session_id = decision["row"]["session_id"] cmd = ["muse-code", "resume", session_id] if args.dry_run: print("would exec: {}".format(" ".join(cmd))) return 0 os.execvp(cmd[0], cmd) return 0 # unreachable def main(argv=None): parser = argparse.ArgumentParser( prog="muse_resume_pool", description="Per-repo, profile-aware Muse Code resume pool and guard.", ) parser.add_argument( "--workspace", help="Workspace root to scope to (default: git top-level or cwd).", ) sub = parser.add_subparsers(dest="command", required=True) p_pool = sub.add_parser("pool", help="List sessions resumable here and now.") p_pool.add_argument("--all", action="store_true", help="Also show profile-blocked sessions.") p_pool.add_argument("--json", action="store_true", help="Machine-readable output.") p_pool.set_defaults(func=cmd_pool) p_check = sub.add_parser("check", help="Explain whether a resume would succeed.") p_check.add_argument("ref", help="Session UUID, UUID prefix, or session name.") p_check.add_argument("--json", action="store_true", help="Machine-readable output.") p_check.set_defaults(func=cmd_check) p_resume = sub.add_parser("resume", help="Guard then exec muse-code resume.") p_resume.add_argument("ref", help="Session UUID, prefix, name, or --last.") p_resume.add_argument("--dry-run", action="store_true", help="Print the resume command instead of exec'ing.") p_resume.set_defaults(func=cmd_resume) args = parser.parse_args(argv) return args.func(args, default_paths()) if __name__ == "__main__": sys.exit(main())