Files
box/bin/muse_resume_pool.py

546 lines
19 KiB
Python
Executable File

#!/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())