546 lines
19 KiB
Python
Executable File
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())
|