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.
This commit is contained in:
2026-10-07 06:17:52 +02:00
parent 99d7cb21b1
commit 85ab7d0b89
3 changed files with 754 additions and 47 deletions
+27 -1
View File
@@ -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 |
+321 -46
View File
@@ -2,6 +2,7 @@
# retoor <retoor@molodetz.nl>
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_<NAME> 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_<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", "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)<script.*?</script>|<style.*?</style>", " ", 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)<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()
@@ -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:
+406
View File
@@ -1,4 +1,5 @@
# retoor <retoor@molodetz.nl>
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": "<kept>"}'
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("<kept>", result)
def test_fetch_html_strips(self):
with mock.patch("urllib.request.urlopen", side_effect=self.fake_open(b"<html><body><h1>Hi</h1><script>var x = 1;</script></body></html>")):
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("<h1>", 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"<p>Hi</p>")):
result = self.tools.dispatch("web_fetch", json.dumps({"url": "https://example.test/", "raw": True}))
self.assertIn("<p>Hi</p>", 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()