- install_skill: adopt an existing skill from a local dir, SKILL.md, .zip, or git URL; symlinks when the source already lives in a recognized external skills folder (.claude/skills, .agents/skills) instead of vendoring a duplicate, with install provenance recorded per skill. Skill discovery now also reads those external folders read-only, so skills placed by other tools are visible with no install step at all. - export_terminal_log: full tmux scrollback export, visible only inside a live tmux session (TOOL_ENV_GATES, a new env-conditional layer on top of the existing keyword-based lazy tool loading). - Chunked lexical RAG (SQLite FTS5 + bm25, no vectors/embeddings): new chunks table wired into record_save, shell/subagent output spill, and terminal log export, searchable via the existing search tool (kind chunk) with full chunk text returned, not just a snippet. delete_record cleans up a record's chunks with it. - Fix a regression where create_skill's default scope silently became context-dependent instead of always "project". - Rewrite README.md to document all of the above plus the existing context/tool-selection and search/memory internals in depth. Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
7595 lines
335 KiB
Python
Executable File
7595 lines
335 KiB
Python
Executable File
#!/usr/bin/env python3
|
||
# retoor <retoor@molodetz.nl>
|
||
import argparse
|
||
import base64
|
||
import collections
|
||
import difflib
|
||
import getpass
|
||
import glob
|
||
import hashlib
|
||
import importlib.util
|
||
import hmac
|
||
import html
|
||
import json
|
||
import os
|
||
import queue
|
||
import random
|
||
import re
|
||
import shutil
|
||
import shlex
|
||
import socket
|
||
import sqlite3
|
||
import ssl
|
||
import struct
|
||
import subprocess
|
||
import sys
|
||
import tempfile
|
||
import threading
|
||
import time
|
||
import urllib.error
|
||
import urllib.parse
|
||
import urllib.request
|
||
import uuid
|
||
import zipfile
|
||
from datetime import datetime, timedelta, timezone
|
||
|
||
try:
|
||
import readline
|
||
except ImportError:
|
||
readline = None
|
||
|
||
try:
|
||
import fcntl
|
||
except ImportError:
|
||
fcntl = None
|
||
|
||
# Python 3.6 compatibility notes:
|
||
# - subprocess 'text' and 'capture_output' kwargs were added in 3.7, so this
|
||
# module only uses 'universal_newlines' plus explicit PIPEs (works on 3.6+).
|
||
# - datetime.fromisoformat was added in 3.7, so all parsing goes through
|
||
# _fromisoformat below (native when available, strptime fallback on 3.6).
|
||
|
||
|
||
def _fromisoformat(text):
|
||
"""datetime.fromisoformat backport for Python 3.6 (handles Z and +HH:MM)."""
|
||
native = getattr(datetime, "fromisoformat", None)
|
||
if native is not None:
|
||
raw = str(text or "").strip()
|
||
if raw[-1:] in ("Z", "z"):
|
||
raw = raw[:-1] + "+00:00"
|
||
return native(raw)
|
||
raw = str(text or "").strip()
|
||
if not raw:
|
||
raise ValueError("empty datetime")
|
||
if raw[-1:] in ("Z", "z"):
|
||
raw = raw[:-1] + "+00:00"
|
||
tzinfo = None
|
||
core = raw
|
||
match = re.match(r"^(\d{4}-\d{2}-\d{2}(?:[T ]\d{2}:\d{2}(?::\d{2}(?:\.\d{1,6})?)?)?)([+-]\d{2}:?\d{2}(?::\d{2})?|[+-]\d{4}|[+-]\d{2})$", raw)
|
||
if match is not None and match.group(2):
|
||
core, tzpart = match.group(1), match.group(2)
|
||
sign = 1 if tzpart[0] == "+" else -1
|
||
digits = tzpart[1:].replace(":", "")
|
||
hours = int(digits[:2])
|
||
minutes = int(digits[2:4]) if len(digits) >= 4 else 0
|
||
tzinfo = timezone(sign * timedelta(hours=hours, minutes=minutes))
|
||
for fmt in ("%Y-%m-%dT%H:%M:%S.%f", "%Y-%m-%dT%H:%M:%S", "%Y-%m-%dT%H:%M",
|
||
"%Y-%m-%d %H:%M:%S.%f", "%Y-%m-%d %H:%M:%S", "%Y-%m-%d %H:%M",
|
||
"%Y-%m-%d"):
|
||
try:
|
||
parsed = datetime.strptime(core, fmt)
|
||
except ValueError:
|
||
continue
|
||
return parsed.replace(tzinfo=tzinfo)
|
||
raise ValueError("invalid datetime, use ISO like 2026-10-08T09:00")
|
||
|
||
VERSION = "1.22.0"
|
||
WORKER_STEPS = 12
|
||
CREATE_SKILL_STEPS = 40
|
||
FORK_MAX_DEPTH = 2
|
||
AGENTS = {}
|
||
AGENTS_LOCK = threading.Lock()
|
||
AGENTS_NEXT = [1]
|
||
MUSE_LABEL = "muse"
|
||
MUSE_CHAT_URL = "muse://exec"
|
||
MUSE_MODELS_URL = "muse://models"
|
||
GROK_LABEL = "grok"
|
||
GROK_CHAT_URL = "grok://exec"
|
||
GROK_MODELS_URL = "grok://models"
|
||
CLAUDE_LABEL = "claude"
|
||
CLAUDE_CHAT_URL = "claude://exec"
|
||
CLAUDE_MODELS_URL = "claude://models"
|
||
CODEX_LABEL = "codex"
|
||
CODEX_CHAT_URL = "codex://exec"
|
||
CODEX_MODELS_URL = "codex://models"
|
||
GEMINI_LABEL = "gemini"
|
||
GEMINI_CHAT_URL = "gemini://exec"
|
||
GEMINI_MODELS_URL = "gemini://models"
|
||
OLLAMA_LABEL = "ollama"
|
||
OLLAMA_BASE = "http://localhost:11434/v1"
|
||
OLLAMA_LIST_TTL = 600
|
||
CLI_CHAT_URLS = (MUSE_CHAT_URL, GROK_CHAT_URL, CLAUDE_CHAT_URL, CODEX_CHAT_URL, GEMINI_CHAT_URL)
|
||
MUSE_MAX_STEPS = 1
|
||
MUSE_SESSIONS_FILE = "muse_sessions.json"
|
||
MUSE_LOCKS_DIR = "muse_locks"
|
||
MUSE_WORKER_TTL = 86400
|
||
MUSE_REAP_GRACE = 30
|
||
OPENROUTER_LABEL = "openrouter"
|
||
OPENROUTER_BASE = "https://openrouter.ai/api/v1"
|
||
OPENROUTER_MODEL = "deepseek/deepseek-v4.1-flash"
|
||
OPENCODE_LABEL = "opencode"
|
||
OPENCODE_CHAT_URL = "opencode://run"
|
||
OPENCODE_MODELS_URL = "opencode://models"
|
||
OPENCODE_AGENT = "plan"
|
||
DEVPLACE_LABEL = "devplace"
|
||
DEVPLACE_BASE = "https://devplace.net/openai/v1"
|
||
DEVPLACE_FREE_ROUTES = ("free", "devii")
|
||
ZEN_LABEL = "zen"
|
||
ZEN_BASE = "https://opencode.ai/zen/v1"
|
||
ZEN_FREE_MODELS = ("mimo-v2.6-flash-free", "nemotron-3.5-lightning-free", "nemotron-3-ultra-free", "muse-spark-1.3-contributor-free", "muse-spark-1.2-contributor-free", "jev-1.13-free", "exo-free", "space-bunny-free", "longcat-2.5-preview-free", "ling-3.1-flash-free", "ling-3.0-flash-fin-free", "fledge-alpha-free")
|
||
POLLINATIONS_LABEL = "pollinations"
|
||
POLLINATIONS_BASE = "https://text.pollinations.ai/openai/v1"
|
||
POLLINATIONS_MODELS = ("openai",)
|
||
OPENCODE_USER_AGENT = "opencode/1.18.29 ai-sdk/provider-utils/4.0.23 runtime/bun/1.3.15"
|
||
OPENCODE_CLIENT_NAME = "cli"
|
||
MODEL_HEALTH_FILE = "model_health.json"
|
||
MODEL_NEUTRAL_REWARD = 0.5
|
||
MODEL_WEIGHT_PRIOR = 2.0
|
||
MODEL_SPEED_REF_TPS = 25.0
|
||
MODEL_LATENCY_REF_MS = 8000.0
|
||
MODEL_CIRCUIT_THRESHOLD = 3
|
||
MODEL_CIRCUIT_COOLDOWN = 300.0
|
||
RETRY_BACKOFF = (2, 5, 10, 20, 30, 60)
|
||
MODEL_LIST_TIMEOUT = 8
|
||
OPENROUTER_FREE_CACHE_TTL = 3600
|
||
RSEARCH_BASE = "https://rsearch.app.molodetz.nl"
|
||
FETCH_METHODS = ("GET", "POST", "PUT", "PATCH", "DELETE", "HEAD", "OPTIONS")
|
||
CURL_HINT = "hint: use the web_fetch tool for HTTP(S) (methods, headers, bodies, status) instead of curl/wget via shell"
|
||
|
||
|
||
def curl_hint(command):
|
||
if re.search(r"\bcurl\b|\bwget\b", command or "") and "http" in (command or "").lower():
|
||
return "\n" + CURL_HINT
|
||
return ""
|
||
DEFAULT_MODEL = "muse"
|
||
DEFAULT_VOICE = "en-US-EmmaMultilingualNeural"
|
||
DEFAULT_PASSPHRASE = "tai-default-insecure-change-me"
|
||
BOX_IMAGE = "tai-box:latest"
|
||
BOX_NAME = "tai-box"
|
||
TELEGRAM_API = "https://api.telegram.org/bot"
|
||
TELEGRAM_FILE_API = "https://api.telegram.org/file/bot"
|
||
SKILL_NAME_RE = re.compile(r"^[a-z0-9](?:[a-z0-9-]{0,62}[a-z0-9])?$")
|
||
BASHRC_MARK_BEGIN = "# tai command-not-found hook - start"
|
||
BASHRC_MARK_END = "# tai command-not-found hook - end"
|
||
BASHRC_BLOCK = BASHRC_MARK_BEGIN + "\ncommand_not_found_handle() {\n \"$HOME/.local/bin/tai.py\" \"$@\"\n return $?\n}\n" + BASHRC_MARK_END + "\n"
|
||
BOX_CONTAINERFILE = """FROM python:3.12-slim
|
||
ENV DEBIAN_FRONTEND=noninteractive PIP_NO_CACHE_DIR=1
|
||
RUN apt-get update && apt-get install -y --no-install-recommends ffmpeg espeak-ng python3-venv && rm -rf /var/lib/apt/lists/*
|
||
RUN python -m venv /box/venv && /box/venv/bin/pip install --no-cache-dir faster-whisper edge-tts
|
||
RUN /box/venv/bin/python -c "from faster_whisper import WhisperModel; WhisperModel('tiny', device='cpu', compute_type='int8')"
|
||
COPY stt.py /box/stt.py
|
||
COPY tts.py /box/tts.py
|
||
CMD ["sleep", "infinity"]
|
||
"""
|
||
BOX_PYTHON = "/box/venv/bin/python"
|
||
_BOX_PYTHON_CACHE = {}
|
||
BOX_STT = """import sys
|
||
from faster_whisper import WhisperModel
|
||
data = sys.stdin.buffer.read()
|
||
with open("/tmp/in.audio", "wb") as handle:
|
||
handle.write(data)
|
||
model = WhisperModel("tiny", device="cpu", compute_type="int8")
|
||
segments, _info = model.transcribe("/tmp/in.audio")
|
||
print(" ".join(segment.text for segment in segments).strip(), flush=True)
|
||
"""
|
||
BOX_TTS = """import asyncio
|
||
import sys
|
||
import edge_tts
|
||
async def speak(text, voice):
|
||
talk = edge_tts.Communicate(text, voice or "en-US-EmmaMultilingualNeural")
|
||
await talk.save("/tmp/out.mp3")
|
||
text = sys.stdin.read().strip()
|
||
voice = sys.argv[1] if len(sys.argv) > 1 else ""
|
||
asyncio.run(speak(text, voice))
|
||
sys.stdout.buffer.write(open("/tmp/out.mp3", "rb").read())
|
||
"""
|
||
TELEGRAM_UNIT = """[Unit]
|
||
Description=tai telegram bot
|
||
After=network-online.target
|
||
Wants=network-online.target
|
||
|
||
[Service]
|
||
Type=simple
|
||
WorkingDirectory=%%h
|
||
ExecStart=%s --telegram
|
||
EnvironmentFile=%s
|
||
Restart=on-failure
|
||
RestartSec=10
|
||
|
||
[Install]
|
||
WantedBy=default.target
|
||
"""
|
||
SCHEDULER_UNIT = """[Unit]
|
||
Description=tai scheduler
|
||
After=network-online.target
|
||
Wants=network-online.target
|
||
|
||
[Service]
|
||
Type=simple
|
||
WorkingDirectory=%%h
|
||
ExecStart=%s --scheduler
|
||
Restart=on-failure
|
||
RestartSec=10
|
||
|
||
[Install]
|
||
WantedBy=default.target
|
||
"""
|
||
INSTALL_TARGETS = ("binary", "bash-hook", "venv", "scheduler-service", "telegram-service", "container")
|
||
CONTEXT_CAP = 32000
|
||
COMPACT_RATIO = 0.8
|
||
KEEP_TURNS = 6
|
||
MAX_STEPS = 25
|
||
TOTAL_STEP_CAP = 500
|
||
LOOP_NUDGE_AT = 3
|
||
LOOP_STOP_AT = 5
|
||
SYSTEM_MAX_CHARS = 12000
|
||
SHELL_TIMEOUT = 120
|
||
SCHEDULER_INTERVAL = 30
|
||
HTTP_TIMEOUT = 60
|
||
STREAM_TIMEOUT = 300
|
||
|
||
EDGE_HOST = "speech.platform.bing.com"
|
||
EDGE_WS_PATH = "/consumer/speech/synthesize/readaloud/edge/v1"
|
||
EDGE_TRUSTED_TOKEN = "6A5AA1D4EAFF4E9FB37E23D68491D6F4"
|
||
EDGE_CHROMIUM = "143.0.3650.75"
|
||
|
||
PROFILE_RE = re.compile(r"^[A-Za-z0-9][A-Za-z0-9_-]{0,31}$")
|
||
BOT_RE = re.compile(r"^[a-z0-9][a-z0-9_-]{0,31}$")
|
||
BOT_MENTION_RE = re.compile(r"^@([A-Za-z0-9][A-Za-z0-9_-]{0,31})\s+(.*)$", re.DOTALL)
|
||
SECRET_RE = re.compile(r"^[A-Za-z0-9][A-Za-z0-9_.-]{0,63}$")
|
||
LEGACY_PASSWORD_NOTE = "especially passwords (collect and keep every password the user shares), preferences, or standing behavior changes"
|
||
VAULT_MEMORY_NOTE = "preferences, or standing behavior changes. Store passwords, tokens, and secrets with store_secret and reference them by name; never write secret values into memory"
|
||
SAFE_COMMANDS = ("ls", "pwd", "echo", "cat", "head", "tail", "grep", "find", "wc", "sort", "uniq", "diff", "file", "stat", "date", "whoami", "uname", "lsb_release")
|
||
RECORDERS = (("arecord", ("arecord", "-q", "-d", "{seconds}", "-f", "cd", "-t", "wav", "{path}")), ("rec", ("rec", "-q", "{path}", "trim", "0", "{seconds}")), ("ffmpeg", ("ffmpeg", "-y", "-v", "quiet", "-f", "alsa", "-i", "default", "-t", "{seconds}", "{path}")))
|
||
TRANSCRIBERS = ("whisper-cpp", "whisper-cli", "whisper", "faster-whisper")
|
||
SYSINFO_TIMEOUT = 10
|
||
SYSINFO_BINARIES = ("tmux", "git", "ffmpeg", "arecord", "sox", "whisper-cpp", "whisper", "systemctl", "podman", "docker", "muse", "grok", "opencode", "claude", "codex", "gemini", "ollama")
|
||
|
||
DEFAULT_SYSTEM = (
|
||
"You are tai, a professional autonomous assistant. You act with tools, verify results, and report concisely. "
|
||
"Rules: prefer the smallest change that solves the task; never silently ignore errors; ask when requirements are ambiguous; "
|
||
"keep answers short and factual. Tools: shell and file tools for the local machine, web_search and web_fetch for the internet, "
|
||
"web_fetch handles every HTTP method with headers and bodies so never curl or wget through shell, "
|
||
"speak for voice output, recall to search past session memory. Memory: call remember whenever you learn durable facts, "
|
||
"preferences, or standing behavior changes. Store passwords, tokens, and secrets with store_secret and reference them by name; "
|
||
"never write secret values into memory. Name and value suffice for secrets; never pester the user for more. Call remember with a forget instruction to drop outdated knowledge. Delegation: do quick, interactive, or memory-changing work yourself; fork background subagents for independent, long, or context-heavy work and collect them with poll; schedule future work instead of waiting on it. "
|
||
"Records: save big or durable findings with record_save and tag generously; large tool outputs spill to mem:<hash> ids automatically, page them with record_read, connect related nodes with graph_link. "
|
||
"Tags: strict rules, lowercase singular only, shortest common word for the subject, one canonical tag per subject, at most 5 per item, always reuse existing tags from the tags tool instead of inventing synonyms. Content words matching known tags attach automatically, and new records link to recent same-tag records. "
|
||
"Lazy tools: only everyday tools load by default; name any other tool or topic in your reply text and it becomes callable on your next step. "
|
||
"Bots: one identity holds many bots; a bot is only a system message plus its own history, all vault data is shared. @name routes one turn to another bot and files the exchange in both histories; /bot switches, create_bot makes bots. "
|
||
"Retrieval: use search first for anything stored: one ranked full-text query across records, events, and audit plus graph expansion beats paging and guessing. Chain tool calls until the task is done."
|
||
)
|
||
MERGE_SYSTEM = (
|
||
"You maintain an AI assistant system prompt. You receive the current system prompt and one memory instruction. "
|
||
"Rewrite the system prompt to incorporate the instruction: add new facts, update changed behavior, or remove forgotten items. "
|
||
"Never persist passwords, tokens, or other secret values; if the instruction holds one, note only that a secret exists and must be stored with store_secret. "
|
||
"Preserve everything unrelated. Keep it organized with short sections. Output ONLY the rewritten system prompt, no explanation, no code fences."
|
||
)
|
||
COMPACT_SYSTEM = (
|
||
"You compress assistant session history into a dense handoff summary. Cover: goal, completed work with outcomes, current state, "
|
||
"key decisions, files touched, errors and fixes, pending next steps. Output only the summary, no preamble."
|
||
)
|
||
AUTONOMOUS_NOTE = (
|
||
"Autonomous mode: never ask the user anything this session. When requirements are unclear, blocked, or need input, "
|
||
"research deeply with web_search, web_fetch, shell, and subagents, then decide and proceed. State assumptions briefly in your reply."
|
||
)
|
||
|
||
|
||
HELP_TEXT = (
|
||
"commands:\n"
|
||
" /profile [name] show current profile or switch to it\n"
|
||
" /profiles list all profiles, current marked with *\n"
|
||
" /bots list bots, current marked with *\n"
|
||
" /bot <name> switch bot, history resumes\n"
|
||
" /env [home|sandbox] show or switch execution environment\n"
|
||
" /skills list loaded skills with provenance\n"
|
||
" /skills install ... install an existing skill (copy or symlink)\n"
|
||
" /secret manage sealed secrets (set|list|delete)\n"
|
||
" /sysinfo show host environment checks\n"
|
||
" /models show model roster with speed health\n"
|
||
" /schedule at|every schedule a prompt for later\n"
|
||
" /schedules list scheduled prompts\n"
|
||
" /unschedule <id> delete a scheduled prompt\n"
|
||
" /records [query] search saved records\n"
|
||
" /record <id> read one record page\n"
|
||
" /graph <node> show node neighborhood\n"
|
||
" /tags [prefix] list tags with usage counts\n"
|
||
" /tools [name] show core and lazy tools\n"
|
||
" /install ... install status, install, upgrade, reinstall\n"
|
||
" /search <query> ranked search over everything\n"
|
||
" /audit [path] show file audit trail\n"
|
||
" /restore <id> restore a file from an audit row\n"
|
||
" /release part msg bump version, back up, log message\n"
|
||
" /fork <task> spawn a background subagent, REPL stays free\n"
|
||
" /agents list background subagents\n"
|
||
" /agent <id|clear> show one result, or purge finished agents\n"
|
||
" /compact compress history into a summary\n"
|
||
" /clear drop history, keep system message\n"
|
||
" /help this text\n"
|
||
" /quit exit\n"
|
||
"anything else is sent to the agent."
|
||
)
|
||
|
||
|
||
class Ansi:
|
||
RESET = "\033[0m"
|
||
BOLD = "\033[1m"
|
||
DIM = "\033[2m"
|
||
UNDERLINE = "\033[4m"
|
||
STRIKE = "\033[9m"
|
||
RED = "\033[31m"
|
||
GREEN = "\033[32m"
|
||
YELLOW = "\033[33m"
|
||
BLUE = "\033[34m"
|
||
MAGENTA = "\033[35m"
|
||
CYAN = "\033[36m"
|
||
GRAY = "\033[90m"
|
||
|
||
|
||
def paint(text, *codes):
|
||
if os.environ.get("NO_COLOR") or not sys.stdout.isatty():
|
||
return text
|
||
return "".join(codes) + text + Ansi.RESET
|
||
|
||
|
||
def now_iso():
|
||
return datetime.now(timezone.utc).isoformat(timespec="seconds")
|
||
|
||
|
||
def short_error(exc):
|
||
text = str(exc) or type(exc).__name__
|
||
return text[:300]
|
||
|
||
|
||
def truncate(text, head=1500, tail=500):
|
||
if len(text) <= head + tail + 50:
|
||
return text
|
||
skipped = len(text) - head - tail
|
||
return text[:head] + "\n...[%d chars truncated]...\n" % skipped + text[-tail:]
|
||
|
||
|
||
SPILL_LIMIT = 6000
|
||
RECORD_KINDS = ("note", "output", "file", "research", "transcript", "termlog")
|
||
AUDIT_MAX_CHARS = 20000
|
||
FILE_RECORD_MAX = 50000
|
||
SHELL_SNAPSHOT_MAX_FILES = 25
|
||
BACKUP_KEEP = 10
|
||
_UNSET = object()
|
||
|
||
|
||
def capped_text(text):
|
||
if text is None:
|
||
return None, None
|
||
if len(text) <= AUDIT_MAX_CHARS:
|
||
return text, len(text.encode("utf-8"))
|
||
return text[:AUDIT_MAX_CHARS] + "\n...[truncated, full size %d bytes]..." % len(text.encode("utf-8")), len(text.encode("utf-8"))
|
||
|
||
|
||
def trunc_marker(true_bytes):
|
||
return "\n...[truncated, full size %d bytes]..." % true_bytes
|
||
|
||
|
||
def read_capped(path, cap=AUDIT_MAX_CHARS):
|
||
try:
|
||
size = os.path.getsize(path)
|
||
except OSError:
|
||
return None, None
|
||
try:
|
||
with open(path, "r", encoding="utf-8", errors="replace") as handle:
|
||
text = handle.read(cap + 1)
|
||
except OSError:
|
||
return None, None
|
||
if len(text) > cap:
|
||
return text[:cap], size
|
||
return text, None
|
||
|
||
|
||
SHELL_SEGMENT = (";", "&&", "||", "|", "&", "(", ")")
|
||
SHELL_WRAPPERS = ("sudo", "command", "env", "nice", "nohup", "setsid")
|
||
|
||
|
||
def shell_target_paths(command, workdir):
|
||
spaced = re.sub(r"([;&|()<>])", r" \1 ", command)
|
||
try:
|
||
tokens = shlex.split(spaced, posix=True)
|
||
except ValueError:
|
||
return []
|
||
found = []
|
||
segment = []
|
||
|
||
def flush():
|
||
found.extend(segment_targets(segment))
|
||
|
||
for token in tokens:
|
||
if token in SHELL_SEGMENT:
|
||
flush()
|
||
segment = []
|
||
else:
|
||
segment.append(token)
|
||
flush()
|
||
resolved = []
|
||
for item in found:
|
||
if len(resolved) >= SHELL_SNAPSHOT_MAX_FILES:
|
||
break
|
||
base = item if os.path.isabs(item) else os.path.join(workdir, item)
|
||
if any(mark in item for mark in ("*", "?", "[")):
|
||
try:
|
||
hits = sorted(glob.glob(base, recursive=True))
|
||
except (OSError, ValueError):
|
||
continue
|
||
else:
|
||
hits = [base]
|
||
for hit in hits:
|
||
if len(resolved) >= SHELL_SNAPSHOT_MAX_FILES:
|
||
break
|
||
if os.path.isfile(hit):
|
||
resolved.append(os.path.abspath(hit))
|
||
elif os.path.isdir(hit):
|
||
for root, _dirs, files in os.walk(hit):
|
||
for entry in sorted(files):
|
||
if len(resolved) >= SHELL_SNAPSHOT_MAX_FILES:
|
||
break
|
||
full = os.path.join(root, entry)
|
||
if os.path.isfile(full):
|
||
resolved.append(os.path.abspath(full))
|
||
if len(resolved) >= SHELL_SNAPSHOT_MAX_FILES:
|
||
break
|
||
seen = []
|
||
for item in resolved:
|
||
if item not in seen:
|
||
seen.append(item)
|
||
return seen
|
||
|
||
|
||
def segment_targets(tokens):
|
||
targets = []
|
||
pos = 0
|
||
while pos < len(tokens):
|
||
token = tokens[pos]
|
||
if "=" in token and not token.startswith((">", "<", "-", "/")):
|
||
pos += 1
|
||
elif token in SHELL_WRAPPERS:
|
||
pos += 1
|
||
elif token == "timeout" and pos + 1 < len(tokens) and re.fullmatch(r"[0-9.]+[smhd]?", tokens[pos + 1] or ""):
|
||
pos += 2
|
||
else:
|
||
break
|
||
rest = tokens[pos:]
|
||
for index, token in enumerate(rest):
|
||
if token in (">", ">>") and index + 1 < len(rest):
|
||
cand = rest[index + 1]
|
||
if cand and not cand.startswith("&") and not cand.startswith("/dev/"):
|
||
targets.append(cand)
|
||
if not rest:
|
||
return targets
|
||
name = os.path.basename(rest[0])
|
||
operands = [item for item in rest[1:] if not item.startswith("-") or item in ("-", "--")]
|
||
operands = [item for item in operands if item not in (">", ">>", "<", "--")]
|
||
if name == "rm":
|
||
targets.extend(operands)
|
||
elif name in ("mv", "cp") and operands:
|
||
targets.append(operands[-1])
|
||
elif name == "tee":
|
||
targets.extend(operands)
|
||
elif name == "truncate":
|
||
for item in operands:
|
||
if re.fullmatch(r"[0-9]+[KMGTPE]?[iB]?", item or ""):
|
||
continue
|
||
targets.append(item)
|
||
elif name == "dd":
|
||
for item in operands:
|
||
if item.startswith("of="):
|
||
targets.append(item[3:])
|
||
elif name == "shred":
|
||
targets.extend(operands)
|
||
return [item for item in targets if item and item != "-"]
|
||
|
||
|
||
def spill_output(store, kind, title, text, tags=(), profile=None):
|
||
profile = profile if profile is not None else store.profile
|
||
record_id = store.add_record(kind, title, text, tags, profile)
|
||
store.add_chunks("record", record_id, chunk_lines(text), profile)
|
||
return record_id
|
||
|
||
|
||
def spilled_view(text, record_id, size, head=1500, tail=500):
|
||
return text[:head] + "\n...[%d chars spilled to %s, record_read pages the rest]...\n" % (size, record_id) + text[-tail:]
|
||
|
||
|
||
def spill_or_truncate(store, kind, title, text, tags=(), profile=None):
|
||
if store is None or len(text) <= SPILL_LIMIT:
|
||
return truncate(text)
|
||
record_id = spill_output(store, kind, title, text, tags, profile)
|
||
return spilled_view(text, record_id, len(text))
|
||
|
||
|
||
CHUNK_MAX_LINES = 60
|
||
CHUNK_OVERLAP_LINES = 3
|
||
CHUNK_MAX_CHARS = 4000
|
||
|
||
|
||
def chunk_lines(text, max_lines=CHUNK_MAX_LINES, overlap_lines=CHUNK_OVERLAP_LINES, max_chars=CHUNK_MAX_CHARS):
|
||
"""Generic RAG chunker: bounded windows of whole lines with small overlap, never splitting mid-line.
|
||
|
||
Used for any content that gets indexed for retrieval (terminal captures, long records, ...),
|
||
so there is one chunking rule in the codebase instead of one per content type.
|
||
"""
|
||
lines = text.splitlines()
|
||
n = len(lines)
|
||
if n == 0:
|
||
return []
|
||
chunks = []
|
||
start = 0
|
||
while start < n:
|
||
end = start
|
||
size = 0
|
||
count = 0
|
||
while end < n and count < max_lines and size < max_chars:
|
||
size += len(lines[end]) + 1
|
||
count += 1
|
||
end += 1
|
||
piece = "\n".join(lines[start:end]).strip()
|
||
if piece:
|
||
chunks.append(piece)
|
||
if end >= n:
|
||
break
|
||
start = max(start + 1, end - overlap_lines)
|
||
return chunks
|
||
|
||
|
||
def collapse_repeated_lines(text, threshold=3):
|
||
"""Collapses runs of 3+ identical lines (progress bars, retry spam) to one line + a count note."""
|
||
lines = text.splitlines()
|
||
out = []
|
||
i = 0
|
||
n = len(lines)
|
||
while i < n:
|
||
j = i
|
||
while j < n and lines[j] == lines[i]:
|
||
j += 1
|
||
run = j - i
|
||
if run >= threshold:
|
||
out.append(lines[i])
|
||
out.append("... (repeated %d more times) ..." % (run - 1))
|
||
else:
|
||
out.extend(lines[i:j])
|
||
i = j
|
||
return "\n".join(out)
|
||
|
||
|
||
def tmux_pane_header():
|
||
try:
|
||
info = subprocess.run(["tmux", "display-message", "-p", "#{session_name}:#{window_index}.#{pane_index} #{pane_current_command}"], stdout=subprocess.PIPE, stderr=subprocess.PIPE, universal_newlines=True, timeout=10)
|
||
except (OSError, subprocess.SubprocessError):
|
||
return ""
|
||
if info.returncode == 0 and info.stdout.strip():
|
||
return info.stdout.strip()
|
||
return ""
|
||
|
||
|
||
VERSION_RE = re.compile(r"^VERSION = \"(\d+)\.(\d+)\.(\d+)\"$", re.MULTILINE)
|
||
|
||
|
||
def next_version(current, part):
|
||
major, minor, patch = (int(item) for item in current.split("."))
|
||
if part == "major":
|
||
return "%d.0.0" % (major + 1)
|
||
if part == "minor":
|
||
return "%d.%d.0" % (major, minor + 1)
|
||
if part == "patch":
|
||
return "%d.%d.%d" % (major, minor, patch + 1)
|
||
raise ValueError("part must be major, minor, or patch")
|
||
|
||
|
||
def atomic_write(path, content):
|
||
tmp = "%s.tmp-%d" % (path, os.getpid())
|
||
with open(tmp, "w", encoding="utf-8") as handle:
|
||
handle.write(content)
|
||
handle.flush()
|
||
os.fsync(handle.fileno())
|
||
os.replace(tmp, path)
|
||
|
||
|
||
def sealed_write(path, content):
|
||
atomic_write(path, content)
|
||
try:
|
||
os.chmod(path, 0o600)
|
||
except OSError:
|
||
pass
|
||
|
||
|
||
def quarantine(path):
|
||
try:
|
||
if os.path.exists(path):
|
||
os.replace(path, "%s.corrupt-%d" % (path, int(time.time())))
|
||
except OSError:
|
||
pass
|
||
|
||
|
||
def release_bump(script_file, part):
|
||
with open(script_file, "r", encoding="utf-8") as handle:
|
||
text = handle.read()
|
||
match = VERSION_RE.search(text)
|
||
if match is None:
|
||
raise ValueError("no VERSION line in %s" % script_file)
|
||
old = ".".join(match.groups())
|
||
new = next_version(old, part)
|
||
atomic_write(script_file, text[:match.start()] + "VERSION = \"%s\"" % new + text[match.end():])
|
||
return old, new
|
||
|
||
|
||
def backups_dir(home):
|
||
return os.path.join(home, "backups")
|
||
|
||
|
||
def prune_backups(folder):
|
||
try:
|
||
names = sorted((entry for entry in os.listdir(folder) if entry.startswith("tai-") and entry.endswith(".py")), key=lambda entry: os.path.getmtime(os.path.join(folder, entry)))
|
||
except OSError:
|
||
return 0
|
||
pruned = 0
|
||
while len(names) > BACKUP_KEEP:
|
||
try:
|
||
os.remove(os.path.join(folder, names.pop(0)))
|
||
pruned += 1
|
||
except OSError:
|
||
break
|
||
return pruned
|
||
|
||
|
||
def snapshot_self(script_file, folder, version):
|
||
with open(script_file, "rb") as handle:
|
||
data = handle.read()
|
||
digest = hashlib.sha256(data).hexdigest()[:8]
|
||
stamp = datetime.now(timezone.utc).strftime("%Y%m%dT%H%M%SZ")
|
||
name = "tai-%s-%s-%s.py" % (version, stamp, digest)
|
||
os.makedirs(folder, exist_ok=True)
|
||
dest = os.path.join(folder, name)
|
||
with open(dest, "wb") as handle:
|
||
handle.write(data)
|
||
prune_backups(folder)
|
||
return dest, digest
|
||
|
||
|
||
def backup_has_digest(folder, digest):
|
||
try:
|
||
return any(entry.endswith("-%s.py" % digest) for entry in os.listdir(folder))
|
||
except OSError:
|
||
return False
|
||
|
||
|
||
def ensure_self_backup(config, store, script_file=None, actor="main"):
|
||
script_file = script_file or os.path.abspath(__file__)
|
||
folder = backups_dir(config.home)
|
||
with open(script_file, "rb") as handle:
|
||
digest = hashlib.sha256(handle.read()).hexdigest()[:8]
|
||
if backup_has_digest(folder, digest):
|
||
return None
|
||
dest, _digest = snapshot_self(script_file, folder, VERSION)
|
||
try:
|
||
store.audit_event(actor, "snapshot", script_file, "backed up as %s" % os.path.basename(dest), tags=["snapshot", "v" + VERSION])
|
||
text, override = read_capped(script_file, FILE_RECORD_MAX)
|
||
if text is not None:
|
||
store.upsert_file_record(script_file, text, override)
|
||
except (sqlite3.Error, OSError):
|
||
pass
|
||
return dest
|
||
|
||
|
||
def do_release(store, script_file, folder, part, message, actor="main"):
|
||
if not message or not message.strip():
|
||
raise ValueError("release message is required")
|
||
old, new = release_bump(script_file, part)
|
||
dest, _digest = snapshot_self(script_file, folder, new)
|
||
row_id = store.audit_event(actor, "release", script_file, message.strip()[:500], old=old, new=new, tags=["release", "v" + new])
|
||
return old, new, dest, row_id
|
||
|
||
|
||
def estimate_tokens(text):
|
||
return max(1, len(text or "") // 4)
|
||
|
||
|
||
def as_int(value, default=0):
|
||
try:
|
||
return int(value)
|
||
except (TypeError, ValueError):
|
||
return default
|
||
|
||
|
||
def action_signature(name, raw_arguments):
|
||
try:
|
||
flat = json.dumps(json.loads(raw_arguments or "{}"), sort_keys=True, default=str)
|
||
except ValueError:
|
||
flat = str(raw_arguments or "")
|
||
return "%s %s" % (name, flat[:500])
|
||
|
||
|
||
def result_digest(result):
|
||
return hashlib.sha1((result or "").encode("utf-8", "replace")).hexdigest()
|
||
|
||
|
||
class Spinner:
|
||
FRAMES = ("⠋", "⠙", "⠹", "⠸", "⠼", "⠴", "⠦", "⠧", "⠇", "⠏")
|
||
|
||
def __init__(self, label):
|
||
self.label = label
|
||
self.stop_flag = threading.Event()
|
||
self.worker = None
|
||
|
||
def start(self):
|
||
if not sys.stderr.isatty():
|
||
return
|
||
self.worker = threading.Thread(target=self.spin, daemon=True)
|
||
self.worker.start()
|
||
|
||
def spin(self):
|
||
pos = 0
|
||
while not self.stop_flag.is_set():
|
||
frame = self.FRAMES[pos % len(self.FRAMES)]
|
||
sys.stderr.write("\r%s %s" % (frame, self.label))
|
||
sys.stderr.flush()
|
||
pos += 1
|
||
time.sleep(0.08)
|
||
sys.stderr.write("\r%s\r" % (" " * (len(self.label) + 2)))
|
||
sys.stderr.flush()
|
||
|
||
def stop(self):
|
||
self.stop_flag.set()
|
||
if self.worker is not None:
|
||
self.worker.join(timeout=1)
|
||
|
||
|
||
STREAM_HEIGHT = 4
|
||
|
||
|
||
def human_size(num):
|
||
num = max(0, int(num))
|
||
if num < 1024:
|
||
return "%d B" % num
|
||
if num < 1048576:
|
||
return "%.1f KB" % (num / 1024)
|
||
return "%.1f MB" % (num / 1048576)
|
||
|
||
|
||
def stream_status(code, output, elapsed, timed_out):
|
||
detail = "%d lines · %s" % (len(output.splitlines()), human_size(len(output.encode("utf-8", "replace"))))
|
||
if timed_out:
|
||
return "timed out after %ds (killed) · %s" % (SHELL_TIMEOUT, detail)
|
||
return "exit %d · %.1fs · %s" % (code, elapsed, detail)
|
||
|
||
|
||
def truncate_ansi(text, width):
|
||
plain = ANSI_RE.sub("", text)
|
||
if len(plain) <= width:
|
||
return text
|
||
out = []
|
||
count = 0
|
||
pos = 0
|
||
for match in ANSI_RE.finditer(text):
|
||
chunk = text[pos:match.start()]
|
||
take = min(len(chunk), width - count)
|
||
out.append(chunk[:take])
|
||
count += take
|
||
if count >= width:
|
||
pos = match.start()
|
||
break
|
||
out.append(match.group(0))
|
||
pos = match.end()
|
||
if count < width:
|
||
out.append(ANSI_RE.sub("", text[pos:])[:width - count])
|
||
return "".join(out) + Ansi.RESET
|
||
|
||
|
||
class LiveStream:
|
||
def __init__(self, height=STREAM_HEIGHT, width=0, file=None, enabled=True):
|
||
self.height = max(1, height)
|
||
self.width = width or shutil.get_terminal_size().columns
|
||
self.file = file if file is not None else sys.stdout
|
||
self.enabled = enabled
|
||
self.lines = collections.deque(maxlen=self.height)
|
||
self.shown = 0
|
||
|
||
def feed(self, line):
|
||
if not self.enabled:
|
||
return
|
||
cleaned = truncate_ansi(line.split("\r")[-1].rstrip("\n").expandtabs(8), max(8, self.width - 2))
|
||
self.lines.append(cleaned)
|
||
if self.shown < self.height:
|
||
self.file.write("│ %s\n" % cleaned)
|
||
self.shown += 1
|
||
else:
|
||
self.file.write("\033[%dA" % self.height)
|
||
for kept in self.lines:
|
||
self.file.write("\r\033[K│ %s\n" % kept)
|
||
self.file.flush()
|
||
|
||
def close(self, status):
|
||
if not self.enabled:
|
||
return
|
||
if self.shown == 0:
|
||
self.file.write("│ (no output)\n")
|
||
self.file.write("│ %s\n" % truncate_ansi(status, max(8, self.width - 2)))
|
||
self.file.flush()
|
||
|
||
|
||
def run_live(argv, timeout=SHELL_TIMEOUT, shell=False, cwd=None, env=None, stream=None):
|
||
started = time.time()
|
||
proc = subprocess.Popen(argv, shell=shell, cwd=cwd, env=env, stdin=subprocess.DEVNULL, stdout=subprocess.PIPE, stderr=subprocess.STDOUT, universal_newlines=True, errors="replace", bufsize=1)
|
||
chunks = []
|
||
pending = queue.Queue()
|
||
|
||
def pump():
|
||
try:
|
||
for line in proc.stdout:
|
||
chunks.append(line)
|
||
pending.put(line)
|
||
finally:
|
||
pending.put(None)
|
||
|
||
reader = threading.Thread(target=pump, daemon=True)
|
||
reader.start()
|
||
timed_out = False
|
||
try:
|
||
deadline = started + timeout
|
||
while True:
|
||
remaining = deadline - time.time()
|
||
if remaining <= 0:
|
||
timed_out = True
|
||
break
|
||
try:
|
||
item = pending.get(timeout=min(0.05, remaining))
|
||
except queue.Empty:
|
||
continue
|
||
if item is None:
|
||
break
|
||
if stream is not None:
|
||
stream.feed(item)
|
||
except BaseException:
|
||
proc.kill()
|
||
proc.wait()
|
||
raise
|
||
finally:
|
||
if timed_out:
|
||
proc.kill()
|
||
try:
|
||
proc.wait(timeout=10)
|
||
except subprocess.TimeoutExpired:
|
||
proc.kill()
|
||
proc.wait()
|
||
reader.join(timeout=10)
|
||
if not timed_out:
|
||
while not pending.empty():
|
||
item = pending.get_nowait()
|
||
if item is None:
|
||
break
|
||
if stream is not None:
|
||
stream.feed(item)
|
||
return proc.returncode, "".join(chunks), time.time() - started, timed_out
|
||
|
||
|
||
DIFF_CONTEXT = 3
|
||
DIFF_MAX_LINES = 120
|
||
BG_ADD = "\033[48;5;28m"
|
||
BG_DEL = "\033[48;5;88m"
|
||
FG_ORANGE = "\033[38;5;214m"
|
||
|
||
PY_KEYWORDS = ("False", "None", "True", "and", "as", "assert", "async", "await", "break", "class", "continue", "def", "del", "elif", "else", "except", "finally", "for", "from", "global", "if", "import", "in", "is", "lambda", "nonlocal", "not", "or", "pass", "raise", "return", "try", "while", "with", "yield", "match", "case")
|
||
|
||
PY_TOKEN_RE = re.compile(r"(?P<string>\"\"\"(?:\\.|[^\\])*?\"\"\"|'''(?:\\.|[^\\])*?'''|\"(?:\\.|[^\"\\\n])*\"|'(?:\\.|[^'\\\n])*')|(?P<comment>#[^\n]*)|(?P<number>\b\d[\d._]*(?:[eE][+-]?\d+)?[jJ]?\b)|(?P<keyword>\b(?:%s)\b)|(?P<decorator>@[A-Za-z_][\w.]*)|(?P<defname>(?<=def )[A-Za-z_]\w*|(?<=class )[A-Za-z_]\w*)" % "|".join(PY_KEYWORDS))
|
||
|
||
|
||
def color_enabled():
|
||
return not os.environ.get("NO_COLOR") and sys.stdout.isatty()
|
||
|
||
|
||
def highlight_python(line):
|
||
parts = []
|
||
pos = 0
|
||
for match in PY_TOKEN_RE.finditer(line):
|
||
parts.append(line[pos:match.start()])
|
||
kind = match.lastgroup
|
||
text = match.group(0)
|
||
if kind == "string":
|
||
parts.append(Ansi.YELLOW + text + Ansi.RESET)
|
||
elif kind == "comment":
|
||
parts.append(Ansi.GRAY + text + Ansi.RESET)
|
||
elif kind == "number":
|
||
parts.append(FG_ORANGE + text + Ansi.RESET)
|
||
elif kind == "keyword":
|
||
parts.append(Ansi.MAGENTA + text + Ansi.RESET)
|
||
else:
|
||
parts.append(Ansi.CYAN + text + Ansi.RESET)
|
||
pos = match.end()
|
||
parts.append(line[pos:])
|
||
return "".join(parts)
|
||
|
||
|
||
def split_diff_lines(text):
|
||
if not text:
|
||
return []
|
||
lines = text.split("\n")
|
||
if lines and lines[-1] == "":
|
||
lines.pop()
|
||
return lines
|
||
|
||
|
||
def render_diff(path, old, new, width=0, colors=True):
|
||
if new is not None and (old or "") == new:
|
||
return ""
|
||
old_lines = split_diff_lines(old or "")
|
||
new_lines = split_diff_lines(new) if new is not None else []
|
||
if not width:
|
||
try:
|
||
width = shutil.get_terminal_size().columns
|
||
except OSError:
|
||
width = 80
|
||
matcher = difflib.SequenceMatcher(None, old_lines, new_lines, autojunk=False)
|
||
rows = []
|
||
for group in matcher.get_grouped_opcodes(DIFF_CONTEXT):
|
||
if rows:
|
||
rows.append(("gap", 0, 0, ""))
|
||
for tag, i1, i2, j1, j2 in group:
|
||
if tag == "equal":
|
||
for offset in range(i2 - i1):
|
||
rows.append(("ctx", i1 + offset + 1, j1 + offset + 1, old_lines[i1 + offset]))
|
||
elif tag == "delete":
|
||
for offset in range(i2 - i1):
|
||
rows.append(("del", i1 + offset + 1, 0, old_lines[i1 + offset]))
|
||
elif tag == "insert":
|
||
for offset in range(j2 - j1):
|
||
rows.append(("add", 0, j1 + offset + 1, new_lines[j1 + offset]))
|
||
else:
|
||
for offset in range(i2 - i1):
|
||
rows.append(("del", i1 + offset + 1, 0, old_lines[i1 + offset]))
|
||
for offset in range(j2 - j1):
|
||
rows.append(("add", 0, j1 + offset + 1, new_lines[j1 + offset]))
|
||
added = sum(1 for kind, _a, _b, _t in rows if kind == "add")
|
||
removed = sum(1 for kind, _a, _b, _t in rows if kind == "del")
|
||
if new is None:
|
||
head = "── %s (deleted, %d lines)" % (path, len(old_lines))
|
||
elif not old_lines:
|
||
head = "── %s (new file, %d lines)" % (path, len(new_lines))
|
||
else:
|
||
head = "── %s (+%d -%d)" % (path, added, removed)
|
||
python = str(path).lower().endswith(".py")
|
||
num_width = max(4, len(str(max(len(old_lines), len(new_lines), 1))))
|
||
code_width = max(20, width - num_width - 4)
|
||
fence = None
|
||
out = []
|
||
out.append(Ansi.DIM + "│ " + head + Ansi.RESET if colors else "│ " + head)
|
||
for kind, old_no, new_no, text in rows[:DIFF_MAX_LINES]:
|
||
if kind == "gap":
|
||
out.append(Ansi.DIM + "···" + Ansi.RESET if colors else "···")
|
||
continue
|
||
code = strip_ansi(text).expandtabs(8)
|
||
if fence is not None:
|
||
if python and colors:
|
||
code = Ansi.YELLOW + code + Ansi.RESET
|
||
if fence in text:
|
||
fence = None
|
||
elif python and colors:
|
||
for mark in ('"""', "'''"):
|
||
if text.count(mark) % 2 == 1:
|
||
fence = mark
|
||
code = Ansi.YELLOW + code + Ansi.RESET
|
||
break
|
||
else:
|
||
code = highlight_python(code)
|
||
number = old_no if kind in ("ctx", "del") else new_no
|
||
gutter = str(number).rjust(num_width)
|
||
if not colors:
|
||
sign = {"ctx": " ", "del": "-", "add": "+"}[kind]
|
||
out.append("%s %s %s" % (gutter, sign, code[:code_width]))
|
||
continue
|
||
dimmed = Ansi.DIM + gutter + Ansi.RESET
|
||
if kind == "ctx":
|
||
out.append("%s %s" % (dimmed, truncate_ansi(code, code_width)))
|
||
elif kind == "del":
|
||
row = "%s %s %s" % (dimmed, Ansi.RED + Ansi.BOLD + "-" + Ansi.RESET, truncate_ansi(code, code_width))
|
||
out.append(pad_tinted(style_span(row, BG_DEL), width))
|
||
else:
|
||
row = "%s %s %s" % (dimmed, Ansi.GREEN + Ansi.BOLD + "+" + Ansi.RESET, truncate_ansi(code, code_width))
|
||
out.append(pad_tinted(style_span(row, BG_ADD), width))
|
||
if len(rows) > DIFF_MAX_LINES:
|
||
note = "··· %d more lines hidden" % (len(rows) - DIFF_MAX_LINES)
|
||
out.append(Ansi.DIM + note + Ansi.RESET if colors else note)
|
||
return "\n".join(out)
|
||
|
||
|
||
def pad_tinted(line, width):
|
||
plain = len(ANSI_RE.sub("", line))
|
||
if plain < width:
|
||
fill = " " * (width - plain)
|
||
if line.endswith(Ansi.RESET):
|
||
return line[:-len(Ansi.RESET)] + fill + Ansi.RESET
|
||
return line + fill
|
||
return line
|
||
|
||
|
||
ANSI_RE = re.compile(r"\033\[[0-9;]*m")
|
||
MD_BULLETS = ("•", "◦", "▪")
|
||
|
||
|
||
def strip_ansi(text):
|
||
return ANSI_RE.sub("", text)
|
||
|
||
|
||
def style_text(text, *codes):
|
||
return "".join(codes) + text + Ansi.RESET
|
||
|
||
|
||
def style_span(text, *codes):
|
||
prefix = "".join(codes)
|
||
return prefix + text.replace(Ansi.RESET, Ansi.RESET + prefix) + Ansi.RESET
|
||
|
||
|
||
def wrap_ansi(text, width, first_prefix="", next_prefix=""):
|
||
parts = re.findall(r"\033\[[0-9;]*m|[^\s\033]+|\s+", text)
|
||
lines = []
|
||
current = first_prefix
|
||
length = len(strip_ansi(first_prefix))
|
||
base = length
|
||
active = ""
|
||
for part in parts:
|
||
if not part:
|
||
continue
|
||
if part.startswith("\033"):
|
||
current += part
|
||
if part == Ansi.RESET:
|
||
active = ""
|
||
else:
|
||
active += part
|
||
continue
|
||
if part.isspace():
|
||
if length < width and length > base:
|
||
current += " "
|
||
length += 1
|
||
continue
|
||
if length + len(part) > width and length > base:
|
||
if active:
|
||
current += Ansi.RESET
|
||
lines.append(current.rstrip())
|
||
current = next_prefix + active
|
||
length = len(strip_ansi(next_prefix))
|
||
base = length
|
||
current += part
|
||
length += len(part)
|
||
if active:
|
||
current += Ansi.RESET
|
||
if current.strip():
|
||
lines.append(current.rstrip())
|
||
return lines
|
||
|
||
|
||
def md_inline(text):
|
||
stashed = []
|
||
|
||
def stash(value):
|
||
stashed.append(value)
|
||
return "\x00%d\x00" % (len(stashed) - 1)
|
||
|
||
text = re.sub(r"`([^`\n]+)`", lambda match: stash(style_text(match.group(1), Ansi.CYAN)), text)
|
||
text = re.sub(r"\\([!\"#$%&'()*+,\-./:;<=>?@\[\\\]^_`{|}~])", lambda match: stash(match.group(1)), text)
|
||
text = re.sub(r"!\[([^\]\n]*)\]\(([^)\s\n]+)(?:\s+\"[^\"]*\")?\)", lambda match: stash("[image: %s]" % (match.group(1) or match.group(2))), text)
|
||
text = re.sub(r"\[([^\]\n]+)\]\(([^)\s\n]+)(?:\s+\"[^\"]*\")?\)", lambda match: style_span(md_inline(match.group(1)), Ansi.BLUE, Ansi.UNDERLINE) + style_text(" (%s)" % match.group(2), Ansi.DIM), text)
|
||
text = re.sub(r"\*\*([^*\n]+)\*\*", lambda match: style_text(match.group(1), Ansi.BOLD), text)
|
||
text = re.sub(r"(?<!\w)__([^_\n]+)__(?!\w)", lambda match: style_text(match.group(1), Ansi.BOLD), text)
|
||
text = re.sub(r"~~([^~\n]+)~~", lambda match: style_text(match.group(1), Ansi.STRIKE), text)
|
||
text = re.sub(r"\*([^*\n]+)\*", lambda match: style_text(match.group(1), Ansi.DIM), text)
|
||
text = re.sub(r"(?<!\w)_([^_\n]+)_(?!\w)", lambda match: style_text(match.group(1), Ansi.DIM), text)
|
||
for pos, value in enumerate(stashed):
|
||
text = text.replace("\x00%d\x00" % pos, value)
|
||
return text
|
||
|
||
|
||
def md_split_row(line):
|
||
line = line.strip()
|
||
if line.startswith("|"):
|
||
line = line[1:]
|
||
if line.endswith("|") and not line.endswith("\\|"):
|
||
line = line[:-1]
|
||
return [cell.replace("\\|", "|").strip() for cell in re.split(r"(?<!\\)\|", line)]
|
||
|
||
|
||
def md_is_separator(line):
|
||
cells = md_split_row(line)
|
||
return bool(cells) and all(re.fullmatch(r":?-{1,}:?", cell) for cell in cells)
|
||
|
||
|
||
def md_align_cell(cell, width, align):
|
||
pad = max(0, width - len(strip_ansi(cell)))
|
||
if align == "right":
|
||
return " " * pad + cell
|
||
if align == "center":
|
||
left = pad // 2
|
||
return " " * left + cell + " " * (pad - left)
|
||
return cell + " " * pad
|
||
|
||
|
||
def md_table_block(header, aligns, rows):
|
||
span = max([len(header)] + [len(row) for row in rows])
|
||
header = (header + [""] * span)[:span]
|
||
rows = [(row + [""] * span)[:span] for row in rows]
|
||
head = [md_inline(cell) for cell in header]
|
||
body = [[md_inline(cell) for cell in row] for row in rows]
|
||
widths = [0] * span
|
||
for pos in range(span):
|
||
widths[pos] = max([len(strip_ansi(head[pos]))] + [len(strip_ansi(row[pos])) for row in body])
|
||
bar = style_text("│", Ansi.DIM)
|
||
lines = [style_text("┌" + "┬".join("─" * (width + 2) for width in widths) + "┐", Ansi.DIM)]
|
||
lines.append(bar + bar.join(" " + style_span(md_align_cell(cell, widths[pos], aligns[pos]), Ansi.BOLD) + " " for pos, cell in enumerate(head)) + bar)
|
||
lines.append(style_text("├" + "┼".join("─" * (width + 2) for width in widths) + "┤", Ansi.DIM))
|
||
for row in body:
|
||
lines.append(bar + bar.join(" " + md_align_cell(cell, widths[pos], aligns[pos]) + " " for pos, cell in enumerate(row)) + bar)
|
||
lines.append(style_text("└" + "┴".join("─" * (width + 2) for width in widths) + "┘", Ansi.DIM))
|
||
return lines
|
||
|
||
|
||
def render_markdown(text, width=None, color=None):
|
||
if color is None:
|
||
color = not os.environ.get("NO_COLOR") and sys.stdout.isatty()
|
||
if not color:
|
||
return text
|
||
if width is None:
|
||
try:
|
||
width = shutil.get_terminal_size().columns
|
||
except OSError:
|
||
width = 80
|
||
text = (text or "").replace("\r\n", "\n").expandtabs(8)
|
||
text = re.sub(r"<br\s*/?>", " \n", text, flags=re.IGNORECASE)
|
||
raw = text.split("\n")
|
||
out = []
|
||
pos = 0
|
||
total = len(raw)
|
||
while pos < total:
|
||
line = raw[pos]
|
||
stripped = line.strip()
|
||
if not stripped:
|
||
pos += 1
|
||
continue
|
||
fence = re.match(r"^\s*(`{3,}|~{3,})\s*\S*\s*$", line)
|
||
if fence:
|
||
char = fence.group(1)[0]
|
||
pos += 1
|
||
while pos < total:
|
||
closer = raw[pos].strip()
|
||
if closer and set(closer) == {char} and len(closer) >= 3:
|
||
pos += 1
|
||
break
|
||
out.append(" " + style_text(raw[pos], Ansi.DIM) if raw[pos].strip() else "")
|
||
pos += 1
|
||
out.append("")
|
||
continue
|
||
if "|" in line and pos + 1 < total and md_is_separator(raw[pos + 1]):
|
||
header = md_split_row(line)
|
||
aligns = []
|
||
for cell in md_split_row(raw[pos + 1]):
|
||
if cell.startswith(":") and cell.endswith(":") and len(cell) > 2:
|
||
aligns.append("center")
|
||
elif cell.endswith(":"):
|
||
aligns.append("right")
|
||
else:
|
||
aligns.append("left")
|
||
pos += 2
|
||
rows = []
|
||
while pos < total and "|" in raw[pos] and raw[pos].strip():
|
||
rows.append(md_split_row(raw[pos]))
|
||
pos += 1
|
||
while len(aligns) < max([len(header)] + [len(row) for row in rows]):
|
||
aligns.append("left")
|
||
out.extend(md_table_block(header, aligns, rows))
|
||
out.append("")
|
||
continue
|
||
heading = re.match(r"^(#{1,6})\s+(.*\S)\s*$", stripped)
|
||
if heading:
|
||
level = len(heading.group(1))
|
||
body = md_inline(heading.group(2))
|
||
if level == 1:
|
||
out.append(style_span(body, Ansi.BOLD, Ansi.CYAN))
|
||
elif level == 2:
|
||
out.append(style_span(body, Ansi.BOLD))
|
||
else:
|
||
out.append(style_span(body, Ansi.BOLD, Ansi.DIM))
|
||
out.append("")
|
||
pos += 1
|
||
continue
|
||
if re.fullmatch(r"\*{3,}|-{3,}|_{3,}", stripped.replace(" ", "")):
|
||
out.append(style_text("─" * width, Ansi.DIM))
|
||
out.append("")
|
||
pos += 1
|
||
continue
|
||
if stripped.startswith(">"):
|
||
paras = []
|
||
current = []
|
||
while pos < total and raw[pos].strip().startswith(">"):
|
||
inner = re.sub(r"^\s*(?:>\s?)+", "", raw[pos])
|
||
if inner.strip():
|
||
current.append(inner.strip())
|
||
elif current:
|
||
paras.append(current)
|
||
current = []
|
||
pos += 1
|
||
if current:
|
||
paras.append(current)
|
||
for index, para in enumerate(paras):
|
||
if index:
|
||
out.append(style_text("│", Ansi.DIM))
|
||
prefix = style_text("│ ", Ansi.DIM)
|
||
for wrapped in wrap_ansi(md_inline(" ".join(para)), width, prefix, prefix):
|
||
out.append(wrapped)
|
||
out.append("")
|
||
continue
|
||
bullet = re.match(r"^(\s*)[-*+]\s+(.*\S)\s*$", line)
|
||
numbered = re.match(r"^(\s*)\d+[.)]\s+(.*\S)\s*$", line)
|
||
if bullet or numbered:
|
||
counters = {}
|
||
while pos < total:
|
||
again_bullet = re.match(r"^(\s*)[-*+]\s+(.*\S)\s*$", raw[pos])
|
||
again_numbered = re.match(r"^(\s*)\d+[.)]\s+(.*\S)\s*$", raw[pos])
|
||
if not again_bullet and not again_numbered:
|
||
break
|
||
if again_bullet:
|
||
level = len(again_bullet.group(1)) // 2
|
||
marker = MD_BULLETS[level % len(MD_BULLETS)]
|
||
body = again_bullet.group(2)
|
||
check = re.match(r"^\[([ xX])\]\s+(.*\S)\s*$", body)
|
||
if check:
|
||
marker = "☑" if check.group(1).lower() == "x" else "☐"
|
||
body = check.group(2)
|
||
else:
|
||
level = len(again_numbered.group(1)) // 2
|
||
counters[level] = counters.get(level, 0) + 1
|
||
for deeper in [key for key in counters if key > level]:
|
||
del counters[deeper]
|
||
marker = "%d." % counters[level]
|
||
body = again_numbered.group(2)
|
||
pad = " " * level
|
||
for wrapped in wrap_ansi(md_inline(body), width, pad + marker + " ", pad + " " * (len(marker) + 1)):
|
||
out.append(wrapped)
|
||
pos += 1
|
||
out.append("")
|
||
continue
|
||
para = []
|
||
while pos < total and raw[pos].strip():
|
||
para.append(raw[pos])
|
||
pos += 1
|
||
segments = []
|
||
current = []
|
||
for chunk in para:
|
||
current.append(chunk.strip())
|
||
if chunk.endswith(" "):
|
||
segments.append(" ".join(current))
|
||
current = []
|
||
if current:
|
||
segments.append(" ".join(current))
|
||
for segment in segments:
|
||
for wrapped in wrap_ansi(md_inline(segment), width):
|
||
out.append(wrapped)
|
||
out.append("")
|
||
while out and not out[-1]:
|
||
out.pop()
|
||
return "\n".join(out)
|
||
|
||
|
||
class SealError(Exception):
|
||
pass
|
||
|
||
|
||
class Seal:
|
||
PREFIX = "tai1$"
|
||
ROUNDS = 200000
|
||
|
||
def __init__(self, home, passphrase):
|
||
self.enabled = bool(passphrase)
|
||
self.default_key = False
|
||
self.master = b""
|
||
self.enc_key = b""
|
||
self.mac_key = b""
|
||
if self.enabled:
|
||
self.master = self.load_master(home, passphrase)
|
||
self.enc_key = hmac.new(self.master, b"tai-enc-1", hashlib.sha256).digest()
|
||
self.mac_key = hmac.new(self.master, b"tai-mac-1", hashlib.sha256).digest()
|
||
|
||
def load_master(self, home, passphrase):
|
||
path = os.path.join(home, ".seal")
|
||
raw = ""
|
||
if os.path.exists(path):
|
||
with open(path, "r", encoding="utf-8") as handle:
|
||
raw = handle.read()
|
||
saved = {}
|
||
if raw:
|
||
try:
|
||
saved = json.loads(raw)
|
||
salt = base64.b64decode(saved["salt"])
|
||
except (ValueError, KeyError):
|
||
raise SealError("seal file is corrupt")
|
||
else:
|
||
salt = os.urandom(16)
|
||
master = hashlib.pbkdf2_hmac("sha256", passphrase.encode("utf-8"), salt, self.ROUNDS, 32)
|
||
check = hmac.new(master, b"tai-check-1", hashlib.sha256).hexdigest()
|
||
if raw:
|
||
if not hmac.compare_digest(saved.get("check", ""), check):
|
||
raise SealError("wrong TAI_PASSPHRASE")
|
||
return master
|
||
with open(path, "w", encoding="utf-8") as handle:
|
||
handle.write(json.dumps({"salt": base64.b64encode(salt).decode("ascii"), "check": check}))
|
||
try:
|
||
os.chmod(path, 0o600)
|
||
except OSError:
|
||
pass
|
||
return master
|
||
|
||
def stream(self, nonce, length):
|
||
out = bytearray()
|
||
counter = 0
|
||
while len(out) < length:
|
||
out += hmac.new(self.enc_key, nonce + struct.pack(">I", counter), hashlib.sha256).digest()
|
||
counter += 1
|
||
return bytes(out[:length])
|
||
|
||
def lock(self, text):
|
||
if not self.enabled or text.startswith(self.PREFIX):
|
||
return text
|
||
nonce = os.urandom(16)
|
||
body = text.encode("utf-8")
|
||
keystream = self.stream(nonce, len(body))
|
||
cipher = bytes(piece ^ keystream[pos] for pos, piece in enumerate(body))
|
||
mac = hmac.new(self.mac_key, nonce + cipher, hashlib.sha256).digest()
|
||
parts = [base64.b64encode(blob).decode("ascii") for blob in (nonce, cipher, mac)]
|
||
return self.PREFIX + "$".join(parts)
|
||
|
||
def unlock(self, text):
|
||
if not text.startswith(self.PREFIX):
|
||
raise SealError("not sealed")
|
||
try:
|
||
_tag, nonce64, cipher64, mac64 = text.split("$", 3)
|
||
nonce = base64.b64decode(nonce64)
|
||
cipher = base64.b64decode(cipher64)
|
||
mac = base64.b64decode(mac64)
|
||
except ValueError:
|
||
raise SealError("sealed value is corrupt")
|
||
expect = hmac.new(self.mac_key, nonce + cipher, hashlib.sha256).digest()
|
||
if not hmac.compare_digest(mac, expect):
|
||
raise SealError("seal check failed")
|
||
keystream = self.stream(nonce, len(cipher))
|
||
plain = bytes(piece ^ keystream[pos] for pos, piece in enumerate(cipher))
|
||
return plain.decode("utf-8")
|
||
|
||
def safe_unlock(self, text):
|
||
try:
|
||
return self.unlock(text)
|
||
except SealError:
|
||
return "[sealed: decrypt failed]"
|
||
|
||
|
||
def resolve_passphrase():
|
||
raw = os.environ.get("TAI_PASSPHRASE")
|
||
if raw is None:
|
||
return DEFAULT_PASSPHRASE, True
|
||
return raw, False
|
||
|
||
|
||
def rotate_seal(config, old_seal, new_passphrase):
|
||
db = sqlite3.connect(config.db_path)
|
||
rows = db.execute("SELECT id, text FROM events").fetchall()
|
||
tables = [row[0] for row in db.execute("SELECT name FROM sqlite_master WHERE type = 'table'").fetchall()]
|
||
secret_cols = [row[1] for row in db.execute("PRAGMA table_info(secrets)").fetchall()] if "secrets" in tables else []
|
||
has_profile = "profile" in secret_cols
|
||
secret_rows = db.execute("SELECT profile, name, value FROM secrets").fetchall() if has_profile else [(None, name, value) for name, value in db.execute("SELECT name, value FROM secrets").fetchall()] if "secrets" in tables else []
|
||
meta_rows = db.execute("SELECT profile, name, meta FROM secrets WHERE meta IS NOT NULL").fetchall() if has_profile and "meta" in secret_cols else [(None, name, meta) for name, meta in db.execute("SELECT name, meta FROM secrets WHERE meta IS NOT NULL").fetchall()] if "meta" in secret_cols else []
|
||
planned_rows = db.execute("SELECT id, prompt FROM schedules").fetchall() if "schedules" in tables else []
|
||
db.close()
|
||
texts = {}
|
||
for row_id, text in rows:
|
||
texts[row_id] = old_seal.unlock(text) if text.startswith(Seal.PREFIX) else text
|
||
secret_texts = {}
|
||
for profile, name, value in secret_rows:
|
||
secret_texts[(profile, name)] = old_seal.unlock(value) if value.startswith(Seal.PREFIX) else value
|
||
secret_metas = {}
|
||
for profile, name, meta in meta_rows:
|
||
secret_metas[(profile, name)] = old_seal.unlock(meta) if meta.startswith(Seal.PREFIX) else meta
|
||
planned_texts = {}
|
||
for row_id, prompt in planned_rows:
|
||
planned_texts[row_id] = old_seal.unlock(prompt) if prompt.startswith(Seal.PREFIX) else prompt
|
||
files = {}
|
||
for entry in os.listdir(config.profiles_dir):
|
||
if not entry.endswith((".sys.md", ".session.json", ".bots.json")):
|
||
continue
|
||
path = os.path.join(config.profiles_dir, entry)
|
||
with open(path, "r", encoding="utf-8") as handle:
|
||
raw = handle.read()
|
||
files[path] = old_seal.unlock(raw) if raw.startswith(Seal.PREFIX) else raw
|
||
try:
|
||
os.remove(os.path.join(config.home, ".seal"))
|
||
except OSError:
|
||
pass
|
||
fresh = Seal(config.home, new_passphrase)
|
||
db = sqlite3.connect(config.db_path)
|
||
for row_id, plain in texts.items():
|
||
db.execute("UPDATE events SET text = ? WHERE id = ?", (fresh.lock(plain), row_id))
|
||
for (profile, name), plain in secret_texts.items():
|
||
if profile is None:
|
||
db.execute("UPDATE secrets SET value = ? WHERE name = ?", (fresh.lock(plain), name))
|
||
else:
|
||
db.execute("UPDATE secrets SET value = ? WHERE profile = ? AND name = ?", (fresh.lock(plain), profile, name))
|
||
for (profile, name), plain in secret_metas.items():
|
||
if profile is None:
|
||
db.execute("UPDATE secrets SET meta = ? WHERE name = ?", (fresh.lock(plain), name))
|
||
else:
|
||
db.execute("UPDATE secrets SET meta = ? WHERE profile = ? AND name = ?", (fresh.lock(plain), profile, name))
|
||
for row_id, plain in planned_texts.items():
|
||
db.execute("UPDATE schedules SET prompt = ? WHERE id = ?", (fresh.lock(plain), row_id))
|
||
db.commit()
|
||
db.close()
|
||
for path, plain in files.items():
|
||
with open(path, "w", encoding="utf-8") as handle:
|
||
handle.write(fresh.lock(plain))
|
||
return fresh
|
||
|
||
|
||
def parse_skill_front(lines):
|
||
front = {}
|
||
pos = 0
|
||
while pos < len(lines):
|
||
line = lines[pos]
|
||
if ":" in line and not line.startswith((" ", "\t")):
|
||
key, value = line.split(":", 1)
|
||
value = value.strip()
|
||
if value in (">", "|"):
|
||
joined = []
|
||
pos += 1
|
||
while pos < len(lines) and lines[pos].startswith((" ", "\t")):
|
||
joined.append(lines[pos].strip())
|
||
pos += 1
|
||
front[key.strip()] = (" " if value == ">" else "\n").join(joined)
|
||
continue
|
||
front[key.strip()] = value.strip("'\"")
|
||
pos += 1
|
||
return front
|
||
|
||
|
||
def parse_skill_file(path):
|
||
with open(path, "r", encoding="utf-8") as handle:
|
||
raw = handle.read()
|
||
if not raw.startswith("---"):
|
||
return None
|
||
end = raw.find("\n---", 3)
|
||
if end < 0:
|
||
return None
|
||
front = parse_skill_front(raw[3:end].strip().splitlines())
|
||
name = front.get("name", "")
|
||
description = front.get("description", "")
|
||
if not name or not description or not SKILL_NAME_RE.match(name):
|
||
return None
|
||
return {"name": name, "description": description, "body": raw[end + 4:].lstrip("\n")}
|
||
|
||
|
||
EXTERNAL_SKILL_DIRNAMES = (os.path.join(".claude", "skills"), os.path.join(".agents", "skills"))
|
||
SKILL_INSTALL_META_DIRNAME = ".installs"
|
||
|
||
|
||
def skill_scope_dir(scope, home_dir, project_dir, name=""):
|
||
"""Shared by create_skill and install_skill so both land skills in the same place."""
|
||
base = os.path.join(project_dir, ".tai", "skills") if scope == "project" else os.path.join(home_dir, "skills")
|
||
return os.path.join(base, name) if name else base
|
||
|
||
|
||
def default_skill_scope(project_dir):
|
||
for marker in (".git", ".tai") + EXTERNAL_SKILL_DIRNAMES:
|
||
if os.path.exists(os.path.join(project_dir, marker)):
|
||
return "project"
|
||
return "home"
|
||
|
||
|
||
def skill_extra_files(root):
|
||
extras = []
|
||
for sub in ("scripts", "references", "assets"):
|
||
folder = os.path.join(root, sub)
|
||
if os.path.isdir(folder):
|
||
for entry in sorted(os.listdir(folder)):
|
||
extras.append(os.path.join(root, sub, entry))
|
||
return extras
|
||
|
||
|
||
def skill_install_meta_path(scope_dir, entry):
|
||
return os.path.join(scope_dir, SKILL_INSTALL_META_DIRNAME, entry + ".json")
|
||
|
||
|
||
def load_skill_install_meta(scope_dir, entry):
|
||
try:
|
||
with open(skill_install_meta_path(scope_dir, entry), "r", encoding="utf-8") as handle:
|
||
return json.load(handle)
|
||
except (OSError, ValueError):
|
||
return None
|
||
|
||
|
||
def save_skill_install_meta(scope_dir, entry, meta):
|
||
path = skill_install_meta_path(scope_dir, entry)
|
||
os.makedirs(os.path.dirname(path), exist_ok=True)
|
||
with open(path, "w", encoding="utf-8") as handle:
|
||
json.dump(meta, handle)
|
||
|
||
|
||
def is_external_skill_root(path):
|
||
return os.path.dirname(os.path.normpath(path)).endswith(EXTERNAL_SKILL_DIRNAMES)
|
||
|
||
|
||
def skill_provenance_note(skill):
|
||
install = skill.get("install")
|
||
if install:
|
||
return "[installed %s, %s scope, from %s]" % (install.get("mode", "copy"), install.get("scope", "?"), install.get("source", "?"))
|
||
if not skill.get("managed", True):
|
||
return "[external, not managed by tai: %s]" % skill["root"]
|
||
return "[native]"
|
||
|
||
|
||
def discover_skills(home_dir, project_dir):
|
||
found = {}
|
||
tai_roots = (os.path.join(home_dir, "skills"), os.path.join(project_dir, ".tai", "skills"))
|
||
external_roots = (os.path.join(os.path.expanduser("~"), ".claude", "skills"),) + tuple(os.path.join(project_dir, dirname) for dirname in EXTERNAL_SKILL_DIRNAMES)
|
||
# External roots scan first so tai's own dirs always win on a name collision (dict assignment below overwrites).
|
||
for base in external_roots + tai_roots:
|
||
if not os.path.isdir(base):
|
||
continue
|
||
managed = base in tai_roots
|
||
for entry in sorted(os.listdir(base)):
|
||
path = os.path.join(base, entry, "SKILL.md")
|
||
if not os.path.isfile(path):
|
||
continue
|
||
try:
|
||
skill = parse_skill_file(path)
|
||
except OSError:
|
||
continue
|
||
if skill is None:
|
||
continue
|
||
skill["root"] = os.path.join(base, entry)
|
||
skill["managed"] = managed
|
||
skill["install"] = load_skill_install_meta(base, entry) if managed else None
|
||
found[skill["name"]] = skill
|
||
return found
|
||
|
||
|
||
SKILL_BLUEPRINTS = {
|
||
"bot-creator": {
|
||
"hint": "build a dedicated chat bot from a tai.py variant",
|
||
"scope": "project",
|
||
"brief": "Deep-research preferred chat-bot features first: transport (Telegram long-poll versus webhook), command routing, per-user allowlists, rate limits, and safe restarts. Also research safe techniques for scripting huge single-file programs: exact-match edits, atomic temp-plus-rename writes, and versioned backups before every reshape. Then copy tai.py to a variant path and reshape the copy into a dedicated bot, editing efficiently and keeping the audit trail intact.",
|
||
},
|
||
"api-client": {
|
||
"hint": "build a resilient REST/OpenAPI client",
|
||
"scope": "project",
|
||
"brief": "Deep-research resilient HTTP client design first: retries with backoff, timeouts, pagination, auth header handling, and error classification. Then write a skill that builds small stdlib-only API clients from a base URL and an endpoint list.",
|
||
},
|
||
"web-researcher": {
|
||
"hint": "deep research with triangulated sources",
|
||
"scope": "project",
|
||
"brief": "Deep-research professional open-source research methodology first: query expansion, source triangulation, recency checks, and claim grading. Then write a skill that turns a question into a sourced brief with confidence levels.",
|
||
},
|
||
"pdf-forms": {
|
||
"hint": "fill PDF forms and extract field data",
|
||
"scope": "project",
|
||
"brief": "Deep-research PDF form handling with freely available tooling first: AcroForm field discovery, filling, flattening, and text extraction. Then write a skill that fills forms and extracts field data from this machine's installed tools.",
|
||
},
|
||
"data-wrangler": {
|
||
"hint": "reshape CSV, JSON, and SQLite data",
|
||
"scope": "project",
|
||
"brief": "Deep-research tabular data reshaping first: CSV dialects, JSON normalization, SQLite import ergonomics, and streaming for large files. Then write a skill that reshapes CSV, JSON, and SQLite inputs using stdlib-only scripts.",
|
||
},
|
||
"home-sysadmin": {
|
||
"hint": "routine Linux host care and triage",
|
||
"scope": "home",
|
||
"brief": "Deep-research routine single-host Linux care first: disk, memory, service health, log triage, and backup verification. Then write a skill that runs a host checkup and reports findings with suggested fixes.",
|
||
},
|
||
}
|
||
|
||
|
||
def skill_catalog(skills):
|
||
lines = []
|
||
if skills:
|
||
lines += ["", "", "## Available skills"]
|
||
for name in sorted(skills):
|
||
lines.append("- %s: %s" % (name, skills[name]["description"][:300]))
|
||
wanted = sorted(name for name in SKILL_BLUEPRINTS if name not in (skills or {}))
|
||
if wanted:
|
||
lines += ["", "", "## Buildable skill blueprints (load_skill builds one on demand)"]
|
||
for name in wanted:
|
||
lines.append("- %s: %s" % (name, SKILL_BLUEPRINTS[name]["hint"]))
|
||
return "\n".join(lines)
|
||
lines.append("Load full instructions with load_skill before using one.")
|
||
return "\n".join(lines)
|
||
|
||
|
||
def container_engine():
|
||
if shutil.which("podman"):
|
||
return "podman"
|
||
if shutil.which("docker"):
|
||
return "docker"
|
||
return ""
|
||
|
||
|
||
def box_state(engine):
|
||
try:
|
||
done = subprocess.run([engine, "inspect", "-f", "{{.State.Running}}", BOX_NAME], stdout=subprocess.PIPE, stderr=subprocess.PIPE, universal_newlines=True, timeout=15)
|
||
except (OSError, subprocess.SubprocessError):
|
||
return "missing"
|
||
if done.returncode != 0:
|
||
return "missing"
|
||
return "running" if done.stdout.strip() == "true" else "stopped"
|
||
|
||
|
||
def ensure_box():
|
||
engine = container_engine()
|
||
if not engine:
|
||
raise BackendError("no container engine found, install podman")
|
||
state = box_state(engine)
|
||
if state == "running":
|
||
return engine
|
||
if state == "stopped":
|
||
done = subprocess.run([engine, "start", BOX_NAME], stdout=subprocess.PIPE, stderr=subprocess.PIPE, universal_newlines=True, timeout=60)
|
||
if done.returncode != 0:
|
||
raise BackendError("cannot start box: " + (done.stderr or "").strip()[:200])
|
||
return engine
|
||
print("building sandbox image %s (one-time, takes minutes)..." % BOX_IMAGE)
|
||
with tempfile.TemporaryDirectory() as build:
|
||
for filename, content in (("Containerfile", BOX_CONTAINERFILE), ("stt.py", BOX_STT), ("tts.py", BOX_TTS)):
|
||
with open(os.path.join(build, filename), "w", encoding="utf-8") as handle:
|
||
handle.write(content)
|
||
done = subprocess.run([engine, "build", "-t", BOX_IMAGE, build], timeout=1800)
|
||
if done.returncode != 0:
|
||
raise BackendError("box image build failed")
|
||
done = subprocess.run([engine, "run", "-d", "--name", BOX_NAME, "--restart", "unless-stopped", BOX_IMAGE], stdout=subprocess.PIPE, stderr=subprocess.PIPE, universal_newlines=True, timeout=120)
|
||
if done.returncode != 0:
|
||
raise BackendError("cannot start box: " + (done.stderr or "").strip()[:200])
|
||
return engine
|
||
|
||
|
||
def box_exec(engine, argv, extra=(), input_data=None, timeout=180):
|
||
return subprocess.run([engine, "exec", *extra, "-i", BOX_NAME, *argv], input=input_data, stdout=subprocess.PIPE, stderr=subprocess.PIPE, timeout=timeout)
|
||
|
||
|
||
def box_python(engine):
|
||
if engine not in _BOX_PYTHON_CACHE:
|
||
try:
|
||
probe = box_exec(engine, ["test", "-x", BOX_PYTHON], timeout=30)
|
||
_BOX_PYTHON_CACHE[engine] = BOX_PYTHON if probe.returncode == 0 else "python3"
|
||
except (OSError, subprocess.SubprocessError):
|
||
_BOX_PYTHON_CACHE[engine] = "python3"
|
||
return _BOX_PYTHON_CACHE[engine]
|
||
|
||
|
||
def box_transcribe(engine, audio, timeout=300):
|
||
done = box_exec(engine, [box_python(engine), "/box/stt.py"], input_data=audio, timeout=timeout)
|
||
if done.returncode != 0:
|
||
raise BackendError("box stt failed: " + done.stderr.decode("utf-8", "replace")[:300])
|
||
return done.stdout.decode("utf-8", "replace").strip()
|
||
|
||
|
||
def box_speak(engine, text, voice=""):
|
||
argv = [box_python(engine), "/box/tts.py"] + ([voice] if voice else [])
|
||
done = box_exec(engine, argv, input_data=text.encode("utf-8"), timeout=180)
|
||
if done.returncode != 0 or not done.stdout:
|
||
raise BackendError("box tts failed: " + done.stderr.decode("utf-8", "replace")[:300])
|
||
return done.stdout
|
||
|
||
|
||
def split_models(spec):
|
||
return [item.strip() for item in (spec or "").replace("\n", ",").split(",") if item.strip()]
|
||
|
||
|
||
def parse_backends(spec):
|
||
found = []
|
||
for entry in (spec or "").replace("\n", ";").split(";"):
|
||
fields = [field.strip() for field in entry.split("|")]
|
||
if len(fields) < 3:
|
||
continue
|
||
label, base, key = fields[0], fields[1].rstrip("/"), fields[2]
|
||
if key.startswith("$"):
|
||
key = os.environ.get(key[1:], "")
|
||
model = fields[3].strip() if len(fields) > 3 else ""
|
||
profile = fields[4].strip().lower() if len(fields) > 4 else ""
|
||
if profile not in ("", "opencode"):
|
||
profile = ""
|
||
if not label or not base:
|
||
continue
|
||
found.append({"label": label, "base": base, "key": key, "model": model or "", "profile": profile})
|
||
return found
|
||
|
||
|
||
class Config:
|
||
def __init__(self, args):
|
||
self.home = os.environ.get("TAI_HOME") or os.path.join(os.path.expanduser("~"), ".tai")
|
||
self.profiles_dir = os.path.join(self.home, "profiles")
|
||
self.audio_dir = os.path.join(self.home, "audio")
|
||
self.db_path = os.path.join(self.home, "memory.db")
|
||
self.history_path = os.path.join(self.home, "history")
|
||
self.model = os.environ.get("TAI_MODEL") or DEFAULT_MODEL
|
||
self.explicit_model = os.environ.get("TAI_MODEL") or ""
|
||
self.backends = parse_backends(os.environ.get("TAI_BACKENDS"))
|
||
self.openrouter_key = os.environ.get("OPENROUTER_API_KEY") or ""
|
||
self.openrouter_models = split_models(os.environ.get("TAI_OPENROUTER_MODELS"))
|
||
self.retry_seconds = float(os.environ.get("TAI_RETRY_SECONDS") or 0) if re.match(r"^\d+(\.\d+)?$", (os.environ.get("TAI_RETRY_SECONDS") or "0").strip()) else 0.0
|
||
self.opencode_enabled = (os.environ.get("TAI_OPENCODE") or "1").strip().lower() not in ("0", "false", "no", "off")
|
||
self.opencode_models = split_models(os.environ.get("TAI_OPENCODE_MODELS"))
|
||
self.devplace_key = os.environ.get("DEVPLACE_API_KEY") or ""
|
||
self.devplace_base = (os.environ.get("TAI_DEVPLACE_BASE") or DEVPLACE_BASE).rstrip("/")
|
||
self.devplace_models = split_models(os.environ.get("TAI_DEVPLACE_MODELS"))
|
||
self.openrouter_enabled = (os.environ.get("TAI_OPENROUTER") or "1").strip().lower() not in ("0", "false", "no", "off")
|
||
self.pollinations_enabled = (os.environ.get("TAI_POLLINATIONS") or "1").strip().lower() not in ("0", "false", "no", "off")
|
||
self.pollinations_models = split_models(os.environ.get("TAI_POLLINATIONS_MODELS")) or list(POLLINATIONS_MODELS)
|
||
self.zen_enabled = (os.environ.get("TAI_ZEN") or "").strip().lower() in ("1", "true", "yes", "on")
|
||
self.zen_models = split_models(os.environ.get("TAI_ZEN_MODELS")) or list(ZEN_FREE_MODELS)
|
||
self.voice = os.environ.get("TAI_VOICE") or DEFAULT_VOICE
|
||
self.auto_approve = bool(args.yes)
|
||
self.profile = args.profile or "default"
|
||
os.makedirs(self.profiles_dir, exist_ok=True)
|
||
os.makedirs(self.audio_dir, exist_ok=True)
|
||
|
||
|
||
def ensure_column(db, table, column, ddl):
|
||
names = [row[1] for row in db.execute("PRAGMA table_info(%s)" % table).fetchall()]
|
||
if column not in names:
|
||
db.execute("ALTER TABLE %s ADD COLUMN %s %s" % (table, column, ddl))
|
||
|
||
|
||
def fts_escape(query):
|
||
terms = re.findall(r"[a-z0-9]+", (query or "").lower())
|
||
if not terms:
|
||
return None
|
||
return ['"%s"' % term for term in terms]
|
||
|
||
|
||
def fts_match(terms):
|
||
return " AND ".join("{title body} : %s" % term for term in terms)
|
||
|
||
|
||
MEM_EVENTS_CAP = 10000
|
||
|
||
|
||
def open_mem_events():
|
||
mem = sqlite3.connect(":memory:")
|
||
try:
|
||
mem.execute("CREATE VIRTUAL TABLE mem_events USING fts5(item, kind, sub, title, body, profile, tokenize='porter unicode61 remove_diacritics 2')")
|
||
except sqlite3.Error:
|
||
mem.execute("CREATE VIRTUAL TABLE mem_events USING fts5(item, kind, sub, title, body, profile)")
|
||
return mem
|
||
|
||
|
||
def setup_fts(db):
|
||
try:
|
||
exists = db.execute("SELECT 1 FROM sqlite_master WHERE type = 'table' AND name = 'fts_docs'").fetchone()
|
||
if not exists:
|
||
try:
|
||
db.execute("CREATE VIRTUAL TABLE fts_docs USING fts5(item, kind, sub, title, body, profile, tokenize='porter unicode61 remove_diacritics 2')")
|
||
except sqlite3.Error:
|
||
db.execute("CREATE VIRTUAL TABLE fts_docs USING fts5(item, kind, sub, title, body, profile)")
|
||
db.execute("""CREATE TRIGGER IF NOT EXISTS trg_records_ai AFTER INSERT ON records BEGIN
|
||
INSERT INTO fts_docs(item, kind, sub, title, body, profile) VALUES(new.id, 'record', new.kind, new.title, new.content, new.profile);
|
||
END""")
|
||
db.execute("""CREATE TRIGGER IF NOT EXISTS trg_records_au AFTER UPDATE ON records WHEN old.title != new.title OR old.content != new.content OR old.profile != new.profile BEGIN
|
||
DELETE FROM fts_docs WHERE item = old.id AND kind = 'record';
|
||
INSERT INTO fts_docs(item, kind, sub, title, body, profile) VALUES(new.id, 'record', new.kind, new.title, new.content, new.profile);
|
||
END""")
|
||
db.execute("""CREATE TRIGGER IF NOT EXISTS trg_records_ad AFTER DELETE ON records BEGIN
|
||
DELETE FROM fts_docs WHERE item = old.id AND kind = 'record';
|
||
END""")
|
||
db.execute("""CREATE TRIGGER IF NOT EXISTS trg_events_ai AFTER INSERT ON events BEGIN
|
||
INSERT INTO fts_docs(item, kind, sub, title, body, profile) VALUES(new.id, 'event', new.role || '/' || new.kind, new.ts || ' ' || new.role || '/' || new.kind, new.text, new.profile);
|
||
END""")
|
||
db.execute("""CREATE TRIGGER IF NOT EXISTS trg_events_au AFTER UPDATE ON events WHEN old.text != new.text OR old.profile != new.profile BEGIN
|
||
DELETE FROM fts_docs WHERE item = old.id AND kind = 'event';
|
||
INSERT INTO fts_docs(item, kind, sub, title, body, profile) VALUES(new.id, 'event', new.role || '/' || new.kind, new.ts || ' ' || new.role || '/' || new.kind, new.text, new.profile);
|
||
END""")
|
||
db.execute("""CREATE TRIGGER IF NOT EXISTS trg_audit_ai AFTER INSERT ON audit BEGIN
|
||
INSERT INTO fts_docs(item, kind, sub, title, body, profile) VALUES('audit:' || new.id, 'audit', new.action, new.path, new.message, new.profile);
|
||
END""")
|
||
db.execute("""CREATE TRIGGER IF NOT EXISTS trg_chunks_ai AFTER INSERT ON chunks BEGIN
|
||
INSERT INTO fts_docs(item, kind, sub, title, body, profile) VALUES('chunk:' || new.id, 'chunk', new.parent_kind, new.parent_id, new.text, new.profile);
|
||
END""")
|
||
db.execute("""CREATE TRIGGER IF NOT EXISTS trg_chunks_ad AFTER DELETE ON chunks BEGIN
|
||
DELETE FROM fts_docs WHERE item = 'chunk:' || old.id AND kind = 'chunk';
|
||
END""")
|
||
if db.execute("SELECT COUNT(*) FROM fts_docs").fetchone()[0] == 0:
|
||
db.execute("INSERT INTO fts_docs(item, kind, sub, title, body, profile) SELECT id, 'record', kind, title, content, profile FROM records")
|
||
db.execute("INSERT INTO fts_docs(item, kind, sub, title, body, profile) SELECT id, 'event', role || '/' || kind, ts || ' ' || role || '/' || kind, text, profile FROM events")
|
||
db.execute("INSERT INTO fts_docs(item, kind, sub, title, body, profile) SELECT 'audit:' || id, 'audit', action, path, message, profile FROM audit")
|
||
db.commit()
|
||
return True
|
||
except sqlite3.Error:
|
||
return False
|
||
|
||
|
||
def table_pk_columns(db, table):
|
||
for row in db.execute("PRAGMA index_list(%s)" % table).fetchall():
|
||
if row[3] == "pk":
|
||
return [info[2] for info in db.execute("PRAGMA index_info(%s)" % row[1]).fetchall()]
|
||
return []
|
||
|
||
|
||
def rebuild_with_profile(db, table, columns, pk):
|
||
if table_pk_columns(db, table) == list(pk):
|
||
return False
|
||
names = [column.split()[0] for column in columns]
|
||
have = {row[1] for row in db.execute("PRAGMA table_info(%s)" % table).fetchall()}
|
||
keep = [name for name in names if name in have]
|
||
db.execute("CREATE TABLE %s_new (%s, PRIMARY KEY (%s))" % (table, ", ".join(columns), ", ".join(pk)))
|
||
if keep:
|
||
db.execute("INSERT OR IGNORE INTO %s_new (%s) SELECT %s FROM %s" % (table, ", ".join(keep), ", ".join(keep), table))
|
||
db.execute("DROP TABLE %s" % table)
|
||
db.execute("ALTER TABLE %s_new RENAME TO %s" % (table, table))
|
||
return True
|
||
|
||
|
||
def migrate_profiles(db):
|
||
for table in ("secrets", "records", "tags", "edges", "audit", "events", "schedules"):
|
||
ensure_column(db, table, "profile", "TEXT")
|
||
db.execute("UPDATE %s SET profile = 'default' WHERE profile IS NULL" % table)
|
||
rebuild_with_profile(db, "secrets", ("profile TEXT", "name TEXT", "value TEXT", "meta TEXT", "updated TEXT"), ("profile", "name"))
|
||
rebuild_with_profile(db, "tags", ("item TEXT", "tag TEXT", "profile TEXT"), ("item", "tag", "profile"))
|
||
rebuild_with_profile(db, "edges", ("src TEXT", "dst TEXT", "relation TEXT", "profile TEXT", "created TEXT"), ("src", "dst", "relation", "profile"))
|
||
|
||
|
||
SCHEMA = (
|
||
("events", ("id INTEGER PRIMARY KEY", "profile TEXT", "ts TEXT", "role TEXT", "kind TEXT", "text TEXT", "tags TEXT"), ()),
|
||
("secrets", ("profile TEXT", "name TEXT", "value TEXT", "meta TEXT", "updated TEXT"), ("profile", "name")),
|
||
("schedules", ("id INTEGER PRIMARY KEY", "name TEXT", "prompt TEXT", "profile TEXT", "every_sec INTEGER", "next_run TEXT", "timeout INTEGER", "status TEXT", "last_status TEXT", "last_result TEXT", "created TEXT", "updated TEXT"), ()),
|
||
("records", ("id TEXT PRIMARY KEY", "kind TEXT", "profile TEXT", "title TEXT", "content TEXT", "size INTEGER", "reads INTEGER", "created TEXT", "updated TEXT"), ()),
|
||
("tags", ("item TEXT", "tag TEXT", "profile TEXT"), ("item", "tag", "profile")),
|
||
("edges", ("src TEXT", "dst TEXT", "relation TEXT", "profile TEXT", "created TEXT"), ("src", "dst", "relation", "profile")),
|
||
("audit", ("id INTEGER PRIMARY KEY", "profile TEXT", "ts TEXT", "actor TEXT", "action TEXT", "path TEXT", "message TEXT", "old_size INTEGER", "new_size INTEGER", "old TEXT", "new TEXT", "tags TEXT"), ()),
|
||
("chunks", ("id INTEGER PRIMARY KEY", "parent_kind TEXT", "parent_id TEXT", "idx INTEGER", "text TEXT", "profile TEXT", "created TEXT"), ()),
|
||
)
|
||
|
||
SCHEMA_INDEXES = (
|
||
("idx_events_profile", "events(profile)"),
|
||
("idx_tags_tag", "tags(tag)"),
|
||
("idx_edges_src", "edges(src)"),
|
||
("idx_edges_dst", "edges(dst)"),
|
||
("idx_audit_path", "audit(path)"),
|
||
("idx_tags_profile", "tags(profile)"),
|
||
("idx_edges_profile", "edges(profile)"),
|
||
("idx_records_profile", "records(profile)"),
|
||
("idx_audit_profile", "audit(profile)"),
|
||
("idx_chunks_parent", "chunks(parent_kind, parent_id)"),
|
||
("idx_chunks_profile", "chunks(profile)"),
|
||
)
|
||
|
||
|
||
def ensure_schema(db):
|
||
notes = []
|
||
for table, columns, pk in SCHEMA:
|
||
exists = db.execute("SELECT 1 FROM sqlite_master WHERE type = 'table' AND name = ?", (table,)).fetchone()
|
||
if not exists:
|
||
extra = ", PRIMARY KEY (%s)" % ", ".join(pk) if pk else ""
|
||
db.execute("CREATE TABLE %s (%s%s)" % (table, ", ".join(columns), extra))
|
||
continue
|
||
have = {row[1] for row in db.execute("PRAGMA table_info(%s)" % table).fetchall()}
|
||
for column in columns:
|
||
parts = column.split()
|
||
if parts[0] not in have:
|
||
db.execute("ALTER TABLE %s ADD COLUMN %s" % (table, " ".join(parts[:2])))
|
||
notes.append("added %s.%s" % (table, parts[0]))
|
||
for index, target in SCHEMA_INDEXES:
|
||
db.execute("CREATE INDEX IF NOT EXISTS %s ON %s" % (index, target))
|
||
db.commit()
|
||
return notes
|
||
|
||
|
||
class Store:
|
||
def __init__(self, config, seal):
|
||
self.config = config
|
||
self.seal = seal
|
||
self.db = sqlite3.connect(config.db_path, timeout=30)
|
||
self.db.execute("PRAGMA journal_mode=WAL")
|
||
self.db.execute("PRAGMA synchronous=NORMAL")
|
||
self.db.execute("PRAGMA busy_timeout=30000")
|
||
self.schema_notes = ensure_schema(self.db)
|
||
migrate_profiles(self.db)
|
||
self.db.commit()
|
||
self.profile = config.profile
|
||
self.secret_cache = {}
|
||
self.fts_ok = setup_fts(self.db)
|
||
if seal.enabled:
|
||
self.db.create_function("tai_enc", 1, seal.lock)
|
||
self.db.create_function("tai_dec", 1, seal.safe_unlock)
|
||
self.db.execute("UPDATE events SET text = tai_enc(text) WHERE text NOT LIKE 'tai1$%'")
|
||
self.db.execute("UPDATE secrets SET value = tai_enc(value) WHERE value NOT LIKE 'tai1$%'")
|
||
self.db.execute("UPDATE secrets SET meta = tai_enc(COALESCE(meta, '{}')) WHERE meta IS NULL OR meta NOT LIKE 'tai1$%'")
|
||
self.db.execute("UPDATE schedules SET prompt = tai_enc(prompt) WHERE prompt NOT LIKE 'tai1$%'")
|
||
self.db.commit()
|
||
self.seal_files()
|
||
else:
|
||
sealed = self.db.execute("SELECT COUNT(*) FROM events WHERE text LIKE 'tai1$%'").fetchone()[0]
|
||
if sealed:
|
||
raise SealError("memory is sealed, set TAI_PASSPHRASE")
|
||
locked = self.db.execute("SELECT COUNT(*) FROM secrets WHERE value LIKE 'tai1$%' OR meta LIKE 'tai1$%'").fetchone()[0]
|
||
if locked:
|
||
raise SealError("secrets are sealed, set TAI_PASSPHRASE")
|
||
planned = self.db.execute("SELECT COUNT(*) FROM schedules WHERE prompt LIKE 'tai1$%'").fetchone()[0]
|
||
if planned:
|
||
raise SealError("schedules are sealed, set TAI_PASSPHRASE")
|
||
self.memdb = None
|
||
self.mem_events_max = 0
|
||
if seal.enabled and self.fts_ok:
|
||
try:
|
||
self.memdb = open_mem_events()
|
||
floor = self.db.execute("SELECT COALESCE(MAX(id), 0) FROM events").fetchone()[0] - MEM_EVENTS_CAP
|
||
self.mem_events_max = max(0, floor)
|
||
self._sync_mem_events()
|
||
except sqlite3.Error:
|
||
self.memdb = None
|
||
|
||
def _sync_mem_events(self):
|
||
if self.memdb is None:
|
||
return
|
||
try:
|
||
rows = self.db.execute("SELECT id, profile, ts, role, kind, tai_dec(text) FROM events WHERE id > ? ORDER BY id ASC", (self.mem_events_max,)).fetchall()
|
||
except sqlite3.Error:
|
||
return
|
||
if not rows:
|
||
return
|
||
try:
|
||
self.memdb.executemany("INSERT INTO mem_events(item, kind, sub, title, body, profile) VALUES (?, 'event', ? || '/' || ?, ? || ' ' || ? || '/' || ?, ?, ?)", [(row[0], row[3], row[4], row[2], row[3], row[4], row[5], row[1]) for row in rows])
|
||
self.memdb.commit()
|
||
except sqlite3.Error:
|
||
return
|
||
self.mem_events_max = rows[-1][0]
|
||
try:
|
||
self.memdb.execute("DELETE FROM mem_events WHERE CAST(item AS INTEGER) <= ?", (self.mem_events_max - MEM_EVENTS_CAP,))
|
||
self.memdb.commit()
|
||
except sqlite3.Error:
|
||
pass
|
||
|
||
def seal_files(self):
|
||
for entry in os.listdir(self.config.profiles_dir):
|
||
if not entry.endswith((".sys.md", ".session.json", ".bots.json")):
|
||
continue
|
||
path = os.path.join(self.config.profiles_dir, entry)
|
||
try:
|
||
with open(path, "r", encoding="utf-8") as handle:
|
||
raw = handle.read()
|
||
except OSError:
|
||
continue
|
||
if raw.startswith(Seal.PREFIX):
|
||
continue
|
||
try:
|
||
with open(path, "w", encoding="utf-8") as handle:
|
||
handle.write(self.seal.lock(raw))
|
||
except OSError:
|
||
pass
|
||
|
||
def log_event(self, profile, role, kind, text, tags=()):
|
||
value = (text or "")[:2000]
|
||
stamped = " " + " ".join([role, kind] + normalize_tags(tags)) + " "
|
||
try:
|
||
if self.seal.enabled:
|
||
self.db.execute("INSERT INTO events (profile, ts, role, kind, text, tags) VALUES (?, ?, ?, ?, tai_enc(?), ?)", (profile, now_iso(), role, kind, value, stamped))
|
||
else:
|
||
self.db.execute("INSERT INTO events (profile, ts, role, kind, text, tags) VALUES (?, ?, ?, ?, ?, ?)", (profile, now_iso(), role, kind, value, stamped))
|
||
self.db.commit()
|
||
self._sync_mem_events()
|
||
except sqlite3.Error:
|
||
pass
|
||
|
||
def fts_search(self, query, kinds=(), profile=None, limit=10):
|
||
profile = self._scope(profile)
|
||
if not self.fts_ok:
|
||
return []
|
||
terms = fts_escape(query)
|
||
if terms is None:
|
||
return []
|
||
sql = "SELECT item, kind, sub, title, snippet(fts_docs, 4, '[', ']', '...', 12), bm25(fts_docs) FROM fts_docs WHERE fts_docs MATCH ? AND profile = ?"
|
||
params = [fts_match(terms), profile]
|
||
if kinds:
|
||
marks = ", ".join("?" * len(kinds))
|
||
sql += " AND kind IN (%s)" % marks
|
||
params += list(kinds)
|
||
if self.seal.enabled:
|
||
sql += " AND kind != 'event'"
|
||
sql += " ORDER BY bm25(fts_docs) LIMIT ?"
|
||
capped = max(1, min(limit, 100))
|
||
params.append(capped)
|
||
try:
|
||
rows = self.db.execute(sql, tuple(params)).fetchall()
|
||
except sqlite3.Error:
|
||
return []
|
||
hits = [{"item": row[0], "kind": row[1], "sub": row[2], "title": row[3], "snippet": row[4], "rank": row[5]} for row in rows]
|
||
if self.seal.enabled and self.memdb is not None and (not kinds or "event" in kinds):
|
||
self._sync_mem_events()
|
||
try:
|
||
mem_rows = self.memdb.execute("SELECT item, kind, sub, title, snippet(mem_events, 4, '[', ']', '...', 12), bm25(mem_events) FROM mem_events WHERE mem_events MATCH ? AND profile = ? ORDER BY bm25(mem_events) LIMIT ?", (fts_match(terms), profile, capped)).fetchall()
|
||
except sqlite3.Error:
|
||
mem_rows = []
|
||
hits += [{"item": row[0], "kind": row[1], "sub": row[2], "title": row[3], "snippet": row[4], "rank": row[5]} for row in mem_rows]
|
||
hits.sort(key=lambda hit: hit["rank"])
|
||
hits = hits[:capped]
|
||
return hits
|
||
|
||
def audit_search(self, query, profile=None, limit=20):
|
||
profile = self._scope(profile)
|
||
hits = self.fts_search(query, ("audit",), profile, limit)
|
||
found = []
|
||
for hit in hits:
|
||
try:
|
||
row_id = int(str(hit["item"]).split(":", 1)[1])
|
||
except (IndexError, ValueError):
|
||
continue
|
||
row = self.audit_get(row_id, profile)
|
||
if row is not None:
|
||
row["snippet"] = hit["snippet"]
|
||
found.append(row)
|
||
return found
|
||
|
||
def search_events(self, profile, query, limit=8, tags=()):
|
||
clean = normalize_tags(tags)
|
||
if self.fts_ok:
|
||
hits = self.fts_search(query, ("event",), profile, max(limit * 5, 20))
|
||
if hits:
|
||
ids = [int(hit["item"]) for hit in hits]
|
||
marks = ", ".join("?" * len(ids))
|
||
text_col = "tai_dec(text)" if self.seal.enabled else "text"
|
||
rows = {row[0]: row[1:] for row in self.db.execute("SELECT id, ts, role, kind, %s, tags FROM events WHERE id IN (%s)" % (text_col, marks), tuple(ids)).fetchall()}
|
||
found = []
|
||
for hit in hits:
|
||
row = rows.get(int(hit["item"]))
|
||
if row is None:
|
||
continue
|
||
if clean and not all((" %s " % tag) in (" %s " % (row[4] or "")) for tag in clean):
|
||
continue
|
||
found.append(row[:4])
|
||
if len(found) >= limit:
|
||
break
|
||
if found:
|
||
return found
|
||
return self.events_like(profile, query, limit, clean)
|
||
|
||
def events_like(self, profile, query, limit=8, tags=()):
|
||
clean = normalize_tags(tags)
|
||
extra = "".join(" AND tags LIKE ?" for _tag in clean)
|
||
wild = tuple("% %s %%" % tag for tag in clean)
|
||
if self.seal.enabled:
|
||
return self.db.execute("SELECT ts, role, kind, text FROM (SELECT id, ts, role, kind, tags, tai_dec(text) AS text FROM events WHERE profile = ? ORDER BY id DESC LIMIT 5000) WHERE text LIKE ?" + extra + " ORDER BY id DESC LIMIT ?", (profile, "%" + query + "%") + wild + (limit,)).fetchall()
|
||
return self.db.execute("SELECT ts, role, kind, text FROM events WHERE profile = ? AND text LIKE ?" + extra + " ORDER BY id DESC LIMIT ?", (profile, "%" + query + "%") + wild + (limit,)).fetchall()
|
||
|
||
def _scope(self, profile):
|
||
return profile or self.profile
|
||
|
||
def save_secret(self, name, value, meta=None, tags=None, profile=None):
|
||
profile = self._scope(profile)
|
||
self.secret_cache.pop(profile, None)
|
||
stored = self.seal.lock(value) if self.seal.enabled else value
|
||
blob = json.dumps(meta or {})
|
||
locked_meta = self.seal.lock(blob) if self.seal.enabled else blob
|
||
self.db.execute("INSERT OR REPLACE INTO secrets (profile, name, value, meta, updated) VALUES (?, ?, ?, ?, ?)", (profile, name, stored, locked_meta, now_iso()))
|
||
if tags is None:
|
||
self.tag_item("secret:" + name, ["secret"], profile)
|
||
else:
|
||
self.set_item_tags("secret:" + name, ["secret"] + list(tags), profile)
|
||
self.db.commit()
|
||
|
||
def secret_meta(self, name, profile=None):
|
||
profile = self._scope(profile)
|
||
row = self.db.execute("SELECT meta FROM secrets WHERE profile = ? AND name = ?", (profile, name)).fetchone()
|
||
if row is None or not row[0]:
|
||
return {}
|
||
raw = self.seal.safe_unlock(row[0]) if row[0].startswith(Seal.PREFIX) else row[0]
|
||
try:
|
||
data = json.loads(raw)
|
||
except ValueError:
|
||
return {}
|
||
return data if isinstance(data, dict) else {}
|
||
|
||
def list_secret_infos(self, profile=None):
|
||
profile = self._scope(profile)
|
||
rows = self.db.execute("SELECT name, meta, updated FROM secrets WHERE profile = ? ORDER BY name", (profile,)).fetchall()
|
||
items = []
|
||
for name, meta, updated in rows:
|
||
raw = None
|
||
if meta:
|
||
raw = self.seal.safe_unlock(meta) if meta.startswith(Seal.PREFIX) else meta
|
||
try:
|
||
data = json.loads(raw or "{}")
|
||
except ValueError:
|
||
data = {}
|
||
items.append({"name": name, "meta": data if isinstance(data, dict) else {}, "tags": self.item_tags("secret:" + name, profile), "updated": updated})
|
||
return items
|
||
|
||
def known_tags(self, profile=None):
|
||
profile = self._scope(profile)
|
||
return {row[0] for row in self.db.execute("SELECT DISTINCT tag FROM tags WHERE profile = ?", (profile,)).fetchall()}
|
||
|
||
def canonical_tag(self, tag, vocab=None, profile=None):
|
||
if vocab is None:
|
||
vocab = self.known_tags(profile)
|
||
if tag in vocab:
|
||
return tag
|
||
single = singular_noun(tag)
|
||
if single != tag and single in vocab:
|
||
return single
|
||
return tag
|
||
|
||
def canonical_tags(self, tags, vocab=None, profile=None):
|
||
profile = self._scope(profile)
|
||
if vocab is None:
|
||
vocab = self.known_tags(profile)
|
||
seen = []
|
||
for tag in normalize_tags(tags):
|
||
clean = self.canonical_tag(tag, vocab)
|
||
if clean not in seen:
|
||
seen.append(clean)
|
||
return seen
|
||
|
||
def tag_counts(self, prefix="", limit=50, profile=None):
|
||
profile = self._scope(profile)
|
||
clean = normalize_tag(prefix)
|
||
sql = "SELECT tag, COUNT(*) FROM tags WHERE profile = ?"
|
||
params = [profile]
|
||
if clean:
|
||
sql += " AND tag LIKE ? ESCAPE '\\'"
|
||
params.append(clean.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_") + "%")
|
||
sql += " GROUP BY tag ORDER BY COUNT(*) DESC, tag LIMIT ?"
|
||
params.append(max(1, min(limit, 200)))
|
||
return [(row[0], row[1]) for row in self.db.execute(sql, tuple(params)).fetchall()]
|
||
|
||
def auto_tags_for(self, text, explicit=(), vocab=None):
|
||
if vocab is None:
|
||
vocab = self.known_tags()
|
||
skip = set(explicit) | set(TAG_AUTO_SKIP)
|
||
words = re.findall(r"[a-z0-9]+", (text or "").lower())
|
||
found = []
|
||
for pos, word in enumerate(words):
|
||
if len(found) >= TAG_AUTO_MAX:
|
||
break
|
||
if pos + 1 < len(words):
|
||
pair = word + "-" + words[pos + 1]
|
||
if pair in vocab and pair not in skip and pair not in found:
|
||
found.append(pair)
|
||
continue
|
||
if word in vocab and word not in skip and word not in found:
|
||
found.append(word)
|
||
continue
|
||
single = singular_noun(word)
|
||
if single != word and single in vocab and single not in skip and single not in found:
|
||
found.append(single)
|
||
return found
|
||
|
||
def tag_neighbors(self, record_id, limit=None, profile=None):
|
||
profile = self._scope(profile)
|
||
rows = self.db.execute("SELECT r.id, MIN(t.tag) FROM records r JOIN tags t ON t.item = r.id WHERE t.tag IN (SELECT tag FROM tags WHERE item = ? AND profile = ? AND tag NOT IN ('record', 'file')) AND r.id != ? AND r.profile = ? AND t.profile = ? GROUP BY r.id ORDER BY r.updated DESC LIMIT ?", (record_id, profile, record_id, profile, profile, max(1, limit or TAG_LINK_MAX))).fetchall()
|
||
return [(row[0], row[1]) for row in rows]
|
||
|
||
def tag_item(self, item, tags, profile=None):
|
||
profile = self._scope(profile)
|
||
for tag in self.canonical_tags(tags, None, profile):
|
||
self.db.execute("INSERT OR IGNORE INTO tags (item, tag, profile) VALUES (?, ?, ?)", (item, tag, profile))
|
||
self.db.commit()
|
||
|
||
def set_item_tags(self, item, tags, profile=None):
|
||
profile = self._scope(profile)
|
||
self.db.execute("DELETE FROM tags WHERE item = ? AND profile = ?", (item, profile))
|
||
for tag in self.canonical_tags(tags, None, profile):
|
||
self.db.execute("INSERT OR IGNORE INTO tags (item, tag, profile) VALUES (?, ?, ?)", (item, tag, profile))
|
||
self.db.commit()
|
||
|
||
def item_tags(self, item, profile=None):
|
||
profile = self._scope(profile)
|
||
return [row[0] for row in self.db.execute("SELECT tag FROM tags WHERE item = ? AND profile = ? ORDER BY tag", (item, profile)).fetchall()]
|
||
|
||
def tagged_items(self, tags, profile=None):
|
||
profile = self._scope(profile)
|
||
clean = normalize_tags(tags)
|
||
if not clean:
|
||
return None
|
||
marks = ", ".join("?" * len(clean))
|
||
rows = self.db.execute("SELECT item FROM tags WHERE profile = ? AND tag IN (%s) GROUP BY item HAVING COUNT(DISTINCT tag) = ?" % marks, (profile,) + tuple(clean) + (len(clean),)).fetchall()
|
||
return [row[0] for row in rows]
|
||
|
||
def has_node(self, kind, key, profile=None):
|
||
profile = self._scope(profile)
|
||
if kind == "mem":
|
||
row = self.db.execute("SELECT 1 FROM records WHERE profile = ? AND id = ?", (profile, "mem:" + key,)).fetchone()
|
||
elif kind == "secret":
|
||
row = self.db.execute("SELECT 1 FROM secrets WHERE profile = ? AND name = ?", (profile, key,)).fetchone()
|
||
else:
|
||
row = self.db.execute("SELECT 1 FROM schedules WHERE profile = ? AND id = ?", (profile, int(key),)).fetchone()
|
||
return row is not None
|
||
|
||
def node_title(self, node, profile=None):
|
||
profile = self._scope(profile)
|
||
parsed = parse_node_id(node)
|
||
if parsed is None:
|
||
return node
|
||
kind, key = parsed
|
||
if kind == "mem":
|
||
row = self.db.execute("SELECT title FROM records WHERE profile = ? AND id = ?", (profile, node,)).fetchone()
|
||
return row[0] or node if row else node
|
||
if kind == "secret":
|
||
return "secret:" + key
|
||
row = self.db.execute("SELECT name, prompt FROM schedules WHERE profile = ? AND id = ?", (profile, int(key),)).fetchone()
|
||
if not row:
|
||
return node
|
||
return row[0] or "schedule #%s" % key
|
||
|
||
def add_record(self, kind, title, content, tags=(), profile=None):
|
||
profile = self._scope(profile)
|
||
record_id = "mem:" + uuid.uuid4().hex[:16]
|
||
while self.db.execute("SELECT 1 FROM records WHERE id = ?", (record_id,)).fetchone():
|
||
record_id = "mem:" + uuid.uuid4().hex[:16]
|
||
stamp = datetime.now(timezone.utc).isoformat()
|
||
size = len(content.encode("utf-8"))
|
||
self.db.execute("INSERT INTO records(id, kind, profile, title, content, size, reads, created, updated) VALUES(?, ?, ?, ?, ?, ?, 0, ?, ?)", (record_id, kind, profile, title, content, size, stamp, stamp))
|
||
vocab = self.known_tags(profile)
|
||
explicit = self.canonical_tags(tags, vocab, profile)
|
||
scan = "%s\n%s" % (title, content.split("\n...[truncated", 1)[0])
|
||
auto = self.auto_tags_for(scan, explicit, vocab)
|
||
self.set_item_tags(record_id, ["record"] + explicit + auto, profile)
|
||
for other, shared in self.tag_neighbors(record_id, None, profile):
|
||
self.db.execute("INSERT OR IGNORE INTO edges(src, dst, relation, profile, created) VALUES(?, ?, ?, ?, ?)", (record_id, other, "shares-" + shared, profile, stamp))
|
||
self.db.commit()
|
||
return record_id
|
||
|
||
def add_chunks(self, parent_kind, parent_id, texts, profile=None):
|
||
"""Replaces any prior chunks for this parent, so re-indexing is idempotent."""
|
||
profile = self._scope(profile)
|
||
self.db.execute("DELETE FROM chunks WHERE parent_kind = ? AND parent_id = ? AND profile = ?", (parent_kind, parent_id, profile))
|
||
stamp = datetime.now(timezone.utc).isoformat()
|
||
for idx, text in enumerate(texts):
|
||
self.db.execute("INSERT INTO chunks(parent_kind, parent_id, idx, text, profile, created) VALUES (?, ?, ?, ?, ?, ?)", (parent_kind, parent_id, idx, text, profile, stamp))
|
||
self.db.commit()
|
||
|
||
def get_chunks_by_ids(self, ids, profile=None):
|
||
if not ids:
|
||
return {}
|
||
profile = self._scope(profile)
|
||
marks = ", ".join("?" * len(ids))
|
||
rows = self.db.execute("SELECT id, parent_kind, parent_id, idx, text FROM chunks WHERE profile = ? AND id IN (%s)" % marks, tuple([profile] + list(ids))).fetchall()
|
||
return {row[0]: {"parent_kind": row[1], "parent_id": row[2], "idx": row[3], "text": row[4]} for row in rows}
|
||
|
||
def get_record(self, record_id, profile=None):
|
||
profile = self._scope(profile)
|
||
row = self.db.execute("SELECT id, kind, title, content, size, reads, created, updated FROM records WHERE profile = ? AND id = ?", (profile, record_id,)).fetchone()
|
||
if row is None:
|
||
return None
|
||
self.db.execute("UPDATE records SET reads = reads + 1 WHERE id = ?", (record_id,))
|
||
self.db.commit()
|
||
return {"id": row[0], "kind": row[1], "title": row[2], "content": row[3], "size": as_int(row[4]), "reads": as_int(row[5]) + 1, "created": row[6], "updated": row[7], "tags": self.item_tags(row[0], profile)}
|
||
|
||
def read_record(self, record_id, offset=0, limit=4000, profile=None):
|
||
profile = self._scope(profile)
|
||
row = self.db.execute("SELECT content FROM records WHERE profile = ? AND id = ?", (profile, record_id,)).fetchone()
|
||
if row is None:
|
||
return None
|
||
content = row[0]
|
||
total = len(content)
|
||
start = max(0, min(offset, total))
|
||
end = min(total, start + max(1, limit))
|
||
self.db.execute("UPDATE records SET reads = reads + 1 WHERE id = ?", (record_id,))
|
||
self.db.commit()
|
||
return {"id": record_id, "slice": content[start:end], "start": start, "end": end, "total": total, "tags": self.item_tags(record_id, profile)}
|
||
|
||
def search_records(self, query="", kind=None, tags=(), limit=8, profile=None):
|
||
profile = self._scope(profile)
|
||
clean = normalize_tags(tags)
|
||
terms = fts_escape(query) if query else None
|
||
if terms is not None and self.fts_ok:
|
||
ranked = self.search_records_fts(terms, kind, clean, limit, profile)
|
||
if ranked:
|
||
return ranked
|
||
sql = "SELECT id, kind, title, size, reads, updated FROM records"
|
||
clauses = ["profile = ?"]
|
||
params = [profile]
|
||
if kind:
|
||
clauses.append("kind = ?")
|
||
params.append(kind)
|
||
if query:
|
||
clauses.append("(title LIKE ? OR content LIKE ?)")
|
||
params += ["%" + query + "%", "%" + query + "%"]
|
||
if clean:
|
||
for tag in clean:
|
||
options = tag_variants(tag)
|
||
marks = ", ".join("?" * len(options))
|
||
clauses.append("id IN (SELECT item FROM tags WHERE profile = ? AND tag IN (%s))" % marks)
|
||
params += [profile] + options
|
||
sql += " WHERE " + " AND ".join(clauses)
|
||
sql += " ORDER BY updated DESC LIMIT ?"
|
||
params.append(max(1, min(limit, 50)))
|
||
rows = self.db.execute(sql, tuple(params)).fetchall()
|
||
return [{"id": row[0], "kind": row[1], "title": row[2], "size": row[3], "reads": row[4], "updated": row[5], "tags": self.item_tags(row[0], profile)} for row in rows]
|
||
|
||
def search_records_fts(self, terms, kind, clean, limit, profile):
|
||
sql = "SELECT r.id, r.kind, r.title, r.size, r.reads, r.updated FROM records r JOIN fts_docs f ON f.item = r.id AND f.kind = 'record' WHERE r.profile = ? AND f MATCH ?"
|
||
params = [profile, fts_match(terms)]
|
||
if kind:
|
||
sql += " AND r.kind = ?"
|
||
params.append(kind)
|
||
for tag in clean:
|
||
options = tag_variants(tag)
|
||
marks = ", ".join("?" * len(options))
|
||
sql += " AND r.id IN (SELECT item FROM tags WHERE profile = ? AND tag IN (%s))" % marks
|
||
params += [profile] + options
|
||
sql += " ORDER BY bm25(f) LIMIT ?"
|
||
params.append(max(1, min(limit, 50)))
|
||
try:
|
||
rows = self.db.execute(sql, tuple(params)).fetchall()
|
||
except sqlite3.Error:
|
||
return []
|
||
return [{"id": row[0], "kind": row[1], "title": row[2], "size": row[3], "reads": row[4], "updated": row[5], "tags": self.item_tags(row[0], profile)} for row in rows]
|
||
|
||
def delete_record(self, record_id, profile=None):
|
||
profile = self._scope(profile)
|
||
done = self.db.execute("DELETE FROM records WHERE profile = ? AND id = ?", (profile, record_id,))
|
||
if done.rowcount > 0:
|
||
self.db.execute("DELETE FROM tags WHERE item = ? AND profile = ?", (record_id, profile))
|
||
self.db.execute("DELETE FROM edges WHERE profile = ? AND (src = ? OR dst = ?)", (profile, record_id, record_id))
|
||
self.db.execute("DELETE FROM chunks WHERE parent_kind = 'record' AND parent_id = ? AND profile = ?", (record_id, profile))
|
||
self.db.commit()
|
||
return True
|
||
self.db.commit()
|
||
return False
|
||
|
||
def add_edge(self, src, dst, relation="linked", profile=None):
|
||
profile = self._scope(profile)
|
||
for node in (src, dst):
|
||
parsed = parse_node_id(node)
|
||
if parsed is None or not self.has_node(*parsed, profile):
|
||
raise ValueError("unknown node " + node)
|
||
clean = re.sub(r"[^a-z0-9]+", "-", str(relation or "").lower()).strip("-")[:32] or "linked"
|
||
self.db.execute("INSERT OR IGNORE INTO edges(src, dst, relation, profile, created) VALUES(?, ?, ?, ?, ?)", (src, dst, clean, profile, datetime.now(timezone.utc).isoformat()))
|
||
self.db.commit()
|
||
return clean
|
||
|
||
def edges_for(self, node, profile=None):
|
||
profile = self._scope(profile)
|
||
rows = self.db.execute("SELECT src, dst, relation FROM edges WHERE profile = ? AND (src = ? OR dst = ?)", (profile, node, node)).fetchall()
|
||
return [{"relation": row[2], "other": row[1] if row[0] == node else row[0], "direction": "out" if row[0] == node else "in"} for row in rows]
|
||
|
||
def remove_edges(self, node, profile=None):
|
||
profile = self._scope(profile)
|
||
done = self.db.execute("DELETE FROM edges WHERE profile = ? AND (src = ? OR dst = ?)", (profile, node, node))
|
||
self.db.commit()
|
||
return done.rowcount
|
||
|
||
def traverse(self, start, depth=2, limit=50, profile=None):
|
||
profile = self._scope(profile)
|
||
seen = {start}
|
||
found = [{"node": start, "title": self.node_title(start, profile), "depth": 0, "via": ""}]
|
||
frontier = [(start, 0, "")]
|
||
capped = max(1, min(depth, 4))
|
||
while frontier and len(found) < max(1, min(limit, 200)):
|
||
node, level, _via = frontier.pop(0)
|
||
if level >= capped:
|
||
continue
|
||
for edge in self.edges_for(node, profile):
|
||
other = edge["other"]
|
||
if other in seen:
|
||
continue
|
||
seen.add(other)
|
||
hop = "%s -[%s]-> %s" % (node, edge["relation"], other) if edge["direction"] == "out" else "%s <-[%s]- %s" % (node, edge["relation"], other)
|
||
found.append({"node": other, "title": self.node_title(other, profile), "depth": level + 1, "via": hop})
|
||
frontier.append((other, level + 1, hop))
|
||
if len(found) >= max(1, min(limit, 200)):
|
||
break
|
||
return found
|
||
|
||
def audit_event(self, actor, action, path, message="", old=None, new=None, tags=(), old_size=None, new_size=None, profile=None):
|
||
profile = self._scope(profile)
|
||
old_text, computed_old = capped_text(old)
|
||
new_text, computed_new = capped_text(new)
|
||
if old is not None and old_size is not None and old_size > len(old.encode("utf-8")):
|
||
old_text = old[:AUDIT_MAX_CHARS] + trunc_marker(old_size)
|
||
computed_old = old_size
|
||
if new is not None and new_size is not None and new_size > len(new.encode("utf-8")):
|
||
new_text = new[:AUDIT_MAX_CHARS] + trunc_marker(new_size)
|
||
computed_new = new_size
|
||
cursor = self.db.execute("INSERT INTO audit(profile, ts, actor, action, path, message, old_size, new_size, old, new, tags) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)", (profile, datetime.now(timezone.utc).isoformat(), actor, action, path, str(message or "")[:2000], computed_old, computed_new, old_text, new_text, " ".join(normalize_tags(tags))))
|
||
self.db.commit()
|
||
return cursor.lastrowid
|
||
|
||
def audit_history(self, path=None, tag=None, limit=20, profile=None):
|
||
profile = self._scope(profile)
|
||
sql = "SELECT id, ts, actor, action, path, message, old_size, new_size, tags FROM audit"
|
||
clauses = ["profile = ?"]
|
||
params = [profile]
|
||
if path:
|
||
clauses.append("path = ?")
|
||
params.append(path)
|
||
if tag:
|
||
clauses.append("tags LIKE ?")
|
||
params.append("%" + tag + "%")
|
||
sql += " WHERE " + " AND ".join(clauses)
|
||
sql += " ORDER BY id DESC LIMIT ?"
|
||
params.append(max(1, min(limit, 100)))
|
||
rows = self.db.execute(sql, tuple(params)).fetchall()
|
||
return [{"id": row[0], "ts": row[1], "actor": row[2], "action": row[3], "path": row[4], "message": row[5], "old_size": row[6], "new_size": row[7], "tags": row[8]} for row in rows]
|
||
|
||
def audit_get(self, row_id, profile=None):
|
||
profile = self._scope(profile)
|
||
row = self.db.execute("SELECT id, ts, actor, action, path, message, old_size, new_size, old, new, tags FROM audit WHERE profile = ? AND id = ?", (profile, row_id,)).fetchone()
|
||
if row is None:
|
||
return None
|
||
return {"id": row[0], "ts": row[1], "actor": row[2], "action": row[3], "path": row[4], "message": row[5], "old_size": row[6], "new_size": row[7], "old": row[8], "new": row[9], "tags": row[10]}
|
||
|
||
def file_record_id(self, path, profile=None):
|
||
profile = self._scope(profile)
|
||
row = self.db.execute("SELECT id FROM records WHERE profile = ? AND kind = 'file' AND title = ?", (profile, path,)).fetchone()
|
||
return row[0] if row else None
|
||
|
||
def upsert_file_record(self, path, content, true_size=None, profile=None):
|
||
profile = self._scope(profile)
|
||
full_bytes = len(content.encode("utf-8"))
|
||
if true_size is not None and true_size > full_bytes:
|
||
text = content + trunc_marker(true_size)
|
||
elif len(content) <= FILE_RECORD_MAX:
|
||
text = content
|
||
else:
|
||
text = content[:FILE_RECORD_MAX] + trunc_marker(full_bytes)
|
||
ext = os.path.splitext(path)[1].lstrip(".").lower()
|
||
explicit = ["file"] + ([ext] if ext and ext.isalnum() and len(ext) <= 5 else [])
|
||
vocab = self.known_tags(profile)
|
||
tags = explicit + self.auto_tags_for("%s\n%s" % (path, content), explicit, vocab)
|
||
record_id = self.file_record_id(path, profile)
|
||
stamp = datetime.now(timezone.utc).isoformat()
|
||
if record_id is None:
|
||
return self.add_record("file", path, text, tags, profile)
|
||
self.db.execute("UPDATE records SET content = ?, size = ?, updated = ? WHERE id = ?", (text, len(text.encode("utf-8")), stamp, record_id))
|
||
self.set_item_tags(record_id, tags, profile)
|
||
self.db.commit()
|
||
return record_id
|
||
|
||
def delete_file_record(self, path, profile=None):
|
||
profile = self._scope(profile)
|
||
record_id = self.file_record_id(path, profile)
|
||
if record_id is None:
|
||
return False
|
||
return self.delete_record(record_id, profile)
|
||
|
||
def load_secret(self, name, profile=None):
|
||
profile = self._scope(profile)
|
||
row = self.db.execute("SELECT value FROM secrets WHERE profile = ? AND name = ?", (profile, name,)).fetchone()
|
||
if row is None:
|
||
return None
|
||
if row[0].startswith(Seal.PREFIX):
|
||
return self.seal.unlock(row[0]) if self.seal.enabled else None
|
||
return row[0]
|
||
|
||
def delete_secret(self, name, profile=None):
|
||
profile = self._scope(profile)
|
||
self.secret_cache.pop(profile, None)
|
||
done = self.db.execute("DELETE FROM secrets WHERE profile = ? AND name = ?", (profile, name,))
|
||
if done.rowcount > 0:
|
||
item = "secret:" + name
|
||
self.db.execute("DELETE FROM tags WHERE item = ? AND profile = ?", (item, profile))
|
||
self.db.execute("DELETE FROM edges WHERE profile = ? AND (src = ? OR dst = ?)", (profile, item, item))
|
||
self.db.commit()
|
||
return True
|
||
self.db.commit()
|
||
return False
|
||
|
||
def secret_values(self, profile=None):
|
||
profile = self._scope(profile)
|
||
if profile not in self.secret_cache:
|
||
found = {}
|
||
for name, value in self.db.execute("SELECT name, value FROM secrets WHERE profile = ?", (profile,)).fetchall():
|
||
if value.startswith(Seal.PREFIX):
|
||
if not self.seal.enabled:
|
||
continue
|
||
value = self.seal.unlock(value)
|
||
if len(value) >= 4:
|
||
found[name] = value
|
||
self.secret_cache[profile] = found
|
||
return self.secret_cache[profile]
|
||
|
||
def redact(self, text, profile=None):
|
||
for name, value in sorted(self.secret_values(profile).items(), key=lambda item: -len(item[1])):
|
||
if value in text:
|
||
text = text.replace(value, "[redacted:%s]" % name)
|
||
return text
|
||
|
||
def add_schedule(self, name, prompt, profile, every_sec, next_run, timeout, tags=None):
|
||
stored = self.seal.lock(prompt) if self.seal.enabled else prompt
|
||
done = self.db.execute("INSERT INTO schedules (name, prompt, profile, every_sec, next_run, timeout, status, last_status, last_result, created, updated) VALUES (?, ?, ?, ?, ?, ?, 'pending', '', '', ?, ?)", (name, stored, profile, every_sec, next_run, timeout, now_iso(), now_iso()))
|
||
item = "sched:%s" % done.lastrowid
|
||
if tags is None:
|
||
self.tag_item(item, ["schedule"], profile)
|
||
else:
|
||
self.set_item_tags(item, ["schedule"] + list(tags), profile)
|
||
self.db.commit()
|
||
return done.lastrowid
|
||
|
||
def remove_schedule(self, row_id, profile=None):
|
||
profile = self._scope(profile)
|
||
done = self.db.execute("DELETE FROM schedules WHERE id = ? AND profile = ?", (row_id, profile))
|
||
if done.rowcount > 0:
|
||
item = "sched:%s" % row_id
|
||
self.db.execute("DELETE FROM tags WHERE item = ? AND profile = ?", (item, profile))
|
||
self.db.execute("DELETE FROM edges WHERE profile = ? AND (src = ? OR dst = ?)", (profile, item, item))
|
||
self.db.commit()
|
||
return True
|
||
self.db.commit()
|
||
return False
|
||
|
||
def list_schedules(self, profile=None):
|
||
profile = self._scope(profile)
|
||
rows = self.db.execute("SELECT id, name, prompt, profile, every_sec, next_run, timeout, status, last_status, last_result, created FROM schedules WHERE profile = ? ORDER BY next_run", (profile,)).fetchall()
|
||
items = []
|
||
for row in rows:
|
||
prompt = self.seal.safe_unlock(row[2]) if row[2].startswith(Seal.PREFIX) else row[2]
|
||
items.append({"id": row[0], "name": row[1], "prompt": prompt, "profile": row[3], "every": row[4], "next_run": row[5], "timeout": row[6], "status": row[7], "last": row[8], "result": row[9], "created": row[10], "tags": self.item_tags("sched:%s" % row[0], row[3])})
|
||
return items
|
||
|
||
def profile_path(self, name):
|
||
return os.path.join(self.config.profiles_dir, name + ".sys.md")
|
||
|
||
def session_path(self, name):
|
||
return os.path.join(self.config.profiles_dir, name + ".session.json")
|
||
|
||
def list_profiles(self):
|
||
names = []
|
||
for entry in sorted(os.listdir(self.config.profiles_dir)):
|
||
if entry.endswith(".sys.md") and ".bot." not in entry:
|
||
names.append(entry[:-7])
|
||
return names
|
||
|
||
def unseal_file(self, raw):
|
||
if raw.startswith(Seal.PREFIX):
|
||
if not self.seal.enabled:
|
||
raise SealError("profile is sealed, set TAI_PASSPHRASE")
|
||
return self.seal.unlock(raw)
|
||
return raw
|
||
|
||
def load_system(self, name):
|
||
path = self.profile_path(name)
|
||
if not os.path.exists(path):
|
||
return None
|
||
try:
|
||
with open(path, "r", encoding="utf-8") as handle:
|
||
return self.unseal_file(handle.read())
|
||
except (OSError, ValueError):
|
||
quarantine(path)
|
||
return None
|
||
except SealError:
|
||
if not self.seal.enabled:
|
||
raise
|
||
quarantine(path)
|
||
return None
|
||
|
||
def save_system(self, name, text):
|
||
try:
|
||
sealed_write(self.profile_path(name), self.seal.lock(text))
|
||
except OSError:
|
||
pass
|
||
|
||
def load_session(self, name, bot="main"):
|
||
path = self.session_path(name) if bot == "main" else self.bot_session_path(name, bot)
|
||
if not os.path.exists(path):
|
||
return []
|
||
try:
|
||
with open(path, "r", encoding="utf-8") as handle:
|
||
raw = handle.read()
|
||
data = json.loads(self.unseal_file(raw))
|
||
return [item for item in data if isinstance(item, dict) and item.get("role") in ("user", "assistant", "tool")]
|
||
except (OSError, ValueError):
|
||
quarantine(path)
|
||
return []
|
||
except SealError:
|
||
if not self.seal.enabled:
|
||
raise
|
||
quarantine(path)
|
||
return []
|
||
|
||
def save_session(self, name, messages, bot="main"):
|
||
path = self.session_path(name) if bot == "main" else self.bot_session_path(name, bot)
|
||
try:
|
||
sealed_write(path, self.seal.lock(json.dumps([item for item in messages if item.get("role") != "system"])))
|
||
except OSError:
|
||
pass
|
||
|
||
def turn_state_path(self, profile, bot="main"):
|
||
if bot == "main":
|
||
return os.path.join(self.config.profiles_dir, "%s.turn.json" % profile)
|
||
return os.path.join(self.config.profiles_dir, "%s.bot.%s.turn.json" % (profile, bot))
|
||
|
||
def load_turn_state(self, profile, bot="main"):
|
||
path = self.turn_state_path(profile, bot)
|
||
if not os.path.exists(path):
|
||
return {}
|
||
try:
|
||
with open(path, "r", encoding="utf-8") as handle:
|
||
data = json.loads(self.unseal_file(handle.read()))
|
||
return data if isinstance(data, dict) else {}
|
||
except (OSError, ValueError):
|
||
quarantine(path)
|
||
return {}
|
||
except SealError:
|
||
if not self.seal.enabled:
|
||
raise
|
||
quarantine(path)
|
||
return {}
|
||
|
||
def save_turn_state(self, profile, bot, state):
|
||
try:
|
||
sealed_write(self.turn_state_path(profile, bot), self.seal.lock(json.dumps(state)))
|
||
except OSError:
|
||
pass
|
||
|
||
def bot_system_path(self, profile, bot):
|
||
return os.path.join(self.config.profiles_dir, "%s.bot.%s.sys.md" % (profile, bot))
|
||
|
||
def bot_session_path(self, profile, bot):
|
||
return os.path.join(self.config.profiles_dir, "%s.bot.%s.session.json" % (profile, bot))
|
||
|
||
def bots_registry_path(self, profile):
|
||
return os.path.join(self.config.profiles_dir, "%s.bots.json" % profile)
|
||
|
||
def load_bots(self, profile):
|
||
path = self.bots_registry_path(profile)
|
||
if not os.path.exists(path):
|
||
return {}
|
||
try:
|
||
with open(path, "r", encoding="utf-8") as handle:
|
||
data = json.loads(self.unseal_file(handle.read()))
|
||
bots = data.get("bots") if isinstance(data, dict) else None
|
||
return bots if isinstance(bots, dict) else {}
|
||
except (OSError, ValueError):
|
||
quarantine(path)
|
||
return {}
|
||
except SealError:
|
||
if not self.seal.enabled:
|
||
raise
|
||
quarantine(path)
|
||
return {}
|
||
|
||
def save_bots(self, profile, bots):
|
||
try:
|
||
sealed_write(self.bots_registry_path(profile), self.seal.lock(json.dumps({"bots": bots})))
|
||
except OSError:
|
||
pass
|
||
|
||
def list_bots(self, profile):
|
||
bots = self.load_bots(profile)
|
||
found = [{"name": "main", "nicknames": []}]
|
||
for name in sorted(bots):
|
||
found.append({"name": name, "nicknames": (bots.get(name) or {}).get("nicknames", [])})
|
||
return found
|
||
|
||
def resolve_bot(self, profile, mention):
|
||
want = str(mention or "").strip().lower()
|
||
if want in ("", "main"):
|
||
return "main"
|
||
bots = self.load_bots(profile)
|
||
if want in bots:
|
||
return want
|
||
for name in sorted(bots):
|
||
if want in (bots.get(name) or {}).get("nicknames", []):
|
||
return name
|
||
return None
|
||
|
||
def load_bot_system(self, profile, bot):
|
||
if bot == "main":
|
||
return self.load_system(profile)
|
||
path = self.bot_system_path(profile, bot)
|
||
if not os.path.exists(path):
|
||
return None
|
||
try:
|
||
with open(path, "r", encoding="utf-8") as handle:
|
||
return self.unseal_file(handle.read())
|
||
except (OSError, ValueError):
|
||
quarantine(path)
|
||
return None
|
||
except SealError:
|
||
if not self.seal.enabled:
|
||
raise
|
||
quarantine(path)
|
||
return None
|
||
|
||
def save_bot_system(self, profile, bot, text):
|
||
try:
|
||
sealed_write(self.bot_system_path(profile, bot), self.seal.lock(text))
|
||
except OSError:
|
||
pass
|
||
|
||
def close(self):
|
||
try:
|
||
self.db.execute("PRAGMA wal_checkpoint(TRUNCATE)")
|
||
except sqlite3.Error:
|
||
pass
|
||
try:
|
||
self.db.close()
|
||
except sqlite3.Error:
|
||
pass
|
||
if getattr(self, "memdb", None) is not None:
|
||
try:
|
||
self.memdb.close()
|
||
except sqlite3.Error:
|
||
pass
|
||
self.memdb = None
|
||
|
||
|
||
class BackendError(Exception):
|
||
def __init__(self, message, status=0):
|
||
super().__init__(message)
|
||
self.status = status
|
||
|
||
|
||
def muse_binary():
|
||
return os.environ.get("MUSE_BIN") or shutil.which("muse")
|
||
|
||
|
||
def opencode_binary():
|
||
return os.environ.get("OPENCODE_BIN") or shutil.which("opencode")
|
||
|
||
|
||
def grok_binary():
|
||
return os.environ.get("GROK_BIN") or shutil.which("grok") or shutil.which("grokcli")
|
||
|
||
|
||
def claude_binary():
|
||
return os.environ.get("CLAUDE_BIN") or shutil.which("claude")
|
||
|
||
|
||
def codex_binary():
|
||
return os.environ.get("CODEX_BIN") or shutil.which("codex")
|
||
|
||
|
||
def gemini_binary():
|
||
return os.environ.get("GEMINI_BIN") or shutil.which("gemini")
|
||
|
||
|
||
def opencode_prompt(messages, tools):
|
||
schema = json.dumps(muse_output_schema())
|
||
rules = (
|
||
"You are the model behind another agent. The transcript above is the whole conversation; lines starting with TOOL <id> are results of tool calls that already ran. "
|
||
"Do not use any of your own tools and do not repeat a tool call whose result is already shown. "
|
||
"If the results answer the user, put the final answer in content with an empty tool_calls list; otherwise request the next tool calls from TOOLS. "
|
||
"Reply with one JSON object only, no code fence, matching this schema: "
|
||
)
|
||
return muse_prompt(messages, tools) + "\n\n" + rules + schema
|
||
|
||
|
||
def opencode_parse(text):
|
||
text = (text or "").strip()
|
||
fenced = re.match(r"^```(?:json)?\s*(.*?)\s*```$", text, re.S)
|
||
if fenced:
|
||
text = fenced.group(1)
|
||
start, end = text.find("{"), text.rfind("}")
|
||
if start != -1 and end > start:
|
||
try:
|
||
answer = json.loads(text[start:end + 1])
|
||
if isinstance(answer, dict) and ("content" in answer or "tool_calls" in answer):
|
||
return answer
|
||
except ValueError:
|
||
pass
|
||
return {"content": text, "tool_calls": []}
|
||
|
||
|
||
def opencode_run(model, prompt, timeout):
|
||
binary = opencode_binary()
|
||
if not binary:
|
||
raise BackendError("opencode CLI not found (OPENCODE_BIN or PATH)")
|
||
argv = [binary, "run", "--pure", "--agent", OPENCODE_AGENT, "--format", "json", "-m", "opencode/" + model]
|
||
with tempfile.TemporaryDirectory(prefix="tai-opencode-") as folder:
|
||
try:
|
||
done = subprocess.run(argv, input=prompt, stdout=subprocess.PIPE, stderr=subprocess.PIPE, universal_newlines=True, cwd=folder, timeout=timeout)
|
||
except subprocess.TimeoutExpired:
|
||
raise BackendError("opencode run timed out after %ds" % timeout)
|
||
except OSError as exc:
|
||
raise BackendError("opencode run failed: %s" % short_error(exc))
|
||
texts = []
|
||
for line in done.stdout.splitlines():
|
||
try:
|
||
event = json.loads(line)
|
||
except ValueError:
|
||
continue
|
||
if not isinstance(event, dict):
|
||
continue
|
||
if event.get("type") == "error":
|
||
data = (event.get("error") or {}).get("data") or {}
|
||
raise BackendError("opencode: %s" % (data.get("message") or json.dumps(event.get("error"))[:300]), data.get("statusCode") or 0)
|
||
part = event.get("part") or {}
|
||
if event.get("type") == "text" and part.get("text"):
|
||
texts.append(part["text"])
|
||
if not texts:
|
||
raise BackendError("opencode returned no text (exit %d): %s" % (done.returncode, (done.stderr or done.stdout)[-300:]))
|
||
return "".join(texts)
|
||
|
||
|
||
def opencode_free_models():
|
||
binary = opencode_binary()
|
||
if not binary:
|
||
return None
|
||
try:
|
||
done = subprocess.run([binary, "models", "opencode"], stdout=subprocess.PIPE, stderr=subprocess.PIPE, universal_newlines=True, timeout=MODEL_LIST_TIMEOUT * 4)
|
||
except (OSError, subprocess.TimeoutExpired):
|
||
return None
|
||
models = [line.strip().split("/", 1)[1] for line in done.stdout.splitlines() if line.strip().startswith("opencode/")]
|
||
return [item for item in models if is_free_model(item)] or None
|
||
|
||
|
||
def _cli_unknown_flag(stderr):
|
||
return bool(re.search(r"unknown (option|flag|command)|unrecognised|unrecognized|invalid flag|not found|no such option", stderr or "", re.IGNORECASE))
|
||
|
||
|
||
def _cli_auth_failure(text):
|
||
return bool(re.search(r"not (logged in|authenticated)|auth|api key|unauthorized|401|login required|please (login|log in|authenticate)", text or "", re.IGNORECASE))
|
||
|
||
|
||
def _cli_run(argv, prompt_input, timeout, label):
|
||
try:
|
||
done = subprocess.run(argv, input=prompt_input, stdin=subprocess.DEVNULL if prompt_input is None else None, stdout=subprocess.PIPE, stderr=subprocess.PIPE, universal_newlines=True, timeout=timeout)
|
||
except subprocess.TimeoutExpired:
|
||
raise BackendError("%s run timed out after %ds" % (label, timeout))
|
||
except OSError as exc:
|
||
raise BackendError("%s run failed: %s" % (label, short_error(exc)))
|
||
return done
|
||
|
||
|
||
def grok_parse_output(stdout):
|
||
try:
|
||
data = json.loads(stdout or "")
|
||
except ValueError:
|
||
return opencode_parse(stdout)
|
||
if isinstance(data, dict):
|
||
if data.get("is_error") or data.get("error"):
|
||
detail = data.get("error") or data.get("result") or data.get("text") or "unknown error"
|
||
raise BackendError("grok: %s" % str(detail)[:300])
|
||
if isinstance(data.get("content"), str) or isinstance(data.get("tool_calls"), list):
|
||
if "content" in data or "tool_calls" in data:
|
||
return data
|
||
for key in ("text", "result", "output", "answer", "content"):
|
||
value = data.get(key)
|
||
if isinstance(value, str) and value.strip():
|
||
return opencode_parse(value)
|
||
return opencode_parse(stdout)
|
||
|
||
|
||
def grok_run(prompt, timeout):
|
||
binary = grok_binary()
|
||
if not binary:
|
||
raise BackendError("grok CLI not found (GROK_BIN or PATH)")
|
||
with tempfile.TemporaryDirectory(prefix="tai-grok-") as folder:
|
||
prompt_path = os.path.join(folder, "prompt.txt")
|
||
with open(prompt_path, "w", encoding="utf-8") as handle:
|
||
handle.write(prompt)
|
||
variants = [
|
||
[binary, "--prompt-file", prompt_path, "--output-format", "json", "--no-auto-update"],
|
||
[binary, "--prompt-file", prompt_path, "--output-format", "json"],
|
||
[binary, "-p", prompt, "--output-format", "json"],
|
||
[binary, "-p", prompt],
|
||
]
|
||
last_error = ""
|
||
for argv in variants:
|
||
done = _cli_run(argv, None, timeout, "grok")
|
||
if done.returncode != 0 and _cli_unknown_flag(done.stderr) and argv is not variants[-1]:
|
||
last_error = (done.stderr or "").strip()[:200]
|
||
continue
|
||
if done.returncode != 0 and not (done.stdout or "").strip():
|
||
blob = ((done.stderr or "") + "\n" + (done.stdout or "")).strip()
|
||
if _cli_auth_failure(blob):
|
||
raise BackendError("grok not authenticated (set XAI_API_KEY or run grok login): %s" % blob[:200])
|
||
raise BackendError("grok exit %d: %s" % (done.returncode, blob[:300] or last_error or "no output"))
|
||
if not (done.stdout or "").strip():
|
||
blob = (done.stderr or "").strip()
|
||
raise BackendError("grok returned no output: %s" % (blob[:300] or last_error or "empty stdout"))
|
||
return done.stdout or ""
|
||
raise BackendError("grok failed: %s" % (last_error or "no usable output"))
|
||
|
||
|
||
def claude_parse_output(stdout):
|
||
try:
|
||
data = json.loads(stdout or "")
|
||
except ValueError:
|
||
return opencode_parse(stdout)
|
||
if not isinstance(data, dict):
|
||
return opencode_parse(stdout)
|
||
if data.get("is_error"):
|
||
raise BackendError("claude: %s" % str(data.get("result") or "unknown error")[:300])
|
||
structured = data.get("structured_output")
|
||
if isinstance(structured, dict) and ("content" in structured or "tool_calls" in structured):
|
||
return structured
|
||
result = data.get("result")
|
||
if isinstance(result, str) and result.strip():
|
||
return opencode_parse(result)
|
||
return opencode_parse(stdout)
|
||
|
||
|
||
def claude_run(prompt, timeout):
|
||
binary = claude_binary()
|
||
if not binary:
|
||
raise BackendError("claude CLI not found (CLAUDE_BIN or PATH)")
|
||
schema = json.dumps(muse_output_schema())
|
||
base = [binary, "-p"]
|
||
if len(prompt.encode("utf-8", "replace")) > 120000:
|
||
variants = [
|
||
(base + ["--output-format", "json", "--json-schema", schema], prompt),
|
||
(base + ["--output-format", "json"], prompt),
|
||
]
|
||
else:
|
||
variants = [
|
||
(base + [prompt, "--output-format", "json", "--json-schema", schema], None),
|
||
(base + [prompt, "--output-format", "json"], None),
|
||
(base + [prompt], None),
|
||
]
|
||
last_error = ""
|
||
for argv, prompt_input in variants:
|
||
done = _cli_run(argv, prompt_input, timeout, "claude")
|
||
if done.returncode != 0 and _cli_unknown_flag(done.stderr) and (argv, prompt_input) != variants[-1]:
|
||
last_error = (done.stderr or "").strip()[:200]
|
||
continue
|
||
if done.returncode != 0 and not (done.stdout or "").strip():
|
||
blob = ((done.stderr or "") + "\n" + (done.stdout or "")).strip()
|
||
if _cli_auth_failure(blob):
|
||
raise BackendError("claude not authenticated (run claude login): %s" % blob[:200])
|
||
raise BackendError("claude exit %d: %s" % (done.returncode, blob[:300] or last_error or "no output"))
|
||
if not (done.stdout or "").strip():
|
||
blob = (done.stderr or "").strip()
|
||
raise BackendError("claude returned no output: %s" % (blob[:300] or last_error or "empty stdout"))
|
||
return done.stdout or ""
|
||
raise BackendError("claude failed: %s" % (last_error or "no usable output"))
|
||
|
||
|
||
def codex_jsonl_texts(stdout):
|
||
texts = []
|
||
last_error = ""
|
||
for line in (stdout or "").splitlines():
|
||
line = line.strip()
|
||
if not line.startswith("{"):
|
||
continue
|
||
try:
|
||
event = json.loads(line)
|
||
except ValueError:
|
||
continue
|
||
if not isinstance(event, dict):
|
||
continue
|
||
etype = event.get("type")
|
||
if etype == "item.completed":
|
||
item = event.get("item") or {}
|
||
if item.get("type") == "agent_message" and item.get("text"):
|
||
texts.append(item["text"])
|
||
elif etype in ("turn.failed", "error"):
|
||
msg = event.get("message") or event.get("error") or ""
|
||
if isinstance(msg, dict):
|
||
msg = msg.get("message") or json.dumps(msg)
|
||
if msg:
|
||
last_error = str(msg)[:300]
|
||
return texts, last_error
|
||
|
||
|
||
def codex_parse_output(stdout, outfile_text="", used_schema=False):
|
||
if isinstance(outfile_text, str) and outfile_text.strip():
|
||
if used_schema:
|
||
try:
|
||
answer = json.loads(outfile_text)
|
||
except ValueError:
|
||
answer = None
|
||
if isinstance(answer, dict) and ("content" in answer or "tool_calls" in answer):
|
||
return answer
|
||
return opencode_parse(outfile_text)
|
||
texts, last_error = codex_jsonl_texts(stdout)
|
||
if texts:
|
||
return opencode_parse("".join(texts))
|
||
if last_error:
|
||
raise BackendError("codex: %s" % last_error)
|
||
return opencode_parse(stdout)
|
||
|
||
|
||
def codex_run(prompt, timeout):
|
||
binary = codex_binary()
|
||
if not binary:
|
||
raise BackendError("codex CLI not found (CODEX_BIN or PATH)")
|
||
schema = json.dumps(muse_output_schema())
|
||
with tempfile.TemporaryDirectory(prefix="tai-codex-") as folder:
|
||
schema_path = os.path.join(folder, "schema.json")
|
||
out_path = os.path.join(folder, "last.txt")
|
||
with open(schema_path, "w", encoding="utf-8") as handle:
|
||
handle.write(schema)
|
||
base = [binary, "exec", "--json", "--ephemeral", "--sandbox", "read-only", "--skip-git-repo-check"]
|
||
if len(prompt.encode("utf-8", "replace")) > 120000:
|
||
variants = [
|
||
(base + ["--output-schema", schema_path, "-o", out_path, "-"], prompt, True),
|
||
(base + ["-o", out_path, "-"], prompt, False),
|
||
]
|
||
else:
|
||
variants = [
|
||
(base + ["--output-schema", schema_path, "-o", out_path, prompt], None, True),
|
||
(base + ["-o", out_path, prompt], None, False),
|
||
(base + [prompt], None, False),
|
||
]
|
||
last_error = ""
|
||
for argv, prompt_input, used_schema in variants:
|
||
try:
|
||
done = subprocess.run(argv, input=prompt_input, stdout=subprocess.PIPE, stderr=subprocess.PIPE, universal_newlines=True, cwd=folder, timeout=timeout)
|
||
except subprocess.TimeoutExpired:
|
||
raise BackendError("codex run timed out after %ds" % timeout)
|
||
except OSError as exc:
|
||
raise BackendError("codex run failed: %s" % short_error(exc))
|
||
if done.returncode != 0 and _cli_unknown_flag(done.stderr) and (argv, prompt_input, used_schema) != variants[-1]:
|
||
last_error = (done.stderr or "").strip()[:200]
|
||
continue
|
||
outfile_text = ""
|
||
if "-o" in argv:
|
||
try:
|
||
with open(out_path, "r", encoding="utf-8", errors="replace") as handle:
|
||
outfile_text = handle.read()
|
||
except OSError:
|
||
outfile_text = ""
|
||
if done.returncode != 0 and not (outfile_text.strip() or (done.stdout or "").strip()):
|
||
blob = ((done.stderr or "") + "\n" + (done.stdout or "")).strip()
|
||
if _cli_auth_failure(blob):
|
||
raise BackendError("codex not authenticated (run codex login): %s" % blob[:200])
|
||
raise BackendError("codex exit %d: %s" % (done.returncode, blob[:300] or last_error or "no output"))
|
||
if not (outfile_text.strip() or (done.stdout or "").strip()):
|
||
blob = (done.stderr or "").strip()
|
||
raise BackendError("codex returned no output: %s" % (blob[:300] or last_error or "empty stdout"))
|
||
return done.stdout or "", outfile_text, used_schema
|
||
raise BackendError("codex failed: %s" % (last_error or "no usable output"))
|
||
|
||
|
||
def gemini_parse_output(stdout):
|
||
try:
|
||
data = json.loads(stdout or "")
|
||
except ValueError:
|
||
return opencode_parse(stdout)
|
||
if not isinstance(data, dict):
|
||
return opencode_parse(stdout)
|
||
err = data.get("error")
|
||
if err:
|
||
raise BackendError("gemini: %s" % (json.dumps(err) if isinstance(err, dict) else str(err))[:300])
|
||
response = data.get("response")
|
||
if isinstance(response, str) and response.strip():
|
||
return opencode_parse(response)
|
||
return opencode_parse(stdout)
|
||
|
||
|
||
def gemini_run(prompt, timeout):
|
||
binary = gemini_binary()
|
||
if not binary:
|
||
raise BackendError("gemini CLI not found (GEMINI_BIN or PATH)")
|
||
base = [binary, "-p", prompt]
|
||
variants = [
|
||
base + ["--output-format", "json", "--approval-mode", "plan"],
|
||
base + ["--output-format", "json"],
|
||
base,
|
||
]
|
||
last_error = ""
|
||
for argv in variants:
|
||
done = _cli_run(argv, None, timeout, "gemini")
|
||
if done.returncode != 0 and _cli_unknown_flag(done.stderr) and argv is not variants[-1]:
|
||
last_error = (done.stderr or "").strip()[:200]
|
||
continue
|
||
if done.returncode != 0 and not (done.stdout or "").strip():
|
||
blob = ((done.stderr or "") + "\n" + (done.stdout or "")).strip()
|
||
if _cli_auth_failure(blob):
|
||
raise BackendError("gemini not authenticated (run gemini login or set GEMINI_API_KEY): %s" % blob[:200])
|
||
hint = {42: "input error", 53: "turn limit exceeded"}.get(done.returncode, "")
|
||
raise BackendError("gemini exit %d%s: %s" % (done.returncode, " (%s)" % hint if hint else "", blob[:300] or last_error or "no output"))
|
||
if not (done.stdout or "").strip():
|
||
blob = (done.stderr or "").strip()
|
||
raise BackendError("gemini returned no output: %s" % (blob[:300] or last_error or "empty stdout"))
|
||
return done.stdout or ""
|
||
raise BackendError("gemini failed: %s" % (last_error or "no usable output"))
|
||
|
||
|
||
def muse_output_schema():
|
||
call = {"type": "object", "additionalProperties": False, "properties": {"name": {"type": "string"}, "arguments": {"type": "string"}}, "required": ["name", "arguments"]}
|
||
return {"type": "object", "additionalProperties": False, "properties": {"content": {"type": "string"}, "tool_calls": {"type": "array", "items": call}}, "required": ["content", "tool_calls"]}
|
||
|
||
|
||
def muse_prompt(messages, tools):
|
||
parts = []
|
||
for message in messages or []:
|
||
role = (message.get("role") or "user").upper()
|
||
body = message.get("content") or ""
|
||
if isinstance(body, list):
|
||
body = " ".join(str(part.get("text") or "") for part in body if isinstance(part, dict))
|
||
if message.get("tool_call_id"):
|
||
parts.append("%s %s:\n%s" % (role, message["tool_call_id"], body))
|
||
else:
|
||
parts.append("%s:\n%s" % (role, body))
|
||
for call in message.get("tool_calls") or []:
|
||
func = call.get("function") or {}
|
||
parts.append("TOOL CALL %s:\n%s %s" % (call.get("id") or "", func.get("name") or "", func.get("arguments") or ""))
|
||
if tools:
|
||
parts.append("TOOLS:\n%s" % json.dumps(tools))
|
||
parts.append("Answer with content text plus tool_calls for every tool to invoke now. Empty tool_calls ends the turn.")
|
||
return "\n\n".join(part for part in parts if part.strip())
|
||
|
||
|
||
def muse_argv(config, schema_path, prompt_path, session_id=None):
|
||
argv = [muse_binary(), "exec", "--json", "--max-model-steps", str(MUSE_MAX_STEPS), "--disable-web-tools", "--disable-reminders", "--allow-workspace-switch", "--output-schema", schema_path, "--prompt-file", prompt_path]
|
||
if session_id is None:
|
||
argv.append("--no-session-log")
|
||
else:
|
||
argv.extend(["--session-id", session_id])
|
||
if config.explicit_model:
|
||
argv.extend(["--model", config.explicit_model])
|
||
return argv
|
||
|
||
|
||
def muse_write_file(path, text):
|
||
with open(path, "w", encoding="utf-8") as handle:
|
||
handle.write(text)
|
||
handle.flush()
|
||
os.fsync(handle.fileno())
|
||
|
||
|
||
def muse_message_hash(message):
|
||
blob = json.dumps(message, sort_keys=True, separators=(",", ":"), default=str)
|
||
return hashlib.sha256(blob.encode("utf-8")).hexdigest()
|
||
|
||
|
||
def muse_map_path(home):
|
||
return os.path.join(home, MUSE_SESSIONS_FILE)
|
||
|
||
|
||
def muse_load_map(home):
|
||
path = muse_map_path(home)
|
||
try:
|
||
with open(path, "r", encoding="utf-8") as handle:
|
||
data = json.load(handle)
|
||
except FileNotFoundError:
|
||
return {}
|
||
except (OSError, ValueError):
|
||
quarantine(path)
|
||
return {}
|
||
if not isinstance(data, dict):
|
||
quarantine(path)
|
||
return {}
|
||
now = time.time()
|
||
pruned = False
|
||
for key in list(data):
|
||
entry = data[key]
|
||
if not isinstance(entry, dict) or not isinstance(entry.get("id"), str) or not isinstance(entry.get("hashes"), list):
|
||
del data[key]
|
||
pruned = True
|
||
elif key.startswith("worker/"):
|
||
try:
|
||
stamp = _fromisoformat(entry.get("updated") or "").timestamp()
|
||
except ValueError:
|
||
stamp = 0
|
||
if now - stamp > MUSE_WORKER_TTL:
|
||
del data[key]
|
||
pruned = True
|
||
if pruned:
|
||
muse_save_map(home, data)
|
||
return data
|
||
|
||
|
||
def muse_save_map(home, data):
|
||
muse_write_file(muse_map_path(home), json.dumps(data, sort_keys=True) + "\n")
|
||
|
||
|
||
_MUSE_THREAD_LOCKS = {}
|
||
_MUSE_THREAD_GUARD = threading.Lock()
|
||
|
||
|
||
def muse_acquire(home, name):
|
||
with _MUSE_THREAD_GUARD:
|
||
thread_lock = _MUSE_THREAD_LOCKS.setdefault(name, threading.Lock())
|
||
thread_lock.acquire()
|
||
handle = None
|
||
if fcntl is not None:
|
||
try:
|
||
folder = os.path.join(home, MUSE_LOCKS_DIR)
|
||
os.makedirs(folder, exist_ok=True)
|
||
handle = open(os.path.join(folder, hashlib.sha256(name.encode("utf-8")).hexdigest() + ".lock"), "w")
|
||
fcntl.flock(handle.fileno(), fcntl.LOCK_EX)
|
||
except OSError:
|
||
if handle is not None:
|
||
try:
|
||
handle.close()
|
||
except OSError:
|
||
pass
|
||
handle = None
|
||
|
||
def release():
|
||
try:
|
||
if handle is not None:
|
||
if fcntl is not None:
|
||
try:
|
||
fcntl.flock(handle.fileno(), fcntl.LOCK_UN)
|
||
except OSError:
|
||
pass
|
||
try:
|
||
handle.close()
|
||
except OSError:
|
||
pass
|
||
finally:
|
||
thread_lock.release()
|
||
|
||
return release
|
||
|
||
|
||
def muse_transact(home, func):
|
||
release = muse_acquire(home, "map")
|
||
try:
|
||
data = muse_load_map(home)
|
||
result = func(data)
|
||
muse_save_map(home, data)
|
||
return result
|
||
finally:
|
||
release()
|
||
|
||
|
||
def muse_session_root():
|
||
base = os.environ.get("XDG_DATA_HOME") or os.path.join(os.path.expanduser("~"), ".local", "share")
|
||
return os.path.join(base, "muse", "sessions")
|
||
|
||
|
||
def muse_session_alive(session_id):
|
||
try:
|
||
pattern = os.path.join(muse_session_root(), "*", "*", "*", session_id, "session.jsonl")
|
||
return bool(glob.glob(pattern))
|
||
except (OSError, ValueError):
|
||
return True
|
||
|
||
|
||
def muse_session_split(messages, hashes):
|
||
if len(hashes) > len(messages):
|
||
return None
|
||
for pos, digest in enumerate(hashes):
|
||
if muse_message_hash(messages[pos]) != digest:
|
||
return None
|
||
return list(messages[len(hashes):])
|
||
|
||
|
||
def muse_valid_answer(text):
|
||
try:
|
||
data = json.loads(text)
|
||
except ValueError:
|
||
return None
|
||
if not isinstance(data, dict):
|
||
return None
|
||
if not isinstance(data.get("content"), str) or not isinstance(data.get("tool_calls"), list):
|
||
return None
|
||
for call in data["tool_calls"]:
|
||
if not isinstance(call, dict) or not isinstance(call.get("name"), str) or not isinstance(call.get("arguments"), str):
|
||
return None
|
||
return data
|
||
|
||
|
||
def muse_gate(func):
|
||
def gated(ok):
|
||
func(ok)
|
||
|
||
gated.armed = threading.Event()
|
||
return gated
|
||
|
||
|
||
def muse_reap(proc, stderr_done, on_reap):
|
||
def target():
|
||
ok = False
|
||
try:
|
||
try:
|
||
for _line in proc.stdout:
|
||
pass
|
||
except (OSError, ValueError):
|
||
pass
|
||
try:
|
||
proc.stdout.close()
|
||
except (OSError, ValueError):
|
||
pass
|
||
try:
|
||
ok = proc.wait(timeout=MUSE_REAP_GRACE) == 0
|
||
except subprocess.TimeoutExpired:
|
||
try:
|
||
proc.kill()
|
||
except OSError:
|
||
pass
|
||
try:
|
||
proc.wait(timeout=10)
|
||
except (OSError, subprocess.SubprocessError):
|
||
pass
|
||
try:
|
||
stderr_done.wait(timeout=10)
|
||
except (OSError, RuntimeError):
|
||
pass
|
||
finally:
|
||
if on_reap is not None:
|
||
try:
|
||
on_reap(ok)
|
||
except Exception:
|
||
pass
|
||
|
||
thread = threading.Thread(target=target, daemon=True)
|
||
thread.start()
|
||
return thread
|
||
|
||
|
||
def muse_run(argv, cwd, stream_sink, timeout, on_reap=None):
|
||
try:
|
||
proc = subprocess.Popen(argv, stdout=subprocess.PIPE, stderr=subprocess.PIPE, universal_newlines=True, cwd=cwd, bufsize=1)
|
||
except OSError as exc:
|
||
raise BackendError("muse exec failed: %s" % short_error(exc))
|
||
gate = getattr(on_reap, "armed", None)
|
||
if gate is not None:
|
||
gate.set()
|
||
stderr_parts = []
|
||
stderr_done = threading.Event()
|
||
|
||
def collect_stderr():
|
||
try:
|
||
for chunk in proc.stderr:
|
||
stderr_parts.append(chunk)
|
||
except (OSError, ValueError):
|
||
pass
|
||
finally:
|
||
stderr_done.set()
|
||
|
||
threading.Thread(target=collect_stderr, daemon=True).start()
|
||
finished = threading.Event()
|
||
expired = threading.Event()
|
||
|
||
def watchdog():
|
||
if not finished.wait(timeout):
|
||
expired.set()
|
||
try:
|
||
proc.kill()
|
||
except OSError:
|
||
pass
|
||
|
||
threading.Thread(target=watchdog, daemon=True).start()
|
||
answer = None
|
||
terminal = None
|
||
buffer = []
|
||
try:
|
||
for line in proc.stdout:
|
||
text = line.strip()
|
||
if not text.startswith("{"):
|
||
continue
|
||
try:
|
||
event = json.loads(text)
|
||
except ValueError:
|
||
continue
|
||
kind = event.get("payload_type") or ""
|
||
payload = event.get("payload") or {}
|
||
if kind == "run.output.delta" and payload.get("text"):
|
||
piece = payload["text"]
|
||
if stream_sink is not None:
|
||
stream_sink(piece)
|
||
buffer.append(piece)
|
||
if answer is None:
|
||
answer = muse_valid_answer("".join(buffer))
|
||
if answer is not None:
|
||
break
|
||
elif kind.startswith("run.terminal."):
|
||
terminal = payload
|
||
break
|
||
finally:
|
||
finished.set()
|
||
muse_reap(proc, stderr_done, on_reap)
|
||
if terminal is not None:
|
||
if terminal.get("terminal") != "completed":
|
||
raise BackendError("muse %s: %s" % (terminal.get("terminal") or "failed", (terminal.get("reason") or "unknown")[:300]))
|
||
parsed = muse_valid_answer(terminal.get("text") or "")
|
||
if parsed is None:
|
||
raise BackendError("muse returned invalid structured answer")
|
||
return parsed
|
||
if answer is not None:
|
||
return answer
|
||
if expired.is_set():
|
||
raise BackendError("muse timed out after %ds" % timeout)
|
||
try:
|
||
code = proc.returncode
|
||
except (OSError, ValueError):
|
||
code = "?"
|
||
raise BackendError("muse exit %s: %s" % (code, "".join(stderr_parts).strip()[-300:] or "no terminal event"))
|
||
|
||
|
||
def muse_health():
|
||
binary = muse_binary()
|
||
if not binary:
|
||
return "missing: muse CLI not found"
|
||
try:
|
||
done = subprocess.run([binary, "--version"], stdout=subprocess.PIPE, stderr=subprocess.PIPE, universal_newlines=True, timeout=15)
|
||
except (OSError, subprocess.SubprocessError) as exc:
|
||
return "fail: %s" % short_error(exc)[:60]
|
||
if done.returncode != 0:
|
||
return "fail: exit %d" % done.returncode
|
||
return "ok, subscription (%s)" % (done.stdout.strip().splitlines()[0][:60] if done.stdout.strip() else "muse")
|
||
|
||
|
||
def grok_health():
|
||
binary = grok_binary()
|
||
if not binary:
|
||
return "missing: grok CLI not found"
|
||
try:
|
||
done = subprocess.run([binary, "--version"], stdout=subprocess.PIPE, stderr=subprocess.PIPE, universal_newlines=True, timeout=15)
|
||
except (OSError, subprocess.SubprocessError) as exc:
|
||
return "fail: %s" % short_error(exc)[:60]
|
||
if done.returncode != 0:
|
||
return "fail: exit %d" % done.returncode
|
||
return "ok, subscription (%s)" % ((done.stdout.strip() or done.stderr.strip()).splitlines()[0][:60] if (done.stdout.strip() or done.stderr.strip()) else "grok")
|
||
|
||
|
||
def claude_health():
|
||
binary = claude_binary()
|
||
if not binary:
|
||
return "missing: claude CLI not found"
|
||
try:
|
||
done = subprocess.run([binary, "--version"], stdout=subprocess.PIPE, stderr=subprocess.PIPE, universal_newlines=True, timeout=15)
|
||
except (OSError, subprocess.SubprocessError) as exc:
|
||
return "fail: %s" % short_error(exc)[:60]
|
||
if done.returncode != 0:
|
||
return "fail: exit %d" % done.returncode
|
||
return "ok, subscription (%s)" % ((done.stdout.strip() or done.stderr.strip()).splitlines()[0][:60] if (done.stdout.strip() or done.stderr.strip()) else "claude")
|
||
|
||
|
||
def codex_health():
|
||
binary = codex_binary()
|
||
if not binary:
|
||
return "missing: codex CLI not found"
|
||
try:
|
||
done = subprocess.run([binary, "--version"], stdout=subprocess.PIPE, stderr=subprocess.PIPE, universal_newlines=True, timeout=15)
|
||
except (OSError, subprocess.SubprocessError) as exc:
|
||
return "fail: %s" % short_error(exc)[:60]
|
||
if done.returncode != 0:
|
||
return "fail: exit %d" % done.returncode
|
||
return "ok, subscription (%s)" % ((done.stdout.strip() or done.stderr.strip()).splitlines()[0][:60] if (done.stdout.strip() or done.stderr.strip()) else "codex")
|
||
|
||
|
||
def gemini_health():
|
||
binary = gemini_binary()
|
||
if not binary:
|
||
return "missing: gemini CLI not found"
|
||
try:
|
||
done = subprocess.run([binary, "--version"], stdout=subprocess.PIPE, stderr=subprocess.PIPE, universal_newlines=True, timeout=15)
|
||
except (OSError, subprocess.SubprocessError) as exc:
|
||
return "fail: %s" % short_error(exc)[:60]
|
||
if done.returncode != 0:
|
||
return "fail: exit %d" % done.returncode
|
||
return "ok, subscription (%s)" % ((done.stdout.strip() or done.stderr.strip()).splitlines()[0][:60] if (done.stdout.strip() or done.stderr.strip()) else "gemini")
|
||
|
||
|
||
_OPENCODE_ID_ALPHABET = "0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz"
|
||
_OPENCODE_SESSION = "ses_" + "".join(random.choice(_OPENCODE_ID_ALPHABET) for _ in range(26))
|
||
|
||
|
||
def opencode_headers():
|
||
return {
|
||
"User-Agent": os.environ.get("OPENCODE_CLIENT_USER_AGENT") or OPENCODE_USER_AGENT,
|
||
"x-opencode-client": os.environ.get("OPENCODE_CLIENT_NAME") or OPENCODE_CLIENT_NAME,
|
||
"x-opencode-project": "global",
|
||
"x-opencode-session": _OPENCODE_SESSION,
|
||
"x-opencode-request": "msg_" + "".join(random.choice(_OPENCODE_ID_ALPHABET) for _ in range(26)),
|
||
}
|
||
|
||
|
||
def model_speed_reward(tokens_per_second):
|
||
if not tokens_per_second or tokens_per_second <= 0:
|
||
return MODEL_NEUTRAL_REWARD
|
||
return tokens_per_second / (tokens_per_second + MODEL_SPEED_REF_TPS)
|
||
|
||
|
||
def model_latency_reward(latency_ms):
|
||
if not latency_ms or latency_ms <= 0:
|
||
return 1.0
|
||
return 1.0 / (1.0 + latency_ms / MODEL_LATENCY_REF_MS)
|
||
|
||
|
||
def model_outcome_reward(latency_ms=None, tokens_per_second=None):
|
||
return model_speed_reward(tokens_per_second) * model_latency_reward(latency_ms)
|
||
|
||
|
||
def model_weight(entry):
|
||
entry = entry or {}
|
||
attempts = (entry.get("s") or 0) + (entry.get("f") or 0)
|
||
return ((entry.get("r") or 0.0) + MODEL_NEUTRAL_REWARD * MODEL_WEIGHT_PRIOR) / (attempts + MODEL_WEIGHT_PRIOR)
|
||
|
||
|
||
class ModelHealth:
|
||
def __init__(self, home):
|
||
self.path = os.path.join(home, MODEL_HEALTH_FILE) if home else ""
|
||
self.data = {"models": {}, "lists": {}}
|
||
if self.path:
|
||
try:
|
||
with open(self.path, "r", encoding="utf-8") as handle:
|
||
raw = json.load(handle)
|
||
if isinstance(raw, dict):
|
||
if isinstance(raw.get("models"), dict):
|
||
self.data["models"] = raw["models"]
|
||
if isinstance(raw.get("lists"), dict):
|
||
self.data["lists"] = raw["lists"]
|
||
except (OSError, ValueError):
|
||
pass
|
||
|
||
def _save(self):
|
||
if not self.path:
|
||
return
|
||
try:
|
||
tmp = self.path + ".tmp"
|
||
with open(tmp, "w", encoding="utf-8") as handle:
|
||
json.dump(self.data, handle)
|
||
os.replace(tmp, self.path)
|
||
except OSError:
|
||
pass
|
||
|
||
def entry(self, provider, model):
|
||
return self.data["models"].get("%s|%s" % (provider or "default", model or "")) or {}
|
||
|
||
def weight(self, provider, model):
|
||
return model_weight(self.entry(provider, model))
|
||
|
||
def circuit_open(self, provider, model):
|
||
return (self.entry(provider, model).get("until") or 0.0) > time.time()
|
||
|
||
def record(self, provider, model, success, latency_ms=None, tokens_per_second=None):
|
||
if not model:
|
||
return
|
||
key = "%s|%s" % (provider or "default", model or "")
|
||
entry = self.data["models"].setdefault(key, {"s": 0, "f": 0, "r": 0.0, "lat": 0.0, "tps": 0.0, "cf": 0, "until": 0.0, "used": 0.0, "n": 0})
|
||
entry["n"] += 1
|
||
entry["used"] = time.time()
|
||
if success:
|
||
reward = MODEL_NEUTRAL_REWARD if latency_ms is None and tokens_per_second is None else model_outcome_reward(latency_ms, tokens_per_second)
|
||
entry["s"] += 1
|
||
entry["cf"] = 0
|
||
entry["until"] = 0.0
|
||
entry["r"] += reward
|
||
if latency_ms is not None:
|
||
entry["lat"] += latency_ms
|
||
if tokens_per_second is not None:
|
||
entry["tps"] += tokens_per_second
|
||
else:
|
||
entry["f"] += 1
|
||
entry["cf"] += 1
|
||
if entry["cf"] >= MODEL_CIRCUIT_THRESHOLD:
|
||
entry["until"] = time.time() + MODEL_CIRCUIT_COOLDOWN
|
||
self._save()
|
||
|
||
def cached_list(self, name, ttl=OPENROUTER_FREE_CACHE_TTL):
|
||
row = self.data["lists"].get(name)
|
||
if not isinstance(row, dict):
|
||
return None
|
||
models = row.get("models")
|
||
if not models or (time.time() - (row.get("at") or 0)) > ttl:
|
||
return None
|
||
return [str(item) for item in models if str(item).strip()]
|
||
|
||
def store_list(self, name, models):
|
||
cleaned = [str(item).strip() for item in models if str(item).strip()]
|
||
if not cleaned:
|
||
return
|
||
self.data["lists"][name] = {"at": time.time(), "models": cleaned}
|
||
self._save()
|
||
|
||
def snapshot(self):
|
||
rows = []
|
||
for key, entry in self.data["models"].items():
|
||
provider, _, model = key.partition("|")
|
||
ok = entry.get("s") or 0
|
||
rows.append({
|
||
"provider": provider,
|
||
"model": model,
|
||
"weight": round(model_weight(entry), 4),
|
||
"ok": ok,
|
||
"fail": entry.get("f") or 0,
|
||
"avg_ms": round((entry.get("lat") or 0.0) / ok) if ok and entry.get("lat") else None,
|
||
"avg_tps": round((entry.get("tps") or 0.0) / ok, 1) if ok and entry.get("tps") else None,
|
||
"circuit_open": (entry.get("until") or 0.0) > time.time(),
|
||
"used": entry.get("used") or 0.0,
|
||
})
|
||
rows.sort(key=lambda row: (-row["weight"], row["provider"], row["model"]))
|
||
return rows
|
||
|
||
|
||
def fetch_model_ids(models_url, key="", extra_headers=None, timeout=MODEL_LIST_TIMEOUT):
|
||
headers = dict(extra_headers or {})
|
||
if key:
|
||
headers["Authorization"] = "Bearer " + key
|
||
try:
|
||
request = urllib.request.Request(models_url, headers=headers, method="GET")
|
||
with urllib.request.urlopen(request, timeout=timeout) as response:
|
||
payload = json.loads(response.read().decode("utf-8", "replace"))
|
||
except (OSError, ValueError):
|
||
return None
|
||
rows = payload.get("data") if isinstance(payload, dict) else None
|
||
if not isinstance(rows, list):
|
||
return None
|
||
ids = [str(row.get("id")).strip() for row in rows if isinstance(row, dict) and str(row.get("id") or "").strip()]
|
||
return ids or None
|
||
|
||
|
||
def embedded_failure_message(data):
|
||
if not isinstance(data, dict):
|
||
return ""
|
||
error = data.get("error")
|
||
if isinstance(error, dict):
|
||
message = error.get("message") or error.get("type") or ""
|
||
return str(message).strip()
|
||
if isinstance(error, str) and error.strip():
|
||
return error.strip()
|
||
if data.get("type") == "error":
|
||
return str(data.get("message") or "upstream error").strip()
|
||
return ""
|
||
|
||
|
||
def is_free_model(model):
|
||
model = (model or "").lower()
|
||
return model.endswith(":free") or model.endswith("-free")
|
||
|
||
|
||
def prefer_muse(models):
|
||
return sorted(models, key=lambda item: 0 if MUSE_LABEL in item.lower() else 1)
|
||
|
||
|
||
def openrouter_free_tier(cand):
|
||
model = (cand.get("model") or "").lower()
|
||
muse = 0 if MUSE_LABEL in model else 1
|
||
if cand.get("provider") == OPENCODE_LABEL:
|
||
return muse
|
||
if cand.get("provider") == DEVPLACE_LABEL:
|
||
return 2 + muse
|
||
if cand.get("provider") == OPENROUTER_LABEL and is_free_model(model):
|
||
return 4 + muse
|
||
return 6
|
||
|
||
|
||
def models_report(chat):
|
||
lines = ["label/model | weight ok/fail avg-ms avg-tps status"]
|
||
seen = set()
|
||
for cand in chat.ordered_candidates():
|
||
key = (cand["provider"], cand["model"])
|
||
if key in seen:
|
||
continue
|
||
seen.add(key)
|
||
entry = chat.health.entry(cand["provider"], cand["model"])
|
||
ok = entry.get("s") or 0
|
||
fail = entry.get("f") or 0
|
||
avg_ms = str(round((entry.get("lat") or 0.0) / ok)) if ok and entry.get("lat") else "-"
|
||
avg_tps = str(round((entry.get("tps") or 0.0) / ok, 1)) if ok and entry.get("tps") else "-"
|
||
if cand["chat_url"] == MUSE_CHAT_URL:
|
||
status = muse_health()
|
||
elif chat.health.circuit_open(cand["provider"], cand["model"]):
|
||
status = "circuit open"
|
||
elif not ok and not fail:
|
||
status = "untested"
|
||
else:
|
||
status = "ready"
|
||
lines.append("%s/%s | %.4f %d/%d %s %s %s" % (cand["provider"], cand["model"], model_weight(entry), ok, fail, avg_ms, avg_tps, status))
|
||
return "\n".join(lines)
|
||
|
||
|
||
class ChatClient:
|
||
def __init__(self, config):
|
||
self.config = config
|
||
self.scope = ("default", "main")
|
||
self.fixed_key = None
|
||
self.health = ModelHealth(getattr(config, "home", "") or "")
|
||
|
||
def muse_key(self):
|
||
if self.fixed_key is not None:
|
||
return self.fixed_key
|
||
return "%s/%s" % (self.scope[0], self.scope[1])
|
||
|
||
def reset_muse(self):
|
||
home = getattr(self.config, "home", None)
|
||
if home is None:
|
||
return
|
||
key = self.muse_key()
|
||
|
||
def drop(data):
|
||
data.pop(key, None)
|
||
|
||
muse_transact(home, drop)
|
||
|
||
def openrouter_models(self):
|
||
models = [item for item in getattr(self.config, "openrouter_models", []) if item]
|
||
explicit = getattr(self.config, "explicit_model", "")
|
||
if explicit and explicit not in models:
|
||
models.insert(0, explicit)
|
||
free = self.health.cached_list("openrouter:free")
|
||
if free is None:
|
||
fetched = fetch_model_ids(OPENROUTER_BASE + "/models", self.config.openrouter_key)
|
||
free = [item for item in fetched if item.endswith(":free")] if fetched else []
|
||
if free:
|
||
self.health.store_list("openrouter:free", free)
|
||
free = prefer_muse(free or [])
|
||
models = [item for item in free if item not in models] + models
|
||
if not models:
|
||
models = [OPENROUTER_MODEL]
|
||
return models
|
||
|
||
def opencode_models(self):
|
||
models = list(getattr(self.config, "opencode_models", []) or [])
|
||
if not models:
|
||
models = self.health.cached_list("opencode:free")
|
||
if models is None:
|
||
models = opencode_free_models() or [item for item in ZEN_FREE_MODELS if is_free_model(item)]
|
||
self.health.store_list("opencode:free", models)
|
||
return prefer_muse(models)
|
||
|
||
def ollama_models(self):
|
||
if getattr(self.config, "ollama_enabled", True) is False:
|
||
return []
|
||
models = list(getattr(self.config, "ollama_models", []) or [])
|
||
if not models:
|
||
models = self.health.cached_list("ollama:local", OLLAMA_LIST_TTL)
|
||
if models is None:
|
||
models = fetch_model_ids(OLLAMA_BASE + "/models", "", timeout=3) or []
|
||
if models:
|
||
self.health.store_list("ollama:local", models)
|
||
return models
|
||
|
||
def devplace_models(self):
|
||
models = list(getattr(self.config, "devplace_models", []) or [])
|
||
if models:
|
||
return prefer_muse(models)
|
||
free = self.health.cached_list("devplace:free")
|
||
if free is None:
|
||
fetched = fetch_model_ids(self.config.devplace_base + "/models", self.config.devplace_key)
|
||
if fetched is None:
|
||
return list(DEVPLACE_FREE_ROUTES)
|
||
free = [item for item in fetched if item in DEVPLACE_FREE_ROUTES or (MUSE_LABEL in item.lower() and is_free_model(item))]
|
||
if free:
|
||
self.health.store_list("devplace:free", free)
|
||
return prefer_muse(free or [])
|
||
|
||
def candidates(self):
|
||
found = []
|
||
if getattr(self.config, "opencode_enabled", True) and opencode_binary():
|
||
for model in self.opencode_models():
|
||
found.append({"provider": OPENCODE_LABEL, "chat_url": OPENCODE_CHAT_URL, "models_url": OPENCODE_MODELS_URL, "key": "", "model": model, "profile": ""})
|
||
for model in self.ollama_models():
|
||
found.append({"provider": OLLAMA_LABEL, "chat_url": OLLAMA_BASE + "/chat/completions", "models_url": OLLAMA_BASE + "/models", "key": "", "model": model, "profile": ""})
|
||
if getattr(self.config, "devplace_key", ""):
|
||
for model in self.devplace_models():
|
||
found.append({"provider": DEVPLACE_LABEL, "chat_url": self.config.devplace_base + "/chat/completions", "models_url": self.config.devplace_base + "/models", "key": self.config.devplace_key, "model": model, "profile": ""})
|
||
for entry in self.config.backends:
|
||
for model in split_models(entry.get("model")) or [self.config.model]:
|
||
found.append({"provider": entry["label"], "chat_url": entry["base"] + "/chat/completions", "models_url": entry["base"] + "/models", "key": entry["key"], "model": model, "profile": entry.get("profile") or ""})
|
||
if getattr(self.config, "openrouter_enabled", True):
|
||
for model in self.openrouter_models():
|
||
found.append({"provider": OPENROUTER_LABEL, "chat_url": OPENROUTER_BASE + "/chat/completions", "models_url": OPENROUTER_BASE + "/models", "key": self.config.openrouter_key, "model": model, "profile": ""})
|
||
if getattr(self.config, "pollinations_enabled", True):
|
||
for model in getattr(self.config, "pollinations_models", []) or list(POLLINATIONS_MODELS):
|
||
found.append({"provider": POLLINATIONS_LABEL, "chat_url": POLLINATIONS_BASE + "/chat/completions", "models_url": POLLINATIONS_BASE + "/models", "key": "", "model": model, "profile": ""})
|
||
if getattr(self.config, "zen_enabled", False):
|
||
for model in getattr(self.config, "zen_models", []) or list(ZEN_FREE_MODELS):
|
||
found.append({"provider": ZEN_LABEL, "chat_url": ZEN_BASE + "/chat/completions", "models_url": ZEN_BASE + "/models", "key": "", "model": model, "profile": "opencode"})
|
||
found.append({"provider": MUSE_LABEL, "chat_url": MUSE_CHAT_URL, "models_url": MUSE_MODELS_URL, "key": "", "model": self.config.model, "profile": ""})
|
||
if grok_binary():
|
||
found.append({"provider": GROK_LABEL, "chat_url": GROK_CHAT_URL, "models_url": GROK_MODELS_URL, "key": "", "model": self.config.model, "profile": ""})
|
||
if claude_binary():
|
||
found.append({"provider": CLAUDE_LABEL, "chat_url": CLAUDE_CHAT_URL, "models_url": CLAUDE_MODELS_URL, "key": "", "model": self.config.model, "profile": ""})
|
||
if codex_binary():
|
||
found.append({"provider": CODEX_LABEL, "chat_url": CODEX_CHAT_URL, "models_url": CODEX_MODELS_URL, "key": "", "model": self.config.model, "profile": ""})
|
||
if gemini_binary():
|
||
found.append({"provider": GEMINI_LABEL, "chat_url": GEMINI_CHAT_URL, "models_url": GEMINI_MODELS_URL, "key": "", "model": self.config.model, "profile": ""})
|
||
return found
|
||
|
||
def ordered_candidates(self):
|
||
http = [cand for cand in self.candidates() if cand["chat_url"] not in CLI_CHAT_URLS]
|
||
cli = [cand for cand in self.candidates() if cand["chat_url"] in CLI_CHAT_URLS]
|
||
ranked = sorted(enumerate(http), key=lambda pair: (openrouter_free_tier(pair[1]), -self.health.weight(pair[1]["provider"], pair[1]["model"]), pair[0]))
|
||
usable = [cand for _, cand in ranked if not self.health.circuit_open(cand["provider"], cand["model"])]
|
||
return (usable or [cand for _, cand in ranked]) + cli
|
||
|
||
def backends(self):
|
||
rows = []
|
||
seen = set()
|
||
for cand in self.candidates():
|
||
key = (cand["provider"], cand["chat_url"])
|
||
if key in seen:
|
||
continue
|
||
seen.add(key)
|
||
rows.append((cand["provider"], cand["chat_url"], cand["models_url"], cand["key"], cand["model"]))
|
||
return rows
|
||
|
||
def headers(self, key, extra=None):
|
||
headers = {"Content-Type": "application/json"}
|
||
if key:
|
||
headers["Authorization"] = "Bearer " + key
|
||
if extra:
|
||
headers.update(extra)
|
||
return headers
|
||
|
||
def complete(self, messages, tools=None, stream_sink=None, ephemeral=False, deadline=None):
|
||
payload = {"model": self.config.model, "messages": messages}
|
||
if tools:
|
||
payload["tools"] = tools
|
||
payload["tool_choice"] = "auto"
|
||
limit = getattr(self.config, "retry_seconds", 0)
|
||
give_up = time.time() + limit if limit > 0 else None
|
||
if deadline is not None:
|
||
give_up = deadline if give_up is None else min(give_up, deadline)
|
||
round_no = 0
|
||
while True:
|
||
errors = []
|
||
cands = self.ordered_candidates() if round_no == 0 else self.all_candidates()
|
||
for cand in cands:
|
||
reply = self.try_candidate(cand, messages, tools, payload, stream_sink, ephemeral, errors)
|
||
if reply is not None:
|
||
return reply
|
||
pause = RETRY_BACKOFF[min(round_no, len(RETRY_BACKOFF) - 1)]
|
||
round_no += 1
|
||
if give_up is not None and time.time() + pause > give_up:
|
||
raise BackendError("all backends failed after %d rounds (%s)" % (round_no, "; ".join(errors[-6:])))
|
||
print(paint("all %d backends failed (round %d), retrying in %ds: %s" % (len(cands), round_no, pause, (errors[-1] if errors else "")[:160]), Ansi.DIM), file=sys.stderr)
|
||
time.sleep(pause)
|
||
|
||
def all_candidates(self):
|
||
http = [cand for cand in self.candidates() if cand["chat_url"] not in CLI_CHAT_URLS]
|
||
cli = [cand for cand in self.candidates() if cand["chat_url"] in CLI_CHAT_URLS]
|
||
ranked = sorted(enumerate(http), key=lambda pair: (openrouter_free_tier(pair[1]), -self.health.weight(pair[1]["provider"], pair[1]["model"]), pair[0]))
|
||
return [cand for _, cand in ranked] + cli
|
||
|
||
def call_candidate(self, cand, messages, tools, attempt, stream_sink, ephemeral):
|
||
label, chat_url, key, model = cand["provider"], cand["chat_url"], cand["key"], cand["model"]
|
||
extra = opencode_headers() if cand["profile"] == "opencode" else None
|
||
timeout = STREAM_TIMEOUT if stream_sink is not None else HTTP_TIMEOUT
|
||
if chat_url == MUSE_CHAT_URL:
|
||
return self.muse_turn(messages, tools, stream_sink, label, timeout=timeout, ephemeral=ephemeral)
|
||
if chat_url == GROK_CHAT_URL:
|
||
return self.grok_turn(messages, tools, stream_sink, label, timeout=timeout)
|
||
if chat_url == CLAUDE_CHAT_URL:
|
||
return self.claude_turn(messages, tools, stream_sink, label, timeout=timeout)
|
||
if chat_url == CODEX_CHAT_URL:
|
||
return self.codex_turn(messages, tools, stream_sink, label, timeout=timeout)
|
||
if chat_url == GEMINI_CHAT_URL:
|
||
return self.gemini_turn(messages, tools, stream_sink, label, timeout=timeout)
|
||
if chat_url == OPENCODE_CHAT_URL:
|
||
return self.opencode_turn(messages, tools, stream_sink, label, model, timeout=timeout)
|
||
if stream_sink is None:
|
||
return self.single_shot(chat_url, key, attempt, label, extra)
|
||
return self.streaming(chat_url, key, attempt, label, stream_sink, extra)
|
||
|
||
def try_candidate(self, cand, messages, tools, payload, stream_sink, ephemeral, errors):
|
||
label, model = cand["provider"], cand["model"]
|
||
attempt = dict(payload, model=model)
|
||
tries = [(tools, attempt)]
|
||
if attempt.get("tools"):
|
||
tries.append((None, {name: value for name, value in attempt.items() if name not in ("tools", "tool_choice")}))
|
||
for pos, (use_tools, body) in enumerate(tries):
|
||
started = time.time()
|
||
try:
|
||
reply = self.call_candidate(cand, messages, use_tools, body, stream_sink, ephemeral)
|
||
if not isinstance(reply, dict):
|
||
raise BackendError("backend returned no reply")
|
||
self.health.record(label, model, True, (time.time() - started) * 1000, max(1, len(reply.get("content") or "") // 4) / max(time.time() - started, 0.001))
|
||
return reply
|
||
except BackendError as exc:
|
||
if pos == 0 and exc.status == 400 and len(tries) > 1:
|
||
continue
|
||
self.health.record(label, model, False)
|
||
errors.append("%s/%s: %s" % (label, model, exc))
|
||
return None
|
||
except Exception as exc:
|
||
self.health.record(label, model, False)
|
||
errors.append("%s/%s: %s" % (label, model, short_error(exc)))
|
||
return None
|
||
return None
|
||
|
||
def single_shot(self, url, key, payload, label, extra=None):
|
||
body = json.dumps(payload).encode("utf-8")
|
||
request = urllib.request.Request(url, data=body, headers=self.headers(key, extra), method="POST")
|
||
try:
|
||
with urllib.request.urlopen(request, timeout=HTTP_TIMEOUT) as response:
|
||
data = json.loads(response.read().decode("utf-8", "replace"))
|
||
except urllib.error.HTTPError as exc:
|
||
raise BackendError("HTTP %d: %s" % (exc.code, exc.read().decode("utf-8", "replace")[:300]), exc.code)
|
||
except urllib.error.URLError as exc:
|
||
raise BackendError(short_error(exc))
|
||
failure = embedded_failure_message(data)
|
||
if failure:
|
||
raise BackendError("upstream error: %s" % failure[:300])
|
||
try:
|
||
message = data["choices"][0]["message"]
|
||
except (KeyError, IndexError, TypeError):
|
||
raise BackendError("upstream returned no choices: %s" % json.dumps(data)[:300])
|
||
return self.normalize(message, label)
|
||
|
||
def normalize(self, message, label):
|
||
calls = []
|
||
for call in message.get("tool_calls") or []:
|
||
func = call.get("function") or {}
|
||
calls.append({"id": call.get("id") or uuid.uuid4().hex, "name": func.get("name") or "", "arguments": func.get("arguments") or "{}"})
|
||
return {"role": "assistant", "content": message.get("content") or "", "reasoning": message.get("reasoning") or message.get("reasoning_content") or "", "tool_calls": calls, "backend": label}
|
||
|
||
def streaming(self, url, key, payload, label, sink, extra=None):
|
||
body = json.dumps(dict(payload, stream=True)).encode("utf-8")
|
||
request = urllib.request.Request(url, data=body, headers=self.headers(key, extra), method="POST")
|
||
try:
|
||
response = urllib.request.urlopen(request, timeout=STREAM_TIMEOUT)
|
||
except urllib.error.HTTPError as exc:
|
||
raise BackendError("HTTP %d: %s" % (exc.code, exc.read().decode("utf-8", "replace")[:300]), exc.code)
|
||
except urllib.error.URLError as exc:
|
||
raise BackendError(short_error(exc))
|
||
content_parts = []
|
||
reasoning_parts = []
|
||
calls = {}
|
||
try:
|
||
for raw in response:
|
||
line = raw.decode("utf-8", "replace").strip()
|
||
if not line.startswith("data:"):
|
||
continue
|
||
data = line[5:].strip()
|
||
if data == "[DONE]":
|
||
break
|
||
try:
|
||
event = json.loads(data)
|
||
except ValueError:
|
||
continue
|
||
choices = event.get("choices") or []
|
||
if not choices:
|
||
continue
|
||
delta = choices[0].get("delta") or {}
|
||
piece = delta.get("content")
|
||
if piece:
|
||
content_parts.append(piece)
|
||
sink(piece)
|
||
think = delta.get("reasoning") or delta.get("reasoning_content")
|
||
if think:
|
||
reasoning_parts.append(think)
|
||
for call in delta.get("tool_calls") or []:
|
||
slot = calls.setdefault(call.get("index", 0), {"id": "", "name": "", "arguments": ""})
|
||
if call.get("id"):
|
||
slot["id"] = call["id"]
|
||
func = call.get("function") or {}
|
||
if func.get("name"):
|
||
slot["name"] += func["name"]
|
||
if func.get("arguments"):
|
||
slot["arguments"] += func["arguments"]
|
||
except OSError as exc:
|
||
if not content_parts and not calls:
|
||
raise BackendError(short_error(exc))
|
||
finally:
|
||
response.close()
|
||
ordered = [calls[pos] for pos in sorted(calls)]
|
||
for slot in ordered:
|
||
if not slot["id"]:
|
||
slot["id"] = uuid.uuid4().hex
|
||
return {"role": "assistant", "content": "".join(content_parts), "reasoning": "".join(reasoning_parts), "tool_calls": ordered, "backend": label}
|
||
|
||
def muse_turn(self, messages, tools, stream_sink, label, timeout=HTTP_TIMEOUT, ephemeral=False):
|
||
binary = muse_binary()
|
||
if not binary:
|
||
raise BackendError("muse CLI not found (MUSE_BIN or PATH)")
|
||
home = getattr(self.config, "home", None)
|
||
if ephemeral or home is None:
|
||
with tempfile.TemporaryDirectory(prefix="tai-muse-") as folder:
|
||
schema_path = os.path.join(folder, "schema.json")
|
||
prompt_path = os.path.join(folder, "prompt.txt")
|
||
muse_write_file(schema_path, json.dumps(muse_output_schema()))
|
||
muse_write_file(prompt_path, muse_prompt(messages, tools))
|
||
answer = muse_run(muse_argv(self.config, schema_path, prompt_path), folder, stream_sink, timeout)
|
||
return self.muse_reply(answer, label)
|
||
key = self.muse_key()
|
||
release = muse_acquire(home, "session:" + key)
|
||
|
||
def on_reap(ok):
|
||
try:
|
||
if not ok:
|
||
self.reset_muse()
|
||
finally:
|
||
release()
|
||
|
||
gated = muse_gate(on_reap)
|
||
try:
|
||
return self.muse_persistent(messages, tools, stream_sink, label, timeout, home, key, gated)
|
||
except BaseException:
|
||
if not gated.armed.is_set():
|
||
release()
|
||
raise
|
||
|
||
def muse_persistent(self, messages, tools, stream_sink, label, timeout, home, key, gated):
|
||
def plan(data):
|
||
entry = data.get(key)
|
||
pending = None
|
||
session_id = None
|
||
if isinstance(entry, dict) and isinstance(entry.get("id"), str) and isinstance(entry.get("hashes"), list):
|
||
pending = muse_session_split(messages, entry["hashes"])
|
||
if pending is not None and not muse_session_alive(entry["id"]):
|
||
pending = None
|
||
if pending is not None:
|
||
session_id = entry["id"]
|
||
if pending is None or not pending:
|
||
session_id = str(uuid.uuid4())
|
||
pending = list(messages)
|
||
return session_id, pending
|
||
|
||
session_id, pending = muse_transact(home, plan)
|
||
try:
|
||
with tempfile.TemporaryDirectory(prefix="tai-muse-") as folder:
|
||
schema_path = os.path.join(folder, "schema.json")
|
||
prompt_path = os.path.join(folder, "prompt.txt")
|
||
muse_write_file(schema_path, json.dumps(muse_output_schema()))
|
||
muse_write_file(prompt_path, muse_prompt(pending, tools))
|
||
answer = muse_run(muse_argv(self.config, schema_path, prompt_path, session_id), folder, stream_sink, timeout, gated)
|
||
except BackendError:
|
||
self.reset_muse()
|
||
raise
|
||
|
||
def commit(data):
|
||
data[key] = {"id": session_id, "hashes": [muse_message_hash(item) for item in messages], "updated": now_iso()}
|
||
|
||
muse_transact(home, commit)
|
||
return self.muse_reply(answer, label)
|
||
|
||
def opencode_turn(self, messages, tools, stream_sink, label, model, timeout=HTTP_TIMEOUT):
|
||
answer = opencode_parse(opencode_run(model, opencode_prompt(messages, tools), timeout))
|
||
if stream_sink is not None and answer.get("content"):
|
||
stream_sink(answer["content"])
|
||
return self.muse_reply(answer, label)
|
||
|
||
def grok_turn(self, messages, tools, stream_sink, label, timeout=HTTP_TIMEOUT):
|
||
answer = grok_parse_output(grok_run(opencode_prompt(messages, tools), timeout))
|
||
if stream_sink is not None and answer.get("content"):
|
||
stream_sink(answer["content"])
|
||
return self.muse_reply(answer, label)
|
||
|
||
def claude_turn(self, messages, tools, stream_sink, label, timeout=HTTP_TIMEOUT):
|
||
answer = claude_parse_output(claude_run(opencode_prompt(messages, tools), timeout))
|
||
if stream_sink is not None and answer.get("content"):
|
||
stream_sink(answer["content"])
|
||
return self.muse_reply(answer, label)
|
||
|
||
def codex_turn(self, messages, tools, stream_sink, label, timeout=HTTP_TIMEOUT):
|
||
stdout, outfile_text, used_schema = codex_run(opencode_prompt(messages, tools), timeout)
|
||
answer = codex_parse_output(stdout, outfile_text, used_schema)
|
||
if stream_sink is not None and answer.get("content"):
|
||
stream_sink(answer["content"])
|
||
return self.muse_reply(answer, label)
|
||
|
||
def gemini_turn(self, messages, tools, stream_sink, label, timeout=HTTP_TIMEOUT):
|
||
answer = gemini_parse_output(gemini_run(opencode_prompt(messages, tools), timeout))
|
||
if stream_sink is not None and answer.get("content"):
|
||
stream_sink(answer["content"])
|
||
return self.muse_reply(answer, label)
|
||
|
||
def muse_reply(self, answer, label):
|
||
calls = []
|
||
for call in answer.get("tool_calls") or []:
|
||
calls.append({"id": uuid.uuid4().hex, "name": call.get("name") or "", "arguments": call.get("arguments") or "{}"})
|
||
return {"role": "assistant", "content": answer.get("content") or "", "reasoning": "", "tool_calls": calls, "backend": label}
|
||
|
||
|
||
def probe_backend(models_url, key):
|
||
if models_url == MUSE_MODELS_URL:
|
||
return muse_health()
|
||
if models_url == GROK_MODELS_URL:
|
||
return grok_health()
|
||
if models_url == CLAUDE_MODELS_URL:
|
||
return claude_health()
|
||
if models_url == CODEX_MODELS_URL:
|
||
return codex_health()
|
||
if models_url == GEMINI_MODELS_URL:
|
||
return gemini_health()
|
||
if models_url == OPENCODE_MODELS_URL:
|
||
return "ok, %s" % opencode_binary() if opencode_binary() else "opencode CLI not found"
|
||
request = urllib.request.Request(models_url, headers={"Authorization": "Bearer " + key})
|
||
started = time.time()
|
||
try:
|
||
with urllib.request.urlopen(request, timeout=6) as response:
|
||
try:
|
||
payload = json.loads(response.read().decode("utf-8", "replace"))
|
||
except ValueError:
|
||
return "bad response"
|
||
data = payload.get("data", []) if isinstance(payload, dict) else []
|
||
count = len(data) if isinstance(data, list) else 0
|
||
return "ok, %d models, %dms" % (count, (time.time() - started) * 1000)
|
||
except urllib.error.HTTPError as exc:
|
||
if exc.code == 401:
|
||
return "needs key"
|
||
return "http %d" % exc.code
|
||
except OSError as exc:
|
||
return "fail: %s" % short_error(exc)[:60]
|
||
except Exception as exc:
|
||
return "fail: %s" % short_error(exc)[:60]
|
||
|
||
|
||
def edge_token(skew=0):
|
||
ticks = int(time.time()) + skew + 11644473600
|
||
ticks -= ticks % 300
|
||
ticks *= 10000000
|
||
return hashlib.sha256(("%d%s" % (ticks, EDGE_TRUSTED_TOKEN)).encode("ascii")).hexdigest().upper()
|
||
|
||
|
||
def edge_timestamp():
|
||
return time.strftime("%a %b %d %Y %H:%M:%S GMT+0000 (Coordinated Universal Time)", time.gmtime())
|
||
|
||
|
||
def edge_ssml(text, voice):
|
||
safe = html.escape(text, quote=False).replace("\r", " ").replace("\n", " ")
|
||
return "<speak version='1.0' xmlns='http://www.w3.org/2001/10/synthesis' xml:lang='en-US'><voice name='%s'><prosody pitch='+0Hz' rate='+0%%' volume='+0%%'>%s</prosody></voice></speak>" % (voice, safe)
|
||
|
||
|
||
def edge_chunks(text, limit=4000):
|
||
sentences = re.split(r"(?<=[.!?])\s+", text.strip())
|
||
parts = []
|
||
current = ""
|
||
for sentence in sentences:
|
||
trial = (current + " " + sentence).strip()
|
||
if len(trial.encode("utf-8")) > limit and current:
|
||
parts.append(current)
|
||
current = sentence
|
||
else:
|
||
current = trial
|
||
if current:
|
||
parts.append(current)
|
||
return parts or [text]
|
||
|
||
|
||
def ws_send(sock, payload):
|
||
if isinstance(payload, str):
|
||
payload = payload.encode("utf-8")
|
||
mask = os.urandom(4)
|
||
header = bytes([0x81])
|
||
length = len(payload)
|
||
if length < 126:
|
||
header += bytes([0x80 | length])
|
||
elif length < 65536:
|
||
header += bytes([0x80 | 126]) + struct.pack(">H", length)
|
||
else:
|
||
header += bytes([0x80 | 127]) + struct.pack(">Q", length)
|
||
masked = bytes(piece ^ mask[pos % 4] for pos, piece in enumerate(payload))
|
||
sock.sendall(header + mask + masked)
|
||
|
||
|
||
def ws_recv_exact(sock, count):
|
||
data = b""
|
||
while len(data) < count:
|
||
piece = sock.recv(count - len(data))
|
||
if not piece:
|
||
raise BackendError("edge connection closed")
|
||
data += piece
|
||
return data
|
||
|
||
|
||
def ws_recv(sock):
|
||
opcode = 0
|
||
payload = b""
|
||
first = True
|
||
while True:
|
||
head = ws_recv_exact(sock, 2)
|
||
done = head[0] & 0x80
|
||
part = head[0] & 0x0F
|
||
masked = head[1] & 0x80
|
||
length = head[1] & 0x7F
|
||
if length == 126:
|
||
length = struct.unpack(">H", ws_recv_exact(sock, 2))[0]
|
||
elif length == 127:
|
||
length = struct.unpack(">Q", ws_recv_exact(sock, 8))[0]
|
||
key = ws_recv_exact(sock, 4) if masked else b""
|
||
body = ws_recv_exact(sock, length) if length else b""
|
||
if masked:
|
||
body = bytes(piece ^ key[pos % 4] for pos, piece in enumerate(body))
|
||
if part == 0x8:
|
||
raise BackendError("edge closed stream")
|
||
if part == 0x9:
|
||
sock.sendall(bytes([0x8A, 0x00]))
|
||
continue
|
||
if part in (0x1, 0x2) and first:
|
||
opcode = part
|
||
first = False
|
||
payload += body
|
||
if done:
|
||
return opcode, payload
|
||
|
||
|
||
def edge_skew(header_text):
|
||
match = re.search(r"Date:\s*(.+?)\r?$", header_text, re.MULTILINE | re.IGNORECASE)
|
||
if not match:
|
||
return 0
|
||
try:
|
||
stamp = time.mktime(time.strptime(match.group(1).strip(), "%a, %d %b %Y %H:%M:%S %Z"))
|
||
return int(stamp - time.time())
|
||
except (ValueError, OverflowError):
|
||
return 0
|
||
|
||
|
||
def edge_connect():
|
||
skew = 0
|
||
for attempt in range(2):
|
||
path = "%s?TrustedClientToken=%s&ConnectionId=%s&Sec-MS-GEC=%s&Sec-MS-GEC-Version=1-%s" % (EDGE_WS_PATH, EDGE_TRUSTED_TOKEN, uuid.uuid4().hex, edge_token(skew), EDGE_CHROMIUM)
|
||
raw = socket.create_connection((EDGE_HOST, 443), timeout=20)
|
||
sock = ssl.create_default_context().wrap_socket(raw, server_hostname=EDGE_HOST)
|
||
sock.settimeout(30)
|
||
key = base64.b64encode(os.urandom(16)).decode("ascii")
|
||
muid = "".join(random.choice("0123456789ABCDEF") for _ in range(32))
|
||
handshake = "\r\n".join([
|
||
"GET %s HTTP/1.1" % path,
|
||
"Host: %s" % EDGE_HOST,
|
||
"Upgrade: websocket",
|
||
"Connection: Upgrade",
|
||
"Sec-WebSocket-Version: 13",
|
||
"Sec-WebSocket-Key: %s" % key,
|
||
"Pragma: no-cache",
|
||
"Cache-Control: no-cache",
|
||
"Origin: chrome-extension://jdiccldimpdaibmpdkjnbmckianbfold",
|
||
"User-Agent: Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/143.0.0.0 Safari/537.36 Edg/143.0.0.0",
|
||
"Accept-Encoding: gzip, deflate, br, zstd",
|
||
"Accept-Language: en-US,en;q=0.9",
|
||
"Cookie: muid=%s;" % muid,
|
||
"", "",
|
||
])
|
||
sock.sendall(handshake.encode("ascii"))
|
||
answer = b""
|
||
while b"\r\n\r\n" not in answer:
|
||
piece = sock.recv(4096)
|
||
if not piece:
|
||
break
|
||
answer += piece
|
||
header_text = answer.decode("latin-1", "replace")
|
||
if " 101 " in header_text.split("\r\n", 1)[0]:
|
||
return sock
|
||
sock.close()
|
||
if " 403 " in header_text and attempt == 0:
|
||
skew = edge_skew(header_text)
|
||
continue
|
||
raise BackendError("edge handshake failed: %s" % header_text.split("\r\n", 1)[0][:200])
|
||
raise BackendError("edge handshake failed after retry")
|
||
|
||
|
||
def edge_synthesize(text, voice):
|
||
sock = edge_connect()
|
||
audio = bytearray()
|
||
try:
|
||
config = "X-Timestamp:%s\r\nContent-Type:application/json; charset=utf-8\r\nPath:speech.config\r\n\r\n{\"context\":{\"synthesis\":{\"audio\":{\"metadataoptions\":{\"sentenceBoundaryEnabled\":\"false\",\"wordBoundaryEnabled\":\"true\"},\"outputFormat\":\"audio-24khz-48kbitrate-mono-mp3\"}}}}" % edge_timestamp()
|
||
ws_send(sock, config)
|
||
for chunk in edge_chunks(text):
|
||
message = "X-RequestId:%s\r\nContent-Type:application/ssml+xml\r\nX-Timestamp:%sZ\r\nPath:ssml\r\n\r\n%s" % (uuid.uuid4().hex, edge_timestamp(), edge_ssml(chunk, voice))
|
||
ws_send(sock, message)
|
||
while True:
|
||
opcode, payload = ws_recv(sock)
|
||
if opcode == 0x2 and len(payload) > 2:
|
||
head_len = int.from_bytes(payload[0:2], "big")
|
||
audio += payload[2 + head_len:]
|
||
elif opcode == 0x1 and "Path:turn.end" in payload.decode("utf-8", "replace"):
|
||
break
|
||
return bytes(audio)
|
||
finally:
|
||
try:
|
||
sock.close()
|
||
except OSError:
|
||
pass
|
||
|
||
|
||
def play_audio(path):
|
||
if sys.platform == "darwin" and shutil.which("afplay"):
|
||
runner = ["afplay", path]
|
||
elif shutil.which("paplay"):
|
||
runner = ["paplay", path]
|
||
elif shutil.which("aplay"):
|
||
runner = ["aplay", "-q", path]
|
||
elif shutil.which("pw-play"):
|
||
runner = ["pw-play", path]
|
||
elif sys.platform == "win32":
|
||
runner = ["powershell", "-c", "(New-Object Media.SoundPlayer '%s').PlaySync();" % path]
|
||
else:
|
||
return False
|
||
try:
|
||
subprocess.run(runner, timeout=300, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, check=False)
|
||
return True
|
||
except (OSError, subprocess.SubprocessError):
|
||
return False
|
||
|
||
|
||
def listen_audio(seconds, workdir):
|
||
recorder = next((item for item in RECORDERS if shutil.which(item[0])), None)
|
||
transcriber = next((name for name in TRANSCRIBERS if shutil.which(name)), "")
|
||
if recorder is None or not transcriber:
|
||
missing = []
|
||
if recorder is None:
|
||
missing.append("audio recorder (install arecord, sox, or ffmpeg)")
|
||
if not transcriber:
|
||
missing.append("transcriber (install whisper-cpp or whisper)")
|
||
return "listening unavailable, missing: " + ", ".join(missing)
|
||
path = os.path.join(workdir, "listen-%d.wav" % int(time.time()))
|
||
command = [part.replace("{seconds}", str(seconds)).replace("{path}", path) for part in recorder[1]]
|
||
try:
|
||
subprocess.run(command, timeout=seconds + 15, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, check=False)
|
||
done = subprocess.run([transcriber, path], stdout=subprocess.PIPE, stderr=subprocess.PIPE, universal_newlines=True, timeout=180)
|
||
except (OSError, subprocess.SubprocessError) as exc:
|
||
return "listening failed: " + short_error(exc)
|
||
finally:
|
||
try:
|
||
os.remove(path)
|
||
except OSError:
|
||
pass
|
||
if done.returncode != 0:
|
||
return "transcriber error: " + (done.stderr or "")[:500]
|
||
text = (done.stdout or "").strip()
|
||
return truncate(text, 3000, 500) if text else "heard nothing"
|
||
|
||
|
||
def sysinfo_check_os():
|
||
try:
|
||
uname = os.uname()
|
||
base = "%s %s %s" % (uname.sysname, uname.release, uname.machine)
|
||
except (AttributeError, OSError):
|
||
base = sys.platform
|
||
try:
|
||
with open("/etc/os-release", "r", encoding="utf-8") as handle:
|
||
fields = dict(line.strip().split("=", 1) for line in handle if "=" in line)
|
||
pretty = fields.get("PRETTY_NAME", "").strip('"')
|
||
if pretty:
|
||
return "%s (%s)" % (pretty, base)
|
||
except OSError:
|
||
pass
|
||
return base
|
||
|
||
|
||
def sysinfo_check_python():
|
||
return "%s (%s)" % (sys.version.split()[0], sys.executable)
|
||
|
||
|
||
def sysinfo_check_venv():
|
||
env_path = os.environ.get("VIRTUAL_ENV") or ""
|
||
if env_path:
|
||
return "yes (%s)" % env_path
|
||
if sys.prefix != sys.base_prefix:
|
||
return "yes (%s)" % sys.prefix
|
||
return "no (system python at %s)" % sys.prefix
|
||
|
||
|
||
def sysinfo_check_root():
|
||
geteuid = getattr(os, "geteuid", None)
|
||
if geteuid is None:
|
||
return "unknown (no geteuid on %s)" % sys.platform
|
||
uid = geteuid()
|
||
return "yes (uid 0)" if uid == 0 else "no (uid %d)" % uid
|
||
|
||
|
||
def sysinfo_check_container():
|
||
engine = container_engine()
|
||
if not engine:
|
||
return "no engine (install podman or docker)"
|
||
return "%s, box %s" % (engine, box_state(engine))
|
||
|
||
|
||
def sysinfo_check_binaries():
|
||
found = []
|
||
lost = []
|
||
for name in SYSINFO_BINARIES:
|
||
(found if shutil.which(name) else lost).append(name)
|
||
return "present: %s; missing: %s" % (", ".join(found) or "none", ", ".join(lost) or "none")
|
||
|
||
|
||
def sysinfo_check_cpu():
|
||
count = os.cpu_count() or 0
|
||
try:
|
||
one, five, fifteen = os.getloadavg()
|
||
return "%d cores, load %.2f/%.2f/%.2f" % (count, one, five, fifteen)
|
||
except OSError:
|
||
return "%d cores" % count
|
||
|
||
|
||
def sysinfo_check_disk():
|
||
usage = shutil.disk_usage(os.getcwd())
|
||
return "%.1f GiB free of %.1f GiB at %s" % (usage.free / 2 ** 30, usage.total / 2 ** 30, os.getcwd())
|
||
|
||
|
||
SYSINFO_CHECKS = (
|
||
("os", sysinfo_check_os),
|
||
("python", sysinfo_check_python),
|
||
("venv", sysinfo_check_venv),
|
||
("root", sysinfo_check_root),
|
||
("container", sysinfo_check_container),
|
||
("binaries", sysinfo_check_binaries),
|
||
("cpu", sysinfo_check_cpu),
|
||
("disk", sysinfo_check_disk),
|
||
)
|
||
|
||
|
||
def collect_sysinfo(names=None):
|
||
wanted = [pair for pair in SYSINFO_CHECKS if names is None or pair[0] in names]
|
||
results = {}
|
||
|
||
def target(name, func):
|
||
started = time.time()
|
||
try:
|
||
value = func()
|
||
except Exception as exc:
|
||
value = "error: " + short_error(exc)
|
||
results[name] = (value, int((time.time() - started) * 1000))
|
||
|
||
workers = [threading.Thread(target=target, args=(name, func), daemon=True) for name, func in wanted]
|
||
for worker in workers:
|
||
worker.start()
|
||
deadline = time.time() + SYSINFO_TIMEOUT
|
||
for worker in workers:
|
||
worker.join(timeout=max(0, deadline - time.time()))
|
||
lines = []
|
||
for name, _func in wanted:
|
||
if name in results:
|
||
value, elapsed = results[name]
|
||
lines.append("%s: %s (%dms)" % (name, value, elapsed))
|
||
else:
|
||
lines.append("%s: timed out after %ds" % (name, SYSINFO_TIMEOUT))
|
||
return "\n".join(lines)
|
||
|
||
|
||
def tool_schema(name, description, properties, required):
|
||
return {"type": "function", "function": {"name": name, "description": description, "parameters": {"type": "object", "properties": properties, "required": required}}}
|
||
|
||
|
||
TOOL_SCHEMAS = [
|
||
tool_schema("shell", "Run a shell command. Output is truncated. On home, destructive commands ask the user first. On sandbox, commands run isolated in the container and need no approval. Pass secrets to expose named vault secrets as TAI_SECRET_<NAME> variables for that command only; using secrets asks the user first. For local commands only: use web_fetch for HTTP(S), never curl or wget.", {"command": {"type": "string"}, "workdir": {"type": "string"}, "secrets": {"type": "array"}}, ["command"]),
|
||
tool_schema("read_file", "Read a text file. Large files are truncated.", {"path": {"type": "string"}}, ["path"]),
|
||
tool_schema("write_file", "Write content to a file, creating parent directories. Overwrites existing files.", {"path": {"type": "string"}, "content": {"type": "string"}}, ["path", "content"]),
|
||
tool_schema("edit_file", "Replace one exact text match in a file. Fails unless the match is unique.", {"path": {"type": "string"}, "find": {"type": "string"}, "replace": {"type": "string"}}, ["path", "find", "replace"]),
|
||
tool_schema("web_search", "Search the web. Set images for image search, content to include fetched page text.", {"query": {"type": "string"}, "images": {"type": "boolean"}, "content": {"type": "boolean"}}, ["query"]),
|
||
tool_schema("web_fetch", "HTTP client: fetch a URL with any method and return status, content type, size, and content. GET by default; pass method (POST, PUT, PATCH, DELETE, HEAD, OPTIONS), headers, and a string body for APIs. HTML becomes readable text, JSON stays raw, raw forces the raw body. Use this for all HTTP(S), never curl or wget via shell. Pass auth_secret to send a vault secret as an Authorization header; the value stays sealed and secret use asks the user first.", {"url": {"type": "string"}, "method": {"type": "string"}, "headers": {"type": "object"}, "body": {"type": "string"}, "timeout": {"type": "integer"}, "raw": {"type": "boolean"}, "auth_secret": {"type": "string"}, "auth_header": {"type": "string"}, "auth_prefix": {"type": "string"}}, ["url"]),
|
||
tool_schema("speak", "Synthesize text to speech with a free voice, save the MP3, and play it when a player exists.", {"text": {"type": "string"}, "voice": {"type": "string"}}, ["text"]),
|
||
tool_schema("listen", "Record from the microphone for some seconds and transcribe it. Needs OS audio tools.", {"seconds": {"type": "integer"}}, []),
|
||
tool_schema("remember", "Update your own system message with new knowledge or behavior. The instruction is merged into the existing system message, unrelated parts stay intact. Use it to forget as well by instructing what to remove. Execute by default when behavior, preferences, or durable facts need updating. Never pass passwords, tokens, or secret values here; store those with store_secret instead.", {"instruction": {"type": "string"}}, ["instruction"]),
|
||
tool_schema("recall", "Search past session memory of the current profile by keyword, optionally filtered by tags.", {"query": {"type": "string"}, "tags": {"type": "array"}}, ["query"]),
|
||
tool_schema("load_skill", "Load a skill by name. Returns full instructions plus bundled file paths. Missing skills with a blueprint are researched and built on demand.", {"name": {"type": "string"}}, ["name"]),
|
||
tool_schema("get_current_terminal_content", "Capture visible text of the current tmux pane including scrollback. Works inside tmux or against a running tmux server.", {"lines": {"type": "integer"}}, []),
|
||
tool_schema("export_terminal_log", "Only available inside a live tmux session. Exports the full pane scrollback (tmux capture-pane -pS -, the entire history, not just what's visible) to a log file under tai's home dir, then chunks it and indexes the chunks for retrieval with SQLite FTS5 + bm25 (lexical only, no embeddings). From then on, use the search tool (kind chunk) to answer questions about anything that happened in this terminal; it returns the actual chunk text, not just a snippet.", {"tags": {"type": "array"}}, []),
|
||
tool_schema("fork", "Spawn a background subagent with its own context that works while you continue. Returns an agent id immediately. Collect its summarized result with poll. Subagents get a smaller step budget and a time limit.", {"task": {"type": "string"}, "timeout": {"type": "integer"}, "profile": {"type": "string"}}, ["task"]),
|
||
tool_schema("poll", "Collect a background subagent result by id. Waits up to wait seconds, then reports running or the result.", {"id": {"type": "integer"}, "wait": {"type": "integer"}}, ["id"]),
|
||
tool_schema("sysinfo", "Inspect the host machine where tai runs. Runs os, python, venv, root, container, binaries, cpu, and disk checks in parallel and reports each with timing. Pass checks to run a named subset.", {"checks": {"type": "array"}}, []),
|
||
tool_schema("create_skill", "Create a new agent skill by name from a brief. Runs a dedicated deep-research worker on this machine (it calls sysinfo first, then researches with web search and shell until it has verified information) and writes SKILL.md plus supporting files into the skill directory. Runs synchronously and can take many minutes. Scope project writes under ./.tai/skills (the default), scope home under ~/.tai/skills.", {"name": {"type": "string"}, "brief": {"type": "string"}, "scope": {"type": "string"}, "timeout": {"type": "integer"}}, ["name", "brief"]),
|
||
tool_schema("install_skill", "Install a skill that already exists somewhere else, instead of building one from scratch. source can be a local directory containing SKILL.md (or one subdirectory that does), a bare SKILL.md file, a local .zip, or a git URL (needs git installed; most sources need nothing beyond the standard library). Validates the SKILL.md frontmatter, then lands it under tai's own skill directories: copy vendors a private snapshot, symlink points at the source in place (the default when the source already sits in a recognized skills folder like .claude/skills or .agents/skills, so there is one canonical copy instead of a drifting duplicate). Scope defaults like create_skill: project when the cwd looks like a project, else home. Overwriting an existing skill of the same name asks the user first unless force is set.", {"source": {"type": "string"}, "name": {"type": "string"}, "scope": {"type": "string"}, "mode": {"type": "string"}, "force": {"type": "boolean"}}, ["source"]),
|
||
tool_schema("store_secret", "Store a password, token, or secret in the sealed vault under a name. Values are encrypted at rest, never shown back, and only usable by reference; prefer /secret set in the REPL so the value never enters the conversation. Optional username, host, port, notes, expires (ISO datetime), and tags describe it; name and value suffice, never pester the user for more.", {"name": {"type": "string"}, "value": {"type": "string"}, "username": {"type": "string"}, "host": {"type": "string"}, "port": {"type": "integer"}, "notes": {"type": "string"}, "expires": {"type": "string"}, "tags": {"type": "array"}}, ["name", "value"]),
|
||
tool_schema("list_secrets", "List vault secret names with metadata and tags. Values are never revealed.", {}, []),
|
||
tool_schema("delete_secret", "Delete a vault secret by name. Asks the user first.", {"name": {"type": "string"}}, ["name"]),
|
||
tool_schema("schedule", "Schedule a prompt to run at a future time or on a repeating interval. Exactly one of at (ISO datetime, naive means local time) or every (interval seconds, minimum 60) is required. When due, a background subagent runs the prompt without blocking anyone; inspect outcomes with schedules. Asks the user first. Optional tags label it.", {"at": {"type": "string"}, "every": {"type": "integer"}, "prompt": {"type": "string"}, "name": {"type": "string"}, "timeout": {"type": "integer"}, "profile": {"type": "string"}, "tags": {"type": "array"}}, ["prompt"]),
|
||
tool_schema("unschedule", "Delete a scheduled prompt by id. Asks the user first.", {"id": {"type": "integer"}}, ["id"]),
|
||
tool_schema("schedules", "List scheduled prompts with due times and last outcomes.", {}, []),
|
||
tool_schema("record_save", "Save text as a tagged record in the vault and get back a mem:<hash> id. Use for anything too big or too durable for chat: research notes, command output, transcripts. Kind is one of note, output, file, research, transcript. Tag rules: lowercase singular, reuse tags from the tags tool, at most 5; known words in the text attach automatically and the record links to recent same-tag records.", {"title": {"type": "string"}, "content": {"type": "string"}, "kind": {"type": "string"}, "tags": {"type": "array"}}, ["content"]),
|
||
tool_schema("record_read", "Read a slice of a record by mem:<hash> id. Offset and limit page through huge records without loading them whole.", {"id": {"type": "string"}, "offset": {"type": "integer"}, "limit": {"type": "integer"}}, ["id"]),
|
||
tool_schema("record_search", "Search records by text, kind, and tags. Returns mem:<hash> ids with sizes so you can page in only what you need.", {"query": {"type": "string"}, "kind": {"type": "string"}, "tags": {"type": "array"}, "limit": {"type": "integer"}}, []),
|
||
tool_schema("record_delete", "Delete a record by mem:<hash> id. Asks the user first.", {"id": {"type": "string"}}, ["id"]),
|
||
tool_schema("graph_link", "Link two vault nodes with a relation: mem:<hash> records, secret:<name> secrets, sched:<id> schedules. Both endpoints must exist.", {"src": {"type": "string"}, "dst": {"type": "string"}, "relation": {"type": "string"}}, ["src", "dst"]),
|
||
tool_schema("graph_query", "Show the neighborhood of a vault node (mem:<hash>, secret:<name>, sched:<id>) by breadth-first traversal. Depth caps at 4.", {"node": {"type": "string"}, "depth": {"type": "integer"}, "limit": {"type": "integer"}}, ["node"]),
|
||
tool_schema("delete_file", "Delete a file. Requires reading it first; the full pre-image stays in the audit trail so restore can undelete. Asks the user first.", {"path": {"type": "string"}}, ["path"]),
|
||
tool_schema("audit", "Show the file audit trail for time travel: every write, edit, delete, shell snapshot, release, and restore with sizes and messages. Filter by exact path or tag, or rank by a full-text query over paths and messages.", {"path": {"type": "string"}, "tag": {"type": "string"}, "query": {"type": "string"}, "limit": {"type": "integer"}}, []),
|
||
tool_schema("restore", "Restore a file to an audit row: the post-image for writes and edits, the pre-image for deletes and shell snapshots. Refuses truncated images. Asks the user first.", {"id": {"type": "integer"}}, ["id"]),
|
||
tool_schema("release", "Cut a release of the running script: bump VERSION by part (patch for fixes, minor for features, major for breaking changes), keep a versioned backup, log the message in the audit trail. Asks the user first.", {"part": {"type": "string"}, "message": {"type": "string"}}, ["part", "message"]),
|
||
tool_schema("tags", "List vault tags with usage counts, most used first. Consult before tagging so new items reuse established tags instead of inventing synonyms.", {"prefix": {"type": "string"}, "limit": {"type": "integer"}}, []),
|
||
tool_schema("install", "Manage tai installations: binary, bash command-not-found hook, venv, scheduler service, telegram service, container. Status reports what exists. Install adds missing pieces, upgrade refreshes in place and restarts services, reinstall rebuilds artifacts, uninstall removes them. Service changes take effect immediately. The vault is never touched by any action. Non-status actions ask the user first.", {"action": {"type": "string"}, "targets": {"type": "array"}}, ["action"]),
|
||
tool_schema("create_bot", "Create a bot in the current profile: a name plus its own system message and history. Give at least one of description, rules, or behavior; optional nicknames register as @mention aliases (a short form of the name registers automatically). Names are lowercase.", {"name": {"type": "string"}, "description": {"type": "string"}, "rules": {"type": "string"}, "behavior": {"type": "string"}, "nicknames": {"type": "array"}}, ["name"]),
|
||
tool_schema("search", "Search everything at once with ranked full-text matching (SQLite FTS5 + bm25, no vectors): records, episodic events, the audit trail, and RAG chunks. Chunks are the retrieval-grade granular hits (e.g. one slice of an exported terminal log) and come back with their actual text, not just a snippet; use the chunk's parent id with record_read to see more surrounding context. One call replaces paging through each store separately. expand adds one graph hop per record hit.", {"query": {"type": "string"}, "kinds": {"type": "array"}, "limit": {"type": "integer"}, "expand": {"type": "boolean"}}, ["query"]),
|
||
]
|
||
|
||
|
||
CORE_TOOLS = ("shell", "read_file", "write_file", "remember", "recall", "load_skill", "fork", "poll")
|
||
LAZY_NAME_STOP = ("get", "set", "list", "add")
|
||
|
||
# Tools gated on top of the lazy keyword match: a gate must pass before a tool can even be
|
||
# offered to the model, regardless of relevance. Used for tools that only make sense in a
|
||
# specific environment (e.g. a tmux pane to capture).
|
||
TOOL_ENV_GATES = {
|
||
"export_terminal_log": lambda: bool(os.environ.get("TMUX")),
|
||
}
|
||
|
||
|
||
def tool_env_ok(name):
|
||
gate = TOOL_ENV_GATES.get(name)
|
||
return gate is None or gate()
|
||
|
||
TOOL_TAGS = {
|
||
"shell": ("run", "execute", "command", "bash", "terminal", "script"),
|
||
"read_file": ("read", "open", "view", "contents"),
|
||
"write_file": ("write", "create", "overwrite"),
|
||
"edit_file": ("edit", "change", "modify", "replace", "patch", "fix", "alter"),
|
||
"web_search": ("search", "web", "internet", "lookup", "google"),
|
||
"web_fetch": ("fetch", "download", "url", "link", "page", "curl", "wget"),
|
||
"speak": ("speak", "say", "voice", "talk", "tts"),
|
||
"listen": ("listen", "hear", "microphone", "mic", "transcribe", "stt", "dictate"),
|
||
"remember": ("remember", "memorize", "preference"),
|
||
"recall": ("recall", "memory", "history", "past"),
|
||
"load_skill": ("skill", "capability"),
|
||
"get_current_terminal_content": ("terminal", "tmux", "pane", "scrollback", "screen"),
|
||
"export_terminal_log": ("terminal", "tmux", "export", "log", "capture", "history", "rag", "index"),
|
||
"fork": ("fork", "subagent", "background", "parallel", "delegate", "spawn"),
|
||
"poll": ("poll", "collect"),
|
||
"sysinfo": ("sysinfo", "system", "host", "machine", "hardware", "specs", "installed", "python", "container"),
|
||
"create_skill": ("skill", "create", "author", "blueprint"),
|
||
"install_skill": ("skill", "install", "import", "adopt", "copy", "symlink"),
|
||
"store_secret": ("secret", "password", "token", "credential", "passwd", "vault", "apikey"),
|
||
"list_secrets": ("secrets", "vault", "credentials"),
|
||
"delete_secret": ("secret", "remove", "revoke"),
|
||
"schedule": ("schedule", "cron", "appointment", "reminder", "recurring", "periodic", "later", "daily", "hourly", "interval"),
|
||
"unschedule": ("unschedule", "cancel"),
|
||
"schedules": ("schedules", "upcoming", "planned", "agenda"),
|
||
"record_save": ("record", "note", "document", "memo"),
|
||
"record_read": ("record", "page"),
|
||
"record_search": ("record", "find"),
|
||
"record_delete": ("record", "remove"),
|
||
"graph_link": ("graph", "link", "connect", "relate", "edge", "relation"),
|
||
"graph_query": ("graph", "neighbors", "neighborhood", "traverse", "related"),
|
||
"delete_file": ("delete", "remove", "erase"),
|
||
"audit": ("audit", "trail", "changes", "blame", "history", "version"),
|
||
"restore": ("restore", "undo", "revert", "recover", "rollback", "undelete", "previous"),
|
||
"release": ("release", "version", "bump", "changelog", "publish"),
|
||
"tags": ("tags", "tag", "label", "categorize", "taxonomy", "vocabulary"),
|
||
"install": ("install", "uninstall", "reinstall", "upgrade", "setup", "service", "systemd", "venv", "hook", "deploy"),
|
||
"create_bot": ("bot", "bots", "persona", "character", "mention"),
|
||
"search": ("search", "find", "lookup", "everything"),
|
||
}
|
||
|
||
|
||
def tool_catalog():
|
||
lines = ["", "", "## Tool catalog (core loads always; name a lazy tool or topic to load it)"]
|
||
for schema in TOOL_SCHEMAS:
|
||
name = schema["function"]["name"]
|
||
if not tool_env_ok(name):
|
||
continue
|
||
if name in CORE_TOOLS:
|
||
continue
|
||
first = schema["function"]["description"].split(".")[0][:80]
|
||
lines.append("- %s: %s" % (name, first))
|
||
return "\n".join(lines)
|
||
|
||
|
||
def conversation_text(messages):
|
||
parts = []
|
||
for item in messages:
|
||
content = item.get("content")
|
||
if item.get("role") in ("user", "assistant", "tool") and isinstance(content, str) and content:
|
||
parts.append(content)
|
||
return "\n".join(parts)[-6000:]
|
||
|
||
|
||
def select_tools(text):
|
||
words = set(re.findall(r"[a-z0-9]+", (text or "").lower()))
|
||
picked = []
|
||
for schema in TOOL_SCHEMAS:
|
||
name = schema["function"]["name"]
|
||
if not tool_env_ok(name):
|
||
continue
|
||
if name in CORE_TOOLS:
|
||
picked.append(schema)
|
||
continue
|
||
tokens = [item for item in name.split("_") if len(item) >= 4 and item not in LAZY_NAME_STOP]
|
||
if tokens and set(tokens) & words:
|
||
picked.append(schema)
|
||
continue
|
||
if set(TOOL_TAGS.get(name, ())) & words:
|
||
picked.append(schema)
|
||
continue
|
||
desc = set(word for word in re.findall(r"[a-z0-9]+", schema["function"]["description"].lower()) if len(word) >= 6)
|
||
if desc & words:
|
||
picked.append(schema)
|
||
return picked
|
||
|
||
|
||
def secret_env_name(name):
|
||
return "TAI_SECRET_" + re.sub(r"[^A-Z0-9_]", "_", name.upper())
|
||
|
||
|
||
def normalize_tag(raw):
|
||
return re.sub(r"[^a-z0-9]+", "-", str(raw or "").strip().lower()).strip("-")[:32]
|
||
|
||
|
||
def normalize_tags(items):
|
||
seen = []
|
||
for item in items:
|
||
cleaned = normalize_tag(item)
|
||
if cleaned and cleaned not in seen:
|
||
seen.append(cleaned)
|
||
return seen[:20]
|
||
|
||
|
||
TAG_AUTO_SKIP = ("record", "file")
|
||
TAG_AUTO_MAX = 5
|
||
TAG_LINK_MAX = 3
|
||
SINGULAR_KEEP = ("news", "means", "series", "species", "physics")
|
||
|
||
|
||
def singular_noun(word):
|
||
if word in SINGULAR_KEEP or len(word) <= 3:
|
||
return word
|
||
if word.endswith("ies"):
|
||
return word[:-3] + "y"
|
||
if word.endswith(("ses", "xes", "zes", "ches", "shes")):
|
||
return word[:-2]
|
||
if word.endswith("s") and not word.endswith(("ss", "us", "is", "os")):
|
||
return word[:-1]
|
||
return word
|
||
|
||
|
||
def tag_variants(tag):
|
||
found = {tag, singular_noun(tag)}
|
||
if not tag.endswith("s"):
|
||
found.add(tag + "s")
|
||
return sorted(found)
|
||
|
||
|
||
def parse_node_id(raw):
|
||
match = re.fullmatch(r"(mem|secret|sched):(.+)", str(raw or "").strip())
|
||
if not match:
|
||
return None
|
||
kind, key = match.group(1), match.group(2)
|
||
if kind == "mem" and re.fullmatch(r"[0-9a-f]{16}", key):
|
||
return kind, key
|
||
if kind == "secret" and SECRET_RE.match(key):
|
||
return kind, key
|
||
if kind == "sched" and key.isdigit():
|
||
return kind, key
|
||
return None
|
||
|
||
|
||
def create_skill_prompt(name, brief, skill_dir):
|
||
return (
|
||
"Create a new agent skill named '%s' in %s.\n"
|
||
"Brief: %s\n\n"
|
||
"Follow this process exactly:\n"
|
||
"1. Call sysinfo first so the skill matches this machine (os, installed tools, root, container, venv).\n"
|
||
"2. Research the subject deeply with web_search, web_fetch, and shell. Keep researching until you literally "
|
||
"have enough verified information to write complete instructions. Verify factual claims against at least two "
|
||
"independent sources. Never invent commands, paths, versions, or API details.\n"
|
||
"3. Write SKILL.md in the skill directory with this exact frontmatter:\n"
|
||
"---\nname: %s\ndescription: <one sentence saying what the skill does and when to use it>\n---\n"
|
||
"followed by clear markdown instructions grounded in your research.\n"
|
||
"4. Add scripts/, references/, or assets/ files only when they carry real weight; keep every file focused.\n"
|
||
"5. Read back every file you wrote, fix mistakes, and confirm the frontmatter holds a name plus a description.\n"
|
||
"6. Reply with a short summary: what the skill covers, which sources verified it, and which files you wrote.\n"
|
||
"Do not stop early: an incomplete or unverified skill is a failure. Do not ask questions; decide and record assumptions in the skill body."
|
||
% (name, skill_dir, brief, name)
|
||
)
|
||
|
||
|
||
class Tools:
|
||
def __init__(self, app):
|
||
self.app = app
|
||
self.handlers = {
|
||
"shell": self.run_shell,
|
||
"read_file": self.run_read_file,
|
||
"write_file": self.run_write_file,
|
||
"edit_file": self.run_edit_file,
|
||
"web_search": self.run_web_search,
|
||
"web_fetch": self.run_web_fetch,
|
||
"speak": self.run_speak,
|
||
"listen": self.run_listen,
|
||
"remember": self.run_remember,
|
||
"recall": self.run_recall,
|
||
"load_skill": self.run_load_skill,
|
||
"get_current_terminal_content": self.run_terminal_content,
|
||
"export_terminal_log": self.run_export_terminal_log,
|
||
"fork": self.run_fork,
|
||
"poll": self.run_poll,
|
||
"sysinfo": self.run_sysinfo,
|
||
"create_skill": self.run_create_skill,
|
||
"install_skill": self.run_install_skill,
|
||
"store_secret": self.run_store_secret,
|
||
"list_secrets": self.run_list_secrets,
|
||
"delete_secret": self.run_delete_secret,
|
||
"schedule": self.run_schedule,
|
||
"unschedule": self.run_unschedule,
|
||
"schedules": self.run_schedules,
|
||
"record_save": self.run_record_save,
|
||
"record_read": self.run_record_read,
|
||
"record_search": self.run_record_search,
|
||
"record_delete": self.run_record_delete,
|
||
"graph_link": self.run_graph_link,
|
||
"graph_query": self.run_graph_query,
|
||
"delete_file": self.run_delete_file,
|
||
"audit": self.run_audit,
|
||
"restore": self.run_restore,
|
||
"release": self.run_release,
|
||
"tags": self.run_tags,
|
||
"install": self.run_install,
|
||
"create_bot": self.run_create_bot,
|
||
"search": self.run_search,
|
||
}
|
||
self.secret_grants = set()
|
||
self.read_files = set()
|
||
self.live = False
|
||
self.diff_preview = None
|
||
|
||
def note_diff(self, path, old, new):
|
||
self.diff_preview = {"path": path, "old": old, "new": new}
|
||
|
||
def box_engine(self):
|
||
try:
|
||
return ensure_box(), ""
|
||
except BackendError as exc:
|
||
return "", "error: " + short_error(exc)
|
||
|
||
def dispatch(self, name, raw_arguments):
|
||
try:
|
||
args = json.loads(raw_arguments or "{}")
|
||
except ValueError:
|
||
return "error: arguments are not valid JSON"
|
||
handler = self.handlers.get(name)
|
||
if handler is None:
|
||
return "error: unknown tool " + name
|
||
try:
|
||
result = handler(args)
|
||
except Denied:
|
||
raise
|
||
except Exception as exc:
|
||
return "error: " + short_error(exc)
|
||
if self.app.store is not None:
|
||
return self.app.store.redact(result, self.active_profile)
|
||
return result
|
||
|
||
def secret_grant(self, kind, key, names):
|
||
token = (kind, key, tuple(sorted(names)))
|
||
if token in self.secret_grants:
|
||
return True
|
||
if not self.app.ask_approval("share secrets (%s) with %s" % (", ".join(sorted(names)), kind)):
|
||
return False
|
||
self.secret_grants.add(token)
|
||
return True
|
||
|
||
def resolve_secrets(self, names):
|
||
vault = {}
|
||
for name in names:
|
||
value = self.app.store.load_secret(name, self.active_profile)
|
||
if value is None:
|
||
return None, "error: unknown secret '%s'" % name
|
||
vault[name] = value
|
||
return vault, ""
|
||
|
||
def shell_is_safe(self, command):
|
||
parts = command.split()
|
||
if not parts:
|
||
return False
|
||
if parts[0] == "git":
|
||
return len(parts) > 1 and parts[1] in ("status", "diff", "log", "show", "branch", "remote")
|
||
return parts[0] in SAFE_COMMANDS
|
||
|
||
def snapshot_shell_targets(self, command, workdir):
|
||
if self.app.store is None:
|
||
return 0
|
||
profile = self.active_profile
|
||
count = 0
|
||
for found in shell_target_paths(command, workdir):
|
||
if not os.path.isfile(found):
|
||
continue
|
||
text, override = read_capped(found)
|
||
if text is None:
|
||
continue
|
||
try:
|
||
self.app.store.audit_event(self.audit_actor(), "shell-snapshot", found, command[:200], old=text, tags=["shell"], old_size=override, profile=profile)
|
||
self.app.store.upsert_file_record(found, text, override, profile)
|
||
count += 1
|
||
except (sqlite3.Error, OSError):
|
||
continue
|
||
return count
|
||
|
||
def run_shell(self, args):
|
||
command = str(args.get("command") or "").strip()
|
||
if not command:
|
||
return "error: empty command"
|
||
secret_names = args.get("secrets") or []
|
||
if not isinstance(secret_names, list) or any(not isinstance(item, str) for item in secret_names):
|
||
return "error: secrets must be a list of names"
|
||
vault, failure = self.resolve_secrets(secret_names) if secret_names else ({}, "")
|
||
if failure:
|
||
return failure
|
||
if vault and not self.secret_grant("shell", "", sorted(vault)):
|
||
return "denied by user"
|
||
if self.app.env == "sandbox":
|
||
return self.run_shell_box(command, str(args.get("workdir") or "/"), vault)
|
||
workdir = str(args.get("workdir") or os.getcwd())
|
||
if not self.app.config.auto_approve and not self.shell_is_safe(command) and not self.app.ask_approval(command):
|
||
return "denied by user"
|
||
self.snapshot_shell_targets(command, workdir)
|
||
env = dict(os.environ)
|
||
for name, value in vault.items():
|
||
env[secret_env_name(name)] = value
|
||
stream = LiveStream() if self.live and sys.stdout.isatty() else None
|
||
code, output, elapsed, timed_out = run_live(command, shell=True, cwd=workdir, env=env, stream=stream)
|
||
if stream is not None:
|
||
stream.close(stream_status(code, output, elapsed, timed_out))
|
||
if timed_out:
|
||
return "error: timed out after %d seconds" % SHELL_TIMEOUT
|
||
return "exit %d\n%s%s" % (code, spill_or_truncate(self.app.store, "output", "shell: " + command[:80], output.strip() or "(no output)", ["shell"], self.active_profile), curl_hint(command))
|
||
|
||
def run_shell_box(self, command, workdir, vault):
|
||
engine, failure = self.box_engine()
|
||
if not engine:
|
||
return failure
|
||
extra = ["--workdir", workdir]
|
||
for name, value in vault.items():
|
||
extra += ["-e", "%s=%s" % (secret_env_name(name), value)]
|
||
try:
|
||
stream = LiveStream() if self.live and sys.stdout.isatty() else None
|
||
code, output, elapsed, timed_out = run_live([engine, "exec", *extra, "-i", BOX_NAME, "sh", "-c", command], stream=stream)
|
||
except (OSError, subprocess.SubprocessError) as exc:
|
||
return "error: " + short_error(exc)
|
||
if stream is not None:
|
||
stream.close(stream_status(code, output, elapsed, timed_out))
|
||
if timed_out:
|
||
return "error: timed out after %d seconds" % SHELL_TIMEOUT
|
||
return "exit %d\n%s%s" % (code, spill_or_truncate(self.app.store, "output", "shell: " + command[:80], output.strip() or "(no output)", ["shell"], self.active_profile), curl_hint(command))
|
||
|
||
def run_read_file(self, args):
|
||
path = str(args.get("path") or "")
|
||
if self.app.env == "sandbox":
|
||
result = self.box_read(path)
|
||
if not result.startswith("error:"):
|
||
self.read_files.add(("box", path))
|
||
return result
|
||
if not path or not os.path.isfile(path):
|
||
return "error: no such file"
|
||
if os.path.getsize(path) > 200000:
|
||
return "error: file too large"
|
||
try:
|
||
with open(path, "r", encoding="utf-8", errors="replace") as handle:
|
||
result = truncate(handle.read(), 8000, 2000)
|
||
except OSError as exc:
|
||
return "error: " + short_error(exc)
|
||
self.read_files.add(("home", os.path.abspath(path)))
|
||
return result
|
||
|
||
def audit_actor(self):
|
||
if getattr(self.app, "persist", True) and getattr(self.app, "depth", 0) == 0:
|
||
return "main"
|
||
return "worker"
|
||
|
||
@property
|
||
def active_profile(self):
|
||
profile = getattr(self.app, "profile", None)
|
||
if profile:
|
||
return profile
|
||
store = getattr(self.app, "store", None)
|
||
return getattr(store, "profile", None) or "default"
|
||
|
||
def audit_file(self, action, path, old, new, message="", tags=(), old_size=None, new_size=None):
|
||
if self.app.store is None:
|
||
return None
|
||
profile = self.active_profile
|
||
try:
|
||
row_id = self.app.store.audit_event(self.audit_actor(), action, path, message, old, new, ["file", action] + list(tags), old_size, new_size, profile)
|
||
except (sqlite3.Error, OSError):
|
||
return False
|
||
try:
|
||
if action == "delete":
|
||
self.app.store.delete_file_record(path, profile)
|
||
elif new is not None:
|
||
self.app.store.upsert_file_record(path, new, new_size, profile)
|
||
elif old is not None:
|
||
self.app.store.upsert_file_record(path, old, old_size, profile)
|
||
except (sqlite3.Error, OSError):
|
||
pass
|
||
return row_id
|
||
|
||
def box_read_raw(self, path, limit=200000):
|
||
engine, failure = self.box_engine()
|
||
if not engine:
|
||
return None, failure
|
||
try:
|
||
done = box_exec(engine, ["sh", "-c", "wc -c < \"$1\" && head -c \"$2\" \"$1\"", "box", path, str(limit + 1)], timeout=60)
|
||
except (OSError, subprocess.SubprocessError) as exc:
|
||
return None, "error: " + short_error(exc)
|
||
if done.returncode != 0:
|
||
return None, "error: no such file"
|
||
head, _, rest = done.stdout.decode("utf-8", "replace").partition("\n")
|
||
try:
|
||
size = int(head.strip())
|
||
except ValueError:
|
||
return None, "error: unreadable file"
|
||
return size, rest
|
||
|
||
def box_read(self, path):
|
||
if not path:
|
||
return "error: empty path"
|
||
engine, failure = self.box_engine()
|
||
if not engine:
|
||
return failure
|
||
try:
|
||
done = box_exec(engine, ["head", "-c", "200000", path], timeout=60)
|
||
except (OSError, subprocess.SubprocessError) as exc:
|
||
return "error: " + short_error(exc)
|
||
if done.returncode != 0:
|
||
return "error: no such file"
|
||
return truncate(done.stdout.decode("utf-8", "replace"), 8000, 2000)
|
||
|
||
def box_write(self, path, content, old=_UNSET, action="write", message="", preview=True):
|
||
engine, failure = self.box_engine()
|
||
if not engine:
|
||
return failure
|
||
if ("box", path) not in self.read_files:
|
||
try:
|
||
probe = box_exec(engine, ["test", "-f", path], timeout=30)
|
||
except (OSError, subprocess.SubprocessError):
|
||
probe = None
|
||
if probe is not None and probe.returncode == 0:
|
||
return "error: write denied for '%s', read it first with read_file" % path
|
||
if old is _UNSET:
|
||
old, old_size = self.box_old_image(path)
|
||
else:
|
||
old_size = None
|
||
parent = os.path.dirname(path) or "."
|
||
try:
|
||
done = box_exec(engine, ["sh", "-c", "mkdir -p \"$1\" && cat > \"$2\"", "box", parent, path], input_data=content.encode("utf-8"), timeout=60)
|
||
except (OSError, subprocess.SubprocessError) as exc:
|
||
return "error: " + short_error(exc)
|
||
if done.returncode != 0:
|
||
return "error: " + done.stderr.decode("utf-8", "replace")[:200]
|
||
self.read_files.add(("box", path))
|
||
audited = self.audit_file(action, "sandbox:" + path, old, content, message, (), old_size, None)
|
||
note = "" if audited is not False else " (audit failed)"
|
||
if preview and old_size is None:
|
||
self.note_diff(path, old, content)
|
||
return "wrote %s%s" % (path, note)
|
||
|
||
def box_old_image(self, path):
|
||
size, raw = self.box_read_raw(path, AUDIT_MAX_CHARS)
|
||
if size is None:
|
||
return None, None
|
||
if size <= AUDIT_MAX_CHARS:
|
||
return raw, None
|
||
return raw[:AUDIT_MAX_CHARS], size
|
||
|
||
def run_write_file(self, args):
|
||
path = str(args.get("path") or "")
|
||
if not path:
|
||
return "error: empty path"
|
||
if self.app.env == "sandbox":
|
||
return self.box_write(path, str(args.get("content") or ""))
|
||
key = ("home", os.path.abspath(path))
|
||
if os.path.isfile(path) and key not in self.read_files:
|
||
return "error: write denied for '%s', read it first with read_file" % path
|
||
content = str(args.get("content") or "")
|
||
if os.path.isfile(path):
|
||
old, old_size = read_capped(path)
|
||
else:
|
||
old, old_size = None, None
|
||
try:
|
||
parent = os.path.dirname(os.path.abspath(path))
|
||
os.makedirs(parent, exist_ok=True)
|
||
with open(path, "w", encoding="utf-8") as handle:
|
||
handle.write(content)
|
||
self.read_files.add(key)
|
||
except OSError as exc:
|
||
return "error: " + short_error(exc)
|
||
audited = self.audit_file("write", os.path.abspath(path), old, content, "", (), old_size, None)
|
||
note = "" if audited is not False else " (audit failed)"
|
||
if old_size is None:
|
||
self.note_diff(path, old, content)
|
||
return "wrote %s%s" % (path, note)
|
||
|
||
def run_edit_file(self, args):
|
||
path = str(args.get("path") or "")
|
||
find = str(args.get("find") or "")
|
||
if self.app.env == "sandbox":
|
||
return self.box_edit(path, find, str(args.get("replace") or ""))
|
||
if not path or not os.path.isfile(path):
|
||
return "error: no such file"
|
||
if not find:
|
||
return "error: empty match"
|
||
if ("home", os.path.abspath(path)) not in self.read_files:
|
||
return "error: edit denied for '%s', read it first with read_file" % path
|
||
try:
|
||
with open(path, "r", encoding="utf-8", errors="replace") as handle:
|
||
content = handle.read()
|
||
except OSError as exc:
|
||
return "error: " + short_error(exc)
|
||
if content.count(find) != 1:
|
||
return "error: match is not unique (%d occurrences)" % content.count(find)
|
||
updated = content.replace(find, str(args.get("replace") or ""))
|
||
try:
|
||
with open(path, "w", encoding="utf-8") as handle:
|
||
handle.write(updated)
|
||
except OSError as exc:
|
||
return "error: " + short_error(exc)
|
||
audited = self.audit_file("edit", os.path.abspath(path), content, updated)
|
||
note = "" if audited is not False else " (audit failed)"
|
||
self.note_diff(path, content, updated)
|
||
return "edited %s%s" % (path, note)
|
||
|
||
def run_web_search(self, args):
|
||
query = str(args.get("query") or "").strip()
|
||
if not query:
|
||
return "error: empty query"
|
||
params = {"query": query}
|
||
if args.get("images"):
|
||
params["type"] = "images"
|
||
if args.get("content"):
|
||
params["content"] = "true"
|
||
url = RSEARCH_BASE + "/search?" + urllib.parse.urlencode(params)
|
||
request = urllib.request.Request(url, headers={"Accept": "application/json"})
|
||
try:
|
||
with urllib.request.urlopen(request, timeout=30) as response:
|
||
data = json.loads(response.read().decode("utf-8", "replace"))
|
||
except (OSError, ValueError) as exc:
|
||
return "error: search failed: " + short_error(exc)
|
||
lines = []
|
||
for pos, item in enumerate(data.get("results") or [], 1):
|
||
lines.append("%d. %s\n %s\n %s" % (pos, item.get("title"), item.get("url"), (item.get("description") or "")[:400]))
|
||
if item.get("content"):
|
||
lines.append(" content: " + truncate(str(item["content"]), 2000, 500))
|
||
return "\n".join(lines) if lines else "no results"
|
||
|
||
def run_web_fetch(self, args):
|
||
url = str(args.get("url") or "").strip()
|
||
if not url.startswith(("http://", "https://")):
|
||
return "error: url must start with http:// or https://"
|
||
method = str(args.get("method") or "GET").strip().upper()
|
||
if method not in FETCH_METHODS:
|
||
return "error: method must be one of %s" % ", ".join(FETCH_METHODS)
|
||
headers = {"User-Agent": "tai/%s" % VERSION}
|
||
custom = args.get("headers") or {}
|
||
if not isinstance(custom, dict) or any(not isinstance(name, str) or not isinstance(value, str) for name, value in custom.items()):
|
||
return "error: headers must be an object of string names to string values"
|
||
for name, value in custom.items():
|
||
if not re.fullmatch(r"[A-Za-z0-9-]+", name.strip()) or "\n" in value or "\r" in value:
|
||
return "error: invalid header '%s'" % name[:40]
|
||
headers[name.strip()] = value
|
||
data = None
|
||
if args.get("body") is not None:
|
||
if not isinstance(args.get("body"), str):
|
||
return "error: body must be a string"
|
||
data = args.get("body").encode("utf-8")
|
||
try:
|
||
timeout = int(args.get("timeout") or 30)
|
||
except (TypeError, ValueError):
|
||
return "error: invalid timeout"
|
||
timeout = max(1, min(timeout, 120))
|
||
auth_name = str(args.get("auth_secret") or "").strip()
|
||
if auth_name:
|
||
value = self.app.store.load_secret(auth_name, self.active_profile)
|
||
if value is None:
|
||
return "error: unknown secret '%s'" % auth_name
|
||
header = str(args.get("auth_header") or "Authorization").strip() or "Authorization"
|
||
if not re.fullmatch(r"[A-Za-z0-9-]+", header):
|
||
return "error: invalid auth header name"
|
||
prefix = args.get("auth_prefix") if args.get("auth_prefix") is not None else "Bearer "
|
||
prefix = str(prefix)
|
||
if "\n" in prefix or "\r" in prefix:
|
||
return "error: invalid auth prefix"
|
||
host = urllib.parse.urlparse(url).hostname or ""
|
||
if not self.secret_grant("web_fetch", host, [auth_name]):
|
||
return "denied by user"
|
||
headers[header] = prefix + value
|
||
request = urllib.request.Request(url, data=data, headers=headers, method=method)
|
||
try:
|
||
with urllib.request.urlopen(request, timeout=timeout) as response:
|
||
status = getattr(response, "status", 200)
|
||
fields = getattr(response, "headers", None)
|
||
content_type = fields.get("Content-Type", "") if fields is not None else ""
|
||
payload = response.read(200000)
|
||
except urllib.error.HTTPError as exc:
|
||
try:
|
||
detail = exc.read(2000).decode("utf-8", "replace").strip()
|
||
except OSError:
|
||
detail = ""
|
||
if detail:
|
||
return "error: HTTP %d: %s" % (exc.code, truncate(detail, 1000, 300))
|
||
return "error: HTTP %d: %s" % (exc.code, short_error(exc))
|
||
except OSError as exc:
|
||
return "error: fetch failed: " + short_error(exc)
|
||
text = payload.decode("utf-8", "replace")
|
||
if not args.get("raw") and ("html" in content_type or not content_type):
|
||
stripped = re.sub(r"(?s)<script.*?</script>|<style.*?</style>", " ", text)
|
||
stripped = re.sub(r"<[^>]+>", " ", stripped)
|
||
text = re.sub(r"\s+", " ", html.unescape(stripped)).strip()
|
||
text = truncate(text.strip(), 5000, 1000) or "empty page"
|
||
head = "HTTP %d · %s · %s" % (status, content_type.split(";")[0].strip() or "unknown", human_size(len(payload)))
|
||
return "%s\n%s" % (head, text)
|
||
|
||
def run_speak(self, args):
|
||
text = str(args.get("text") or "").strip()
|
||
if not text:
|
||
return "error: empty text"
|
||
voice = str(args.get("voice") or self.app.config.voice)
|
||
try:
|
||
audio = edge_synthesize(text[:8000], voice)
|
||
except (BackendError, OSError) as exc:
|
||
return "error: speech failed: " + short_error(exc)
|
||
if not audio:
|
||
return "error: speech returned no audio"
|
||
path = os.path.join(self.app.config.audio_dir, "speech-%d.mp3" % int(time.time()))
|
||
try:
|
||
with open(path, "wb") as handle:
|
||
handle.write(audio)
|
||
except OSError as exc:
|
||
return "error: " + short_error(exc)
|
||
return "saved %s (%d bytes), played: %s" % (path, len(audio), play_audio(path))
|
||
|
||
def run_listen(self, args):
|
||
try:
|
||
seconds = max(1, min(30, int(args.get("seconds") or 5)))
|
||
except (TypeError, ValueError):
|
||
seconds = 5
|
||
return listen_audio(seconds, self.app.config.audio_dir)
|
||
|
||
def run_remember(self, args):
|
||
instruction = str(args.get("instruction") or "").strip()
|
||
if not instruction:
|
||
return "error: empty instruction"
|
||
current = self.app.system_message
|
||
prompt = [
|
||
{"role": "system", "content": MERGE_SYSTEM},
|
||
{"role": "user", "content": "CURRENT SYSTEM PROMPT:\n" + current + "\n\nMEMORY INSTRUCTION:\n" + instruction},
|
||
]
|
||
spinner = Spinner("merging memory")
|
||
spinner.start()
|
||
try:
|
||
merged = self.app.chat.complete(prompt, ephemeral=True)
|
||
except BackendError as exc:
|
||
return "error: memory merge failed: " + short_error(exc)
|
||
finally:
|
||
spinner.stop()
|
||
text = merged["content"].strip()
|
||
text = re.sub(r"^```[a-zA-Z]*\n|\n```$", "", text).strip()
|
||
if len(text) < 20:
|
||
return "error: merged system message too short, rejected"
|
||
if len(text) > SYSTEM_MAX_CHARS:
|
||
return "error: merged system message exceeds %d chars, rejected" % SYSTEM_MAX_CHARS
|
||
self.app.update_system(text)
|
||
return "system message updated for profile '%s' (%d to %d chars)" % (self.app.profile, len(current), len(text))
|
||
|
||
def run_recall(self, args):
|
||
query = str(args.get("query") or "").strip()
|
||
if not query:
|
||
return "error: empty query"
|
||
tags = args.get("tags") or []
|
||
if not isinstance(tags, list) or any(not isinstance(item, str) for item in tags):
|
||
return "error: tags must be a list of names"
|
||
rows = self.app.store.search_events(self.active_profile, query, tags=tags)
|
||
if not rows:
|
||
return "no memories matching " + query
|
||
return "\n".join("[%s %s/%s] %s" % (stamp, role, kind, text[:300]) for stamp, role, kind, text in rows)
|
||
|
||
def box_edit(self, path, find, replace):
|
||
if not path:
|
||
return "error: empty path"
|
||
if not find:
|
||
return "error: empty match"
|
||
if ("box", path) not in self.read_files:
|
||
return "error: edit denied for '%s', read it first with read_file" % path
|
||
size, content = self.box_read_raw(path)
|
||
if size is None:
|
||
return content
|
||
if size > 200000:
|
||
return "error: file too large to edit safely (%d bytes)" % size
|
||
if content.count(find) != 1:
|
||
return "error: match is not unique (%d occurrences)" % content.count(find)
|
||
written = self.box_write(path, content.replace(find, replace), old=content, action="edit")
|
||
if written.startswith("wrote "):
|
||
return "edited %s%s" % (path, written[len("wrote " + path):])
|
||
return written
|
||
|
||
def run_load_skill(self, args):
|
||
name = str(args.get("name") or "").strip()
|
||
skill = self.app.skills.get(name)
|
||
if skill is None:
|
||
blueprint = SKILL_BLUEPRINTS.get(name)
|
||
if blueprint is None:
|
||
known = ", ".join(sorted(self.app.skills)) or "none"
|
||
return "error: unknown skill, known: " + known
|
||
return "building skill '%s' from blueprint, this runs deep research and takes a while:\n%s" % (name, self.run_create_skill({"name": name, "brief": blueprint["brief"], "scope": blueprint["scope"]}))
|
||
parts = [skill["body"].strip(), skill_provenance_note(skill)]
|
||
extras = skill_extra_files(skill["root"])
|
||
if extras:
|
||
parts.append("bundled files:\n" + "\n".join(extras))
|
||
return "\n\n".join(parts)
|
||
|
||
def run_terminal_content(self, args):
|
||
try:
|
||
lines = max(10, min(500, int(args.get("lines") or 100)))
|
||
except (TypeError, ValueError):
|
||
lines = 100
|
||
if not shutil.which("tmux"):
|
||
return "tmux not available"
|
||
pane_header = tmux_pane_header()
|
||
header = [pane_header] if pane_header else []
|
||
try:
|
||
done = subprocess.run(["tmux", "capture-pane", "-p", "-S", "-%d" % lines], stdout=subprocess.PIPE, stderr=subprocess.PIPE, universal_newlines=True, timeout=10)
|
||
except (OSError, subprocess.SubprocessError) as exc:
|
||
return "terminal capture failed: " + short_error(exc)
|
||
if done.returncode != 0:
|
||
return "terminal capture failed: " + (done.stderr or "").strip()[:200]
|
||
if not done.stdout.strip():
|
||
return "terminal pane is empty"
|
||
return "\n".join(header + [truncate(done.stdout, 6000, 2000)])
|
||
|
||
def run_export_terminal_log(self, args):
|
||
if not os.environ.get("TMUX"):
|
||
return "error: tmux is not active (this tool only works inside a live tmux session)"
|
||
if not shutil.which("tmux"):
|
||
return "error: tmux not available"
|
||
extra_tags = args.get("tags") or []
|
||
if not isinstance(extra_tags, list) or any(not isinstance(item, str) for item in extra_tags):
|
||
return "error: tags must be a list of strings"
|
||
try:
|
||
done = subprocess.run(["tmux", "capture-pane", "-p", "-S", "-"], stdout=subprocess.PIPE, stderr=subprocess.PIPE, universal_newlines=True, timeout=30)
|
||
except (OSError, subprocess.SubprocessError) as exc:
|
||
return "error: " + short_error(exc)
|
||
if done.returncode != 0:
|
||
return "error: tmux capture failed: " + (done.stderr or "").strip()[:200]
|
||
raw = done.stdout
|
||
if not raw.strip():
|
||
return "terminal pane history is empty"
|
||
profile = self.active_profile
|
||
pane_header = tmux_pane_header() or "pane"
|
||
stamp = datetime.now(timezone.utc).strftime("%Y%m%dT%H%M%SZ")
|
||
slug = re.sub(r"[^A-Za-z0-9_.-]+", "-", pane_header).strip("-") or "pane"
|
||
folder = os.path.join(self.app.config.home, "termlogs", profile)
|
||
try:
|
||
os.makedirs(folder, exist_ok=True)
|
||
path = os.path.join(folder, "%s-%s.log" % (slug, stamp))
|
||
with open(path, "w", encoding="utf-8") as handle:
|
||
handle.write(raw)
|
||
except OSError as exc:
|
||
return "error: " + short_error(exc)
|
||
if self.app.store is None:
|
||
return "captured %s to %s (%d bytes); vault unavailable, not indexed for search" % (pane_header, path, len(raw.encode("utf-8")))
|
||
title = "tmux %s captured %s" % (pane_header, stamp)
|
||
record_id = self.app.store.add_record("termlog", title, raw, ["termlog"] + extra_tags, profile)
|
||
chunks = chunk_lines(collapse_repeated_lines(raw))
|
||
self.app.store.add_chunks("record", record_id, chunks, profile)
|
||
return "captured %s to %s (%d bytes), indexed as %s in %d chunks -- use search (kind chunk) to answer questions about this terminal session" % (pane_header, path, len(raw.encode("utf-8")), record_id, len(chunks))
|
||
|
||
def run_fork(self, args):
|
||
task = str(args.get("task") or "").strip()
|
||
if not task:
|
||
return "error: empty task"
|
||
if self.app.depth >= FORK_MAX_DEPTH:
|
||
return "error: fork depth limit reached"
|
||
try:
|
||
timeout = max(30, min(3600, int(args.get("timeout") or 600)))
|
||
except (TypeError, ValueError):
|
||
timeout = 600
|
||
asked = str(args.get("profile") or "").strip()
|
||
if asked and not PROFILE_RE.match(asked):
|
||
return "error: invalid profile name"
|
||
if asked and asked != self.active_profile:
|
||
return "error: cannot fork as another profile"
|
||
profile = self.active_profile
|
||
agent_id = spawn_agent(task, profile, timeout, self.app.config, self.app.store.seal, self.app.depth + 1, self.app.runner_override)
|
||
return "agent %d started: %s. Collect its result with poll." % (agent_id, task[:80])
|
||
|
||
def run_poll(self, args):
|
||
try:
|
||
agent_id = int(args.get("id") or 0)
|
||
except (TypeError, ValueError):
|
||
return "error: invalid agent id"
|
||
try:
|
||
wait = max(0, min(120, int(args.get("wait") or 0)))
|
||
except (TypeError, ValueError):
|
||
wait = 0
|
||
status, text = poll_agent(agent_id, wait, self.active_profile)
|
||
if status in ("done", "timeout") and self.app.store is not None:
|
||
with AGENTS_LOCK:
|
||
record = dict(AGENTS.get(agent_id) or {})
|
||
full = record.get("result") or ""
|
||
if len(full) > SPILL_LIMIT:
|
||
spilled = record.get("spilled")
|
||
if not spilled:
|
||
task = str(record.get("task") or "")[:80]
|
||
spilled = self.app.store.add_record("output", "agent %d: %s" % (agent_id, task), full, ["agent-output"], self.active_profile)
|
||
self.app.store.add_chunks("record", spilled, chunk_lines(full), self.active_profile)
|
||
with AGENTS_LOCK:
|
||
if agent_id in AGENTS:
|
||
AGENTS[agent_id]["spilled"] = spilled
|
||
return "[%s] %s" % (status, spilled_view(full, spilled, len(full)))
|
||
return "[%s] %s" % (status, text)
|
||
|
||
def run_sysinfo(self, args):
|
||
names = args.get("checks")
|
||
if names is None:
|
||
return collect_sysinfo()
|
||
if not isinstance(names, list) or not names:
|
||
return "error: checks must be a non-empty list"
|
||
valid = [name for name, _func in SYSINFO_CHECKS]
|
||
unknown = [str(item) for item in names if str(item) not in valid]
|
||
if unknown:
|
||
return "error: unknown checks: %s (valid: %s)" % (", ".join(unknown), ", ".join(valid))
|
||
return collect_sysinfo([str(item) for item in names])
|
||
|
||
def run_create_skill(self, args):
|
||
name = str(args.get("name") or "").strip()
|
||
brief = str(args.get("brief") or "").strip()
|
||
if not name or not SKILL_NAME_RE.match(name):
|
||
return "error: invalid skill name, use lowercase letters, digits, and hyphens"
|
||
if not brief:
|
||
return "error: empty brief"
|
||
if self.app.depth >= FORK_MAX_DEPTH:
|
||
return "error: create_skill depth limit reached"
|
||
scope = str(args.get("scope") or "project").strip().lower()
|
||
if scope not in ("project", "home"):
|
||
return "error: scope must be project or home"
|
||
try:
|
||
timeout = max(60, min(3600, int(args.get("timeout") or 1200)))
|
||
except (TypeError, ValueError):
|
||
timeout = 1200
|
||
skill_dir = skill_scope_dir(scope, self.app.config.home, os.getcwd(), name)
|
||
skill_file = os.path.join(skill_dir, "SKILL.md")
|
||
existed = os.path.isfile(skill_file)
|
||
spinner = Spinner("creating skill %s" % name)
|
||
spinner.start()
|
||
try:
|
||
if self.app.runner_override is None:
|
||
result = default_runner(self.app.config, self.app.store.seal, self.app.profile, self.app.depth + 1, create_skill_prompt(name, brief, skill_dir), timeout, CREATE_SKILL_STEPS)
|
||
else:
|
||
result = self.app.runner_override(create_skill_prompt(name, brief, skill_dir), self.app.profile, timeout)
|
||
finally:
|
||
spinner.stop()
|
||
if os.path.isfile(skill_file):
|
||
self.app.skills = discover_skills(self.app.config.home, os.getcwd())
|
||
self.app.apply_system()
|
||
action = "replaced" if existed else "created"
|
||
return "skill '%s' %s at %s\n%s" % (name, action, skill_file, truncate(result, 2000, 500))
|
||
status = "timeout" if "[time limit reached]" in result else "done"
|
||
return "skill '%s' was not created (%s), worker output:\n%s" % (name, status, truncate(result, 3000, 1000))
|
||
|
||
@staticmethod
|
||
def find_skill_root(base):
|
||
"""A source dir qualifies if it's SKILL.md itself, or has exactly one child that is."""
|
||
if os.path.isfile(os.path.join(base, "SKILL.md")):
|
||
return base
|
||
try:
|
||
entries = sorted(os.listdir(base))
|
||
except OSError:
|
||
return None
|
||
candidates = [entry for entry in entries if os.path.isfile(os.path.join(base, entry, "SKILL.md"))]
|
||
return os.path.join(base, candidates[0]) if len(candidates) == 1 else None
|
||
|
||
def resolve_skill_source(self, source, workdir):
|
||
"""Returns ((path, ephemeral), '') on success, (None, 'error: ...') otherwise.
|
||
|
||
path is a directory containing SKILL.md, or a bare SKILL.md file. ephemeral
|
||
means the path lives under workdir (git clone / zip extraction) and symlinking
|
||
to it would dangle once workdir is cleaned up, so install must copy it.
|
||
"""
|
||
if source.endswith(".git") or source.startswith("git@") or source.startswith("git+"):
|
||
git_url = source[4:] if source.startswith("git+") else source
|
||
if not shutil.which("git"):
|
||
return None, "error: git is not installed, supply a local path or .zip instead"
|
||
clone_dir = os.path.join(workdir, "clone")
|
||
try:
|
||
done = subprocess.run(["git", "clone", "--depth", "1", git_url, clone_dir], stdout=subprocess.PIPE, stderr=subprocess.PIPE, universal_newlines=True, timeout=120)
|
||
except (OSError, subprocess.SubprocessError) as exc:
|
||
return None, "error: git clone failed: " + short_error(exc)
|
||
if done.returncode != 0:
|
||
return None, "error: git clone failed: " + (done.stderr or "").strip()[:300]
|
||
found = self.find_skill_root(clone_dir)
|
||
if found is None:
|
||
return None, "error: no SKILL.md found in " + git_url
|
||
return (found, True), ""
|
||
path = os.path.expanduser(source)
|
||
if not os.path.exists(path):
|
||
return None, "error: no such file or directory: " + source
|
||
if os.path.isfile(path) and path.endswith(".zip"):
|
||
extract_dir = os.path.join(workdir, "zip")
|
||
try:
|
||
with zipfile.ZipFile(path) as archive:
|
||
archive.extractall(extract_dir)
|
||
except (OSError, zipfile.BadZipFile) as exc:
|
||
return None, "error: bad zip file: " + short_error(exc)
|
||
found = self.find_skill_root(extract_dir)
|
||
if found is None:
|
||
return None, "error: no SKILL.md found in " + source
|
||
return (found, True), ""
|
||
if os.path.isfile(path) and path.endswith(".md"):
|
||
return (path, False), ""
|
||
if os.path.isdir(path):
|
||
found = self.find_skill_root(path)
|
||
if found is None:
|
||
return None, "error: no SKILL.md found under " + source
|
||
return (found, False), ""
|
||
return None, "error: source must be a directory, a SKILL.md file, a .zip, or a git URL"
|
||
|
||
def run_install_skill(self, args):
|
||
source = str(args.get("source") or "").strip()
|
||
if not source:
|
||
return "error: empty source"
|
||
scope = str(args.get("scope") or "").strip().lower() or default_skill_scope(os.getcwd())
|
||
if scope not in ("project", "home"):
|
||
return "error: scope must be project or home"
|
||
mode = str(args.get("mode") or "").strip().lower() or None
|
||
if mode is not None and mode not in ("copy", "symlink"):
|
||
return "error: mode must be copy or symlink"
|
||
force = bool(args.get("force"))
|
||
workdir = tempfile.mkdtemp(prefix="tai-skill-")
|
||
try:
|
||
spinner = Spinner("installing skill from %s" % source)
|
||
spinner.start()
|
||
try:
|
||
resolved, failure = self.resolve_skill_source(source, workdir)
|
||
finally:
|
||
spinner.stop()
|
||
if failure:
|
||
return failure
|
||
real_source, ephemeral = resolved
|
||
source_dir = real_source if os.path.isdir(real_source) else None
|
||
skill_file = os.path.join(real_source, "SKILL.md") if source_dir else real_source
|
||
try:
|
||
skill = parse_skill_file(skill_file)
|
||
except OSError as exc:
|
||
return "error: " + short_error(exc)
|
||
if skill is None:
|
||
return "error: %s has no valid SKILL.md (needs name, description, and a name of lowercase letters, digits, hyphens)" % source
|
||
name = str(args.get("name") or "").strip() or skill["name"]
|
||
if not SKILL_NAME_RE.match(name):
|
||
return "error: invalid skill name, use lowercase letters, digits, and hyphens"
|
||
if mode is None:
|
||
mode = "symlink" if (source_dir and not ephemeral and is_external_skill_root(source_dir)) else "copy"
|
||
elif mode == "symlink" and (ephemeral or source_dir is None):
|
||
return "error: symlink mode needs a persistent source directory, not a git/zip source or a bare SKILL.md file"
|
||
scope_dir = skill_scope_dir(scope, self.app.config.home, os.getcwd())
|
||
dest = os.path.join(scope_dir, name)
|
||
existed = os.path.islink(dest) or os.path.exists(dest)
|
||
if existed and not force and not self.app.ask_approval("overwrite existing skill '%s' at %s" % (name, dest)):
|
||
return "denied by user"
|
||
if existed:
|
||
if os.path.islink(dest) or os.path.isfile(dest):
|
||
os.remove(dest)
|
||
else:
|
||
shutil.rmtree(dest)
|
||
os.makedirs(scope_dir, exist_ok=True)
|
||
if mode == "symlink":
|
||
os.symlink(os.path.abspath(source_dir), dest)
|
||
elif source_dir is not None:
|
||
shutil.copytree(source_dir, dest)
|
||
else:
|
||
os.makedirs(dest)
|
||
shutil.copy2(real_source, os.path.join(dest, "SKILL.md"))
|
||
save_skill_install_meta(scope_dir, name, {"source": source, "mode": mode, "scope": scope, "installed_at": datetime.now(timezone.utc).isoformat()})
|
||
if self.app.store is not None:
|
||
try:
|
||
self.app.store.audit_event(self.audit_actor(), "skill-install", dest, "installed skill '%s' from %s (%s, %s scope)" % (name, source, mode, scope), tags=["skill", "install"])
|
||
except (sqlite3.Error, OSError):
|
||
pass
|
||
self.app.skills = discover_skills(self.app.config.home, os.getcwd())
|
||
self.app.apply_system()
|
||
extras = skill_extra_files(dest)
|
||
tail = "\nbundled files: %s" % ", ".join(extras) if extras else ""
|
||
replaced = " (replacing the previous install)" if existed else ""
|
||
return "skill '%s' installed (%s, %s scope) at %s from %s%s\n%s%s" % (name, mode, scope, dest, source, replaced, skill["description"], tail)
|
||
finally:
|
||
shutil.rmtree(workdir, ignore_errors=True)
|
||
|
||
def run_store_secret(self, args):
|
||
name = str(args.get("name") or "").strip()
|
||
value = str(args.get("value") or "")
|
||
if not SECRET_RE.match(name):
|
||
return "error: invalid secret name, use letters, digits, dot, underscore, hyphen"
|
||
if not value:
|
||
return "error: empty value"
|
||
username = str(args.get("username") or "").strip()[:128]
|
||
host = str(args.get("host") or "").strip()[:255]
|
||
port = 0
|
||
if args.get("port") is not None and str(args.get("port")).strip() != "":
|
||
try:
|
||
port = int(args.get("port"))
|
||
except (TypeError, ValueError):
|
||
return "error: port must be a number"
|
||
if not 1 <= port <= 65535:
|
||
return "error: port must be 1-65535"
|
||
notes = str(args.get("notes") or "").strip()[:2000]
|
||
expires = str(args.get("expires") or "").strip()
|
||
if expires:
|
||
try:
|
||
expires = parse_schedule_at(expires)
|
||
except ValueError as exc:
|
||
return "error: " + short_error(exc)
|
||
raw_tags = args.get("tags") or []
|
||
if not isinstance(raw_tags, list) or any(not isinstance(item, str) for item in raw_tags):
|
||
return "error: tags must be a list of names"
|
||
tags = normalize_tags(raw_tags)
|
||
meta = {}
|
||
if username:
|
||
meta["username"] = username
|
||
if host:
|
||
meta["host"] = host
|
||
if port:
|
||
meta["port"] = port
|
||
if notes:
|
||
meta["notes"] = notes
|
||
if expires:
|
||
meta["expires"] = expires
|
||
self.app.store.save_secret(name, value, meta, tags, self.active_profile)
|
||
return "stored secret '%s' (sealed; tags: %s)" % (name, ", ".join(["secret"] + tags))
|
||
|
||
def run_list_secrets(self, args):
|
||
infos = self.app.store.list_secret_infos(self.active_profile)
|
||
if not infos:
|
||
return "secrets: none"
|
||
now = datetime.now(timezone.utc).isoformat()
|
||
lines = []
|
||
for info in infos:
|
||
meta = info["meta"]
|
||
parts = ["- " + info["name"]]
|
||
if info["tags"]:
|
||
parts.append("[%s]" % " ".join(info["tags"]))
|
||
who = ""
|
||
if meta.get("username") and meta.get("host"):
|
||
who = "%s@%s" % (meta["username"], meta["host"])
|
||
elif meta.get("host"):
|
||
who = str(meta["host"])
|
||
elif meta.get("username"):
|
||
who = str(meta["username"])
|
||
if who and meta.get("port"):
|
||
who += ":%s" % meta["port"]
|
||
if who:
|
||
parts.append(who)
|
||
if meta.get("notes"):
|
||
parts.append(str(meta["notes"])[:60])
|
||
if meta.get("expires"):
|
||
parts.append("(expired)" if meta["expires"] <= now else "(expires %s)" % local_display(meta["expires"]))
|
||
lines.append(" ".join(parts))
|
||
return "\n".join(lines)
|
||
|
||
def run_delete_secret(self, args):
|
||
name = str(args.get("name") or "").strip()
|
||
if self.app.store.load_secret(name, self.active_profile) is None:
|
||
return "error: unknown secret '%s'" % name
|
||
if not self.app.ask_approval("delete secret '%s'" % name):
|
||
return "denied by user"
|
||
self.app.store.delete_secret(name, self.active_profile)
|
||
return "deleted secret '%s'" % name
|
||
|
||
def run_schedule(self, args):
|
||
prompt = str(args.get("prompt") or "").strip()
|
||
if not prompt:
|
||
return "error: empty prompt"
|
||
at_raw = args.get("at")
|
||
every_raw = args.get("every")
|
||
if (at_raw is None) == (every_raw is None):
|
||
return "error: pass exactly one of at or every"
|
||
if at_raw is not None:
|
||
try:
|
||
next_run = parse_schedule_at(str(at_raw))
|
||
except ValueError as exc:
|
||
return "error: " + short_error(exc)
|
||
every_sec = 0
|
||
else:
|
||
try:
|
||
every_sec = int(every_raw)
|
||
except (TypeError, ValueError):
|
||
return "error: every must be seconds"
|
||
if every_sec < 60:
|
||
return "error: every must be at least 60 seconds"
|
||
next_run = (datetime.now(timezone.utc) + timedelta(seconds=every_sec)).isoformat()
|
||
name = str(args.get("name") or "").strip()[:80]
|
||
try:
|
||
timeout = max(60, min(3600, int(args.get("timeout") or 600)))
|
||
except (TypeError, ValueError):
|
||
timeout = 600
|
||
asked = str(args.get("profile") or "").strip()
|
||
if asked and not PROFILE_RE.match(asked):
|
||
return "error: invalid profile name"
|
||
if asked and asked != self.active_profile:
|
||
return "error: cannot schedule for another profile"
|
||
profile = self.active_profile
|
||
raw_tags = args.get("tags") or []
|
||
if not isinstance(raw_tags, list) or any(not isinstance(item, str) for item in raw_tags):
|
||
return "error: tags must be a list of names"
|
||
if not self.app.ask_approval("schedule '%s'" % (name or prompt[:60])):
|
||
return "denied by user"
|
||
row_id = self.app.store.add_schedule(name, prompt, profile, every_sec, next_run, timeout, normalize_tags(raw_tags))
|
||
when = "every %s" % format_delay(every_sec) if every_sec else "at %s" % local_display(next_run)
|
||
return "scheduled #%d (%s)" % (row_id, when)
|
||
|
||
def run_unschedule(self, args):
|
||
try:
|
||
row_id = int(args.get("id") or 0)
|
||
except (TypeError, ValueError):
|
||
return "error: invalid schedule id"
|
||
if row_id not in [item["id"] for item in self.app.store.list_schedules(self.active_profile)]:
|
||
return "error: no schedule #%d" % row_id
|
||
if not self.app.ask_approval("delete schedule #%d" % row_id):
|
||
return "denied by user"
|
||
self.app.store.remove_schedule(row_id, self.active_profile)
|
||
return "deleted schedule #%d" % row_id
|
||
|
||
def run_schedules(self, args):
|
||
return format_schedule_lines(self.app.store.list_schedules(self.active_profile), datetime.now(timezone.utc))
|
||
|
||
def record_tag_args(self, args):
|
||
raw_tags = args.get("tags") or []
|
||
if not isinstance(raw_tags, list) or any(not isinstance(item, str) for item in raw_tags):
|
||
return None
|
||
return normalize_tags(raw_tags)
|
||
|
||
def run_record_save(self, args):
|
||
content = str(args.get("content") or "")
|
||
if not content:
|
||
return "error: content is empty"
|
||
title = str(args.get("title") or "").strip()[:200] or "(untitled)"
|
||
kind = str(args.get("kind") or "note").strip().lower() or "note"
|
||
if kind not in RECORD_KINDS:
|
||
return "error: unknown kind %s (valid: %s)" % (kind, ", ".join(RECORD_KINDS))
|
||
tags = self.record_tag_args(args)
|
||
if tags is None:
|
||
return "error: tags must be a list of names"
|
||
record_id = self.app.store.add_record(kind, title, content, tags, self.active_profile)
|
||
chunks = chunk_lines(content)
|
||
self.app.store.add_chunks("record", record_id, chunks, self.active_profile)
|
||
saved = "saved %s (%d chars, kind %s, %d chunks, tags: %s)" % (record_id, len(content), kind, len(chunks), ", ".join(self.app.store.item_tags(record_id, self.active_profile)))
|
||
links = sorted(edge["other"] for edge in self.app.store.edges_for(record_id, self.active_profile) if edge["direction"] == "out")
|
||
if links:
|
||
saved += " [linked: %s]" % ", ".join(links)
|
||
return saved
|
||
|
||
def run_create_bot(self, args):
|
||
name = str(args.get("name") or "").strip().lower()
|
||
if not BOT_RE.match(name):
|
||
return "error: invalid bot name (lowercase letters, digits, - and _)"
|
||
if name == "main":
|
||
return "error: 'main' is the default bot"
|
||
description = str(args.get("description") or "").strip()
|
||
rules = str(args.get("rules") or "").strip()
|
||
behavior = str(args.get("behavior") or "").strip()
|
||
if not description and not rules and not behavior:
|
||
return "error: give at least one of description, rules, or behavior"
|
||
raw_nicks = args.get("nicknames") or []
|
||
if not isinstance(raw_nicks, list) or any(not isinstance(item, str) for item in raw_nicks):
|
||
return "error: nicknames must be a list of names"
|
||
profile = self.active_profile
|
||
taken = {"main"}
|
||
for item in self.app.store.list_bots(profile):
|
||
taken.add(item["name"])
|
||
taken.update(item["nicknames"])
|
||
short = re.split(r"[-_]", name, 1)[0]
|
||
nicks = []
|
||
for cand in list(raw_nicks) + [short]:
|
||
clean = cand.strip().lower()
|
||
if clean and clean != name and BOT_RE.match(clean) and clean not in nicks:
|
||
nicks.append(clean)
|
||
for cand in list(raw_nicks):
|
||
clean = cand.strip().lower()
|
||
if clean and not BOT_RE.match(clean):
|
||
return "error: invalid nickname '%s'" % cand.strip()
|
||
if name in taken:
|
||
return "error: bot name '%s' is taken" % name
|
||
clashes = [nick for nick in nicks if nick in taken]
|
||
if clashes:
|
||
return "error: nickname '%s' is taken" % clashes[0]
|
||
parts = []
|
||
if description:
|
||
parts.append(description)
|
||
if rules:
|
||
parts.append("Rules:\n" + rules)
|
||
if behavior:
|
||
parts.append("Behavior:\n" + behavior)
|
||
self.app.store.save_bot_system(profile, name, "\n\n".join(parts))
|
||
bots = self.app.store.load_bots(profile)
|
||
stamp = now_iso()
|
||
bots[name] = {"nicknames": nicks, "created": stamp, "updated": stamp}
|
||
self.app.store.save_bots(profile, bots)
|
||
extra = " (nicknames: %s)" % ", ".join(nicks) if nicks else ""
|
||
return "created bot '%s'%s. Switch with /bot %s or mention @%s." % (name, extra, name, name)
|
||
|
||
def run_install(self, args):
|
||
action = str(args.get("action") or "").strip().lower()
|
||
if action not in ("status", "install", "upgrade", "reinstall", "uninstall"):
|
||
return "error: action must be status, install, upgrade, reinstall, or uninstall"
|
||
targets = args.get("targets") or list(INSTALL_TARGETS)
|
||
if not isinstance(targets, list) or not targets or any(item not in INSTALL_TARGETS for item in targets):
|
||
return "error: unknown target, valid: %s" % ", ".join(INSTALL_TARGETS)
|
||
if action != "status" and not self.app.ask_approval("install %s: %s (vault always kept)" % (action, ", ".join(targets))):
|
||
return "denied by user"
|
||
return install_report(action, targets, self.app.config.home, self.app.store)
|
||
|
||
def run_search(self, args):
|
||
query = str(args.get("query") or "").strip()
|
||
if not query:
|
||
return "error: empty query"
|
||
kinds = args.get("kinds") or ["record", "event", "audit", "chunk"]
|
||
if not isinstance(kinds, list) or not kinds or any(item not in ("record", "event", "audit", "chunk") for item in kinds):
|
||
return "error: kinds must be a list of record, event, audit, chunk"
|
||
try:
|
||
limit = int(args.get("limit") or 10)
|
||
except (TypeError, ValueError):
|
||
return "error: invalid limit"
|
||
expand = bool(args.get("expand"))
|
||
profile = self.active_profile
|
||
hits = self.app.store.fts_search(query, kinds, profile, limit)
|
||
if self.app.store.seal.enabled and "event" in kinds and not any(hit["kind"] == "event" for hit in hits):
|
||
for ts, _role, kind, text in self.app.store.events_like(profile, query, limit):
|
||
hits.append({"kind": "event", "item": "event %s" % ts, "sub": kind, "title": text[:60], "snippet": text[:160], "rank": 99.0})
|
||
hits.sort(key=lambda hit: hit["rank"])
|
||
if not hits:
|
||
return "no matches for '%s'" % query[:80]
|
||
chunk_ids = []
|
||
for hit in hits:
|
||
if hit["kind"] == "chunk":
|
||
try:
|
||
chunk_ids.append(int(str(hit["item"]).split(":", 1)[1]))
|
||
except (IndexError, ValueError):
|
||
pass
|
||
chunk_rows = self.app.store.get_chunks_by_ids(chunk_ids, profile) if chunk_ids else {}
|
||
lines = []
|
||
for hit in hits:
|
||
if hit["kind"] == "record":
|
||
lines.append("[record/%s] %s :: %s -- %s" % (hit["sub"], hit["item"], hit["title"][:80], hit["snippet"]))
|
||
if expand:
|
||
neighbors = []
|
||
for edge in self.app.store.edges_for(hit["item"], profile):
|
||
if len(neighbors) < 3:
|
||
neighbors.append("%s (%s)" % (self.app.store.node_title(edge["other"], profile)[:60], edge["relation"]))
|
||
if neighbors:
|
||
lines.append(" linked: %s" % "; ".join(neighbors))
|
||
elif hit["kind"] == "event":
|
||
lines.append("[event/%s] #%s %s -- %s" % (hit["sub"], hit["item"], hit["title"][:60], hit["snippet"]))
|
||
elif hit["kind"] == "chunk":
|
||
try:
|
||
chunk_id = int(str(hit["item"]).split(":", 1)[1])
|
||
except (IndexError, ValueError):
|
||
chunk_id = None
|
||
chunk = chunk_rows.get(chunk_id)
|
||
if chunk is None:
|
||
lines.append("[chunk/%s] %s -- %s" % (hit["sub"], hit["title"][:80], hit["snippet"]))
|
||
else:
|
||
preview = chunk["text"][:500] + ("..." if len(chunk["text"]) > 500 else "")
|
||
lines.append("[chunk/%s] %s#%d (parent %s, read full with record_read) -- %s" % (chunk["parent_kind"], chunk["parent_id"][:60], chunk["idx"], chunk["parent_id"], preview))
|
||
else:
|
||
lines.append("[audit/%s] #%s %s -- %s" % (hit["sub"], str(hit["item"]).split(":", 1)[1], hit["title"][:80], hit["snippet"]))
|
||
return "\n".join(lines)
|
||
|
||
def run_tags(self, args):
|
||
prefix = str(args.get("prefix") or "").strip()
|
||
try:
|
||
limit = int(args.get("limit") or 50)
|
||
except (TypeError, ValueError):
|
||
return "error: invalid limit"
|
||
counts = self.app.store.tag_counts(prefix, limit, self.active_profile)
|
||
if not counts:
|
||
return "no tags yet"
|
||
lines = ["%s (%d)" % (tag, count) for tag, count in counts]
|
||
return "%d tags:\n%s" % (len(counts), "\n".join(lines))
|
||
|
||
def run_record_read(self, args):
|
||
record_id = str(args.get("id") or "").strip()
|
||
try:
|
||
offset = max(0, int(args.get("offset") or 0))
|
||
except (TypeError, ValueError):
|
||
return "error: invalid offset"
|
||
try:
|
||
limit = int(args.get("limit") or 4000)
|
||
except (TypeError, ValueError):
|
||
return "error: invalid limit"
|
||
page = self.app.store.read_record(record_id, offset, max(1, min(limit, 20000)), self.active_profile)
|
||
if page is None:
|
||
return "error: unknown record " + record_id
|
||
head = "%s chars %d-%d of %d :: %s" % (record_id, page["start"], page["end"], page["total"], self.app.store.node_title(record_id, self.active_profile))
|
||
if page["tags"]:
|
||
head += " [%s]" % ", ".join(page["tags"])
|
||
return head + "\n" + page["slice"]
|
||
|
||
def run_record_search(self, args):
|
||
query = str(args.get("query") or "")
|
||
kind = str(args.get("kind") or "").strip().lower() or None
|
||
if kind is not None and kind not in RECORD_KINDS:
|
||
return "error: unknown kind %s (valid: %s)" % (kind, ", ".join(RECORD_KINDS))
|
||
tags = self.record_tag_args(args)
|
||
if tags is None:
|
||
return "error: tags must be a list of names"
|
||
try:
|
||
limit = int(args.get("limit") or 8)
|
||
except (TypeError, ValueError):
|
||
return "error: invalid limit"
|
||
found = self.app.store.search_records(query, kind, tags, limit, self.active_profile)
|
||
if not found:
|
||
return "no records match"
|
||
lines = []
|
||
for item in found:
|
||
lines.append("%s [%s] %s (%d chars, %d reads) tags: %s" % (item["id"], item["kind"], item["title"], item["size"], item["reads"], ", ".join(item["tags"])))
|
||
return "\n".join(lines)
|
||
|
||
def run_record_delete(self, args):
|
||
record_id = str(args.get("id") or "").strip()
|
||
if not record_id:
|
||
return "error: record id is required"
|
||
if not self.app.ask_approval("delete record %s" % record_id):
|
||
return "denied by user"
|
||
if self.app.store.delete_record(record_id, self.active_profile):
|
||
return "deleted record " + record_id
|
||
return "error: unknown record " + record_id
|
||
|
||
def run_graph_link(self, args):
|
||
src = str(args.get("src") or "").strip()
|
||
dst = str(args.get("dst") or "").strip()
|
||
if not src or not dst:
|
||
return "error: src and dst are required"
|
||
try:
|
||
relation = self.app.store.add_edge(src, dst, args.get("relation"), self.active_profile)
|
||
except ValueError as exc:
|
||
return "error: " + str(exc)
|
||
return "linked %s -[%s]-> %s" % (src, relation, dst)
|
||
|
||
def run_delete_file(self, args):
|
||
path = str(args.get("path") or "")
|
||
if not path:
|
||
return "error: empty path"
|
||
if self.app.env == "sandbox":
|
||
return self.box_delete(path)
|
||
if not os.path.isfile(path):
|
||
return "error: no such file"
|
||
if ("home", os.path.abspath(path)) not in self.read_files:
|
||
return "error: delete denied for '%s', read it first with read_file" % path
|
||
if not self.app.ask_approval("delete file %s" % path):
|
||
return "denied by user"
|
||
old, old_size = read_capped(path)
|
||
try:
|
||
os.remove(path)
|
||
except OSError as exc:
|
||
return "error: " + short_error(exc)
|
||
self.read_files.discard(("home", os.path.abspath(path)))
|
||
audited = self.audit_file("delete", os.path.abspath(path), old, None, "", (), old_size, None)
|
||
note = "" if audited is not False else " (audit failed)"
|
||
if old_size is None:
|
||
self.note_diff(path, old, None)
|
||
return "deleted %s%s" % (path, note)
|
||
|
||
def box_delete(self, path):
|
||
engine, failure = self.box_engine()
|
||
if not engine:
|
||
return failure
|
||
if ("box", path) not in self.read_files:
|
||
try:
|
||
probe = box_exec(engine, ["test", "-f", path], timeout=30)
|
||
except (OSError, subprocess.SubprocessError):
|
||
probe = None
|
||
if probe is None or probe.returncode != 0:
|
||
return "error: no such file"
|
||
return "error: delete denied for '%s', read it first with read_file" % path
|
||
old, old_size = self.box_old_image(path)
|
||
if old is None:
|
||
return "error: no such file"
|
||
try:
|
||
done = box_exec(engine, ["rm", "-f", path], timeout=60)
|
||
except (OSError, subprocess.SubprocessError) as exc:
|
||
return "error: " + short_error(exc)
|
||
if done.returncode != 0:
|
||
return "error: " + done.stderr.decode("utf-8", "replace")[:200]
|
||
self.read_files.discard(("box", path))
|
||
audited = self.audit_file("delete", "sandbox:" + path, old, None, "", (), old_size, None)
|
||
note = "" if audited is not False else " (audit failed)"
|
||
if old_size is None:
|
||
self.note_diff(path, old, None)
|
||
return "deleted %s%s" % (path, note)
|
||
|
||
def run_audit(self, args):
|
||
path = str(args.get("path") or "").strip() or None
|
||
tag = str(args.get("tag") or "").strip() or None
|
||
query = str(args.get("query") or "").strip() or None
|
||
try:
|
||
limit = int(args.get("limit") or 20)
|
||
except (TypeError, ValueError):
|
||
return "error: invalid limit"
|
||
if query:
|
||
rows = self.app.store.audit_search(query, self.active_profile, limit)
|
||
else:
|
||
rows = self.app.store.audit_history(path, tag, limit, self.active_profile)
|
||
if not rows:
|
||
return "no audit rows match"
|
||
lines = []
|
||
for row in rows:
|
||
old_size = row["old_size"] if row["old_size"] is not None else "-"
|
||
new_size = row["new_size"] if row["new_size"] is not None else "-"
|
||
line = "#%d %s %s %s %s (%s->%s bytes)" % (row["id"], row["ts"][:19], row["actor"], row["action"], row["path"], old_size, new_size)
|
||
if row["message"]:
|
||
line += " :: " + row["message"][:120]
|
||
if row["tags"]:
|
||
line += " [%s]" % row["tags"]
|
||
lines.append(line)
|
||
return "\n".join(lines)
|
||
|
||
def run_restore(self, args):
|
||
try:
|
||
row_id = int(args.get("id") or 0)
|
||
except (TypeError, ValueError):
|
||
return "error: invalid audit id"
|
||
row = self.app.store.audit_get(row_id, self.active_profile)
|
||
if row is None:
|
||
return "error: unknown audit #%d" % row_id
|
||
if row["action"] in ("snapshot", "release"):
|
||
return "error: audit #%d is a %s row, restorable rows are write, edit, delete, shell-snapshot, restore" % (row_id, row["action"])
|
||
if row["new"] is not None:
|
||
target, which, full_size = row["new"], "post", row["new_size"]
|
||
elif row["old"] is not None:
|
||
target, which, full_size = row["old"], "pre", row["old_size"]
|
||
else:
|
||
return "error: audit #%d holds no content" % row_id
|
||
if "[truncated, full size" in target:
|
||
return "error: audit #%d %s-image is truncated (full size %s bytes), cannot restore safely" % (row_id, which, full_size)
|
||
if not self.app.ask_approval("restore %s from audit #%d (%s-image)" % (row["path"], row_id, which)):
|
||
return "denied by user"
|
||
message = "restored from audit #%d (%s-image)" % (row_id, which)
|
||
if row["path"].startswith("sandbox:"):
|
||
return self.box_restore(row["path"][len("sandbox:"):], target, message)
|
||
if os.path.isfile(row["path"]):
|
||
current, capped = read_capped(row["path"])
|
||
else:
|
||
current, capped = None, None
|
||
try:
|
||
parent = os.path.dirname(row["path"])
|
||
if parent:
|
||
os.makedirs(parent, exist_ok=True)
|
||
with open(row["path"], "w", encoding="utf-8") as handle:
|
||
handle.write(target)
|
||
self.read_files.add(("home", os.path.abspath(row["path"])))
|
||
except OSError as exc:
|
||
return "error: " + short_error(exc)
|
||
self.audit_file("restore", row["path"], None, target, message)
|
||
self.app.store.upsert_file_record(row["path"], target, None, self.active_profile)
|
||
if capped is None:
|
||
self.note_diff(row["path"], current, target)
|
||
return "restored %s from audit #%d (%s-image)" % (row["path"], row_id, which)
|
||
|
||
def box_restore(self, path, target, message):
|
||
self.read_files.add(("box", path))
|
||
old, capped = self.box_old_image(path)
|
||
written = self.box_write(path, target, old=None, action="restore", message=message, preview=False)
|
||
if written.startswith("wrote "):
|
||
if capped is None:
|
||
self.note_diff(path, old, target)
|
||
return "restored sandbox:%s (%s)" % (path, message)
|
||
return written
|
||
|
||
def run_release(self, args):
|
||
part = str(args.get("part") or "").strip().lower()
|
||
message = str(args.get("message") or "").strip()
|
||
if part not in ("major", "minor", "patch"):
|
||
return "error: part must be major, minor, or patch"
|
||
if not message:
|
||
return "error: release message is required"
|
||
script_file = os.path.abspath(__file__)
|
||
try:
|
||
with open(script_file, "r", encoding="utf-8") as handle:
|
||
match = VERSION_RE.search(handle.read())
|
||
if match is None:
|
||
return "error: no VERSION line in running script"
|
||
current = ".".join(match.groups())
|
||
upcoming = next_version(current, part)
|
||
except (OSError, ValueError) as exc:
|
||
return "error: " + short_error(exc)
|
||
if not self.app.ask_approval("release %s (%s): %s" % (upcoming, part, message[:120])):
|
||
return "denied by user"
|
||
try:
|
||
old, new, dest, row_id = do_release(self.app.store, script_file, backups_dir(self.app.config.home), part, message, self.audit_actor())
|
||
except (OSError, ValueError, sqlite3.Error) as exc:
|
||
return "error: " + short_error(exc)
|
||
return "released %s -> %s, backup %s, audit #%d" % (old, new, os.path.basename(dest), row_id)
|
||
|
||
def run_graph_query(self, args):
|
||
node = str(args.get("node") or "").strip()
|
||
parsed = parse_node_id(node)
|
||
if parsed is None or not self.app.store.has_node(*parsed, self.active_profile):
|
||
return "error: unknown node " + node
|
||
try:
|
||
depth = int(args.get("depth") or 2)
|
||
except (TypeError, ValueError):
|
||
return "error: invalid depth"
|
||
try:
|
||
limit = int(args.get("limit") or 30)
|
||
except (TypeError, ValueError):
|
||
return "error: invalid limit"
|
||
found = self.app.store.traverse(node, depth, limit, self.active_profile)
|
||
lines = []
|
||
for item in found:
|
||
if item["depth"] == 0:
|
||
lines.append("%s :: %s" % (item["node"], item["title"]))
|
||
else:
|
||
lines.append("%s%s :: %s" % (" " * item["depth"], item["via"], item["title"]))
|
||
return "\n".join(lines)
|
||
|
||
|
||
def conversation_tokens(messages, tools=None):
|
||
total = 0
|
||
for item in messages:
|
||
total += estimate_tokens(item.get("content"))
|
||
for call in item.get("tool_calls") or []:
|
||
total += estimate_tokens(json.dumps(call))
|
||
if tools:
|
||
total += estimate_tokens(json.dumps(tools))
|
||
return total
|
||
|
||
|
||
def compact_messages(messages, chat, keep=KEEP_TURNS):
|
||
if len(messages) <= keep + 1:
|
||
return messages, False
|
||
system = messages[0] if messages[0].get("role") == "system" else {"role": "system", "content": DEFAULT_SYSTEM}
|
||
recent = messages[-keep:]
|
||
while recent and recent[0].get("role") != "user":
|
||
recent = recent[1:]
|
||
if not recent:
|
||
recent = messages[-keep:]
|
||
middle = messages[1:len(messages) - len(recent)]
|
||
if not middle:
|
||
return messages, False
|
||
digest = []
|
||
for item in middle:
|
||
role = item.get("role")
|
||
if role == "assistant" and item.get("tool_calls"):
|
||
names = ", ".join(call.get("function", {}).get("name", "?") for call in item["tool_calls"])
|
||
digest.append("assistant tool calls: " + names)
|
||
digest.append("%s: %s" % (role, (item.get("content") or "")[:1500]))
|
||
request = [
|
||
{"role": "system", "content": COMPACT_SYSTEM},
|
||
{"role": "user", "content": "Compress this history:\n" + "\n".join(digest)},
|
||
]
|
||
summary = chat.complete(request, ephemeral=True)["content"].strip()
|
||
rebuilt = [
|
||
system,
|
||
{"role": "user", "content": "Previous session summary:\n" + summary},
|
||
{"role": "assistant", "content": "Understood. Continuing with summarized context."},
|
||
] + recent
|
||
return rebuilt, True
|
||
|
||
|
||
class Denied(Exception):
|
||
def __init__(self, guidance=""):
|
||
super().__init__("denied by user")
|
||
self.guidance = guidance
|
||
|
||
|
||
class Agent:
|
||
def __init__(self, config, store, persist=True, quiet=False, depth=0, auto=False, yolo=False):
|
||
self.config = config
|
||
self.auto = auto
|
||
self.yolo = yolo
|
||
self.store = store
|
||
self.chat = ChatClient(config)
|
||
if not persist or depth > 0:
|
||
self.chat.fixed_key = "worker/%s" % uuid.uuid4().hex
|
||
self.tools = Tools(self)
|
||
self.profile = ""
|
||
self.bot = "main"
|
||
self.system_message = ""
|
||
self.messages = []
|
||
self.skills = discover_skills(config.home, os.getcwd())
|
||
self.env = "home"
|
||
self.persist = persist
|
||
self.quiet = quiet
|
||
self.depth = depth
|
||
self.deadline = None
|
||
self.timed_out = False
|
||
self.runner_override = None
|
||
self.user_stepped_in = False
|
||
self._turn_open = False
|
||
self._turn_goal = ""
|
||
self._turn_steps = 0
|
||
self._turn_tools = 0
|
||
self.switch_profile(config.profile, silent=True)
|
||
|
||
def apply_system(self):
|
||
content = self.system_message + skill_catalog(self.skills) + tool_catalog()
|
||
if self.auto:
|
||
content += "\n\n" + AUTONOMOUS_NOTE
|
||
self.messages[0] = {"role": "system", "content": content}
|
||
|
||
def sync_chat_scope(self):
|
||
self.chat.scope = (self.profile, self.bot)
|
||
|
||
def reset_history(self):
|
||
self.messages = [{"role": "system", "content": self.system_message}]
|
||
self.apply_system()
|
||
|
||
def switch_profile(self, name, silent=False):
|
||
if self.messages and self.persist:
|
||
self.store.save_session(self.profile, self.messages, self.bot)
|
||
self.profile = name
|
||
self.bot = "main"
|
||
self.store.profile = name
|
||
self.tools.read_files.clear()
|
||
self.tools.secret_grants.clear()
|
||
loaded = self.store.load_system(name)
|
||
if loaded is None:
|
||
loaded = DEFAULT_SYSTEM
|
||
self.store.save_system(name, loaded)
|
||
if LEGACY_PASSWORD_NOTE in loaded:
|
||
loaded = loaded.replace(LEGACY_PASSWORD_NOTE, VAULT_MEMORY_NOTE)
|
||
self.store.save_system(name, loaded)
|
||
self.system_message = loaded
|
||
restored = self.store.load_session(name) if self.persist else []
|
||
self.messages = [{"role": "system", "content": self.system_message}] + restored
|
||
self.apply_system()
|
||
self.sync_chat_scope()
|
||
if not silent:
|
||
print(paint("profile: %s (%d chars knowledge, %d restored messages)" % (name, len(self.system_message), len(restored)), Ansi.GREEN))
|
||
|
||
def update_system(self, text):
|
||
self.system_message = text
|
||
if self.bot == "main":
|
||
self.store.save_system(self.profile, text)
|
||
else:
|
||
self.store.save_bot_system(self.profile, self.bot, text)
|
||
self.apply_system()
|
||
self.store.log_event(self.profile, "system", "remember", "system message updated (%d chars)" % len(text))
|
||
|
||
def switch_bot(self, name, silent=False):
|
||
target = self.store.resolve_bot(self.profile, name)
|
||
if target is None:
|
||
return "unknown bot @%s (available: %s)" % (name, ", ".join(item["name"] for item in self.store.list_bots(self.profile)))
|
||
if self.messages and self.persist:
|
||
self.store.save_session(self.profile, self.messages, self.bot)
|
||
self.bot = target
|
||
loaded = self.store.load_bot_system(self.profile, target)
|
||
if loaded is None:
|
||
loaded = DEFAULT_SYSTEM
|
||
if target == "main":
|
||
self.store.save_system(self.profile, loaded)
|
||
else:
|
||
self.store.save_bot_system(self.profile, target, loaded)
|
||
self.system_message = loaded
|
||
restored = self.store.load_session(self.profile, target) if self.persist else []
|
||
self.messages = [{"role": "system", "content": self.system_message}] + restored
|
||
self.apply_system()
|
||
self.sync_chat_scope()
|
||
if not silent:
|
||
print(paint("bot: %s (%d restored messages)" % (target, len(restored)), Ansi.GREEN))
|
||
return "switched to bot '%s'" % target
|
||
|
||
def scrub_secret_from_history(self):
|
||
for item in self.messages:
|
||
content = item.get("content")
|
||
if isinstance(content, str) and content:
|
||
item["content"] = self.store.redact(content, self.profile)
|
||
for recorded in item.get("tool_calls") or []:
|
||
args = (recorded.get("function") or {}).get("arguments")
|
||
if isinstance(args, str) and args:
|
||
recorded["function"]["arguments"] = self.store.redact(args, self.profile)
|
||
self.chat.reset_muse()
|
||
|
||
def ask_approval(self, command, guidance=True):
|
||
if self.yolo or self.auto:
|
||
return True
|
||
if not self.persist:
|
||
return self.config.auto_approve
|
||
if not sys.stdin.isatty():
|
||
return False
|
||
print(paint("run: %s" % command, Ansi.YELLOW))
|
||
try:
|
||
answer = input(paint("allow? [y]once [Y]always [n]o ", Ansi.YELLOW)).strip()
|
||
except (EOFError, KeyboardInterrupt):
|
||
return False
|
||
self.user_stepped_in = True
|
||
if answer == "Y" or answer.lower() in ("yolo", "always"):
|
||
self.yolo = True
|
||
self.config.auto_approve = True
|
||
print(paint("yolo mode on, approvals off for this session", Ansi.YELLOW))
|
||
return True
|
||
if answer.lower() in ("y", "yes"):
|
||
return True
|
||
if not guidance or self.auto:
|
||
return False
|
||
try:
|
||
typed = input(paint("denied. What should I do instead? (empty aborts this turn) ", Ansi.YELLOW)).strip()
|
||
except (EOFError, KeyboardInterrupt):
|
||
typed = ""
|
||
raise Denied(typed)
|
||
|
||
def show_call(self, call):
|
||
if self.quiet:
|
||
return
|
||
print(paint("┌─ %s" % call["name"], Ansi.MAGENTA))
|
||
preview_args = call["arguments"]
|
||
if call["name"] == "store_secret":
|
||
try:
|
||
masked = json.loads(preview_args or "{}")
|
||
masked["value"] = "[hidden]"
|
||
preview_args = json.dumps(masked)
|
||
except ValueError:
|
||
preview_args = "[hidden]"
|
||
preview = self.store.redact(preview_args, self.profile)[:400].replace("\n", " ")
|
||
print(paint("│ %s" % preview, Ansi.DIM))
|
||
|
||
def show_result(self, result, elapsed_ms):
|
||
if self.quiet:
|
||
return
|
||
self.show_diff()
|
||
print(paint("└─ %d chars · %dms" % (len(result), elapsed_ms), Ansi.DIM))
|
||
|
||
def show_diff(self):
|
||
preview = getattr(self.tools, "diff_preview", None)
|
||
if preview is None:
|
||
return
|
||
self.tools.diff_preview = None
|
||
rendered = render_diff(preview["path"], preview["old"], preview["new"], colors=color_enabled())
|
||
if rendered:
|
||
print(rendered)
|
||
|
||
def checkpoint(self):
|
||
if not self.persist:
|
||
return
|
||
self.store.save_session(self.profile, self.messages, self.bot)
|
||
if self._turn_open:
|
||
self.store.save_turn_state(self.profile, self.bot, {"open": True, "goal": self._turn_goal, "steps": self._turn_steps, "tools": self._turn_tools, "updated": now_iso()})
|
||
|
||
def close_turn(self):
|
||
if not self.persist:
|
||
self._turn_open = False
|
||
return
|
||
self.store.save_session(self.profile, self.messages, self.bot)
|
||
self._turn_open = False
|
||
if self._turn_goal:
|
||
self.store.save_turn_state(self.profile, self.bot, {"open": False, "goal": self._turn_goal, "steps": self._turn_steps, "tools": self._turn_tools, "updated": now_iso()})
|
||
|
||
def run_turn(self, text, capture=False, max_steps=MAX_STEPS, _routed=False):
|
||
if not _routed:
|
||
mentioned = BOT_MENTION_RE.match(text or "")
|
||
if mentioned:
|
||
target = self.store.resolve_bot(self.profile, mentioned.group(1))
|
||
if target is None:
|
||
return "unknown bot @%s (available: %s)" % (mentioned.group(1), ", ".join(item["name"] for item in self.store.list_bots(self.profile)))
|
||
if target != self.bot:
|
||
return self.run_mention(target, mentioned.group(2), text, capture, max_steps)
|
||
text = mentioned.group(2)
|
||
self.messages.append({"role": "user", "content": text})
|
||
self.store.log_event(self.profile, "user", "message", self.store.redact(text, self.profile))
|
||
self._turn_open = True
|
||
self._turn_goal = text[:500]
|
||
self._turn_steps = 0
|
||
self._turn_tools = 0
|
||
self.checkpoint()
|
||
started = time.time()
|
||
backends = []
|
||
pieces = []
|
||
last_text = ""
|
||
sink = pieces.append
|
||
loud = not capture and not self.quiet
|
||
self.tools.live = loud
|
||
seen_actions = collections.deque(maxlen=400)
|
||
seen_results = collections.deque(maxlen=400)
|
||
stall = 0
|
||
last_pair = None
|
||
run_len = 0
|
||
nudged = False
|
||
for _step in range(TOTAL_STEP_CAP):
|
||
if self.deadline is not None and time.time() > self.deadline:
|
||
self.timed_out = True
|
||
last_text = (last_text + "\n[time limit reached]").strip()
|
||
break
|
||
payload = select_tools(conversation_text(self.messages))
|
||
tokens = conversation_tokens(self.messages, payload)
|
||
if tokens > CONTEXT_CAP * COMPACT_RATIO:
|
||
if loud:
|
||
print(paint("compacting context (%d tokens)..." % tokens, Ansi.DIM))
|
||
self.messages, _changed = compact_messages(self.messages, self.chat)
|
||
self.apply_system()
|
||
spinner = None
|
||
if loud:
|
||
print(paint("tai", Ansi.BOLD, Ansi.CYAN))
|
||
spinner = Spinner("thinking")
|
||
spinner.start()
|
||
try:
|
||
reply = self.chat.complete(self.messages, payload, stream_sink=sink, deadline=self.deadline)
|
||
except BackendError as exc:
|
||
if loud:
|
||
print(paint("backend error: %s" % exc, Ansi.RED))
|
||
self.messages.pop()
|
||
self.checkpoint()
|
||
return last_text or "backend error: %s" % exc
|
||
finally:
|
||
if spinner is not None:
|
||
spinner.stop()
|
||
reply["content"] = self.store.redact(reply["content"] or "", self.profile)
|
||
if loud and reply["content"].strip():
|
||
print(render_markdown(reply["content"]))
|
||
if reply["backend"] not in backends:
|
||
backends.append(reply["backend"])
|
||
if loud and (len(backends) > 1 or reply["backend"] != MUSE_LABEL):
|
||
print(paint("(via %s)" % reply["backend"], Ansi.DIM))
|
||
if loud and reply["reasoning"] and not reply["content"] and not reply["tool_calls"]:
|
||
print(paint(reply["reasoning"][:2000], Ansi.GRAY))
|
||
history_calls = [{"id": call["id"], "type": "function", "function": {"name": call["name"], "arguments": call["arguments"]}} for call in reply["tool_calls"]]
|
||
self.messages.append({"role": "assistant", "content": reply["content"] or None, "tool_calls": history_calls or None})
|
||
if not reply["tool_calls"]:
|
||
if reply["content"]:
|
||
last_text = reply["content"]
|
||
elif reply["reasoning"]:
|
||
last_text = reply["reasoning"]
|
||
if last_text:
|
||
self.store.log_event(self.profile, "assistant", "message", last_text)
|
||
break
|
||
aborted = False
|
||
step_pairs = []
|
||
for pos, call in enumerate(reply["tool_calls"]):
|
||
self.show_call(call)
|
||
call_started = time.time()
|
||
try:
|
||
result = self.tools.dispatch(call["name"], call["arguments"])
|
||
except Denied as denied:
|
||
result = "denied by user"
|
||
self.show_result(result, (time.time() - call_started) * 1000)
|
||
self.messages.append({"role": "tool", "tool_call_id": call["id"], "content": result})
|
||
self.store.log_event(self.profile, "tool", call["name"], result)
|
||
step_pairs.append((action_signature(call["name"], call["arguments"]), result_digest(result)))
|
||
for skipped in reply["tool_calls"][pos + 1:]:
|
||
self.messages.append({"role": "tool", "tool_call_id": skipped["id"], "content": "skipped: stopped after denial"})
|
||
if denied.guidance:
|
||
self.messages.append({"role": "user", "content": denied.guidance})
|
||
self.store.log_event(self.profile, "user", "guidance", self.store.redact(denied.guidance, self.profile))
|
||
else:
|
||
aborted = True
|
||
last_text = "stopped by user"
|
||
break
|
||
self.show_result(result, (time.time() - call_started) * 1000)
|
||
self.messages.append({"role": "tool", "tool_call_id": call["id"], "content": result})
|
||
self.store.log_event(self.profile, "tool", call["name"], result)
|
||
step_pairs.append((action_signature(call["name"], call["arguments"]), result_digest(result)))
|
||
if call["name"] == "store_secret":
|
||
self.scrub_secret_from_history()
|
||
if aborted:
|
||
break
|
||
productive = self.user_stepped_in
|
||
self.user_stepped_in = False
|
||
for pair in step_pairs:
|
||
if pair[0] not in seen_actions or pair[1] not in seen_results:
|
||
productive = True
|
||
seen_actions.append(pair[0])
|
||
seen_results.append(pair[1])
|
||
if pair == last_pair:
|
||
run_len += 1
|
||
else:
|
||
last_pair = pair
|
||
run_len = 1
|
||
nudged = False
|
||
self._turn_steps += 1
|
||
self._turn_tools += len(step_pairs)
|
||
self.checkpoint()
|
||
if productive:
|
||
stall = 0
|
||
else:
|
||
stall += 1
|
||
if run_len == LOOP_NUDGE_AT and not nudged:
|
||
nudged = True
|
||
nudge = "loop warning: you called %s %d times in a row with identical results. Change approach, try different arguments, or stop instead of repeating it." % (last_pair[0][:200], run_len)
|
||
self.messages.append({"role": "user", "content": nudge})
|
||
self.store.log_event(self.profile, "user", "guidance", nudge)
|
||
if loud:
|
||
print(paint(nudge, Ansi.YELLOW))
|
||
if run_len >= LOOP_STOP_AT:
|
||
last_text = "stuck in a loop: %s repeated %d times with identical results" % (last_pair[0][:200], run_len)
|
||
if loud:
|
||
print(paint(last_text, Ansi.YELLOW))
|
||
self.store.log_event(self.profile, "system", "stall", last_text)
|
||
break
|
||
if stall >= max_steps:
|
||
last_text = "no progress for %d steps: actions and results keep repeating. Say 'continue' to resume with fresh direction." % max_steps
|
||
if loud:
|
||
print(paint(last_text, Ansi.YELLOW))
|
||
self.store.log_event(self.profile, "system", "stall", last_text)
|
||
break
|
||
else:
|
||
if loud:
|
||
print(paint("step budget exhausted (%d total steps)" % TOTAL_STEP_CAP, Ansi.YELLOW))
|
||
self.close_turn()
|
||
total = time.time() - started
|
||
if loud:
|
||
print(paint("tokens≈%d · %s · %.1fs" % (conversation_tokens(self.messages), "+".join(backends), total), Ansi.DIM))
|
||
return last_text
|
||
|
||
def run_mention(self, target, rest, original, capture=False, max_steps=MAX_STEPS):
|
||
if self.persist:
|
||
self.store.save_session(self.profile, self.messages, self.bot)
|
||
other = self.store.load_session(self.profile, target)
|
||
else:
|
||
other = []
|
||
system = self.store.load_bot_system(self.profile, target) or DEFAULT_SYSTEM
|
||
saved_messages, saved_system, saved_bot = self.messages, self.system_message, self.bot
|
||
self.messages = [{"role": "system", "content": system}] + other
|
||
self.system_message = system
|
||
self.bot = target
|
||
self.apply_system()
|
||
self.sync_chat_scope()
|
||
try:
|
||
answer = self.run_turn(rest, capture=capture, max_steps=max_steps, _routed=True)
|
||
finally:
|
||
if self.persist:
|
||
self.store.save_session(self.profile, self.messages, target)
|
||
self.messages, self.system_message, self.bot = saved_messages, saved_system, saved_bot
|
||
self.apply_system()
|
||
self.sync_chat_scope()
|
||
self.store.log_event(self.profile, "user", "message", self.store.redact(original, self.profile))
|
||
self.messages.append({"role": "user", "content": original})
|
||
self.messages.append({"role": "assistant", "content": answer, "reasoning": "", "tool_calls": [], "backend": "mention"})
|
||
if self.persist:
|
||
self.store.save_session(self.profile, self.messages, self.bot)
|
||
return answer
|
||
|
||
|
||
def default_runner(config, seal, profile, depth, task, timeout, max_steps=WORKER_STEPS):
|
||
store = Store(config, seal)
|
||
try:
|
||
worker = Agent(config, store, persist=False, quiet=True, depth=depth)
|
||
if profile != worker.profile:
|
||
worker.switch_profile(profile, silent=True)
|
||
worker.deadline = time.time() + timeout
|
||
return worker.run_turn(task, capture=True, max_steps=max_steps)
|
||
finally:
|
||
store.close()
|
||
|
||
|
||
def spawn_agent(task, profile, timeout, config, seal, depth=0, runner=None):
|
||
with AGENTS_LOCK:
|
||
agent_id = AGENTS_NEXT[0]
|
||
AGENTS_NEXT[0] += 1
|
||
record = {"id": agent_id, "task": task, "profile": profile, "status": "running", "result": "", "started": time.time(), "ended": None, "thread": None, "spilled": None}
|
||
AGENTS[agent_id] = record
|
||
|
||
def target():
|
||
try:
|
||
if runner is None:
|
||
result = default_runner(config, seal, profile, depth, task, timeout)
|
||
else:
|
||
result = runner(task, profile, timeout)
|
||
status = "timeout" if "[time limit reached]" in result else "done"
|
||
except Exception as exc:
|
||
result = "worker failed: " + short_error(exc)
|
||
status = "error"
|
||
with AGENTS_LOCK:
|
||
record["result"] = result
|
||
record["status"] = status
|
||
record["ended"] = time.time()
|
||
|
||
worker_thread = threading.Thread(target=target, daemon=True)
|
||
with AGENTS_LOCK:
|
||
record["thread"] = worker_thread
|
||
worker_thread.start()
|
||
return agent_id
|
||
|
||
|
||
def poll_agent(agent_id, wait=0, profile=None):
|
||
with AGENTS_LOCK:
|
||
record = AGENTS.get(agent_id)
|
||
if profile is not None and record is not None and record.get("profile") != profile:
|
||
record = None
|
||
worker_thread = record["thread"] if record else None
|
||
if record is None:
|
||
return "missing", "no agent %d" % agent_id
|
||
if worker_thread is not None:
|
||
worker_thread.join(timeout=max(0, wait))
|
||
with AGENTS_LOCK:
|
||
status = record["status"]
|
||
result = record["result"]
|
||
elapsed = (record["ended"] or time.time()) - record["started"]
|
||
if status == "running":
|
||
return status, "agent %d still running (%ds elapsed)" % (agent_id, elapsed)
|
||
if status == "error":
|
||
return status, result
|
||
return status, truncate(result, 6000, 2000) or "(empty result)"
|
||
|
||
|
||
def list_agents(profile=None):
|
||
with AGENTS_LOCK:
|
||
return [{"id": record["id"], "task": record["task"], "status": record["status"], "elapsed": int((record["ended"] or time.time()) - record["started"])} for record in AGENTS.values() if profile is None or record.get("profile") == profile]
|
||
|
||
|
||
def clear_agents(profile=None):
|
||
with AGENTS_LOCK:
|
||
finished = [key for key, record in AGENTS.items() if record["status"] != "running" and (profile is None or record.get("profile") == profile)]
|
||
for key in finished:
|
||
del AGENTS[key]
|
||
return len(finished)
|
||
|
||
|
||
def running_agents():
|
||
with AGENTS_LOCK:
|
||
return sum(1 for record in AGENTS.values() if record["status"] == "running")
|
||
|
||
|
||
def parse_schedule_at(raw):
|
||
text = str(raw or "").strip()
|
||
if not text:
|
||
raise ValueError("empty datetime")
|
||
if text[-1:] in ("Z", "z"):
|
||
text = text[:-1] + "+00:00"
|
||
try:
|
||
moment = _fromisoformat(text)
|
||
except ValueError:
|
||
raise ValueError("invalid datetime, use ISO like 2026-10-08T09:00")
|
||
return moment.astimezone(timezone.utc).isoformat()
|
||
|
||
|
||
def parse_duration(raw):
|
||
match = re.fullmatch(r"(\d+)\s*([smhdw])?", str(raw or "").strip().lower())
|
||
if not match:
|
||
return None
|
||
return int(match.group(1)) * {"s": 1, "m": 60, "h": 3600, "d": 86400, "w": 604800}[match.group(2) or "s"]
|
||
|
||
|
||
def format_delay(seconds):
|
||
seconds = int(seconds)
|
||
if seconds < 0:
|
||
return "overdue"
|
||
if seconds < 60:
|
||
return "%ds" % seconds
|
||
if seconds < 3600:
|
||
return "%dm" % (seconds // 60)
|
||
if seconds < 86400:
|
||
return "%dh" % (seconds // 3600)
|
||
return "%dd" % (seconds // 86400)
|
||
|
||
|
||
def local_display(utc_iso):
|
||
try:
|
||
return _fromisoformat(utc_iso).astimezone().strftime("%Y-%m-%d %H:%M")
|
||
except ValueError:
|
||
return utc_iso
|
||
|
||
|
||
def format_schedule_lines(items, now):
|
||
if not items:
|
||
return "(no schedules)"
|
||
lines = []
|
||
for item in items:
|
||
try:
|
||
due_in = (_fromisoformat(item["next_run"]) - now).total_seconds()
|
||
except ValueError:
|
||
due_in = -1
|
||
when = "every %s" % format_delay(item["every"]) if item["every"] else "at %s" % local_display(item["next_run"])
|
||
marked = " [%s]" % " ".join(item["tags"]) if item["tags"] else ""
|
||
lines.append("#%d [%s] %s (%s) %s%s :: %s" % (item["id"], item["status"], when, format_delay(due_in), item["name"] or "(no name)", marked, item["prompt"][:80]))
|
||
if item["last"]:
|
||
extra = " :: %s" % item["result"][:120] if item["result"] else ""
|
||
lines.append(" last: %s%s" % (item["last"], extra))
|
||
return "\n".join(lines)
|
||
|
||
|
||
def note_schedule_result(db_path, row_id, status, result=""):
|
||
try:
|
||
db = sqlite3.connect(db_path, timeout=30)
|
||
db.execute("PRAGMA journal_mode=WAL")
|
||
db.execute("PRAGMA busy_timeout=30000")
|
||
try:
|
||
db.execute("UPDATE schedules SET last_status = ?, last_result = ?, updated = ? WHERE id = ?", (status, result[:2000], datetime.now(timezone.utc).isoformat(), row_id))
|
||
db.commit()
|
||
finally:
|
||
db.close()
|
||
except sqlite3.Error:
|
||
pass
|
||
|
||
|
||
def watch_schedule(db_path, row_id, agent_id, timeout):
|
||
status, text = poll_agent(agent_id, wait=timeout + 30)
|
||
note_schedule_result(db_path, row_id, status, text)
|
||
|
||
|
||
def scheduler_tick(config, seal, runner=None):
|
||
fired = 0
|
||
now = datetime.now(timezone.utc)
|
||
db = sqlite3.connect(config.db_path, timeout=30)
|
||
db.execute("PRAGMA journal_mode=WAL")
|
||
db.execute("PRAGMA busy_timeout=30000")
|
||
try:
|
||
rows = db.execute("SELECT id, prompt, profile, every_sec, next_run, timeout FROM schedules WHERE status = 'pending' AND next_run <= ? ORDER BY next_run", (now.isoformat(),)).fetchall()
|
||
for row_id, prompt, profile, every_sec, next_run, timeout in rows:
|
||
stamp = datetime.now(timezone.utc).isoformat()
|
||
if every_sec:
|
||
try:
|
||
claimed_next = _fromisoformat(next_run)
|
||
except ValueError:
|
||
continue
|
||
while claimed_next <= datetime.now(timezone.utc):
|
||
claimed_next += timedelta(seconds=every_sec)
|
||
done = db.execute("UPDATE schedules SET next_run = ?, updated = ? WHERE id = ? AND status = 'pending' AND next_run = ?", (claimed_next.isoformat(), stamp, row_id, next_run))
|
||
else:
|
||
done = db.execute("UPDATE schedules SET status = 'done', updated = ? WHERE id = ? AND status = 'pending'", (stamp, row_id))
|
||
db.commit()
|
||
if done.rowcount != 1:
|
||
continue
|
||
if prompt.startswith(Seal.PREFIX):
|
||
try:
|
||
prompt = seal.unlock(prompt)
|
||
except SealError as exc:
|
||
note_schedule_result(config.db_path, row_id, "error: " + short_error(exc))
|
||
continue
|
||
agent_id = spawn_agent(prompt, profile, timeout, config, seal, 0, runner)
|
||
watcher = threading.Thread(target=watch_schedule, args=(config.db_path, row_id, agent_id, timeout), daemon=True)
|
||
watcher.start()
|
||
fired += 1
|
||
finally:
|
||
db.close()
|
||
return fired
|
||
|
||
|
||
def run_scheduler_loop(config, seal, stop_flag):
|
||
while not stop_flag.is_set():
|
||
try:
|
||
scheduler_tick(config, seal)
|
||
except Exception as exc:
|
||
print("scheduler error: %s" % short_error(exc))
|
||
stop_flag.wait(SCHEDULER_INTERVAL)
|
||
|
||
|
||
def start_scheduler(config, seal):
|
||
stop_flag = threading.Event()
|
||
worker = threading.Thread(target=run_scheduler_loop, args=(config, seal, stop_flag), daemon=True)
|
||
worker.start()
|
||
return stop_flag
|
||
|
||
|
||
COMMANDS = ("help", "profile", "profiles", "bots", "bot", "env", "skills", "sysinfo", "models", "secret", "schedule", "schedules", "unschedule", "records", "record", "graph", "tags", "tools", "install", "search", "audit", "restore", "release", "fork", "agents", "agent", "compact", "clear", "quit", "exit")
|
||
|
||
|
||
def complete_command(text, state):
|
||
options = ["/" + name for name in COMMANDS if name.startswith(text.lstrip("/"))]
|
||
return options[state] if state < len(options) else None
|
||
|
||
|
||
def handle_command(agent, text):
|
||
parts = text[1:].split(None, 1)
|
||
name = parts[0].lower()
|
||
arg = parts[1].strip() if len(parts) > 1 else ""
|
||
if name in ("quit", "exit"):
|
||
return False
|
||
if name == "help":
|
||
print(HELP_TEXT)
|
||
elif name == "profiles":
|
||
names = agent.store.list_profiles()
|
||
for item in names:
|
||
print("%s %s" % ("*" if item == agent.profile else " ", item))
|
||
if not names:
|
||
print("(no profiles)")
|
||
elif name == "bots":
|
||
for item in agent.store.list_bots(agent.profile):
|
||
nicks = " aka %s" % ", ".join(item["nicknames"]) if item["nicknames"] else ""
|
||
print("%s %s%s" % ("*" if item["name"] == agent.bot else " ", item["name"], nicks))
|
||
elif name == "bot":
|
||
if not arg:
|
||
print(paint("use /bot <name>, see /bots", Ansi.RED))
|
||
else:
|
||
print(agent.switch_bot(arg))
|
||
elif name == "profile":
|
||
if not arg:
|
||
print("current profile: %s (%d chars)" % (agent.profile, len(agent.system_message)))
|
||
elif not PROFILE_RE.match(arg):
|
||
print(paint("invalid profile name", Ansi.RED))
|
||
else:
|
||
agent.switch_profile(arg)
|
||
elif name == "env":
|
||
if not arg:
|
||
print("current environment: %s" % agent.env)
|
||
elif arg not in ("home", "sandbox"):
|
||
print(paint("use /env home or /env sandbox", Ansi.RED))
|
||
elif arg == "sandbox":
|
||
try:
|
||
ensure_box()
|
||
except BackendError as exc:
|
||
print(paint("sandbox unavailable: %s" % exc, Ansi.RED))
|
||
else:
|
||
agent.env = "sandbox"
|
||
print("environment: sandbox (isolated container)")
|
||
else:
|
||
agent.env = "home"
|
||
print("environment: home")
|
||
elif name == "skills":
|
||
bits = arg.split()
|
||
if bits and bits[0] == "install":
|
||
if len(bits) < 2:
|
||
print(paint("use /skills install <source> [name=x] [scope=project|home] [mode=copy|symlink] [force]", Ansi.RED))
|
||
return True
|
||
call_args = {"source": bits[1]}
|
||
for token in bits[2:]:
|
||
if token == "force":
|
||
call_args["force"] = True
|
||
elif "=" in token:
|
||
key, _, value = token.partition("=")
|
||
if key in ("name", "scope", "mode"):
|
||
call_args[key] = value
|
||
print(agent.tools.dispatch("install_skill", json.dumps(call_args)))
|
||
return True
|
||
agent.skills = discover_skills(agent.config.home, os.getcwd())
|
||
agent.apply_system()
|
||
if not agent.skills:
|
||
print("(no skills)")
|
||
for skill_name in sorted(agent.skills):
|
||
skill = agent.skills[skill_name]
|
||
print("- %s: %s %s" % (skill_name, skill["description"][:200], skill_provenance_note(skill)))
|
||
wanted = sorted(item for item in SKILL_BLUEPRINTS if item not in agent.skills)
|
||
if wanted:
|
||
print("blueprints, built on first load:")
|
||
for skill_name in wanted:
|
||
print("- %s: %s" % (skill_name, SKILL_BLUEPRINTS[skill_name]["hint"]))
|
||
elif name == "sysinfo":
|
||
print(collect_sysinfo())
|
||
elif name == "models":
|
||
print(models_report(agent.chat))
|
||
elif name == "secret":
|
||
sub = arg.split(None, 1)
|
||
action = sub[0].lower() if sub else ""
|
||
rest = sub[1].strip() if len(sub) > 1 else ""
|
||
if action == "set" and rest:
|
||
if not SECRET_RE.match(rest):
|
||
print(paint("invalid secret name, use letters, digits, dot, underscore, hyphen", Ansi.RED))
|
||
else:
|
||
try:
|
||
value = getpass.getpass("value for '%s': " % rest)
|
||
except (EOFError, KeyboardInterrupt):
|
||
print("\naborted")
|
||
return True
|
||
if not value:
|
||
print("empty value, aborted")
|
||
else:
|
||
agent.store.save_secret(rest, value, None, None, agent.profile)
|
||
print("stored secret '%s'" % rest)
|
||
elif action in ("list", "") and not rest:
|
||
print(agent.tools.dispatch("list_secrets", "{}"))
|
||
elif action == "delete" and rest:
|
||
if agent.store.load_secret(rest, agent.profile) is None:
|
||
print(paint("unknown secret '%s'" % rest, Ansi.RED))
|
||
elif agent.ask_approval("delete secret '%s'" % rest, guidance=False):
|
||
agent.store.delete_secret(rest, agent.profile)
|
||
print("deleted secret '%s'" % rest)
|
||
else:
|
||
print("aborted")
|
||
else:
|
||
print(paint("use /secret set <name> or /secret list or /secret delete <name>", Ansi.RED))
|
||
elif name == "schedule":
|
||
parts = arg.split(None, 2)
|
||
if len(parts) != 3 or parts[0] not in ("at", "every"):
|
||
print(paint("use /schedule at <ISO datetime> <prompt> or /schedule every <30s|10m|2h|1d> <prompt>", Ansi.RED))
|
||
else:
|
||
mode, when, prompt = parts
|
||
if mode == "at":
|
||
try:
|
||
next_run = parse_schedule_at(when)
|
||
except ValueError as exc:
|
||
print(paint("error: %s" % exc, Ansi.RED))
|
||
return True
|
||
every_sec = 0
|
||
else:
|
||
every_sec = parse_duration(when)
|
||
if every_sec is None or every_sec < 60:
|
||
print(paint("error: every must be a duration of at least 60s", Ansi.RED))
|
||
return True
|
||
next_run = (datetime.now(timezone.utc) + timedelta(seconds=every_sec)).isoformat()
|
||
row_id = agent.store.add_schedule("", prompt, agent.profile, every_sec, next_run, 600)
|
||
print("scheduled #%d" % row_id)
|
||
elif name == "schedules":
|
||
print(format_schedule_lines(agent.store.list_schedules(agent.profile), datetime.now(timezone.utc)))
|
||
elif name == "unschedule":
|
||
try:
|
||
row_id = int(arg or "0")
|
||
except ValueError:
|
||
print(paint("use /unschedule <id>", Ansi.RED))
|
||
return True
|
||
if agent.store.remove_schedule(row_id, agent.profile):
|
||
print("deleted schedule #%d" % row_id)
|
||
else:
|
||
print(paint("no schedule #%d" % row_id, Ansi.RED))
|
||
elif name == "records":
|
||
print(agent.tools.dispatch("record_search", json.dumps({"query": arg})))
|
||
elif name == "record":
|
||
bits = arg.split()
|
||
if not bits:
|
||
print(paint("use /record <mem:id> [offset]", Ansi.RED))
|
||
else:
|
||
try:
|
||
offset = int(bits[1]) if len(bits) > 1 else 0
|
||
except ValueError:
|
||
print(paint("offset must be a number", Ansi.RED))
|
||
return True
|
||
print(agent.tools.dispatch("record_read", json.dumps({"id": bits[0], "offset": offset})))
|
||
elif name == "graph":
|
||
if not arg:
|
||
print(paint("use /graph <mem:id|secret:name|sched:id>", Ansi.RED))
|
||
else:
|
||
print(agent.tools.dispatch("graph_query", json.dumps({"node": arg})))
|
||
elif name == "tools":
|
||
if arg:
|
||
match = [schema["function"] for schema in TOOL_SCHEMAS if schema["function"]["name"] == arg]
|
||
if not match:
|
||
print(paint("unknown tool '%s'" % arg, Ansi.RED))
|
||
return True
|
||
print("%s [%s]" % (match[0]["name"], "core" if match[0]["name"] in CORE_TOOLS else "lazy"))
|
||
print(match[0]["description"])
|
||
print("trigger tags: %s" % ", ".join(TOOL_TAGS.get(match[0]["name"], ())))
|
||
return True
|
||
names = [schema["function"]["name"] for schema in TOOL_SCHEMAS if tool_env_ok(schema["function"]["name"])]
|
||
cores = sorted(tool for tool in names if tool in CORE_TOOLS)
|
||
lazy = sorted(tool for tool in names if tool not in CORE_TOOLS)
|
||
print("core, always loaded:\n %s" % "\n ".join(cores))
|
||
print("lazy, loads when named:\n %s" % "\n ".join(lazy))
|
||
elif name == "search":
|
||
if not arg:
|
||
print(paint("use /search <query>", Ansi.RED))
|
||
else:
|
||
print(agent.tools.dispatch("search", json.dumps({"query": arg})))
|
||
elif name == "install":
|
||
bits = arg.split()
|
||
action = bits[0] if bits else "status"
|
||
if action not in ("status", "install", "upgrade", "reinstall", "uninstall"):
|
||
print(paint("use /install [status|install|upgrade|reinstall|uninstall] [targets...]", Ansi.RED))
|
||
else:
|
||
print(agent.tools.dispatch("install", json.dumps({"action": action, "targets": bits[1:] or list(INSTALL_TARGETS)})))
|
||
elif name == "tags":
|
||
print(agent.tools.dispatch("tags", json.dumps({"prefix": arg} if arg else {})))
|
||
elif name == "audit":
|
||
print(agent.tools.dispatch("audit", json.dumps({"path": arg} if arg else {})))
|
||
elif name == "restore":
|
||
try:
|
||
row_id = int(arg or "0")
|
||
except ValueError:
|
||
print(paint("use /restore <audit id>", Ansi.RED))
|
||
return True
|
||
print(agent.tools.dispatch("restore", json.dumps({"id": row_id})))
|
||
agent.show_diff()
|
||
elif name == "release":
|
||
bits = arg.split(None, 1)
|
||
if len(bits) != 2 or bits[0] not in ("major", "minor", "patch"):
|
||
print(paint("use /release <major|minor|patch> <message>", Ansi.RED))
|
||
else:
|
||
print(agent.tools.dispatch("release", json.dumps({"part": bits[0], "message": bits[1]})))
|
||
elif name == "fork":
|
||
if not arg:
|
||
print(paint("use /fork <task>", Ansi.RED))
|
||
else:
|
||
agent_id = spawn_agent(arg, agent.profile, 600, agent.config, agent.store.seal, agent.depth + 1, agent.runner_override)
|
||
print("agent %d started, REPL stays free" % agent_id)
|
||
elif name == "agents":
|
||
records = list_agents(agent.profile)
|
||
if not records:
|
||
print("(no agents)")
|
||
for record in records:
|
||
print("#%d [%s] %ds %s" % (record["id"], record["status"], record["elapsed"], record["task"][:60]))
|
||
elif name == "agent":
|
||
if arg == "clear":
|
||
print("purged %d finished agents" % clear_agents(agent.profile))
|
||
else:
|
||
try:
|
||
agent_id = int(arg)
|
||
except ValueError:
|
||
print(paint("use /agent <id> or /agent clear", Ansi.RED))
|
||
return True
|
||
status, text = poll_agent(agent_id, 0, agent.profile)
|
||
print("[%s]\n%s" % (status, text))
|
||
elif name == "compact":
|
||
agent.messages, changed = compact_messages(agent.messages, agent.chat)
|
||
agent.apply_system()
|
||
print("compacted" if changed else "nothing to compact")
|
||
elif name == "clear":
|
||
agent.reset_history()
|
||
print("history cleared")
|
||
else:
|
||
print(paint("unknown command, try /help", Ansi.RED))
|
||
return True
|
||
|
||
|
||
class Telegram:
|
||
def __init__(self, token):
|
||
self.token = token
|
||
|
||
def call(self, method, params, timeout=70):
|
||
body = json.dumps(params).encode("utf-8")
|
||
request = urllib.request.Request(TELEGRAM_API + self.token + "/" + method, data=body, headers={"Content-Type": "application/json"}, method="POST")
|
||
with urllib.request.urlopen(request, timeout=timeout) as response:
|
||
return json.loads(response.read().decode("utf-8", "replace"))
|
||
|
||
|
||
def telegram_env_path(home):
|
||
return os.path.join(home, "telegram.env")
|
||
|
||
|
||
def load_telegram_token(home):
|
||
token = os.environ.get("TELEGRAM_BOT_TOKEN") or ""
|
||
if token:
|
||
return token
|
||
path = telegram_env_path(home)
|
||
if not os.path.exists(path):
|
||
return ""
|
||
with open(path, "r", encoding="utf-8") as handle:
|
||
for line in handle.read().splitlines():
|
||
if line.startswith("TELEGRAM_BOT_TOKEN="):
|
||
return line.split("=", 1)[1].strip().strip("'\"")
|
||
return ""
|
||
|
||
|
||
def telegram_send(bot, chat_id, text):
|
||
for pos in range(0, len(text), 4000):
|
||
bot.call("sendMessage", {"chat_id": chat_id, "text": text[pos:pos + 4000] or "(empty reply)"})
|
||
|
||
|
||
def telegram_transcribe(bot, file_id):
|
||
try:
|
||
info = bot.call("getFile", {"file_id": file_id})
|
||
remote = (info.get("result") or {}).get("file_path", "")
|
||
if not remote:
|
||
return "error: telegram returned no file"
|
||
with urllib.request.urlopen(TELEGRAM_FILE_API + bot.token + "/" + remote, timeout=60) as response:
|
||
audio = response.read(20000000)
|
||
engine = ensure_box()
|
||
return box_transcribe(engine, audio)
|
||
except (OSError, ValueError, BackendError) as exc:
|
||
return "error: transcription failed: " + short_error(exc)
|
||
|
||
|
||
def handle_telegram_update(agent, bot, update):
|
||
message = update.get("message") or {}
|
||
chat_id = (message.get("chat") or {}).get("id")
|
||
if not chat_id:
|
||
return
|
||
text = message.get("text") or ""
|
||
if text == "/start":
|
||
telegram_send(bot, chat_id, "tai online. Send any message, /new for a fresh start.")
|
||
return
|
||
if text == "/new":
|
||
agent.reset_history()
|
||
telegram_send(bot, chat_id, "fresh start.")
|
||
return
|
||
voice = message.get("voice") or message.get("audio")
|
||
if voice:
|
||
bot.call("sendChatAction", {"chat_id": chat_id, "action": "typing"})
|
||
text = telegram_transcribe(bot, voice.get("file_id", ""))
|
||
if text.startswith("error"):
|
||
telegram_send(bot, chat_id, text)
|
||
return
|
||
if not text.strip():
|
||
return
|
||
bot.call("sendChatAction", {"chat_id": chat_id, "action": "typing"})
|
||
reply = agent.run_turn(text, capture=True)
|
||
telegram_send(bot, chat_id, reply)
|
||
|
||
|
||
def run_telegram_bot(agent):
|
||
token = load_telegram_token(agent.config.home)
|
||
if not token:
|
||
print("missing TELEGRAM_BOT_TOKEN", file=sys.stderr)
|
||
return 2
|
||
bot = Telegram(token)
|
||
try:
|
||
me = bot.call("getMe", {})
|
||
except (OSError, ValueError) as exc:
|
||
print("telegram auth failed: %s" % short_error(exc), file=sys.stderr)
|
||
return 2
|
||
print("telegram bot online as @%s" % (me.get("result") or {}).get("username", "?"))
|
||
start_scheduler(agent.config, agent.store.seal)
|
||
offset = 0
|
||
while True:
|
||
try:
|
||
data = bot.call("getUpdates", {"offset": offset, "timeout": 50})
|
||
except (OSError, ValueError) as exc:
|
||
print("poll error: %s, retrying" % short_error(exc))
|
||
time.sleep(5)
|
||
continue
|
||
for update in data.get("result") or []:
|
||
offset = update.get("update_id", offset) + 1
|
||
try:
|
||
handle_telegram_update(agent, bot, update)
|
||
except Exception as exc:
|
||
print("update error: %s" % short_error(exc))
|
||
|
||
|
||
def run_host(argv, timeout=120):
|
||
try:
|
||
done = subprocess.run(argv, stdout=subprocess.PIPE, stderr=subprocess.PIPE, universal_newlines=True, timeout=timeout)
|
||
return done.returncode == 0, (done.stdout or "") + (done.stderr or "")
|
||
except (OSError, subprocess.SubprocessError) as exc:
|
||
return False, short_error(exc)
|
||
|
||
|
||
def venv_available():
|
||
if importlib.util.find_spec("venv") is None:
|
||
return False
|
||
if importlib.util.find_spec("ensurepip") is not None:
|
||
return True
|
||
return shutil.which("pip3") is not None or shutil.which("pip") is not None
|
||
|
||
|
||
def venv_dir(home):
|
||
return os.path.join(home, "venv")
|
||
|
||
|
||
def venv_python(home):
|
||
cand = os.path.join(venv_dir(home), "bin", "python")
|
||
if os.path.isfile(cand) and os.access(cand, os.X_OK):
|
||
return cand
|
||
return None
|
||
|
||
|
||
def ensure_venv(home):
|
||
found = venv_python(home)
|
||
if found:
|
||
return True, "kept %s" % found
|
||
if not venv_available():
|
||
return False, "python venv module unavailable"
|
||
try:
|
||
done = subprocess.run([sys.executable, "-m", "venv", venv_dir(home)], stdout=subprocess.PIPE, stderr=subprocess.PIPE, universal_newlines=True, timeout=300)
|
||
except (OSError, subprocess.SubprocessError) as exc:
|
||
return False, short_error(exc)
|
||
if done.returncode != 0 or not venv_python(home):
|
||
return False, ((done.stderr or "").strip()[:200] or "venv creation failed")
|
||
return True, "created %s" % venv_python(home)
|
||
|
||
|
||
def service_exec(home, script, *args):
|
||
parts = []
|
||
python = venv_python(home)
|
||
if python:
|
||
parts.append(python)
|
||
parts.append(script)
|
||
parts.extend(args)
|
||
return " ".join(parts)
|
||
|
||
|
||
def unit_dir(home):
|
||
return os.path.join(os.path.expanduser("~"), ".config", "systemd", "user")
|
||
|
||
|
||
def service_state(name):
|
||
if not shutil.which("systemctl"):
|
||
return "unknown", "unknown"
|
||
active, _out = run_host(["systemctl", "--user", "is-active", name], timeout=15)
|
||
enabled, _out = run_host(["systemctl", "--user", "is-enabled", name], timeout=15)
|
||
return ("active" if active else "inactive"), ("enabled" if enabled else "disabled")
|
||
|
||
|
||
def activate_service(name):
|
||
if not shutil.which("systemctl"):
|
||
return "systemctl missing, unit written only"
|
||
run_host(["systemctl", "--user", "daemon-reload"])
|
||
active, _enabled = service_state(name)
|
||
if active == "active":
|
||
ok, out = run_host(["systemctl", "--user", "restart", name])
|
||
return "restarted, changes effective immediately" if ok else "restart failed: " + out.strip()[:150]
|
||
ok, out = run_host(["systemctl", "--user", "enable", "--now", name])
|
||
if ok:
|
||
user = os.environ.get("USER") or getpass.getuser()
|
||
run_host(["loginctl", "enable-linger", user])
|
||
return "enabled and started"
|
||
return "enable failed: " + out.strip()[:150]
|
||
|
||
|
||
def remove_service_unit(home, name):
|
||
path = os.path.join(unit_dir(home), name)
|
||
if shutil.which("systemctl"):
|
||
run_host(["systemctl", "--user", "disable", "--now", name])
|
||
run_host(["systemctl", "--user", "daemon-reload"])
|
||
if os.path.exists(path):
|
||
try:
|
||
os.remove(path)
|
||
return True, "removed " + path
|
||
except OSError as exc:
|
||
return False, short_error(exc)
|
||
return True, "not present"
|
||
|
||
|
||
def script_version(path):
|
||
try:
|
||
with open(path, "r", encoding="utf-8") as handle:
|
||
match = VERSION_RE.search(handle.read())
|
||
except OSError:
|
||
return None
|
||
return ".".join(match.groups()) if match else None
|
||
|
||
|
||
def hook_present(path):
|
||
try:
|
||
with open(path, "r", encoding="utf-8") as handle:
|
||
return BASHRC_BLOCK in handle.read()
|
||
except OSError:
|
||
return False
|
||
|
||
|
||
def remove_bashrc_block(path):
|
||
try:
|
||
with open(path, "r", encoding="utf-8") as handle:
|
||
raw = handle.read()
|
||
except OSError:
|
||
return False
|
||
stripped = re.sub(re.escape(BASHRC_MARK_BEGIN) + r".*?" + re.escape(BASHRC_MARK_END) + r"\n?", "", raw, flags=re.DOTALL)
|
||
if stripped == raw:
|
||
return False
|
||
backup = path + ".bak-tai"
|
||
if raw and not os.path.exists(backup):
|
||
with open(backup, "w", encoding="utf-8") as handle:
|
||
handle.write(raw)
|
||
with open(path, "w", encoding="utf-8") as handle:
|
||
handle.write(stripped)
|
||
return True
|
||
|
||
|
||
def box_venv_state(engine):
|
||
try:
|
||
probe = box_exec(engine, ["test", "-x", BOX_PYTHON], timeout=30)
|
||
except (OSError, subprocess.SubprocessError):
|
||
return "unknown"
|
||
return "venv ready" if probe.returncode == 0 else "no venv (reinstall container to rebuild)"
|
||
|
||
|
||
def status_lines(home, store=None):
|
||
lines = []
|
||
binary = os.path.join(os.path.expanduser("~"), ".local", "bin", "tai.py")
|
||
if os.path.isfile(binary):
|
||
lines.append("binary: installed %s (%s, running %s)" % (binary, script_version(binary) or "unknown", VERSION))
|
||
else:
|
||
lines.append("binary: not installed")
|
||
lines.append("bash-hook: %s in ~/.bashrc" % ("installed" if hook_present(os.path.join(os.path.expanduser("~"), ".bashrc")) else "not installed"))
|
||
python = venv_python(home)
|
||
if python:
|
||
lines.append("venv: ready %s" % python)
|
||
elif venv_available():
|
||
lines.append("venv: not created")
|
||
else:
|
||
lines.append("venv: unavailable (no venv module)")
|
||
for name, label in (("tai-scheduler.service", "scheduler-service"), ("tai-telegram.service", "telegram-service")):
|
||
unit = os.path.join(unit_dir(home), name)
|
||
if not os.path.isfile(unit):
|
||
lines.append("%s: not installed" % label)
|
||
continue
|
||
active, enabled = service_state(name)
|
||
lines.append("%s: unit present, %s, %s" % (label, active, enabled))
|
||
engine = container_engine()
|
||
if not engine:
|
||
lines.append("container: no engine (install podman)")
|
||
else:
|
||
state = box_state(engine)
|
||
if state == "running":
|
||
lines.append("container: running, %s" % box_venv_state(engine))
|
||
else:
|
||
lines.append("container: %s" % state)
|
||
lines.append(vault_line(home, store))
|
||
return lines
|
||
|
||
|
||
def vault_line(home, store=None):
|
||
db_path = os.path.join(home, "memory.db")
|
||
if not os.path.isfile(db_path):
|
||
return "vault: empty, install ops never touch data"
|
||
if store is None:
|
||
return "vault: kept %s, install ops never touch data" % db_path
|
||
try:
|
||
secrets = store.db.execute("SELECT COUNT(*) FROM secrets").fetchone()[0]
|
||
schedules = store.db.execute("SELECT COUNT(*) FROM schedules").fetchone()[0]
|
||
records = store.db.execute("SELECT COUNT(*) FROM records").fetchone()[0]
|
||
except sqlite3.Error:
|
||
return "vault: kept %s, install ops never touch data" % db_path
|
||
return "vault: kept (%d secrets, %d schedules, %d records), install ops never touch data" % (secrets, schedules, records)
|
||
|
||
|
||
def op_binary(mode, home):
|
||
target = os.path.join(os.path.expanduser("~"), ".local", "bin", "tai.py")
|
||
src = os.path.abspath(__file__)
|
||
current = script_version(target)
|
||
if mode == "install" and current is not None:
|
||
return "binary: already installed %s (running %s)" % (current, VERSION)
|
||
if mode == "uninstall":
|
||
if not os.path.exists(target):
|
||
return "binary: not present"
|
||
try:
|
||
os.remove(target)
|
||
return "binary: removed %s" % target
|
||
except OSError as exc:
|
||
return "binary: remove failed: %s" % short_error(exc)
|
||
try:
|
||
os.makedirs(os.path.dirname(target), exist_ok=True)
|
||
with open(src, "rb") as handle:
|
||
blob = handle.read()
|
||
with open(target, "wb") as handle:
|
||
handle.write(blob)
|
||
os.chmod(target, 0o755)
|
||
except OSError as exc:
|
||
return "binary: copy failed: %s" % short_error(exc)
|
||
if current is None:
|
||
return "binary: installed %s" % VERSION
|
||
return "binary: refreshed %s -> %s" % (current, VERSION)
|
||
|
||
|
||
def op_hook(mode, home):
|
||
path = os.path.join(os.path.expanduser("~"), ".bashrc")
|
||
if mode == "uninstall":
|
||
return "bash-hook: %s" % ("removed" if remove_bashrc_block(path) else "not present")
|
||
if mode == "reinstall":
|
||
remove_bashrc_block(path)
|
||
try:
|
||
changed = upsert_bashrc_block(path)
|
||
except OSError as exc:
|
||
return "bash-hook: write failed: %s" % short_error(exc)
|
||
if changed:
|
||
return "bash-hook: installed, run: source ~/.bashrc"
|
||
return "bash-hook: already present"
|
||
|
||
|
||
def op_venv(mode, home):
|
||
if mode == "uninstall":
|
||
folder = venv_dir(home)
|
||
if not os.path.isdir(folder):
|
||
return "venv: not present"
|
||
try:
|
||
shutil.rmtree(folder)
|
||
return "venv: removed %s" % folder
|
||
except OSError as exc:
|
||
return "venv: remove failed: %s" % short_error(exc)
|
||
if mode == "reinstall" and os.path.isdir(venv_dir(home)):
|
||
try:
|
||
shutil.rmtree(venv_dir(home))
|
||
except OSError as exc:
|
||
return "venv: rebuild failed: %s" % short_error(exc)
|
||
ok, detail = ensure_venv(home)
|
||
return "venv: %s" % detail
|
||
|
||
|
||
def op_service(mode, home, name, label, body):
|
||
unit = os.path.join(unit_dir(home), name)
|
||
if mode == "uninstall":
|
||
_done, detail = remove_service_unit(home, name)
|
||
return "%s: %s" % (label, detail)
|
||
existed = os.path.isfile(unit)
|
||
if mode == "install" and existed:
|
||
active, enabled = service_state(name)
|
||
return "%s: already installed (%s, %s)" % (label, active, enabled)
|
||
if mode == "reinstall" and shutil.which("systemctl"):
|
||
run_host(["systemctl", "--user", "disable", "--now", name])
|
||
try:
|
||
os.makedirs(os.path.dirname(unit), exist_ok=True)
|
||
with open(unit, "w", encoding="utf-8") as handle:
|
||
handle.write(body)
|
||
except OSError as exc:
|
||
return "%s: unit write failed: %s" % (label, short_error(exc))
|
||
return "%s: unit %s, %s" % (label, "rewritten" if existed else "written", activate_service(name))
|
||
|
||
|
||
def op_scheduler_service(mode, home):
|
||
body = SCHEDULER_UNIT % service_exec(home, os.path.abspath(__file__))
|
||
return op_service(mode, home, "tai-scheduler.service", "scheduler-service", body)
|
||
|
||
|
||
def op_telegram_service(mode, home):
|
||
if mode == "uninstall":
|
||
_done, detail = remove_service_unit(home, "tai-telegram.service")
|
||
return "telegram-service: %s" % detail
|
||
unit = os.path.join(unit_dir(home), "tai-telegram.service")
|
||
if mode == "install" and os.path.isfile(unit):
|
||
active, enabled = service_state("tai-telegram.service")
|
||
return "telegram-service: already installed (%s, %s)" % (active, enabled)
|
||
token = os.environ.get("TELEGRAM_BOT_TOKEN") or load_telegram_token(home)
|
||
if not token:
|
||
return "telegram-service: needs TELEGRAM_BOT_TOKEN in the environment"
|
||
try:
|
||
me = Telegram(token).call("getMe", {})
|
||
except (OSError, ValueError) as exc:
|
||
return "telegram-service: token rejected: %s" % short_error(exc)
|
||
env_path = telegram_env_path(home)
|
||
try:
|
||
with open(env_path, "w", encoding="utf-8") as handle:
|
||
handle.write("TELEGRAM_BOT_TOKEN=%s\n" % token)
|
||
os.chmod(env_path, 0o600)
|
||
except OSError as exc:
|
||
return "telegram-service: env write failed: %s" % short_error(exc)
|
||
body = TELEGRAM_UNIT % (service_exec(home, os.path.abspath(__file__)), env_path)
|
||
if mode == "reinstall" and shutil.which("systemctl"):
|
||
run_host(["systemctl", "--user", "disable", "--now", "tai-telegram.service"])
|
||
try:
|
||
os.makedirs(os.path.dirname(unit), exist_ok=True)
|
||
with open(unit, "w", encoding="utf-8") as handle:
|
||
handle.write(body)
|
||
except OSError as exc:
|
||
return "telegram-service: unit write failed: %s" % short_error(exc)
|
||
who = (me.get("result") or {}).get("username", "?")
|
||
return "telegram-service: @%s, %s" % (who, activate_service("tai-telegram.service"))
|
||
|
||
|
||
def op_container(mode, _home):
|
||
engine = container_engine()
|
||
if not engine:
|
||
return "container: no engine (install podman)"
|
||
if mode == "uninstall":
|
||
run_host([engine, "rm", "-f", BOX_NAME])
|
||
ok, _out = run_host([engine, "rmi", "-f", BOX_IMAGE])
|
||
return "container: removed%s" % (", image purged" if ok else "")
|
||
if mode == "reinstall":
|
||
run_host([engine, "rm", "-f", BOX_NAME])
|
||
run_host([engine, "rmi", "-f", BOX_IMAGE])
|
||
try:
|
||
engine = ensure_box()
|
||
except BackendError as exc:
|
||
return "container: failed: %s" % exc
|
||
state = box_state(engine)
|
||
detail = ", %s" % box_venv_state(engine) if state == "running" else ""
|
||
if mode == "reinstall":
|
||
return "container: rebuilt and %s%s" % (state, detail)
|
||
if mode == "install" and state == "running":
|
||
return "container: already running%s" % detail
|
||
return "container: %s%s" % (state, detail)
|
||
|
||
|
||
def install_report(action, targets, home, store=None):
|
||
if action == "status":
|
||
return "\n".join(status_lines(home, store))
|
||
lines = []
|
||
for target in targets:
|
||
if target == "binary":
|
||
lines.append(op_binary(action, home))
|
||
elif target == "bash-hook":
|
||
lines.append(op_hook(action, home))
|
||
elif target == "venv":
|
||
lines.append(op_venv(action, home))
|
||
elif target == "scheduler-service":
|
||
lines.append(op_scheduler_service(action, home))
|
||
elif target == "telegram-service":
|
||
lines.append(op_telegram_service(action, home))
|
||
elif target == "container":
|
||
lines.append(op_container(action, home))
|
||
lines.append(vault_line(home, store))
|
||
return "\n".join(lines)
|
||
|
||
|
||
def install_telegram():
|
||
token = os.environ.get("TELEGRAM_BOT_TOKEN") or ""
|
||
if token:
|
||
print("using TELEGRAM_BOT_TOKEN from environment")
|
||
else:
|
||
try:
|
||
token = getpass.getpass("Telegram bot token: ").strip()
|
||
except (EOFError, KeyboardInterrupt):
|
||
print("\naborted")
|
||
return 1
|
||
if not token:
|
||
print("empty token, aborted")
|
||
return 1
|
||
try:
|
||
me = Telegram(token).call("getMe", {})
|
||
except (OSError, ValueError) as exc:
|
||
print("token rejected: %s" % short_error(exc))
|
||
return 1
|
||
print("token ok for @%s" % (me.get("result") or {}).get("username", "?"))
|
||
home = os.environ.get("TAI_HOME") or os.path.join(os.path.expanduser("~"), ".tai")
|
||
os.makedirs(home, exist_ok=True)
|
||
env_path = telegram_env_path(home)
|
||
with open(env_path, "w", encoding="utf-8") as handle:
|
||
handle.write("TELEGRAM_BOT_TOKEN=%s\n" % token)
|
||
try:
|
||
os.chmod(env_path, 0o600)
|
||
except OSError:
|
||
pass
|
||
print("wrote %s" % env_path)
|
||
try:
|
||
ensure_box()
|
||
except BackendError as exc:
|
||
print("container failed: %s" % exc)
|
||
return 1
|
||
unit_folder = os.path.join(os.path.expanduser("~"), ".config", "systemd", "user")
|
||
os.makedirs(unit_folder, exist_ok=True)
|
||
unit_path = os.path.join(unit_folder, "tai-telegram.service")
|
||
with open(unit_path, "w", encoding="utf-8") as handle:
|
||
handle.write(TELEGRAM_UNIT % (service_exec(home, os.path.abspath(__file__)), env_path))
|
||
print("wrote %s" % unit_path)
|
||
print("service: %s" % activate_service("tai-telegram.service"))
|
||
return 0
|
||
|
||
|
||
def uninstall_telegram():
|
||
unit_path = os.path.join(os.path.expanduser("~"), ".config", "systemd", "user", "tai-telegram.service")
|
||
if shutil.which("systemctl"):
|
||
run_host(["systemctl", "--user", "disable", "--now", "tai-telegram.service"])
|
||
if os.path.exists(unit_path):
|
||
try:
|
||
os.remove(unit_path)
|
||
print("removed %s" % unit_path)
|
||
except OSError as exc:
|
||
print("unit remove failed: %s" % short_error(exc))
|
||
else:
|
||
print("no service unit found")
|
||
if shutil.which("systemctl"):
|
||
run_host(["systemctl", "--user", "daemon-reload"])
|
||
engine = container_engine()
|
||
if not engine:
|
||
print("no container engine found")
|
||
return 0
|
||
ok, _out = run_host([engine, "rm", "-f", BOX_NAME])
|
||
print("container remove: %s" % ("ok" if ok else "not present"))
|
||
ok, _out = run_host([engine, "rmi", "-f", BOX_IMAGE])
|
||
print("image purge: %s" % ("ok" if ok else "not present"))
|
||
return 0
|
||
|
||
|
||
def upsert_bashrc_block(path):
|
||
try:
|
||
with open(path, "r", encoding="utf-8") as handle:
|
||
raw = handle.read()
|
||
except OSError:
|
||
raw = ""
|
||
if BASHRC_BLOCK in raw:
|
||
return False
|
||
backup = path + ".bak-tai"
|
||
if raw and not os.path.exists(backup):
|
||
with open(backup, "w", encoding="utf-8") as handle:
|
||
handle.write(raw)
|
||
stripped = re.sub(re.escape(BASHRC_MARK_BEGIN) + r".*?" + re.escape(BASHRC_MARK_END) + r"\n?", "", raw, flags=re.DOTALL)
|
||
if stripped and not stripped.endswith("\n"):
|
||
stripped += "\n"
|
||
with open(path, "w", encoding="utf-8") as handle:
|
||
handle.write(stripped + BASHRC_BLOCK)
|
||
return True
|
||
|
||
|
||
def install_self():
|
||
home = os.environ.get("TAI_HOME") or os.path.join(os.path.expanduser("~"), ".tai")
|
||
print(op_binary("upgrade", home))
|
||
print(op_hook("install", home))
|
||
print(op_venv("install", home))
|
||
print("restart your shell or run: source ~/.bashrc")
|
||
return 0
|
||
|
||
|
||
def boot(args):
|
||
args.yes = args.yes or args.yolo or args.auto
|
||
config = Config(args)
|
||
passphrase, default_key = resolve_passphrase()
|
||
try:
|
||
seal = Seal(config.home, passphrase)
|
||
except SealError as exc:
|
||
if default_key:
|
||
raise
|
||
try:
|
||
old = Seal(config.home, DEFAULT_PASSPHRASE)
|
||
except SealError:
|
||
raise exc
|
||
try:
|
||
seal = rotate_seal(config, old, passphrase)
|
||
except (OSError, sqlite3.Error) as fail:
|
||
raise SealError(str(fail))
|
||
print("seal upgraded from default key to TAI_PASSPHRASE")
|
||
seal.default_key = default_key and seal.enabled
|
||
store = Store(config, seal)
|
||
if store.schema_notes:
|
||
print("vault schema upgraded: %s" % ", ".join(store.schema_notes))
|
||
try:
|
||
backed_up = ensure_self_backup(config, store)
|
||
except (OSError, sqlite3.Error) as exc:
|
||
print("warning: self-backup failed: %s" % short_error(exc))
|
||
backed_up = None
|
||
if backed_up is not None:
|
||
print("self-backup: %s" % os.path.basename(backed_up))
|
||
return Agent(config, store, yolo=args.yolo or args.auto, auto=args.auto)
|
||
|
||
|
||
def build_parser():
|
||
parser = argparse.ArgumentParser(prog="tai", description="single-file autonomous agent, standard library only")
|
||
parser.add_argument("--profile", default=None)
|
||
parser.add_argument("--yes", action="store_true")
|
||
parser.add_argument("--yolo", action="store_true")
|
||
parser.add_argument("--auto", action="store_true")
|
||
parser.add_argument("--continue", dest="resume", action="store_true")
|
||
parser.add_argument("--env", default=None)
|
||
parser.add_argument("--version", action="store_true")
|
||
parser.add_argument("--telegram", action="store_true")
|
||
parser.add_argument("--scheduler", action="store_true")
|
||
parser.add_argument("--install", action="store_true")
|
||
parser.add_argument("--install-telegram", action="store_true")
|
||
parser.add_argument("--uninstall-telegram", action="store_true")
|
||
parser.add_argument("prompt", nargs="*")
|
||
return parser
|
||
|
||
|
||
def banner(agent):
|
||
print(paint("tai %s · profile %s" % (VERSION, agent.profile), Ansi.BOLD, Ansi.CYAN))
|
||
if agent.yolo:
|
||
print(paint("YOLO MODE: all approvals off", Ansi.BOLD, Ansi.RED))
|
||
if agent.auto:
|
||
print(paint("AUTO MODE: never asks, researches instead", Ansi.BOLD, Ansi.YELLOW))
|
||
seal_state = "off"
|
||
if agent.store.seal.enabled:
|
||
seal_state = "on (default key)" if agent.store.seal.default_key else "on (TAI_PASSPHRASE)"
|
||
print(" seal: %s" % seal_state)
|
||
for label, _chat_url, models_url, key, _model in agent.chat.backends():
|
||
print(" %s: %s" % (label, probe_backend(models_url, key)))
|
||
if agent.store.load_turn_state(agent.profile, agent.bot).get("open"):
|
||
print(paint("interrupted turn available, restart with --continue to pick it up", Ansi.YELLOW))
|
||
print(paint("type /help for commands", Ansi.DIM))
|
||
|
||
|
||
def resume_turn(agent):
|
||
state = agent.store.load_turn_state(agent.profile, agent.bot)
|
||
if not state.get("open"):
|
||
print("nothing to resume: last turn for %s/%s completed normally" % (agent.profile, agent.bot))
|
||
return ""
|
||
goal = state.get("goal") or ""
|
||
print("resuming interrupted turn for %s/%s: %d steps, %d tool calls, goal: %s" % (agent.profile, agent.bot, state.get("steps", 0), state.get("tools", 0), goal[:200]))
|
||
message = "Continue where you left off%s. Do not repeat steps already completed above; carry on with the next step." % (": the interrupted request was: %s" % goal if goal else "")
|
||
return agent.run_turn(message)
|
||
|
||
|
||
def repl(agent):
|
||
if readline is not None:
|
||
try:
|
||
readline.read_history_file(agent.config.history_path)
|
||
except OSError:
|
||
pass
|
||
readline.set_completer(complete_command)
|
||
readline.parse_and_bind("tab: complete")
|
||
scheduler_stop = start_scheduler(agent.config, agent.store.seal)
|
||
try:
|
||
while True:
|
||
try:
|
||
busy = running_agents()
|
||
base = agent.profile + ("!" if agent.yolo else "")
|
||
tag = "%s+%d" % (base, busy) if busy else base
|
||
line = input(paint("tai[%s]› " % tag, Ansi.BOLD, Ansi.GREEN))
|
||
except EOFError:
|
||
print()
|
||
break
|
||
except KeyboardInterrupt:
|
||
print()
|
||
continue
|
||
text = line.strip()
|
||
if not text:
|
||
continue
|
||
if text.startswith("/"):
|
||
if not handle_command(agent, text):
|
||
break
|
||
continue
|
||
try:
|
||
agent.run_turn(text)
|
||
except KeyboardInterrupt:
|
||
print(paint("\ninterrupted", Ansi.YELLOW))
|
||
agent.checkpoint()
|
||
finally:
|
||
scheduler_stop.set()
|
||
if readline is not None:
|
||
try:
|
||
readline.write_history_file(agent.config.history_path)
|
||
except OSError:
|
||
pass
|
||
agent.store.save_session(agent.profile, agent.messages, agent.bot)
|
||
|
||
|
||
def main(argv=None):
|
||
args = build_parser().parse_args(argv)
|
||
if args.install:
|
||
return install_self()
|
||
if args.install_telegram:
|
||
return install_telegram()
|
||
if args.uninstall_telegram:
|
||
return uninstall_telegram()
|
||
if args.version:
|
||
print("tai %s" % VERSION)
|
||
return 0
|
||
if args.profile is None:
|
||
args.profile = "telegram" if args.telegram else "default"
|
||
if not PROFILE_RE.match(args.profile):
|
||
print("invalid profile name", file=sys.stderr)
|
||
return 1
|
||
try:
|
||
agent = boot(args)
|
||
except KeyboardInterrupt:
|
||
print(paint("\ninterrupted", Ansi.YELLOW))
|
||
return 130
|
||
except SealError as exc:
|
||
print("seal error: %s" % exc, file=sys.stderr)
|
||
return 2
|
||
if args.env is not None:
|
||
if args.env not in ("home", "sandbox"):
|
||
print("use --env home or --env sandbox", file=sys.stderr)
|
||
return 1
|
||
if args.env == "sandbox":
|
||
try:
|
||
ensure_box()
|
||
except BackendError as exc:
|
||
print("sandbox unavailable: %s" % exc, file=sys.stderr)
|
||
return 2
|
||
agent.env = args.env
|
||
if args.scheduler:
|
||
print("scheduler online, press Ctrl-C to stop")
|
||
try:
|
||
run_scheduler_loop(agent.config, agent.store.seal, threading.Event())
|
||
except KeyboardInterrupt:
|
||
print()
|
||
finally:
|
||
agent.store.close()
|
||
return 0
|
||
if args.telegram:
|
||
return run_telegram_bot(agent)
|
||
if args.prompt:
|
||
try:
|
||
agent.run_turn(" ".join(args.prompt))
|
||
except KeyboardInterrupt:
|
||
print(paint("\ninterrupted", Ansi.YELLOW))
|
||
return 130
|
||
finally:
|
||
try:
|
||
agent.checkpoint()
|
||
except Exception:
|
||
pass
|
||
try:
|
||
agent.store.close()
|
||
except Exception:
|
||
pass
|
||
return 0
|
||
if args.resume and not args.telegram and not args.scheduler:
|
||
resume_turn(agent)
|
||
banner(agent)
|
||
try:
|
||
repl(agent)
|
||
finally:
|
||
agent.store.close()
|
||
return 0
|
||
|
||
|
||
if __name__ == "__main__":
|
||
sys.exit(main())
|