From 85ab7d0b89702a5b0c891db7c11d433ee3587238 Mon Sep 17 00:00:00 2001 From: retoor Date: Wed, 7 Oct 2026 06:17:52 +0200 Subject: [PATCH] Full HTTP web_fetch, live shell streaming, progress budget (v1.16.0) web_fetch is now a real HTTP client: methods, custom headers, string bodies, configurable timeout, status plus content-type plus size on every response, content-aware rendering (HTML to text, JSON raw), and server error bodies surfaced on HTTP failures. Shell, schema descriptions, and the system prompt steer all HTTP(S) to web_fetch; curl/wget shell use earns a model-side hint instead of a refusal. Shell output streams live in a rolling 4-line window for the active agent only, with an exit line showing return code, elapsed time, line count, and byte size. Workers, pipes, and capture stay silent. Fixed step budget replaced by a progress budget: productive steps (unseen actions, unseen results, user interaction) reset the stall counter; identical repeats get a loop warning at 3 and a stop at 5; wider stalls stop at the patience budget; 500-step total cap backstops everything. 210 + 16 tests green. --- README.md | 28 +++- tai.py | 367 +++++++++++++++++++++++++++++++++++++++++------ test_tai.py | 406 ++++++++++++++++++++++++++++++++++++++++++++++++++++ 3 files changed, 754 insertions(+), 47 deletions(-) diff --git a/README.md b/README.md index ae499e6..35ed5ed 100644 --- a/README.md +++ b/README.md @@ -121,6 +121,18 @@ commands may run. Workers cannot schedule, cannot prompt for approval, and stop nesting after two levels; scheduled prompts fire as subagents that nobody waits on, with outcomes kept in the vault. +Turns run on a progress budget instead of a fixed step count. A step +counts as productive when it tries an unseen action, returns an +unseen result, or answers a user prompt, and any productive step +resets the stall counter, so a turn doing varied work can run as +long as it keeps moving. Repeating one action with identical +results earns a loop warning at 3 repeats and a stop at 5; wider +stalls (alternating actions, no new information) stop after the +patience budget (25 main, 12 worker, 40 skill builder); a 500-step +total cap backstops everything. Every user approval or typed +guidance resets the stall counter, since a supervised agent is a +safe agent. + ## Tools | Tool | Purpose | @@ -130,7 +142,7 @@ that nobody waits on, with outcomes kept in the vault. | `write_file` | Write content to a file, creating parent directories | | `edit_file` | Replace one unique exact text match in a file | | `web_search` | Search the web, optionally images or page content | -| `web_fetch` | Fetch a URL and return its text content | +| `web_fetch` | HTTP client: methods, headers, bodies, status | | `speak` | Synthesize speech, save MP3, play when possible | | `listen` | Record from the microphone and transcribe it | | `remember` | Merge knowledge into the profile system message | @@ -292,6 +304,12 @@ This construction uses only the standard library and is honest file-theft protection, not audited cryptography; high-value secrets still belong in a dedicated manager. +The vault heals its own schema: every boot creates missing tables and +adds missing columns automatically, reporting what changed +(`vault schema upgraded: added secrets.meta, ...`). Old vaults from +any previous version open without manual steps, and SQLite's dynamic +typing means historic type drift in a column never blocks a boot. + ## Sealed search The seal guards against file theft: an attacker who copies `~/.tai` @@ -457,6 +475,14 @@ rules, and inline bold, italic, code, strikethrough, and links. Long lines wrap to the terminal width without breaking styles. Piped output and `NO_COLOR` stay raw markdown. +Shell output streams live while the command runs: the active agent +shows a rolling 4-line window inside the call box, then an +exit line with the return code, elapsed time, line count, and byte +size. Long lines trim to +the terminal width without breaking colors, progress-style output +keeps its latest segment, and the full text still reaches the model +for feedback. Workers, pipes, and capture mode stay silent. + ## Configuration | Variable | Purpose | Default | diff --git a/tai.py b/tai.py index 09814c5..4d40db0 100755 --- a/tai.py +++ b/tai.py @@ -2,6 +2,7 @@ # retoor import argparse import base64 +import collections import getpass import glob import hashlib @@ -10,6 +11,7 @@ import hmac import html import json import os +import queue import random import re import shutil @@ -34,7 +36,7 @@ try: except ImportError: readline = None -VERSION = "1.14.0" +VERSION = "1.16.0" WORKER_STEPS = 12 CREATE_SKILL_STEPS = 40 FORK_MAX_DEPTH = 2 @@ -46,6 +48,14 @@ PRIMARY_LABEL = "primary" FALLBACK_BASE = "https://devplace.net/openai/v1" FALLBACK_LABEL = "fallback" 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 = "openrouter/free" DEFAULT_VOICE = "en-US-EmmaMultilingualNeural" DEFAULT_PASSPHRASE = "tai-default-insecure-change-me" @@ -124,6 +134,9 @@ 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 @@ -151,6 +164,7 @@ 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. " @@ -505,6 +519,25 @@ 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 = ("⠋", "⠙", "⠹", "⠸", "⠼", "⠴", "⠦", "⠧", "⠇", "⠏") @@ -536,6 +569,134 @@ class Spinner: 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, text=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 + + ANSI_RE = re.compile(r"\033\[[0-9;]*m") MD_BULLETS = ("•", "◦", "▪") @@ -1231,8 +1392,11 @@ 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))) - db.execute("INSERT OR IGNORE INTO %s_new (%s) SELECT %s FROM %s" % (table, ", ".join(names), ", ".join(names), table)) + 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 @@ -1245,13 +1409,49 @@ def migrate_profiles(db): 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")) - db.execute("CREATE INDEX IF NOT EXISTS idx_tags_tag ON tags(tag)") - db.execute("CREATE INDEX IF NOT EXISTS idx_edges_src ON edges(src)") - db.execute("CREATE INDEX IF NOT EXISTS idx_edges_dst ON edges(dst)") - db.execute("CREATE INDEX IF NOT EXISTS idx_tags_profile ON tags(profile)") - db.execute("CREATE INDEX IF NOT EXISTS idx_edges_profile ON edges(profile)") - db.execute("CREATE INDEX IF NOT EXISTS idx_records_profile ON records(profile)") - db.execute("CREATE INDEX IF NOT EXISTS idx_audit_profile ON audit(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"), ()), +) + +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)"), +) + + +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: @@ -1262,22 +1462,8 @@ class Store: self.db.execute("PRAGMA journal_mode=WAL") self.db.execute("PRAGMA synchronous=NORMAL") self.db.execute("PRAGMA busy_timeout=30000") - self.db.execute("CREATE TABLE IF NOT EXISTS events (id INTEGER PRIMARY KEY, profile TEXT, ts TEXT, role TEXT, kind TEXT, text TEXT, tags TEXT)") - self.db.execute("CREATE INDEX IF NOT EXISTS idx_events_profile ON events(profile)") - self.db.execute("CREATE TABLE IF NOT EXISTS secrets (profile TEXT, name TEXT, value TEXT, meta TEXT, updated TEXT, PRIMARY KEY (profile, name))") - self.db.execute("CREATE TABLE IF NOT EXISTS 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)") - self.db.execute("CREATE TABLE IF NOT EXISTS records (id TEXT PRIMARY KEY, kind TEXT, profile TEXT, title TEXT, content TEXT, size INTEGER, reads INTEGER, created TEXT, updated TEXT)") - self.db.execute("CREATE TABLE IF NOT EXISTS tags (item TEXT, tag TEXT, profile TEXT, PRIMARY KEY (item, tag, profile))") - self.db.execute("CREATE INDEX IF NOT EXISTS idx_tags_tag ON tags(tag)") - self.db.execute("CREATE TABLE IF NOT EXISTS edges (src TEXT, dst TEXT, relation TEXT, profile TEXT, created TEXT, PRIMARY KEY (src, dst, relation, profile))") - self.db.execute("CREATE INDEX IF NOT EXISTS idx_edges_src ON edges(src)") - self.db.execute("CREATE INDEX IF NOT EXISTS idx_edges_dst ON edges(dst)") - self.db.execute("CREATE TABLE IF NOT EXISTS 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)") - self.db.execute("CREATE INDEX IF NOT EXISTS idx_audit_path ON audit(path)") - self.db.commit() + self.schema_notes = ensure_schema(self.db) migrate_profiles(self.db) - ensure_column(self.db, "secrets", "meta", "TEXT") - ensure_column(self.db, "events", "tags", "TEXT") self.db.commit() self.profile = config.profile self.secret_cache = {} @@ -1628,7 +1814,7 @@ class Store: 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": row[4], "reads": row[5] + 1, "created": row[6], "updated": row[7], "tags": self.item_tags(row[0], profile)} + 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) @@ -2485,12 +2671,12 @@ def tool_schema(name, description, properties, 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_ variables for that command only; using secrets asks the user first.", {"command": {"type": "string"}, "workdir": {"type": "string"}, "secrets": {"type": "array"}}, ["command"]), + 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_ 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", "Fetch a URL and return its text content. 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"}, "auth_secret": {"type": "string"}, "auth_header": {"type": "string"}, "auth_prefix": {"type": "string"}}, ["url"]), + 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"]), @@ -2727,6 +2913,7 @@ class Tools: } self.secret_grants = set() self.read_files = set() + self.live = False def box_engine(self): try: @@ -2818,12 +3005,13 @@ class Tools: env = dict(os.environ) for name, value in vault.items(): env[secret_env_name(name)] = value - try: - done = subprocess.run(command, shell=True, cwd=workdir, capture_output=True, text=True, timeout=SHELL_TIMEOUT, env=env) - except subprocess.TimeoutExpired: + 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 - output = (done.stdout or "") + (done.stderr or "") - return "exit %d\n%s" % (done.returncode, spill_or_truncate(self.app.store, "output", "shell: " + command[:80], output.strip() or "(no output)", ["shell"], self.active_profile)) + 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() @@ -2833,13 +3021,15 @@ class Tools: for name, value in vault.items(): extra += ["-e", "%s=%s" % (secret_env_name(name), value)] try: - done = box_exec(engine, ["sh", "-c", command], extra=tuple(extra), timeout=SHELL_TIMEOUT) - except subprocess.TimeoutExpired: - return "error: timed out after %d seconds" % SHELL_TIMEOUT + 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) - output = done.stdout.decode("utf-8", "replace") + done.stderr.decode("utf-8", "replace") - return "exit %d\n%s" % (done.returncode, spill_or_truncate(self.app.store, "output", "shell: " + command[:80], output.strip() or "(no output)", ["shell"], self.active_profile)) + 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 "") @@ -3039,7 +3229,27 @@ class Tools: 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) @@ -3056,16 +3266,31 @@ class Tools: if not self.secret_grant("web_fetch", host, [auth_name]): return "denied by user" headers[header] = prefix + value - request = urllib.request.Request(url, headers=headers) + request = urllib.request.Request(url, data=data, headers=headers, method=method) try: - with urllib.request.urlopen(request, timeout=30) as response: - raw = response.read(200000).decode("utf-8", "replace") + 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 = re.sub(r"(?s)|", " ", raw) - text = re.sub(r"<[^>]+>", " ", text) - text = re.sub(r"\s+", " ", html.unescape(text)).strip() - return truncate(text, 5000, 1000) or "empty page" + 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)|", " ", 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() @@ -3853,6 +4078,7 @@ class Agent: self.deadline = None self.timed_out = False self.runner_override = None + self.user_stepped_in = False self.switch_profile(config.profile, silent=True) def apply_system(self): @@ -3940,6 +4166,7 @@ class Agent: 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 @@ -3993,7 +4220,14 @@ class Agent: last_text = "" sink = pieces.append loud = not capture and not self.quiet - for _step in range(max_steps): + 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() @@ -4040,6 +4274,7 @@ class Agent: 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() @@ -4050,6 +4285,7 @@ class Agent: 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: @@ -4062,13 +4298,50 @@ class Agent: 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 + 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", Ansi.YELLOW)) + print(paint("step budget exhausted (%d total steps)" % TOTAL_STEP_CAP, Ansi.YELLOW)) if self.persist: self.store.save_session(self.profile, self.messages, self.bot) total = time.time() - started @@ -5117,6 +5390,8 @@ def boot(args): 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: diff --git a/test_tai.py b/test_tai.py index ac5e871..58b9b83 100644 --- a/test_tai.py +++ b/test_tai.py @@ -1,4 +1,5 @@ # retoor +import io import json import os import re @@ -7,6 +8,7 @@ import sys import tempfile import time import unittest +import urllib.error from datetime import datetime, timedelta, timezone from unittest import mock @@ -1438,6 +1440,54 @@ class MigrationTests(unittest.TestCase): finally: store.close() + def test_sparse_legacy_db_auto_heals(self): + db = sqlite3.connect(self.config.db_path) + db.execute("CREATE TABLE secrets (name TEXT PRIMARY KEY, value TEXT, updated TEXT)") + db.execute("INSERT INTO secrets VALUES ('srv', 'pw', 't')") + db.execute("CREATE TABLE events (id INTEGER PRIMARY KEY, profile TEXT, ts TEXT, role TEXT, kind TEXT, text TEXT)") + db.execute("INSERT INTO events VALUES (1, 'default', 't', 'user', 'message', 'hello world')") + db.execute("CREATE TABLE records (id TEXT PRIMARY KEY, kind TEXT, title TEXT, content TEXT, size TEXT, reads INTEGER, created TEXT, updated TEXT)") + db.execute("INSERT INTO records VALUES ('mem:aaaaaaaaaaaaaaaa', 'note', 't', 'words', '5', 0, '', '')") + db.commit() + db.close() + store = tai.Store(self.config, tai.Seal(self.config.home, "")) + try: + self.assertEqual(store.load_secret("srv"), "pw") + self.assertEqual(store.get_record("mem:aaaaaaaaaaaaaaaa")["content"], "words") + self.assertEqual(store.get_record("mem:aaaaaaaaaaaaaaaa")["size"], 5) + self.assertEqual(len(store.search_events("default", "hello")), 1) + for table in ("secrets", "events", "records", "tags", "edges", "audit", "schedules"): + self.assertTrue(store.db.execute("SELECT 1 FROM sqlite_master WHERE type = 'table' AND name = ?", (table,)).fetchone(), table) + secret_cols = [row[1] for row in store.db.execute("PRAGMA table_info(secrets)").fetchall()] + event_cols = [row[1] for row in store.db.execute("PRAGMA table_info(events)").fetchall()] + record_cols = [row[1] for row in store.db.execute("PRAGMA table_info(records)").fetchall()] + self.assertIn("meta", secret_cols) + self.assertIn("tags", event_cols) + self.assertIn("profile", record_cols) + self.assertIn("added secrets.meta", store.schema_notes) + self.assertIn("added events.tags", store.schema_notes) + self.assertIn("added records.profile", store.schema_notes) + finally: + store.close() + + def test_legacy_db_without_meta_migrates(self): + db = sqlite3.connect(self.config.db_path) + db.execute("CREATE TABLE secrets (name TEXT PRIMARY KEY, value TEXT, updated TEXT)") + db.execute("INSERT INTO secrets VALUES ('srv', 'pw-legacy', '2026-01-01')") + db.commit() + db.close() + store = tai.Store(self.config, tai.Seal(self.config.home, "")) + try: + self.assertEqual(store.load_secret("srv"), "pw-legacy") + self.assertEqual(store.secret_meta("srv"), {}) + self.assertEqual(tai.table_pk_columns(store.db, "secrets"), ["profile", "name"]) + names = [row[1] for row in store.db.execute("PRAGMA table_info(secrets)").fetchall()] + self.assertIn("meta", names) + store.save_secret("srv", "pw-new", {"host": "example.com"}, None, "default") + self.assertEqual(store.secret_meta("srv"), {"host": "example.com"}) + finally: + store.close() + class LazyTests(unittest.TestCase): def setUp(self): @@ -2032,6 +2082,362 @@ class FTSSearchTests(unittest.TestCase): self.store = tai.Store(self.config, tai.Seal(self.config.home, "fts-test-2")) +class StreamTests(unittest.TestCase): + def setUp(self): + self.tmp = tempfile.TemporaryDirectory() + os.environ["TAI_HOME"] = self.tmp.name + + class FakeArgs: + profile = "t" + yes = True + + self.config = tai.Config(FakeArgs()) + self.store = tai.Store(self.config, tai.Seal(self.config.home, "stream-test-1")) + + def tearDown(self): + self.store.close() + self.tmp.cleanup() + os.environ.pop("TAI_HOME", None) + + def test_truncate_ansi_keeps_short_colors(self): + text = "\033[31mhi\033[0m" + self.assertEqual(tai.truncate_ansi(text, 10), text) + + def test_truncate_ansi_cuts_visible_width(self): + cut = tai.truncate_ansi("\033[31m" + "x" * 50 + "\033[0m", 10) + self.assertEqual(tai.strip_ansi(cut), "x" * 10) + self.assertTrue(cut.endswith(tai.Ansi.RESET)) + self.assertIn("\033[31m", cut) + + def test_truncate_ansi_plain(self): + self.assertEqual(tai.truncate_ansi("abcdef", 4), "abcd" + tai.Ansi.RESET) + + def test_human_size(self): + self.assertEqual(tai.human_size(0), "0 B") + self.assertEqual(tai.human_size(512), "512 B") + self.assertEqual(tai.human_size(1024), "1.0 KB") + self.assertEqual(tai.human_size(1536), "1.5 KB") + self.assertEqual(tai.human_size(2097152), "2.0 MB") + + def test_stream_status(self): + self.assertEqual(tai.stream_status(0, "a\nb\n", 1.25, False), "exit 0 · 1.2s · 2 lines · 4 B") + self.assertEqual(tai.stream_status(3, "", 0.0, False), "exit 3 · 0.0s · 0 lines · 0 B") + self.assertEqual(tai.stream_status(-9, "x\n" * 10, 120.0, True), "timed out after 120s (killed) · 10 lines · 20 B") + self.assertEqual(tai.stream_status(0, "z" * 2048, 0.5, False), "exit 0 · 0.5s · 1 lines · 2.0 KB") + + def test_window_rolls_at_height(self): + buf = io.StringIO() + stream = tai.LiveStream(height=2, width=40, file=buf, enabled=True) + stream.feed("one\n") + stream.feed("two\n") + stream.feed("three\n") + stream.close("exit 0") + out = buf.getvalue() + self.assertIn("│ one\n│ two\n", out) + self.assertIn("\033[2A", out) + self.assertIn("\r\033[K│ three\n", out) + self.assertTrue(out.endswith("│ exit 0\n")) + + def test_window_disabled_writes_nothing(self): + buf = io.StringIO() + stream = tai.LiveStream(file=buf, enabled=False) + stream.feed("one\n") + stream.close("exit 0") + self.assertEqual(buf.getvalue(), "") + + def test_window_no_output_status(self): + buf = io.StringIO() + stream = tai.LiveStream(file=buf, enabled=True) + stream.close("exit 3") + self.assertEqual(buf.getvalue(), "│ (no output)\n│ exit 3\n") + + def test_window_keeps_progress_tail(self): + buf = io.StringIO() + stream = tai.LiveStream(height=4, width=40, file=buf, enabled=True) + stream.feed("50%\r100%\n") + self.assertIn("│ 100%\n", buf.getvalue()) + + def test_run_live_captures(self): + code, out, _elapsed, timed = tai.run_live(["echo", "hi"]) + self.assertEqual((code, out, timed), (0, "hi\n", False)) + + def test_run_live_exit_code(self): + code, _out, _elapsed, timed = tai.run_live(["sh", "-c", "exit 3"]) + self.assertEqual((code, timed), (3, False)) + + def test_run_live_merges_stderr(self): + code, out, _elapsed, _timed = tai.run_live(["sh", "-c", "echo out; echo err >&2"]) + self.assertIn("out\n", out) + self.assertIn("err\n", out) + self.assertEqual(code, 0) + + def test_run_live_streams(self): + buf = io.StringIO() + stream = tai.LiveStream(file=buf, enabled=True) + _code, out, _elapsed, timed = tai.run_live(["echo", "streamed"], stream=stream) + stream.close("exit 0") + self.assertEqual(out, "streamed\n") + self.assertFalse(timed) + self.assertIn("│ streamed\n", buf.getvalue()) + + def test_run_live_timeout_kills(self): + started = time.time() + _code, _out, _elapsed, timed = tai.run_live(["sleep", "30"], timeout=1) + self.assertTrue(timed) + self.assertLess(time.time() - started, 10) + + def test_shell_live_flag(self): + app = FakeSecretApp(self.store) + app.profile = "t" + tools = tai.Tools(app) + self.assertFalse(tools.live) + tools.live = True + buf = io.StringIO() + buf.isatty = lambda: True + with mock.patch.object(sys, "stdout", buf): + result = tools.dispatch("shell", json.dumps({"command": "echo live-mark"})) + self.assertTrue(result.startswith("exit 0")) + self.assertIn("live-mark", result) + self.assertIn("│ live-mark\n", buf.getvalue()) + self.assertIn("exit 0 ·", buf.getvalue()) + self.assertIn("1 lines · 10 B", buf.getvalue()) + + def test_shell_silent_without_flag(self): + app = FakeSecretApp(self.store) + app.profile = "t" + tools = tai.Tools(app) + buf = io.StringIO() + buf.isatty = lambda: True + with mock.patch.object(sys, "stdout", buf): + result = tools.dispatch("shell", json.dumps({"command": "echo quiet-mark"})) + self.assertTrue(result.startswith("exit 0")) + self.assertEqual(buf.getvalue(), "") + + def test_agent_enables_live_for_turn(self): + agent = tai.Agent(self.config, self.store, persist=False, quiet=True) + reply = {"role": "assistant", "content": "done", "reasoning": "", "tool_calls": [], "backend": "x"} + with mock.patch.object(agent.chat, "complete", return_value=dict(reply)): + agent.run_turn("hi", capture=True) + self.assertFalse(agent.tools.live) + agent.quiet = False + with mock.patch.object(agent.chat, "complete", return_value=dict(reply)): + with mock.patch("builtins.print"): + agent.run_turn("hi") + self.assertTrue(agent.tools.live) + + +class ProgressTests(unittest.TestCase): + def setUp(self): + self.tmp = tempfile.TemporaryDirectory() + os.environ["TAI_HOME"] = self.tmp.name + + class FakeArgs: + profile = "t" + yes = True + + self.config = tai.Config(FakeArgs()) + self.store = tai.Store(self.config, tai.Seal(self.config.home, "progress-test-1")) + self.agent = tai.Agent(self.config, self.store, persist=False, quiet=True) + + def tearDown(self): + self.store.close() + self.tmp.cleanup() + os.environ.pop("TAI_HOME", None) + + def tool_reply(self, command, call_id="1"): + return {"role": "assistant", "content": "", "reasoning": "", "tool_calls": [{"id": call_id, "name": "shell", "arguments": json.dumps({"command": command})}], "backend": "x"} + + def text_reply(self, text="all done"): + return {"role": "assistant", "content": text, "reasoning": "", "tool_calls": [], "backend": "x"} + + def tool_count(self): + return len([item for item in self.agent.messages if item["role"] == "tool"]) + + def test_novelty_resets_stall(self): + replies = [self.tool_reply("echo novel-%d" % num, str(num)) for num in range(5)] + [self.text_reply()] + with mock.patch.object(self.agent.chat, "complete", side_effect=[dict(item) for item in replies]): + answer = self.agent.run_turn("go", capture=True, max_steps=3) + self.assertEqual(answer, "all done") + self.assertEqual(self.tool_count(), 5) + + def test_loop_nudge_then_stop(self): + replies = [self.tool_reply("echo same", str(num)) for num in range(6)] + [self.text_reply()] + with mock.patch.object(self.agent.chat, "complete", side_effect=[dict(item) for item in replies]): + answer = self.agent.run_turn("go", capture=True, max_steps=10) + self.assertIn("stuck in a loop", answer) + self.assertEqual(self.tool_count(), 5) + nudges = [item for item in self.agent.messages if item["role"] == "user" and "loop warning" in item.get("content", "")] + self.assertEqual(len(nudges), 1) + + def test_nudge_recovery(self): + replies = [self.tool_reply("echo same", str(num)) for num in range(3)] + [self.tool_reply("echo different", "9"), self.text_reply()] + with mock.patch.object(self.agent.chat, "complete", side_effect=[dict(item) for item in replies]): + answer = self.agent.run_turn("go", capture=True, max_steps=10) + self.assertEqual(answer, "all done") + self.assertEqual(self.tool_count(), 4) + nudges = [item for item in self.agent.messages if item["role"] == "user" and "loop warning" in item.get("content", "")] + self.assertEqual(len(nudges), 1) + + def test_stall_oscillation(self): + replies = [] + for num in range(8): + replies.append(self.tool_reply("echo osc-%d" % (num % 2), str(num))) + replies.append(self.text_reply()) + with mock.patch.object(self.agent.chat, "complete", side_effect=[dict(item) for item in replies]): + answer = self.agent.run_turn("go", capture=True, max_steps=3) + self.assertIn("no progress for 3 steps", answer) + self.assertEqual(self.tool_count(), 5) + + def test_user_interaction_resets(self): + replies = [self.tool_reply("echo int-%d" % (num % 2), str(num)) for num in range(8)] + [self.text_reply()] + real_dispatch = self.agent.tools.dispatch + calls = [] + + def approving(name, raw): + calls.append(name) + if len(calls) <= 3: + self.agent.user_stepped_in = True + return real_dispatch(name, raw) + + with mock.patch.object(self.agent.chat, "complete", side_effect=[dict(item) for item in replies]): + with mock.patch.object(self.agent.tools, "dispatch", side_effect=approving): + answer = self.agent.run_turn("go", capture=True, max_steps=2) + self.assertIn("no progress for 2 steps", answer) + self.assertEqual(self.tool_count(), 5) + + def test_total_cap_backstop(self): + replies = [self.tool_reply("echo fresh-%d" % num, str(num)) for num in range(8)] + [self.text_reply()] + with mock.patch.object(tai, "TOTAL_STEP_CAP", 6): + with mock.patch.object(self.agent.chat, "complete", side_effect=[dict(item) for item in replies]): + answer = self.agent.run_turn("go", capture=True, max_steps=100) + self.assertEqual(answer, "") + self.assertEqual(self.tool_count(), 6) + + def test_ask_approval_marks_interaction(self): + agent = tai.Agent(self.config, self.store, persist=True, quiet=True) + self.assertFalse(agent.user_stepped_in) + fake_stdin = io.StringIO("y\n") + fake_stdin.isatty = lambda: True + with mock.patch.object(sys, "stdin", fake_stdin): + with mock.patch("builtins.print"): + self.assertTrue(agent.ask_approval("echo hi")) + self.assertTrue(agent.user_stepped_in) + auto_agent = tai.Agent(self.config, self.store, persist=True, quiet=True, auto=True) + self.assertTrue(auto_agent.ask_approval("echo hi")) + self.assertFalse(auto_agent.user_stepped_in) + + +class WebFetchTests(unittest.TestCase): + def setUp(self): + self.tmp = tempfile.TemporaryDirectory() + os.environ["TAI_HOME"] = self.tmp.name + + class FakeArgs: + profile = "t" + yes = True + + self.config = tai.Config(FakeArgs()) + self.store = tai.Store(self.config, tai.Seal(self.config.home, "fetch-test-1")) + self.app = FakeSecretApp(self.store) + self.app.profile = "t" + self.tools = tai.Tools(self.app) + + def tearDown(self): + self.store.close() + self.tmp.cleanup() + os.environ.pop("TAI_HOME", None) + + def fake_open(self, body, status=200, content_type="text/html; charset=utf-8", seen=None): + class FakeResponse: + def __enter__(self): + return self + + def __exit__(self, *exc): + return False + + def read(self, limit=0): + return body + + def fake_urlopen(request, timeout=30): + if seen is not None: + seen["method"] = request.get_method() + seen["data"] = request.data + seen["headers"] = {key.lower(): value for key, value in request.headers.items()} + seen["timeout"] = timeout + response = FakeResponse() + response.status = status + response.headers = {"Content-Type": content_type} + return response + + return fake_urlopen + + def test_fetch_post_json(self): + seen = {} + payload = b'{"ok": true, "tag": ""}' + with mock.patch("urllib.request.urlopen", side_effect=self.fake_open(payload, 201, "application/json", seen)): + result = self.tools.dispatch("web_fetch", json.dumps({"url": "https://api.example.test/items", "method": "post", "headers": {"Content-Type": "application/json", "X-Trace": "1"}, "body": payload.decode("utf-8")})) + self.assertEqual(seen["method"], "POST") + self.assertEqual(seen["data"], payload) + self.assertEqual(seen["headers"]["content-type"], "application/json") + self.assertEqual(seen["headers"]["x-trace"], "1") + self.assertTrue(result.startswith("HTTP 201 · application/json · 29 B\n")) + self.assertIn("", result) + + def test_fetch_html_strips(self): + with mock.patch("urllib.request.urlopen", side_effect=self.fake_open(b"

Hi

")): + result = self.tools.dispatch("web_fetch", json.dumps({"url": "https://example.test/"})) + self.assertTrue(result.startswith("HTTP 200 · text/html · ")) + self.assertIn("Hi", result) + self.assertNotIn("

", result) + self.assertNotIn("var x", result) + + def test_fetch_raw_keeps_tags(self): + with mock.patch("urllib.request.urlopen", side_effect=self.fake_open(b"

Hi

")): + result = self.tools.dispatch("web_fetch", json.dumps({"url": "https://example.test/", "raw": True})) + self.assertIn("

Hi

", result) + + def test_fetch_validation(self): + base = "https://example.test/" + self.assertIn("method must be", self.tools.dispatch("web_fetch", json.dumps({"url": base, "method": "BREW"}))) + self.assertIn("headers must be", self.tools.dispatch("web_fetch", json.dumps({"url": base, "headers": ["x"]}))) + self.assertIn("headers must be", self.tools.dispatch("web_fetch", json.dumps({"url": base, "headers": {"X-A": 1}}))) + self.assertIn("invalid header", self.tools.dispatch("web_fetch", json.dumps({"url": base, "headers": {"Bad Name": "x"}}))) + self.assertIn("invalid header", self.tools.dispatch("web_fetch", json.dumps({"url": base, "headers": {"X-A": "a\nb"}}))) + self.assertIn("body must be", self.tools.dispatch("web_fetch", json.dumps({"url": base, "method": "POST", "body": 42}))) + self.assertIn("invalid timeout", self.tools.dispatch("web_fetch", json.dumps({"url": base, "timeout": "x"}))) + self.assertIn("url must start", self.tools.dispatch("web_fetch", json.dumps({"url": "ftp://example.test/"}))) + + def test_fetch_http_error_surfaces_body(self): + failure = urllib.error.HTTPError("https://example.test/", 404, "Not Found", None, io.BytesIO(b"no such widget")) + with mock.patch("urllib.request.urlopen", side_effect=failure): + result = self.tools.dispatch("web_fetch", json.dumps({"url": "https://example.test/"})) + self.assertEqual(result, "error: HTTP 404: no such widget") + empty = urllib.error.HTTPError("https://example.test/", 500, "Server Error", None, io.BytesIO(b"")) + with mock.patch("urllib.request.urlopen", side_effect=empty): + result = self.tools.dispatch("web_fetch", json.dumps({"url": "https://example.test/"})) + self.assertTrue(result.startswith("error: HTTP 500:")) + + def test_fetch_connection_error(self): + with mock.patch("urllib.request.urlopen", side_effect=urllib.error.URLError("refused")): + result = self.tools.dispatch("web_fetch", json.dumps({"url": "https://example.test/"})) + self.assertIn("error: fetch failed:", result) + + def test_shell_curl_hint(self): + marked = self.tools.dispatch("shell", json.dumps({"command": "echo curl http://example.test"})) + self.assertIn(tai.CURL_HINT, marked) + plain = self.tools.dispatch("shell", json.dumps({"command": "echo just-local"})) + self.assertNotIn("web_fetch", plain) + self.assertEqual(tai.curl_hint("wget https://example.test/x"), "\n" + tai.CURL_HINT) + self.assertEqual(tai.curl_hint("curl --version"), "") + self.assertEqual(tai.curl_hint("echo hi"), "") + + def test_steering_text(self): + descs = {schema["function"]["name"]: schema["function"]["description"] for schema in tai.TOOL_SCHEMAS} + self.assertIn("never curl or wget", descs["web_fetch"]) + self.assertIn("use web_fetch for HTTP(S)", descs["shell"]) + self.assertIn("never curl or wget", tai.DEFAULT_SYSTEM) + + class DenialTests(unittest.TestCase): def setUp(self): self.tmp = tempfile.TemporaryDirectory()