Files
tai/test_tai.py

3160 lines
152 KiB
Python

# retoor <retoor@molodetz.nl>
import io
import json
import os
import re
import sqlite3
import sys
import tempfile
import time
import unittest
import urllib.error
from datetime import datetime, timedelta, timezone
from unittest import mock
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
import tai
VALID_SKILL = """---
name: pdf-forms
description: >
Fill PDF forms and extract field data.
Use when the user mentions PDF documents.
---
# PDF forms
Do the thing.
"""
class SkillTests(unittest.TestCase):
def write_skill(self, root, entry, body):
folder = os.path.join(root, entry)
os.makedirs(folder, exist_ok=True)
path = os.path.join(folder, "SKILL.md")
with open(path, "w", encoding="utf-8") as handle:
handle.write(body)
return path
def test_parse_valid(self):
with tempfile.TemporaryDirectory() as tmp:
path = self.write_skill(tmp, "x", VALID_SKILL)
skill = tai.parse_skill_file(path)
self.assertEqual(skill["name"], "pdf-forms")
self.assertIn("Fill PDF forms", skill["description"])
self.assertIn("Do the thing.", skill["body"])
def test_parse_rejects_bad(self):
with tempfile.TemporaryDirectory() as tmp:
missing = self.write_skill(tmp, "a", "---\nname: x\n---\nbody\n")
self.assertIsNone(tai.parse_skill_file(missing))
bad_name = self.write_skill(tmp, "b", "---\nname: Bad_Name!\ndescription: d\n---\nbody\n")
self.assertIsNone(tai.parse_skill_file(bad_name))
no_front = self.write_skill(tmp, "c", "just markdown\n")
self.assertIsNone(tai.parse_skill_file(no_front))
def test_discover_project_wins(self):
with tempfile.TemporaryDirectory() as home, tempfile.TemporaryDirectory() as project:
self.write_skill(os.path.join(home, "skills"), "dup", VALID_SKILL)
other = VALID_SKILL.replace("Do the thing.", "Project variant.")
self.write_skill(os.path.join(project, ".tai", "skills"), "dup", other)
found = tai.discover_skills(home, project)
self.assertEqual(list(found), ["pdf-forms"])
self.assertIn("Project variant.", found["pdf-forms"]["body"])
def test_catalog(self):
skills = {"b-skill": {"description": "second"}, "a-skill": {"description": "first"}}
catalog = tai.skill_catalog(skills)
self.assertLess(catalog.index("a-skill"), catalog.index("b-skill"))
self.assertIn("load_skill", catalog)
empty = tai.skill_catalog({})
self.assertIn("Buildable skill blueprints", empty)
self.assertIn("bot-creator", empty)
class FakeApp:
def __init__(self):
self.skills = {}
self.env = "home"
self.store = None
class ToolTests(unittest.TestCase):
def test_load_skill_unknown(self):
tools = tai.Tools(FakeApp())
result = tools.dispatch("load_skill", json.dumps({"name": "nope"}))
self.assertIn("unknown skill", result)
def test_terminal_without_tmux(self):
tools = tai.Tools(FakeApp())
with tempfile.TemporaryDirectory() as empty:
with mock.patch.dict(os.environ, {"PATH": empty}):
self.assertEqual(tools.dispatch("get_current_terminal_content", "{}"), "tmux not available")
def test_box_helpers_present(self):
self.assertIn("faster-whisper", tai.BOX_CONTAINERFILE)
self.assertIn("edge-tts", tai.BOX_CONTAINERFILE)
self.assertIn("sleep", tai.BOX_CONTAINERFILE)
self.assertIn("WhisperModel", tai.BOX_STT)
self.assertIn("Communicate", tai.BOX_TTS)
class SysinfoTests(unittest.TestCase):
def test_collect_all_checks_visible(self):
report = tai.collect_sysinfo()
lines = report.splitlines()
self.assertEqual(len(lines), len(tai.SYSINFO_CHECKS))
for (name, _func), line in zip(tai.SYSINFO_CHECKS, lines):
self.assertTrue(line.startswith(name + ": "), line)
self.assertRegex(line, r"\(\d+ms\)$")
def test_subset(self):
report = tai.collect_sysinfo(["os", "root"])
lines = report.splitlines()
self.assertEqual(len(lines), 2)
self.assertTrue(lines[0].startswith("os: "))
self.assertTrue(lines[1].startswith("root: "))
def test_checks_run_in_parallel(self):
def slow():
time.sleep(0.3)
return "ok"
stubs = tuple(("slow%d" % pos, slow) for pos in range(4))
with mock.patch.object(tai, "SYSINFO_CHECKS", stubs):
started = time.time()
report = tai.collect_sysinfo()
self.assertLess(time.time() - started, 1.0)
self.assertEqual(len(report.splitlines()), 4)
def test_timeout_surfaces(self):
def stuck():
time.sleep(5)
return "never"
with mock.patch.object(tai, "SYSINFO_CHECKS", (("stuck", stuck),)):
with mock.patch.object(tai, "SYSINFO_TIMEOUT", 1):
report = tai.collect_sysinfo()
self.assertIn("stuck: timed out after 1s", report)
def test_check_error_surfaces(self):
def broken():
raise RuntimeError("boom")
with mock.patch.object(tai, "SYSINFO_CHECKS", (("broken", broken),)):
report = tai.collect_sysinfo()
self.assertIn("broken: error: boom", report)
def test_venv_detection(self):
with mock.patch.dict(os.environ, {"VIRTUAL_ENV": "/tmp/fake-venv"}):
self.assertIn("/tmp/fake-venv", tai.sysinfo_check_venv())
def test_root_shape(self):
result = tai.sysinfo_check_root()
self.assertTrue(result.startswith(("yes", "no", "unknown")), result)
def test_tool_dispatch(self):
tools = tai.Tools(FakeApp())
self.assertIn("os: ", tools.dispatch("sysinfo", "{}"))
subset = tools.dispatch("sysinfo", json.dumps({"checks": ["python"]}))
self.assertEqual(len(subset.splitlines()), 1)
self.assertIn("unknown checks: nope", tools.dispatch("sysinfo", json.dumps({"checks": ["nope"]})))
self.assertIn("non-empty list", tools.dispatch("sysinfo", json.dumps({"checks": []})))
def test_sysinfo_command(self):
self.assertIn("sysinfo", tai.COMMANDS)
self.assertIn("/sysinfo", tai.HELP_TEXT)
with mock.patch.object(tai, "collect_sysinfo", return_value="fake-report") as collector:
with mock.patch("builtins.print") as printer:
self.assertTrue(tai.handle_command(mock.Mock(), "/sysinfo"))
collector.assert_called_once_with()
printer.assert_called_once_with("fake-report")
class FakeSkillApp:
def __init__(self, config):
self.depth = 0
self.profile = "t"
self.config = config
self.store = mock.Mock()
self.store.redact = lambda text, profile=None: text
self.skills = {}
self.applied = 0
self.runner_override = None
def apply_system(self):
self.applied += 1
class CreateSkillTests(unittest.TestCase):
def config_in(self, tmp):
old_home = os.environ.get("TAI_HOME")
os.environ["TAI_HOME"] = tmp
class FakeArgs:
profile = "t"
yes = True
try:
return tai.Config(FakeArgs())
finally:
if old_home is None:
os.environ.pop("TAI_HOME", None)
else:
os.environ["TAI_HOME"] = old_home
def test_prompt_is_self_contained(self):
prompt = tai.create_skill_prompt("pdf-forms", "fill pdf forms", "/tmp/x/pdf-forms")
self.assertIn("pdf-forms", prompt)
self.assertIn("fill pdf forms", prompt)
self.assertIn("/tmp/x/pdf-forms", prompt)
self.assertIn("sysinfo", prompt)
self.assertIn("SKILL.md", prompt)
self.assertIn("name: pdf-forms", prompt)
self.assertIn("two independent sources", prompt)
def test_validation(self):
tools = tai.Tools(FakeSkillApp(None))
self.assertIn("invalid skill name", tools.dispatch("create_skill", json.dumps({"name": "Bad_Name!", "brief": "b"})))
self.assertIn("invalid skill name", tools.dispatch("create_skill", json.dumps({"brief": "b"})))
self.assertIn("empty brief", tools.dispatch("create_skill", json.dumps({"name": "ok-name"})))
self.assertIn("scope must be", tools.dispatch("create_skill", json.dumps({"name": "ok-name", "brief": "b", "scope": "moon"})))
deep = FakeSkillApp(None)
deep.depth = 2
self.assertIn("depth limit", tai.Tools(deep).dispatch("create_skill", json.dumps({"name": "ok-name", "brief": "b"})))
def test_creates_and_refreshes(self):
with tempfile.TemporaryDirectory() as tmp:
app = FakeSkillApp(self.config_in(tmp))
seen = {}
def runner(task, profile, timeout):
seen["task"] = task
seen["profile"] = profile
seen["timeout"] = timeout
skill_dir = os.path.join(tmp, ".tai", "skills", "pdf-forms")
os.makedirs(skill_dir, exist_ok=True)
with open(os.path.join(skill_dir, "SKILL.md"), "w", encoding="utf-8") as handle:
handle.write(VALID_SKILL)
return "researched and wrote the skill"
app.runner_override = runner
with mock.patch("os.getcwd", return_value=tmp):
result = tai.Tools(app).dispatch("create_skill", json.dumps({"name": "pdf-forms", "brief": "fill pdf forms"}))
self.assertIn("created at", result)
self.assertIn("pdf-forms", app.skills)
self.assertEqual(app.applied, 1)
self.assertEqual(seen["profile"], "t")
self.assertIn("fill pdf forms", seen["task"])
self.assertIn("SKILL.md", seen["task"])
def test_replace_and_missing(self):
with tempfile.TemporaryDirectory() as tmp:
config = self.config_in(tmp)
skill_dir = os.path.join(tmp, ".tai", "skills", "pdf-forms")
os.makedirs(skill_dir, exist_ok=True)
with open(os.path.join(skill_dir, "SKILL.md"), "w", encoding="utf-8") as handle:
handle.write(VALID_SKILL)
app = FakeSkillApp(config)
app.runner_override = lambda task, profile, timeout: "rewrote it"
with mock.patch("os.getcwd", return_value=tmp):
result = tai.Tools(app).dispatch("create_skill", json.dumps({"name": "pdf-forms", "brief": "b"}))
self.assertIn("replaced at", result)
app2 = FakeSkillApp(config)
app2.runner_override = lambda task, profile, timeout: "gave up"
with mock.patch("os.getcwd", return_value=tmp):
result2 = tai.Tools(app2).dispatch("create_skill", json.dumps({"name": "other-skill", "brief": "b"}))
self.assertIn("was not created (done)", result2)
self.assertNotIn("other-skill", app2.skills)
self.assertEqual(app2.applied, 0)
def test_home_scope(self):
with tempfile.TemporaryDirectory() as tmp:
app = FakeSkillApp(self.config_in(tmp))
def runner(task, profile, timeout):
skill_dir = os.path.join(tmp, "skills", "home-skill")
os.makedirs(skill_dir, exist_ok=True)
with open(os.path.join(skill_dir, "SKILL.md"), "w", encoding="utf-8") as handle:
handle.write(VALID_SKILL.replace("pdf-forms", "home-skill"))
return "done"
app.runner_override = runner
with mock.patch("os.getcwd", return_value="/nonexistent-dir"):
result = tai.Tools(app).dispatch("create_skill", json.dumps({"name": "home-skill", "brief": "b", "scope": "home"}))
self.assertIn("created at", result)
self.assertIn("home-skill", app.skills)
class FakeSecretApp:
def __init__(self, store, approve=True):
self.store = store
self.env = "home"
self.config = mock.Mock(auto_approve=True)
self.approvals = []
self.approve = approve
def ask_approval(self, command):
self.approvals.append(command)
return self.approve
def show_diff(self):
pass
class SecretsTests(unittest.TestCase):
def setUp(self):
self.tmp = tempfile.TemporaryDirectory()
os.environ["TAI_HOME"] = self.tmp.name
class FakeArgs:
profile = "t"
yes = True
config = tai.Config(FakeArgs())
self.store = tai.Store(config, tai.Seal(config.home, "vault-test-1"))
def tearDown(self):
self.store.close()
self.tmp.cleanup()
os.environ.pop("TAI_HOME", None)
def test_store_list_delete_roundtrip(self):
app = FakeSecretApp(self.store)
tools = tai.Tools(app)
self.assertIn("stored secret 'api'", tools.dispatch("store_secret", json.dumps({"name": "api", "value": "token-abc-123"})))
self.assertEqual(self.store.load_secret("api"), "token-abc-123")
raw = self.store.db.execute("SELECT value FROM secrets WHERE name = 'api'").fetchone()[0]
self.assertTrue(raw.startswith("tai1$"))
self.assertNotIn("token-abc", raw)
listed = tools.dispatch("list_secrets", "{}")
self.assertIn("api", listed)
self.assertNotIn("token-abc", listed)
self.assertIn("deleted secret 'api'", tools.dispatch("delete_secret", json.dumps({"name": "api"})))
self.assertIsNone(self.store.load_secret("api"))
self.assertEqual(len(app.approvals), 1)
def test_validation(self):
tools = tai.Tools(FakeSecretApp(self.store))
self.assertIn("invalid secret name", tools.dispatch("store_secret", json.dumps({"name": "Bad Name!", "value": "x"})))
self.assertIn("empty value", tools.dispatch("store_secret", json.dumps({"name": "ok", "value": ""})))
self.assertIn("unknown secret", tools.dispatch("delete_secret", json.dumps({"name": "nope"})))
self.assertIn("unknown tool", tools.dispatch("get_secret", json.dumps({"name": "x"})))
def test_shell_blind_injection(self):
self.store.save_secret("demo", "injected-value-42")
app = FakeSecretApp(self.store)
tools = tai.Tools(app)
result = tools.dispatch("shell", json.dumps({"command": "echo $TAI_SECRET_DEMO", "secrets": ["demo"]}))
self.assertIn("[redacted:demo]", result)
self.assertNotIn("injected-value-42", result)
tools.dispatch("shell", json.dumps({"command": "echo $TAI_SECRET_DEMO", "secrets": ["demo"]}))
self.assertEqual(len(app.approvals), 1)
self.assertIn("unknown secret 'nope'", tools.dispatch("shell", json.dumps({"command": "echo hi", "secrets": ["nope"]})))
self.assertIn("must be a list", tools.dispatch("shell", json.dumps({"command": "echo hi", "secrets": "demo"})))
def test_shell_secret_denied(self):
self.store.save_secret("demo", "injected-value-42")
app = FakeSecretApp(self.store, approve=False)
result = tai.Tools(app).dispatch("shell", json.dumps({"command": "echo hi", "secrets": ["demo"]}))
self.assertEqual(result, "denied by user")
def test_web_fetch_auth(self):
self.store.save_secret("api", "fetch-token-7")
seen = {}
class FakeResponse:
def __enter__(self):
return self
def __exit__(self, *exc):
return False
def read(self, limit=0):
return b"<html><body>ok fetch-token-7 here</body></html>"
def fake_urlopen(request, timeout=30):
seen[request.full_url] = {key.lower(): value for key, value in request.headers.items()}
return FakeResponse()
app = FakeSecretApp(self.store)
tools = tai.Tools(app)
with mock.patch("urllib.request.urlopen", side_effect=fake_urlopen):
result = tools.dispatch("web_fetch", json.dumps({"url": "https://api.example.test/v1", "auth_secret": "api"}))
keyed = tools.dispatch("web_fetch", json.dumps({"url": "https://api.example.test/v2", "auth_secret": "api", "auth_header": "X-Api-Key", "auth_prefix": ""}))
first = seen["https://api.example.test/v1"]
second = seen["https://api.example.test/v2"]
self.assertEqual(first["authorization"], "Bearer fetch-token-7")
self.assertIn("tai/", first["user-agent"])
self.assertEqual(second["x-api-key"], "fetch-token-7")
self.assertIn("[redacted:api]", result)
self.assertNotIn("fetch-token-7", result)
self.assertIn("[redacted:api]", keyed)
self.assertEqual(len(app.approvals), 1)
self.assertIn("unknown secret", tools.dispatch("web_fetch", json.dumps({"url": "https://api.example.test/v1", "auth_secret": "nope"})))
self.assertIn("invalid auth header", tools.dispatch("web_fetch", json.dumps({"url": "https://api.example.test/v1", "auth_secret": "api", "auth_header": "Bad\nHeader"})))
def test_redact(self):
self.store.save_secret("long", "abcdefghij")
self.store.save_secret("short", "abc")
self.store.save_secret("sub", "cdef")
text = self.store.redact("see abcdefghij and abc here")
self.assertEqual(text, "see [redacted:long] and abc here")
def test_run_turn_scrubs_and_redacts(self):
class FakeArgs:
profile = "t"
yes = True
config = tai.Config(FakeArgs())
agent = tai.Agent(config, self.store, persist=False, quiet=True)
call_reply = {"role": "assistant", "content": "storing now hunter2-leak", "reasoning": "", "tool_calls": [{"id": "c9", "name": "store_secret", "arguments": json.dumps({"name": "leak", "value": "hunter2-leak"})}], "backend": "x"}
final_reply = {"role": "assistant", "content": "done", "reasoning": "", "tool_calls": [], "backend": "x"}
with mock.patch.object(agent.chat, "complete", side_effect=[call_reply, final_reply]):
result = agent.run_turn("remember it", capture=True)
self.assertEqual(result, "done")
blob = json.dumps(agent.messages)
self.assertNotIn("hunter2-leak", blob)
self.assertIn("[redacted:leak]", blob)
self.assertEqual(self.store.load_secret("leak"), "hunter2-leak")
self.assertEqual(self.store.search_events("t", "hunter2"), [])
def test_secret_repl(self):
agent = mock.Mock()
agent.store = self.store
agent.profile = "t"
agent.ask_approval = lambda command, guidance=True: True
agent.tools = tai.Tools(agent)
with mock.patch("getpass.getpass", return_value="repl-value-1"):
with mock.patch("builtins.print") as printer:
self.assertTrue(tai.handle_command(agent, "/secret set wifi"))
self.assertEqual(self.store.load_secret("wifi"), "repl-value-1")
printed = "\n".join(str(call.args[0]) for call in printer.call_args_list if call.args)
self.assertIn("stored secret 'wifi'", printed)
self.assertNotIn("repl-value-1", printed)
with mock.patch("builtins.print") as printer:
tai.handle_command(agent, "/secret list")
tai.handle_command(agent, "/secret")
listed = "\n".join(str(call.args[0]) for call in printer.call_args_list if call.args)
self.assertIn("- wifi", listed)
self.assertNotIn("repl-value-1", listed)
with mock.patch("builtins.print") as printer:
tai.handle_command(agent, "/secret delete wifi")
tai.handle_command(agent, "/secret delete nope")
tai.handle_command(agent, "/secret frobnicate")
tai.handle_command(agent, "/secret set Bad Name")
removed = "\n".join(str(call.args[0]) for call in printer.call_args_list if call.args)
self.assertIn("deleted secret 'wifi'", removed)
self.assertIn("unknown secret", removed)
self.assertIn("use /secret set", removed)
self.assertIn("invalid secret name", removed)
self.assertIsNone(self.store.load_secret("wifi"))
def test_legacy_memory_migrates_to_vault(self):
self.assertIn("store_secret", tai.DEFAULT_SYSTEM)
self.assertNotIn("collect and keep", tai.DEFAULT_SYSTEM)
class FakeArgs:
profile = "t"
yes = True
config = tai.Config(FakeArgs())
old_system = "You are tai. Memory: call remember whenever you learn durable facts, " + tai.LEGACY_PASSWORD_NOTE + "; call it with a forget instruction."
self.store.save_system("old", old_system)
agent = tai.Agent(config, self.store, persist=False, quiet=True)
agent.switch_profile("old", silent=True)
self.assertNotIn(tai.LEGACY_PASSWORD_NOTE, agent.system_message)
self.assertIn("store_secret", agent.system_message)
self.assertIn("store_secret", self.store.load_system("old"))
class FakeSchedulerApp:
def __init__(self, store, approve=True):
self.store = store
self.profile = "t"
self.approvals = []
self.approve = approve
def ask_approval(self, command):
self.approvals.append(command)
return self.approve
class SchedulerTests(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.seal = tai.Seal(self.config.home, "sched-test-1")
self.store = tai.Store(self.config, self.seal)
with tai.AGENTS_LOCK:
tai.AGENTS.clear()
tai.AGENTS_NEXT[0] = 1
def tearDown(self):
self.store.close()
self.tmp.cleanup()
os.environ.pop("TAI_HOME", None)
with tai.AGENTS_LOCK:
tai.AGENTS.clear()
tai.AGENTS_NEXT[0] = 1
def test_parse_at_and_duration(self):
parsed = tai.parse_schedule_at("2026-10-08T09:00")
self.assertEqual(datetime.fromisoformat(parsed).utcoffset().total_seconds(), 0)
self.assertEqual(tai.local_display(parsed), "2026-10-08 09:00")
with self.assertRaises(ValueError):
tai.parse_schedule_at("not a date")
self.assertEqual(tai.parse_duration("10m"), 600)
self.assertEqual(tai.parse_duration("2h"), 7200)
self.assertEqual(tai.parse_duration("1d"), 86400)
self.assertEqual(tai.parse_duration("45"), 45)
self.assertIsNone(tai.parse_duration("nope"))
self.assertEqual(tai.format_delay(4000), "1h")
self.assertEqual(tai.format_delay(-5), "overdue")
def test_schedule_validation(self):
tools = tai.Tools(FakeSchedulerApp(self.store))
self.assertIn("empty prompt", tools.dispatch("schedule", json.dumps({"at": "2026-10-08T09:00"})))
self.assertIn("exactly one", tools.dispatch("schedule", json.dumps({"prompt": "x"})))
self.assertIn("exactly one", tools.dispatch("schedule", json.dumps({"prompt": "x", "at": "2026-10-08T09:00", "every": 60})))
self.assertIn("invalid datetime", tools.dispatch("schedule", json.dumps({"prompt": "x", "at": "soon"})))
self.assertIn("at least 60", tools.dispatch("schedule", json.dumps({"prompt": "x", "every": 5})))
self.assertIn("invalid profile", tools.dispatch("schedule", json.dumps({"prompt": "x", "every": 60, "profile": "bad name!"})))
denied = tai.Tools(FakeSchedulerApp(self.store, approve=False))
self.assertEqual(denied.dispatch("schedule", json.dumps({"prompt": "x", "every": 60})), "denied by user")
self.assertEqual(self.store.list_schedules(), [])
def test_schedule_roundtrip(self):
app = FakeSchedulerApp(self.store)
tools = tai.Tools(app)
once = tools.dispatch("schedule", json.dumps({"prompt": "water plants", "name": "plants", "at": "2026-10-08T09:00"}))
self.assertIn("scheduled #1", once)
rep = tools.dispatch("schedule", json.dumps({"prompt": "check mail", "every": 3600}))
self.assertIn("scheduled #2", rep)
raw = self.store.db.execute("SELECT prompt FROM schedules WHERE id = 1").fetchone()[0]
self.assertTrue(raw.startswith("tai1$"))
listed = tools.dispatch("schedules", "{}")
self.assertIn("#1", listed)
self.assertIn("plants", listed)
self.assertIn("every 1h", listed)
self.assertIn("water plants", listed)
self.assertIn("deleted schedule #1", tools.dispatch("unschedule", json.dumps({"id": 1})))
self.assertIn("no schedule #1", tools.dispatch("unschedule", json.dumps({"id": 1})))
self.assertEqual(len(app.approvals), 3)
def test_tick_fires_once_async(self):
self.store.add_schedule("job", "do the thing", "t", 0, "2020-01-01T00:00:00+00:00", 60)
calls = []
def stub(task, profile, timeout):
calls.append((task, profile, timeout))
return "stub-result"
self.assertEqual(tai.scheduler_tick(self.config, self.seal, runner=stub), 1)
self.assertEqual(calls[0][0], "do the thing")
self.assertEqual(self.store.db.execute("SELECT status FROM schedules WHERE id = 1").fetchone()[0], "done")
self.assertEqual(tai.scheduler_tick(self.config, self.seal, runner=stub), 0)
self.assertEqual(len(calls), 1)
deadline = time.time() + 5
row = ("", "")
while time.time() < deadline:
row = self.store.db.execute("SELECT last_status, last_result FROM schedules WHERE id = 1").fetchone()
if row[0]:
break
time.sleep(0.05)
self.assertEqual(row[0], "done")
self.assertIn("stub-result", row[1])
def test_tick_repeat_advances_without_backfill(self):
past = (datetime.now(timezone.utc) - timedelta(days=2)).isoformat()
self.store.add_schedule("hourly", "ping", "t", 3600, past, 60)
calls = []
self.assertEqual(tai.scheduler_tick(self.config, self.seal, runner=lambda task, profile, timeout: calls.append(task) or "ok"), 1)
self.assertEqual(len(calls), 1)
row = self.store.db.execute("SELECT status, next_run FROM schedules WHERE id = 1").fetchone()
self.assertEqual(row[0], "pending")
self.assertGreater(datetime.fromisoformat(row[1]), datetime.now(timezone.utc))
def test_tick_skips_claimed(self):
self.store.add_schedule("job", "do it", "t", 0, "2020-01-01T00:00:00+00:00", 60)
self.store.db.execute("UPDATE schedules SET status = 'done' WHERE id = 1")
self.store.db.commit()
calls = []
self.assertEqual(tai.scheduler_tick(self.config, self.seal, runner=lambda task, profile, timeout: calls.append(task) or "ok"), 0)
self.assertEqual(calls, [])
def test_schedule_repl(self):
agent = mock.Mock()
agent.store = self.store
agent.profile = "t"
with mock.patch("builtins.print") as printer:
tai.handle_command(agent, "/schedule every 1h check mail")
tai.handle_command(agent, "/schedule at 2026-10-08T09:00 water plants")
created = "\n".join(str(call.args[0]) for call in printer.call_args_list if call.args)
self.assertIn("scheduled #1", created)
self.assertIn("scheduled #2", created)
with mock.patch("builtins.print") as printer:
tai.handle_command(agent, "/schedules")
listed = "\n".join(str(call.args[0]) for call in printer.call_args_list if call.args)
self.assertIn("#1", listed)
self.assertIn("check mail", listed)
with mock.patch("builtins.print") as printer:
tai.handle_command(agent, "/schedule soon x")
tai.handle_command(agent, "/schedule every 5s x")
tai.handle_command(agent, "/schedule at nope x")
tai.handle_command(agent, "/unschedule 1")
tai.handle_command(agent, "/unschedule 1")
tai.handle_command(agent, "/unschedule x")
errors = "\n".join(str(call.args[0]) for call in printer.call_args_list if call.args)
self.assertIn("use /schedule at", errors)
self.assertIn("at least 60", errors)
self.assertIn("invalid datetime", errors)
self.assertIn("deleted schedule #1", errors)
self.assertIn("no schedule #1", errors)
self.assertIn("use /unschedule", errors)
def test_delegation_wording(self):
self.assertIn("Delegation:", tai.DEFAULT_SYSTEM)
self.assertIn("fork background subagents", tai.DEFAULT_SYSTEM)
self.assertIn("schedule future work", tai.DEFAULT_SYSTEM)
class RecordGraphTests(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, "record-test-1"))
self.app = FakeSecretApp(self.store)
self.tools = tai.Tools(self.app)
with tai.AGENTS_LOCK:
tai.AGENTS.clear()
tai.AGENTS_NEXT[0] = 1
def tearDown(self):
self.store.close()
self.tmp.cleanup()
os.environ.pop("TAI_HOME", None)
with tai.AGENTS_LOCK:
tai.AGENTS.clear()
tai.AGENTS_NEXT[0] = 1
def mem_id(self, text):
match = re.search(r"mem:[0-9a-f]{16}", text)
self.assertIsNotNone(match, "no mem id in %r" % text)
return match.group(0)
def test_save_read_search_delete_roundtrip(self):
saved = self.tools.dispatch("record_save", json.dumps({"title": "site notes", "content": "alpha beta gamma", "kind": "note", "tags": ["Site Visit"]}))
record_id = self.mem_id(saved)
self.assertIn("16 chars", saved)
self.assertIn("site-visit", saved)
page = self.tools.dispatch("record_read", json.dumps({"id": record_id}))
self.assertIn("chars 0-16 of 16", page)
self.assertIn("alpha beta gamma", page)
found = self.tools.dispatch("record_search", json.dumps({"query": "beta"}))
self.assertIn(record_id, found)
self.assertIn("site notes", found)
deleted = self.tools.dispatch("record_delete", json.dumps({"id": record_id}))
self.assertIn("deleted record", deleted)
self.assertIn("unknown record", self.tools.dispatch("record_read", json.dumps({"id": record_id})))
self.assertEqual(self.app.approvals, ["delete record " + record_id])
def test_read_pages_offsets(self):
saved = self.tools.dispatch("record_save", json.dumps({"title": "paged", "content": "".join("%04d" % num for num in range(250))}))
record_id = self.mem_id(saved)
first = self.tools.dispatch("record_read", json.dumps({"id": record_id, "limit": 100}))
self.assertIn("chars 0-100 of 1000", first)
second = self.tools.dispatch("record_read", json.dumps({"id": record_id, "offset": 100, "limit": 100}))
self.assertIn("chars 100-200 of 1000", second)
self.assertNotEqual(first.splitlines()[-1], second.splitlines()[-1])
record = self.store.get_record(record_id)
self.assertGreaterEqual(record["reads"], 3)
def test_record_validation_errors(self):
self.assertIn("content is empty", self.tools.dispatch("record_save", json.dumps({"title": "x"})))
self.assertIn("unknown kind", self.tools.dispatch("record_save", json.dumps({"content": "x", "kind": "song"})))
self.assertIn("tags must be a list", self.tools.dispatch("record_save", json.dumps({"content": "x", "tags": "nope"})))
self.assertIn("unknown record", self.tools.dispatch("record_read", json.dumps({"id": "mem:0123456789abcdef"})))
self.assertIn("no records match", self.tools.dispatch("record_search", json.dumps({"query": "nothing-here-zzz"})))
def test_delete_denied_keeps_record(self):
saved = self.tools.dispatch("record_save", json.dumps({"content": "keep me"}))
record_id = self.mem_id(saved)
self.app.approve = False
self.assertIn("denied by user", self.tools.dispatch("record_delete", json.dumps({"id": record_id})))
self.assertIsNotNone(self.store.get_record(record_id))
def test_search_filters_kind_and_tags(self):
self.tools.dispatch("record_save", json.dumps({"title": "a", "content": "shared word", "kind": "note", "tags": ["team"]}))
self.tools.dispatch("record_save", json.dumps({"title": "b", "content": "shared word", "kind": "output", "tags": ["team"]}))
by_kind = self.tools.dispatch("record_search", json.dumps({"query": "shared", "kind": "output"}))
self.assertIn("[output]", by_kind)
self.assertNotIn("[note]", by_kind)
by_tag = self.tools.dispatch("record_search", json.dumps({"tags": ["team"]}))
self.assertIn("[output]", by_tag)
self.assertIn("[note]", by_tag)
missing = self.tools.dispatch("record_search", json.dumps({"tags": ["team", "other"]}))
self.assertIn("no records match", missing)
def test_graph_link_and_query(self):
first = self.mem_id(self.tools.dispatch("record_save", json.dumps({"title": "plan", "content": "the plan"})))
second = self.mem_id(self.tools.dispatch("record_save", json.dumps({"title": "bravo-log-title", "content": "the log"})))
self.tools.dispatch("store_secret", json.dumps({"name": "deploy", "value": "token-xyz-9"}))
linked = self.tools.dispatch("graph_link", json.dumps({"src": first, "dst": second, "relation": "Follows Up"}))
self.assertIn("follows-up", linked)
self.tools.dispatch("graph_link", json.dumps({"src": second, "dst": "secret:deploy", "relation": "uses"}))
view = self.tools.dispatch("graph_query", json.dumps({"node": first}))
self.assertIn(first + " :: plan", view)
self.assertIn(second, view)
self.assertIn("secret:deploy", view)
self.assertIn("bravo-log-title", view)
self.assertIn("unknown node", self.tools.dispatch("graph_query", json.dumps({"node": "mem:0123456789abcdef"})))
self.assertIn("unknown node", self.tools.dispatch("graph_link", json.dumps({"src": first, "dst": "mem:0123456789abcdef"})))
def test_traverse_caps_depth(self):
ids = [self.mem_id(self.tools.dispatch("record_save", json.dumps({"title": "n%d" % num, "content": "x"}))) for num in range(6)]
for pos in range(5):
self.store.add_edge(ids[pos], ids[pos + 1], "next")
shallow = self.store.traverse(ids[0], depth=1, limit=50)
self.assertEqual(len(shallow), 2)
deep = self.store.traverse(ids[0], depth=99, limit=50)
self.assertEqual(len(deep), 5)
self.assertEqual(max(item["depth"] for item in deep), 4)
def test_shell_output_spills_to_record(self):
result = self.tools.dispatch("shell", json.dumps({"command": "python3 -c \"print('0123456789' * 700)\""}))
self.assertIn("exit 0", result)
record_id = self.mem_id(result)
self.assertIn("record_read pages the rest", result)
page = self.tools.dispatch("record_read", json.dumps({"id": record_id, "limit": 20}))
self.assertIn("of 7000", page)
found = self.tools.dispatch("record_search", json.dumps({"tags": ["shell"]}))
self.assertIn(record_id, found)
def test_poll_spills_once_and_reuses(self):
big = "z" * 7000
agent_id = tai.spawn_agent("big task", "t", 60, None, None, 0, lambda task, profile, timeout: big)
first = self.tools.dispatch("poll", json.dumps({"id": agent_id, "wait": 5}))
record_id = self.mem_id(first)
second = self.tools.dispatch("poll", json.dumps({"id": agent_id}))
self.assertIn(record_id, second)
rows = self.store.db.execute("SELECT COUNT(*) FROM records").fetchone()[0]
self.assertEqual(rows, 1)
def test_record_repl_commands(self):
record_id = self.mem_id(self.tools.dispatch("record_save", json.dumps({"title": "repl note", "content": "visible words"})))
self.app.tools = self.tools
with mock.patch("builtins.print") as shown:
tai.handle_command(self.app, "/records visible")
tai.handle_command(self.app, "/record " + record_id)
tai.handle_command(self.app, "/graph " + record_id)
printed = "\n".join(str(call.args[0]) for call in shown.call_args_list)
self.assertIn(record_id, printed)
self.assertIn("repl note", printed)
def test_wal_mode_enabled(self):
mode = self.store.db.execute("PRAGMA journal_mode").fetchone()[0]
self.assertEqual(mode.lower(), "wal")
class SecretMetaTests(unittest.TestCase):
def setUp(self):
self.tmp = tempfile.TemporaryDirectory()
os.environ["TAI_HOME"] = self.tmp.name
class FakeArgs:
profile = "t"
yes = True
config = tai.Config(FakeArgs())
self.store = tai.Store(config, tai.Seal(config.home, "meta-test-1"))
self.tools = tai.Tools(FakeSecretApp(self.store))
def tearDown(self):
self.store.close()
self.tmp.cleanup()
os.environ.pop("TAI_HOME", None)
def test_metadata_roundtrip(self):
result = self.tools.dispatch("store_secret", json.dumps({"name": "db", "value": "pw-12345", "username": "ops", "host": "db.internal", "port": 5432, "notes": "primary", "tags": ["Prod DB"]}))
self.assertIn("stored secret 'db'", result)
listed = self.tools.dispatch("list_secrets", "{}")
self.assertIn("ops@db.internal:5432", listed)
self.assertIn("prod-db", listed)
self.assertNotIn("pw-12345", listed)
infos = self.store.list_secret_infos()
self.assertEqual(infos[0]["meta"].get("username"), "ops")
self.assertEqual(infos[0]["meta"].get("port"), 5432)
def test_minimum_info_suffices(self):
result = self.tools.dispatch("store_secret", json.dumps({"name": "plain", "value": "v-abcdef"}))
self.assertIn("stored secret 'plain'", result)
self.assertIn("- plain [secret]", self.tools.dispatch("list_secrets", "{}"))
class ScheduleTagTests(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, "schetag-1"))
self.tools = tai.Tools(FakeSchedulerApp(self.store))
def tearDown(self):
self.store.close()
self.tmp.cleanup()
os.environ.pop("TAI_HOME", None)
def test_schedule_tags_shown(self):
self.tools.dispatch("schedule", json.dumps({"prompt": "water plants", "every": 3600, "tags": ["Home Chores"]}))
listed = self.tools.dispatch("schedules", "{}")
self.assertIn("home-chores", listed)
items = self.store.list_schedules()
self.assertIn("schedule", items[0]["tags"])
class ModeTests(unittest.TestCase):
def setUp(self):
self.tmp = tempfile.TemporaryDirectory()
os.environ["TAI_HOME"] = self.tmp.name
class FakeArgs:
profile = "t"
yes = False
self.config = tai.Config(FakeArgs())
self.store = tai.Store(self.config, tai.Seal(self.config.home, "mode-test-1"))
def tearDown(self):
self.store.close()
self.tmp.cleanup()
os.environ.pop("TAI_HOME", None)
def test_interactive_always_enables_yolo(self):
agent = tai.Agent(self.config, self.store, persist=True, quiet=True)
with mock.patch("sys.stdin") as fake_stdin:
fake_stdin.isatty.return_value = True
with mock.patch("builtins.input", return_value="Y"):
with mock.patch("builtins.print"):
self.assertTrue(agent.ask_approval("do it"))
self.assertTrue(agent.yolo)
with mock.patch("builtins.input", side_effect=AssertionError("must not ask")):
self.assertTrue(agent.ask_approval("do it again"))
def test_yolo_never_asks(self):
agent = tai.Agent(self.config, self.store, persist=True, quiet=True, yolo=True)
with mock.patch("builtins.input", side_effect=AssertionError("must not ask")):
self.assertTrue(agent.ask_approval("anything"))
def test_auto_gets_autonomous_note(self):
agent = tai.Agent(self.config, self.store, persist=True, quiet=True, auto=True)
agent.apply_system()
self.assertIn("Autonomous mode", agent.messages[0]["content"])
class AuditTests(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, "audit-test-1"))
self.app = FakeSecretApp(self.store)
self.app.skills = {}
self.tools = tai.Tools(self.app)
def tearDown(self):
self.store.close()
self.tmp.cleanup()
os.environ.pop("TAI_HOME", None)
def workfile(self, name="note.txt", content="v1"):
path = os.path.join(self.tmp.name, name)
with open(path, "w", encoding="utf-8") as handle:
handle.write(content)
return path
def test_write_audits_and_records(self):
path = os.path.join(self.tmp.name, "fresh.txt")
self.assertIn("wrote", self.tools.dispatch("write_file", json.dumps({"path": path, "content": "hello"})))
rows = self.store.audit_history(path)
self.assertEqual(len(rows), 1)
self.assertEqual(rows[0]["action"], "write")
self.assertIsNone(rows[0]["old_size"])
self.assertEqual(rows[0]["new_size"], 5)
record_id = self.store.file_record_id(path)
self.assertIsNotNone(record_id)
found = self.tools.dispatch("record_search", json.dumps({"tags": ["txt"]}))
self.assertIn(record_id, found)
self.tools.dispatch("read_file", json.dumps({"path": path}))
self.tools.dispatch("write_file", json.dumps({"path": path, "content": "hello again"}))
rows = self.store.audit_history(path)
self.assertEqual(len(rows), 2)
self.assertEqual(rows[0]["old_size"], 5)
self.assertEqual(rows[0]["new_size"], 11)
def test_edit_captures_old_and_new(self):
path = self.workfile()
self.tools.dispatch("read_file", json.dumps({"path": path}))
self.assertIn("edited", self.tools.dispatch("edit_file", json.dumps({"path": path, "find": "v1", "replace": "v2"})))
row = self.store.audit_get(self.store.audit_history(path)[0]["id"])
self.assertEqual(row["action"], "edit")
self.assertEqual(row["old"], "v1")
self.assertEqual(row["new"], "v2")
def test_big_file_truncates_with_marker(self):
path = os.path.join(self.tmp.name, "big.txt")
content = "x" * 60000
self.tools.dispatch("write_file", json.dumps({"path": path, "content": content}))
row = self.store.audit_get(self.store.audit_history(path)[0]["id"])
self.assertEqual(row["new_size"], 60000)
self.assertIn("[truncated, full size 60000 bytes]", row["new"])
record = self.store.get_record(self.store.file_record_id(path))
self.assertIn("[truncated, full size 60000 bytes]", record["content"])
def test_shell_rm_snapshots_before_delete(self):
first = self.workfile("a.txt", "alpha")
second = self.workfile("b.txt", "beta")
result = self.tools.dispatch("shell", json.dumps({"command": "rm a.txt b.txt", "workdir": self.tmp.name}))
self.assertIn("exit 0", result)
self.assertFalse(os.path.exists(first))
snaps = self.store.audit_history(tag="shell")
self.assertEqual(len(snaps), 2)
by_path = {row["path"]: row for row in snaps}
full = self.store.audit_get(by_path[first]["id"])
self.assertEqual(full["old"], "alpha")
self.assertIsNone(full["new"])
self.assertIn("rm a.txt b.txt", full["message"])
def test_shell_redirect_snapshots_target(self):
path = self.workfile("out.txt", "old words")
self.tools.dispatch("shell", json.dumps({"command": "echo new > out.txt", "workdir": self.tmp.name}))
snaps = self.store.audit_history(tag="shell")
self.assertEqual(len(snaps), 1)
full = self.store.audit_get(snaps[0]["id"])
self.assertEqual(full["old"], "old words")
def test_shell_parser_units(self):
self.assertEqual(tai.segment_targets(["rm", "-rf", "a", "b"]), ["a", "b"])
self.assertEqual(tai.segment_targets(["sudo", "rm", "x"]), ["x"])
self.assertEqual(tai.segment_targets(["VAR=1", "cmd", ">", "out"]), ["out"])
self.assertEqual(tai.segment_targets(["tee", "t1", "t2"]), ["t1", "t2"])
self.assertEqual(tai.segment_targets(["mv", "a", "b"]), ["b"])
self.assertEqual(tai.segment_targets(["cmd", "2>", "/dev/null"]), [])
self.assertEqual(tai.segment_targets(["truncate", "-s", "0", "f"]), ["f"])
self.assertEqual(tai.segment_targets(["dd", "if=a", "of=b"]), ["b"])
sub = os.path.join(self.tmp.name, "sub")
os.makedirs(sub)
self.workfile("sub/one.txt", "1")
self.workfile("sub/two.txt", "2")
hits = tai.shell_target_paths("rm -rf sub", self.tmp.name)
self.assertEqual(len(hits), 2)
first = self.workfile("t1.txt", "1")
piped = tai.shell_target_paths("echo x | tee t1.txt missing.txt", self.tmp.name)
self.assertEqual(piped, [first])
def test_delete_roundtrip_and_deny(self):
gone = self.workfile("gone.txt", "bye")
kept = self.workfile("kept.txt", "hi")
self.tools.dispatch("read_file", json.dumps({"path": gone}))
self.tools.dispatch("read_file", json.dumps({"path": kept}))
self.assertIn("deleted", self.tools.dispatch("delete_file", json.dumps({"path": gone})))
self.assertFalse(os.path.exists(gone))
self.assertIsNone(self.store.file_record_id(gone))
row = self.store.audit_get(self.store.audit_history(gone)[0]["id"])
self.assertEqual(row["action"], "delete")
self.assertEqual(row["old"], "bye")
self.app.approve = False
self.assertIn("denied by user", self.tools.dispatch("delete_file", json.dumps({"path": kept})))
self.assertTrue(os.path.exists(kept))
def test_audit_tool_filters(self):
first = self.workfile("f1.txt", "one")
self.workfile("f2.txt", "two")
self.tools.dispatch("read_file", json.dumps({"path": first}))
self.tools.dispatch("edit_file", json.dumps({"path": first, "find": "one", "replace": "uno"}))
by_path = self.tools.dispatch("audit", json.dumps({"path": first}))
self.assertIn("edit", by_path)
self.assertNotIn("f2.txt", by_path)
self.assertIn("invalid limit", self.tools.dispatch("audit", json.dumps({"limit": "zzz"})))
self.assertIn("no audit rows match", self.tools.dispatch("audit", json.dumps({"path": "/nope/nothing"})))
def test_restore_post_pre_and_undelete(self):
path = self.workfile("time.txt", "v1")
self.tools.dispatch("read_file", json.dumps({"path": path}))
self.tools.dispatch("edit_file", json.dumps({"path": path, "find": "v1", "replace": "v2"}))
rows = {row["action"]: row for row in self.store.audit_history(path)}
self.assertIn("restored", self.tools.dispatch("restore", json.dumps({"id": rows["edit"]["id"]})))
with open(path, encoding="utf-8") as handle:
self.assertEqual(handle.read(), "v2")
self.tools.dispatch("shell", json.dumps({"command": "rm time.txt", "workdir": self.tmp.name}))
snap = self.store.audit_history(path, tag="shell")[0]
self.assertIn("restored", self.tools.dispatch("restore", json.dumps({"id": snap["id"]})))
with open(path, encoding="utf-8") as handle:
self.assertEqual(handle.read(), "v2")
history = self.store.audit_history(path)
self.assertEqual(history[0]["action"], "restore")
self.assertIn("restored from audit", history[0]["message"])
def test_restore_refuses_truncated(self):
path = os.path.join(self.tmp.name, "huge.txt")
self.tools.dispatch("write_file", json.dumps({"path": path, "content": "y" * 60000}))
row_id = self.store.audit_history(path)[0]["id"]
result = self.tools.dispatch("restore", json.dumps({"id": row_id}))
self.assertIn("truncated", result)
self.assertIn("cannot restore safely", result)
def test_restore_repl_command(self):
path = self.workfile("repl.txt", "keep")
self.tools.dispatch("read_file", json.dumps({"path": path}))
self.tools.dispatch("edit_file", json.dumps({"path": path, "find": "keep", "replace": "changed"}))
row_id = self.store.audit_history(path)[0]["id"]
self.app.tools = self.tools
with mock.patch("builtins.print") as shown:
tai.handle_command(self.app, "/audit " + path)
tai.handle_command(self.app, "/restore %d" % row_id)
tai.handle_command(self.app, "/restore nope")
printed = "\n".join(str(call.args[0]) for call in shown.call_args_list)
self.assertIn("edit", printed)
self.assertIn("restored", printed)
self.assertIn("use /restore <audit id>", printed)
def test_lazy_blueprint_builds_on_load(self):
with mock.patch.object(self.tools, "run_create_skill", return_value="created") as creator:
result = self.tools.dispatch("load_skill", json.dumps({"name": "bot-creator"}))
self.assertIn("building skill 'bot-creator' from blueprint", result)
creator.assert_called_once()
sent = creator.call_args.args[0]
self.assertEqual(sent["name"], "bot-creator")
self.assertIn("Deep-research", sent["brief"])
self.assertIn("unknown skill", self.tools.dispatch("load_skill", json.dumps({"name": "nope"})))
class ReleaseTests(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, "release-test-1"))
def tearDown(self):
self.store.close()
self.tmp.cleanup()
os.environ.pop("TAI_HOME", None)
def script(self, version="1.2.3"):
path = os.path.join(self.tmp.name, "tai-copy.py")
with open(path, "w", encoding="utf-8") as handle:
handle.write("#!/usr/bin/env python3\nVERSION = \"%s\"\nprint('hi')\n" % version)
return path
def test_bump_parts(self):
self.assertEqual(tai.release_bump(self.script(), "patch"), ("1.2.3", "1.2.4"))
self.assertEqual(tai.release_bump(self.script(), "minor"), ("1.2.3", "1.3.0"))
self.assertEqual(tai.release_bump(self.script(), "major"), ("1.2.3", "2.0.0"))
with open(self.script(), encoding="utf-8") as handle:
self.assertIn('VERSION = "1.2.3"', handle.read())
with self.assertRaises(ValueError):
tai.release_bump(self.script(), "banana")
flat = os.path.join(self.tmp.name, "flat.py")
with open(flat, "w", encoding="utf-8") as handle:
handle.write("no version here\n")
with self.assertRaises(ValueError):
tai.release_bump(flat, "patch")
def test_do_release_end_to_end(self):
path = self.script()
folder = tai.backups_dir(self.tmp.name)
old, new, dest, row_id = tai.do_release(self.store, path, folder, "minor", "add audit trail")
self.assertEqual((old, new), ("1.2.3", "1.3.0"))
self.assertTrue(dest.startswith(folder))
self.assertIn("tai-1.3.0-", os.path.basename(dest))
with open(path, encoding="utf-8") as handle:
text = handle.read()
self.assertIn('VERSION = "1.3.0"', text)
self.assertIn("print('hi')", text)
row = self.store.audit_get(row_id)
self.assertEqual(row["action"], "release")
self.assertEqual(row["message"], "add audit trail")
self.assertEqual(row["old"], "1.2.3")
self.assertEqual(row["new"], "1.3.0")
self.assertIn("release", row["tags"])
with self.assertRaises(ValueError):
tai.do_release(self.store, path, folder, "patch", " ")
def test_backup_prune_and_skip(self):
folder = tai.backups_dir(self.tmp.name)
os.makedirs(folder)
now = time.time()
for pos in range(12):
name = os.path.join(folder, "tai-1.0.0-20200101T00000%dZ-abc%d.py" % (pos, pos))
with open(name, "w", encoding="utf-8") as handle:
handle.write("old")
stamp = now - (12 - pos) * 60
os.utime(name, (stamp, stamp))
path = self.script()
dest, _digest = tai.snapshot_self(path, folder, "1.2.3")
kept = [entry for entry in os.listdir(folder) if entry.endswith(".py")]
self.assertEqual(len(kept), tai.BACKUP_KEEP)
self.assertIn(os.path.basename(dest), kept)
with open(path, "w", encoding="utf-8") as handle:
handle.write("#!/usr/bin/env python3\nVERSION = \"9.9.9\"\nprint('changed')\n")
first = tai.ensure_self_backup(self.config, self.store, script_file=path)
self.assertIsNotNone(first)
again = tai.ensure_self_backup(self.config, self.store, script_file=path)
self.assertIsNone(again)
snaps = self.store.audit_history(tag="snapshot")
self.assertEqual(len(snaps), 1)
self.assertIn("backed up as", snaps[0]["message"])
def test_release_tool_validates_safely(self):
app = FakeSecretApp(self.store, approve=False)
tools = tai.Tools(app)
self.assertIn("part must be", tools.dispatch("release", json.dumps({"part": "banana", "message": "x"})))
self.assertIn("message is required", tools.dispatch("release", json.dumps({"part": "patch", "message": ""})))
denied = tools.dispatch("release", json.dumps({"part": "patch", "message": "try bump"}))
self.assertIn("denied by user", denied)
self.assertEqual(len(app.approvals), 1)
self.assertIn("release %s (patch)" % tai.next_version(tai.VERSION, "patch"), app.approvals[0])
def test_boot_creates_self_backup(self):
import argparse as ap
args = ap.Namespace(profile="t", yes=True, yolo=False, auto=False)
with mock.patch.dict(os.environ, {"TAI_PASSPHRASE": "release-test-1"}):
with mock.patch("builtins.print"):
agent = tai.boot(args)
try:
names = os.listdir(tai.backups_dir(self.tmp.name))
self.assertEqual(len(names), 1)
self.assertTrue(names[0].startswith("tai-%s-" % tai.VERSION))
snaps = agent.store.audit_history(tag="snapshot")
self.assertEqual(len(snaps), 1)
finally:
agent.store.close()
class TaggingTests(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, "tag-test-1"))
self.app = FakeSecretApp(self.store)
self.tools = tai.Tools(self.app)
def tearDown(self):
self.store.close()
self.tmp.cleanup()
os.environ.pop("TAI_HOME", None)
def mem_id(self, text):
return re.search(r"mem:[0-9a-f]{16}", text).group(0)
def test_singular_units(self):
self.assertEqual(tai.singular_noun("servers"), "server")
self.assertEqual(tai.singular_noun("cities"), "city")
self.assertEqual(tai.singular_noun("boxes"), "box")
self.assertEqual(tai.singular_noun("branches"), "branch")
self.assertEqual(tai.singular_noun("glass"), "glass")
self.assertEqual(tai.singular_noun("status"), "status")
self.assertEqual(tai.singular_noun("news"), "news")
self.assertEqual(tai.singular_noun("physics"), "physics")
self.assertEqual(tai.singular_noun("api"), "api")
def test_plural_merges_into_known_singular(self):
first = self.mem_id(self.tools.dispatch("record_save", json.dumps({"title": "one", "content": "about a server", "tags": ["server"]})))
self.assertIn("server", self.store.item_tags(first))
second = self.mem_id(self.tools.dispatch("record_save", json.dumps({"title": "two", "content": "more", "tags": ["Servers"]})))
tags = self.store.item_tags(second)
self.assertIn("server", tags)
self.assertNotIn("servers", tags)
self.assertNotIn("servers", self.store.known_tags())
def test_irregular_words_never_corrupt(self):
record_id = self.mem_id(self.tools.dispatch("record_save", json.dumps({"title": "words", "content": "plain", "tags": ["news", "glass", "status", "physics"]})))
tags = self.store.item_tags(record_id)
for word in ("news", "glass", "status", "physics"):
self.assertIn(word, tags)
def test_known_words_attach_automatically(self):
self.tools.dispatch("record_save", json.dumps({"title": "seed", "content": "nothing yet", "tags": ["deploy", "server"]}))
record_id = self.mem_id(self.tools.dispatch("record_save", json.dumps({"title": "friday", "content": "we deploy the thing friday", "tags": []})))
self.assertIn("deploy", self.store.item_tags(record_id))
self.assertNotIn("friday", self.store.item_tags(record_id))
plural = self.mem_id(self.tools.dispatch("record_save", json.dumps({"title": "fleet", "content": "all servers rebooted", "tags": []})))
self.assertIn("server", self.store.item_tags(plural))
def test_new_records_link_to_same_tag_neighbors(self):
first = self.mem_id(self.tools.dispatch("record_save", json.dumps({"title": "a", "content": "deploy alpha", "tags": ["deploy"]})))
saved = self.tools.dispatch("record_save", json.dumps({"title": "b", "content": "deploy beta", "tags": ["deploy"]}))
second = self.mem_id(saved)
self.assertIn(first, saved)
edges = self.store.edges_for(second)
self.assertEqual(len(edges), 1)
self.assertEqual(edges[0]["relation"], "shares-deploy")
self.assertEqual(edges[0]["other"], first)
def test_baseline_only_records_do_not_link(self):
self.tools.dispatch("record_save", json.dumps({"title": "a", "content": "lorem ipsum", "tags": []}))
second = self.mem_id(self.tools.dispatch("record_save", json.dumps({"title": "b", "content": "dolor sit", "tags": []})))
self.assertEqual(self.store.edges_for(second), [])
def test_auto_link_caps_at_three(self):
for pos in range(5):
self.tools.dispatch("record_save", json.dumps({"title": "n%d" % pos, "content": "shared topic here", "tags": ["topic"]}))
sixth = self.mem_id(self.tools.dispatch("record_save", json.dumps({"title": "n5", "content": "shared topic again", "tags": ["topic"]})))
edges = [edge for edge in self.store.edges_for(sixth) if edge["direction"] == "out"]
self.assertEqual(len(edges), 3)
def test_search_expands_singular_plural(self):
legacy = self.store.add_record("note", "old", "zeta-content", ["servers"])
self.assertIn("servers", self.store.item_tags(legacy))
modern = self.mem_id(self.tools.dispatch("record_save", json.dumps({"title": "new", "content": "zeta-content", "tags": ["server"]})))
found = self.tools.dispatch("record_search", json.dumps({"tags": ["server"]}))
self.assertIn(legacy, found)
self.assertIn(modern, found)
def test_tags_tool_counts_and_prefix(self):
self.tools.dispatch("record_save", json.dumps({"title": "a", "content": "x", "tags": ["deploy", "friday"]}))
self.tools.dispatch("record_save", json.dumps({"title": "b", "content": "y", "tags": ["deploy"]}))
listed = self.tools.dispatch("tags", "{}")
self.assertIn("deploy (2)", listed)
self.assertIn("friday (1)", listed)
self.assertLess(listed.index("deploy (2)"), listed.index("friday (1)"))
prefixed = self.tools.dispatch("tags", json.dumps({"prefix": "fri"}))
self.assertIn("friday (1)", prefixed)
self.assertNotIn("deploy", prefixed)
self.app.tools = self.tools
with mock.patch("builtins.print") as shown:
tai.handle_command(self.app, "/tags")
printed = "\n".join(str(call.args[0]) for call in shown.call_args_list)
self.assertIn("deploy (2)", printed)
class ProfileTests(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, "profile-test-1"))
self.app = FakeSecretApp(self.store)
self.tools = tai.Tools(self.app)
with tai.AGENTS_LOCK:
tai.AGENTS.clear()
tai.AGENTS_NEXT[0] = 1
def tearDown(self):
self.store.close()
self.tmp.cleanup()
os.environ.pop("TAI_HOME", None)
with tai.AGENTS_LOCK:
tai.AGENTS.clear()
tai.AGENTS_NEXT[0] = 1
def mem_id(self, text):
return re.search(r"mem:[0-9a-f]{16}", text).group(0)
def test_records_isolated(self):
record_id = self.mem_id(self.tools.dispatch("record_save", json.dumps({"title": "mine", "content": "t-only words", "tags": ["t-tag"]})))
self.assertIsNone(self.store.get_record(record_id, "other"))
self.assertEqual(self.store.search_records("t-only", profile="other"), [])
self.assertEqual(len(self.store.search_records("t-only", profile="t")), 1)
self.app.profile = "other"
self.assertIn("no records match", self.tools.dispatch("record_search", json.dumps({"query": "t-only"})))
self.assertIn("unknown record", self.tools.dispatch("record_read", json.dumps({"id": record_id})))
def test_secrets_isolated(self):
self.tools.dispatch("store_secret", json.dumps({"name": "api", "value": "t-value-123"}))
self.assertIsNone(self.store.load_secret("api", "other"))
self.assertEqual(self.store.list_secret_infos("other"), [])
self.store.save_secret("api", "other-value-456", None, None, "other")
self.assertEqual(self.store.load_secret("api", "t"), "t-value-123")
self.assertEqual(self.store.load_secret("api", "other"), "other-value-456")
self.assertIn("[redacted:api]", self.store.redact("leak t-value-123 here", "t"))
self.assertIn("t-value-123", self.store.redact("leak t-value-123 here", "other"))
self.app.profile = "other"
listed = self.tools.dispatch("list_secrets", "{}")
self.assertIn("api", listed)
self.assertNotIn("t-value-123", listed)
def test_tags_vocab_isolated(self):
self.tools.dispatch("record_save", json.dumps({"title": "seed", "content": "nothing", "tags": ["deploy"]}))
self.app.profile = "other"
record_id = self.mem_id(self.tools.dispatch("record_save", json.dumps({"title": "fresh", "content": "we deploy friday", "tags": []})))
self.assertNotIn("deploy", self.store.item_tags(record_id, "other"))
self.assertEqual(self.store.tag_counts(profile="other"), [("record", 1)])
self.assertIn("deploy", self.store.known_tags("t"))
def test_edges_isolated(self):
first = self.mem_id(self.tools.dispatch("record_save", json.dumps({"title": "a", "content": "x", "tags": ["linkable"]})))
second = self.mem_id(self.tools.dispatch("record_save", json.dumps({"title": "b", "content": "y", "tags": ["unlinked"]})))
self.tools.dispatch("graph_link", json.dumps({"src": first, "dst": second, "relation": "uses"}))
self.app.profile = "other"
self.assertIn("unknown node", self.tools.dispatch("graph_query", json.dumps({"node": first})))
self.assertEqual(self.store.edges_for(first, "other"), [])
self.app.profile = "t"
view = self.tools.dispatch("graph_query", json.dumps({"node": first}))
self.assertIn("uses", view)
def test_audit_isolated(self):
path = os.path.join(self.tmp.name, "aud.txt")
self.tools.dispatch("write_file", json.dumps({"path": path, "content": "t-data"}))
self.assertEqual(len(self.store.audit_history(profile="t")), 1)
self.assertEqual(self.store.audit_history(profile="other"), [])
row_id = self.store.audit_history(profile="t")[0]["id"]
self.assertIsNone(self.store.audit_get(row_id, "other"))
self.app.profile = "other"
self.assertIn("no audit rows match", self.tools.dispatch("audit", "{}"))
self.assertIn("unknown audit", self.tools.dispatch("restore", json.dumps({"id": row_id})))
def test_schedules_isolated(self):
app = FakeSchedulerApp(self.store)
tools = tai.Tools(app)
tools.dispatch("schedule", json.dumps({"prompt": "t-job", "every": 3600}))
self.assertEqual(len(self.store.list_schedules("t")), 1)
self.assertEqual(self.store.list_schedules("other"), [])
row_id = self.store.list_schedules("t")[0]["id"]
app.profile = "other"
self.assertIn("no schedule", tools.dispatch("unschedule", json.dumps({"id": row_id})))
self.assertFalse(self.store.remove_schedule(row_id, "other"))
self.assertTrue(self.store.remove_schedule(row_id, "t"))
self.assertIn("another profile", tools.dispatch("schedule", json.dumps({"prompt": "x", "every": 60, "profile": "t"})))
def test_fork_and_poll_scoped(self):
first = tai.spawn_agent("task-t", "t", 60, None, None, 0, lambda task, profile, timeout: "done-t")
second = tai.spawn_agent("task-other", "other", 60, None, None, 0, lambda task, profile, timeout: "done-other")
self.assertEqual(tai.poll_agent(first, wait=5, profile="t")[0], "done")
self.assertEqual(tai.poll_agent(second, profile="t")[0], "missing")
self.assertEqual([item["id"] for item in tai.list_agents("t")], [first])
app = FakePrincipal()
app.store = self.store
tools = tai.Tools(app)
self.assertIn("done-t", tools.dispatch("poll", json.dumps({"id": first})))
self.assertIn("missing", tools.dispatch("poll", json.dumps({"id": second})))
self.assertIn("another profile", tools.dispatch("fork", json.dumps({"task": "x", "profile": "other"})))
def test_switch_resets_identity(self):
agent = tai.Agent(self.config, self.store, persist=True, quiet=True)
agent.tools.read_files.add(("home", "/tmp/x"))
agent.tools.secret_grants.add(("shell", "", ("api",)))
with mock.patch("builtins.print"):
agent.switch_profile("newbie")
self.assertEqual(agent.profile, "newbie")
self.assertEqual(self.store.profile, "newbie")
self.assertEqual(agent.tools.read_files, set())
self.assertEqual(agent.tools.secret_grants, set())
self.assertEqual([item["role"] for item in agent.messages], ["system"])
self.assertIn("newbie", self.store.list_profiles())
self.assertIn("t", self.store.list_profiles())
def test_recall_scoped(self):
self.store.log_event("t", "user", "message", "t recall marker")
self.store.log_event("other", "user", "message", "other recall marker")
found = self.tools.dispatch("recall", json.dumps({"query": "recall marker"}))
self.assertIn("t recall marker", found)
self.assertNotIn("other recall marker", found)
def test_profile_list_is_global(self):
self.store.save_system("alpha", "system a")
self.store.save_system("beta", "system b")
self.assertIn("alpha", self.store.list_profiles())
self.assertIn("beta", self.store.list_profiles())
class MigrationTests(unittest.TestCase):
def setUp(self):
self.tmp = tempfile.TemporaryDirectory()
os.environ["TAI_HOME"] = self.tmp.name
class FakeArgs:
profile = "default"
yes = True
self.config = tai.Config(FakeArgs())
def tearDown(self):
self.tmp.cleanup()
os.environ.pop("TAI_HOME", None)
def test_legacy_db_migrates_to_default(self):
db = sqlite3.connect(self.config.db_path)
db.execute("CREATE TABLE secrets (name TEXT PRIMARY KEY, value TEXT, meta TEXT, updated TEXT)")
db.execute("INSERT INTO secrets VALUES ('k', 'v-legacy', '{}', '2026-01-01')")
db.execute("CREATE TABLE records (id TEXT PRIMARY KEY, kind TEXT, title TEXT, content TEXT, size INTEGER, reads INTEGER, created TEXT, updated TEXT)")
db.execute("INSERT INTO records VALUES ('mem:0123456789abcdef', 'note', 'old', 'old words', 9, 0, '', '')")
db.execute("CREATE TABLE tags (item TEXT, tag TEXT, PRIMARY KEY (item, tag))")
db.execute("INSERT INTO tags VALUES ('mem:0123456789abcdef', 'legacy')")
db.execute("CREATE TABLE edges (src TEXT, dst TEXT, relation TEXT, created TEXT, PRIMARY KEY (src, dst, relation))")
db.execute("CREATE TABLE audit (id INTEGER PRIMARY KEY, ts TEXT, actor TEXT, action TEXT, path TEXT, message TEXT, old_size INTEGER, new_size INTEGER, old TEXT, new TEXT, tags TEXT)")
db.execute("INSERT INTO audit VALUES (1, 't', 'main', 'write', '/x', '', 0, 9, NULL, 'old words', '')")
db.execute("CREATE TABLE events (id INTEGER PRIMARY KEY, profile TEXT, ts TEXT, role TEXT, kind TEXT, text TEXT, tags TEXT)")
db.execute("CREATE TABLE 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)")
db.commit()
db.close()
store = tai.Store(self.config, tai.Seal(self.config.home, ""))
try:
self.assertEqual(store.load_secret("k"), "v-legacy")
store.save_secret("k", "v-work", None, None, "work")
self.assertEqual(store.load_secret("k", "work"), "v-work")
self.assertEqual(store.load_secret("k"), "v-legacy")
self.assertEqual(store.get_record("mem:0123456789abcdef")["content"], "old words")
self.assertIn("legacy", store.item_tags("mem:0123456789abcdef"))
self.assertEqual(len(store.audit_history()), 1)
self.assertEqual(tai.table_pk_columns(store.db, "secrets"), ["profile", "name"])
self.assertEqual(tai.table_pk_columns(store.db, "tags"), ["item", "tag", "profile"])
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):
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, "lazy-test-1"))
def tearDown(self):
self.store.close()
self.tmp.cleanup()
os.environ.pop("TAI_HOME", None)
def names(self, text):
return {schema["function"]["name"] for schema in tai.select_tools(text)}
def test_core_minimal(self):
self.assertEqual(self.names("hello"), set(tai.CORE_TOOLS))
self.assertEqual(len(tai.CORE_TOOLS), 8)
def test_synonym_triggers(self):
pairs = [
("remove that file", "delete_file"),
("set a cron reminder", "schedule"),
("my password", "store_secret"),
("undo that change", "restore"),
("show version history", "audit"),
("what tags exist", "tags"),
("publish a release", "release"),
("search the web", "web_search"),
("transcribe this", "listen"),
("check the host specs", "sysinfo"),
("upcoming agenda", "schedules"),
("connect these nodes", "graph_link"),
]
for text, tool in pairs:
self.assertIn(tool, self.names(text), "missing %s for %r" % (tool, text))
def test_name_mention_loads_family(self):
found = self.names("use record_save for this")
for tool in ("record_save", "record_read", "record_search", "record_delete"):
self.assertIn(tool, found)
def test_no_overtrigger(self):
self.assertEqual(self.names("hello, how are you today"), set(tai.CORE_TOOLS))
self.assertNotIn("edit_file", self.names("tell me about this"))
def test_catalog_lists_lazy(self):
catalog = tai.tool_catalog()
self.assertIn("## Tool catalog", catalog)
for schema in tai.TOOL_SCHEMAS:
name = schema["function"]["name"]
if name in tai.CORE_TOOLS:
continue
self.assertIn("- %s:" % name, catalog)
def test_tool_results_feed_selection(self):
messages = [
{"role": "user", "content": "run it"},
{"role": "assistant", "content": None, "tool_calls": []},
{"role": "tool", "content": "exit 0\n...[7000 chars spilled to mem:0123456789abcdef, record_read pages the rest]..."},
]
text = tai.conversation_text(messages)
self.assertIn("record_read", text)
self.assertIn("record_read", self.names(text))
def test_run_turn_uses_filtered_payload(self):
agent = tai.Agent(self.config, self.store, persist=False, quiet=True)
seen = {}
reply = {"role": "assistant", "content": "hi", "reasoning": "", "tool_calls": [], "backend": "x"}
def fake_complete(messages, tools, stream_sink=None):
seen["tools"] = {schema["function"]["name"] for schema in tools}
return dict(reply)
with mock.patch.object(agent.chat, "complete", side_effect=fake_complete):
agent.run_turn("hi", capture=True)
self.assertEqual(seen["tools"], set(tai.CORE_TOOLS))
def test_tools_repl(self):
agent = mock.Mock()
with mock.patch("builtins.print") as shown:
tai.handle_command(agent, "/tools")
tai.handle_command(agent, "/tools schedule")
tai.handle_command(agent, "/tools nope")
printed = "\n".join(str(call.args[0]) for call in shown.call_args_list)
self.assertIn("core, always loaded", printed)
self.assertIn("schedule", printed)
self.assertIn("cron", printed)
self.assertIn("unknown tool", printed)
class InstallOpsTests(unittest.TestCase):
def setUp(self):
self.tmp = tempfile.TemporaryDirectory()
self.home = os.path.join(self.tmp.name, "home")
os.makedirs(self.home)
self.env = mock.patch.dict(os.environ, {"HOME": self.home, "TAI_HOME": os.path.join(self.home, ".tai")})
self.env.start()
class FakeArgs:
profile = "t"
yes = True
self.config = tai.Config(FakeArgs())
self.store = tai.Store(self.config, tai.Seal(self.config.home, "install-test-1"))
self.app = FakeSecretApp(self.store)
self.app.config = self.config
self.tools = tai.Tools(self.app)
def tearDown(self):
self.store.close()
self.env.stop()
self.tmp.cleanup()
os.environ.pop("TAI_HOME", None)
tai._BOX_PYTHON_CACHE.clear()
def test_status_reports_all_targets(self):
with mock.patch.object(tai, "container_engine", return_value=None):
report = tai.install_report("status", list(tai.INSTALL_TARGETS), self.config.home, self.store)
for label in ("binary:", "bash-hook:", "venv:", "scheduler-service:", "telegram-service:", "container:", "vault:"):
self.assertIn(label, report)
self.assertIn("not installed", report)
self.assertIn("never touch data", report)
def test_binary_install_upgrade_uninstall(self):
target = os.path.join(self.home, ".local", "bin", "tai.py")
self.assertIn("installed", tai.op_binary("install", self.config.home))
self.assertTrue(os.access(target, os.X_OK))
with open(os.path.abspath(tai.__file__), "rb") as handle:
self.assertEqual(open(target, "rb").read(), handle.read())
self.assertIn("already installed", tai.op_binary("install", self.config.home))
self.assertIn("refreshed", tai.op_binary("upgrade", self.config.home))
self.assertIn("removed", tai.op_binary("uninstall", self.config.home))
self.assertFalse(os.path.exists(target))
self.assertIn("not present", tai.op_binary("uninstall", self.config.home))
def test_hook_install_remove_idempotent(self):
path = os.path.join(self.home, ".bashrc")
self.assertIn("source ~/.bashrc", tai.op_hook("install", self.config.home))
with open(path, encoding="utf-8") as handle:
self.assertIn(tai.BASHRC_MARK_BEGIN, handle.read())
self.assertIn("already present", tai.op_hook("install", self.config.home))
self.assertIn("removed", tai.op_hook("uninstall", self.config.home))
with open(path, encoding="utf-8") as handle:
self.assertNotIn(tai.BASHRC_MARK_BEGIN, handle.read())
self.assertTrue(os.path.isfile(path + ".bak-tai"))
self.assertIn("not present", tai.op_hook("uninstall", self.config.home))
def test_venv_kept_and_removed(self):
folder = tai.venv_dir(self.config.home)
python = os.path.join(folder, "bin", "python")
os.makedirs(os.path.dirname(python))
with open(python, "w", encoding="utf-8") as handle:
handle.write("#!/bin/sh\n")
os.chmod(python, 0o755)
self.assertIn("kept", tai.op_venv("install", self.config.home))
self.assertIn(python, tai.service_exec(self.config.home, "/x/tai.py"))
self.assertIn("removed", tai.op_venv("uninstall", self.config.home))
self.assertFalse(os.path.isdir(folder))
self.assertEqual(tai.service_exec(self.config.home, "/x/tai.py"), "/x/tai.py")
def test_venv_create_and_unavailable(self):
with mock.patch.object(tai, "venv_available", return_value=False):
self.assertIn("unavailable", tai.op_venv("install", self.config.home))
def fake_run(argv, capture_output=True, text=True, timeout=300):
folder = argv[-1]
python = os.path.join(folder, "bin", "python")
os.makedirs(os.path.dirname(python))
with open(python, "w", encoding="utf-8") as handle:
handle.write("#!/bin/sh\n")
os.chmod(python, 0o755)
outcome = mock.Mock()
outcome.returncode = 0
outcome.stderr = ""
return outcome
with mock.patch.object(tai, "venv_available", return_value=True):
with mock.patch("subprocess.run", side_effect=fake_run):
self.assertIn("created", tai.op_venv("install", self.config.home))
def test_scheduler_unit_write_and_remove(self):
with mock.patch("shutil.which", return_value=None):
self.assertIn("unit written", tai.op_scheduler_service("install", self.config.home))
unit = os.path.join(self.home, ".config", "systemd", "user", "tai-scheduler.service")
with open(unit, encoding="utf-8") as handle:
body = handle.read()
self.assertIn("--scheduler", body)
self.assertIn("already installed", tai.op_scheduler_service("install", self.config.home))
self.assertIn("removed", tai.op_scheduler_service("uninstall", self.config.home))
self.assertFalse(os.path.exists(unit))
def test_telegram_needs_token(self):
with mock.patch.dict(os.environ, {}, clear=False):
os.environ.pop("TELEGRAM_BOT_TOKEN", None)
self.assertIn("TELEGRAM_BOT_TOKEN", tai.op_telegram_service("install", self.config.home))
self.assertIn("not present", tai.op_telegram_service("uninstall", self.config.home))
def test_reinstall_keeps_vault(self):
marker = b"vault-data-marker-7"
with open(os.path.join(self.config.home, "memory.db"), "wb") as handle:
handle.write(marker)
with mock.patch.object(tai, "ensure_venv", return_value=(True, "created fake")):
report = tai.install_report("reinstall", ["binary", "bash-hook", "venv"], self.config.home, self.store)
with open(os.path.join(self.config.home, "memory.db"), "rb") as handle:
self.assertIn(marker, handle.read())
self.assertIn("never touch data", report)
def test_install_tool_validation_and_approval(self):
self.assertIn("action must be", self.tools.dispatch("install", json.dumps({"action": "explode"})))
self.assertIn("unknown target", self.tools.dispatch("install", json.dumps({"action": "status", "targets": ["nope"]})))
with mock.patch.object(tai, "container_engine", return_value=None):
status = self.tools.dispatch("install", json.dumps({"action": "status"}))
self.assertIn("scheduler-service:", status)
self.app.approve = False
self.assertIn("denied by user", self.tools.dispatch("install", json.dumps({"action": "upgrade", "targets": ["venv"]})))
self.assertIn("vault always kept", self.app.approvals[0])
self.app.approve = True
done = self.tools.dispatch("install", json.dumps({"action": "install", "targets": ["bash-hook"]}))
self.assertIn("bash-hook:", done)
def test_install_mentions_load_tool(self):
found = {schema["function"]["name"] for schema in tai.select_tools("how do I install tai as a service")}
self.assertIn("install", found)
found = {schema["function"]["name"] for schema in tai.select_tools("set up a venv hook")}
self.assertIn("install", found)
def test_box_python_cache_and_containerfile(self):
self.assertIn("/box/venv", tai.BOX_CONTAINERFILE)
self.assertIn("python3-venv", tai.BOX_CONTAINERFILE)
probe = mock.Mock()
probe.returncode = 0
with mock.patch.object(tai, "box_exec", return_value=probe) as runner:
self.assertEqual(tai.box_python("podman"), tai.BOX_PYTHON)
self.assertEqual(tai.box_python("podman"), tai.BOX_PYTHON)
self.assertEqual(runner.call_count, 1)
tai._BOX_PYTHON_CACHE.clear()
probe.returncode = 1
with mock.patch.object(tai, "box_exec", return_value=probe):
self.assertEqual(tai.box_python("podman"), "python3")
def test_unit_templates_format(self):
body = tai.TELEGRAM_UNIT % ("/x/tai.py", "/x/telegram.env")
self.assertIn("WorkingDirectory=%h", body)
self.assertIn("/x/tai.py --telegram", body)
body = tai.SCHEDULER_UNIT % "/x/tai.py"
self.assertIn("/x/tai.py --scheduler", body)
def test_install_repl(self):
self.app.tools = self.tools
with mock.patch.object(tai, "container_engine", return_value=None):
with mock.patch("builtins.print") as shown:
tai.handle_command(self.app, "/install status")
tai.handle_command(self.app, "/install frobnicate")
printed = "\n".join(str(call.args[0]) for call in shown.call_args_list)
self.assertIn("bash-hook:", printed)
self.assertIn("use /install", printed)
class BotsTests(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, "bots-test-1"))
self.agent = tai.Agent(self.config, self.store, persist=True, quiet=True)
self.tools = self.agent.tools
def tearDown(self):
self.store.close()
self.tmp.cleanup()
os.environ.pop("TAI_HOME", None)
def test_create_bot_roundtrip(self):
result = self.tools.dispatch("create_bot", json.dumps({"name": "Helper-Bot", "description": "helps a lot", "rules": "be kind", "behavior": "concise", "nicknames": ["help"]}))
self.assertIn("created bot 'helper-bot'", result)
self.assertIn("help", result)
self.assertIn("helper", result)
bots = {item["name"]: item for item in self.store.list_bots("t")}
self.assertIn("helper-bot", bots)
self.assertIn("help", bots["helper-bot"]["nicknames"])
self.assertIn("helper", bots["helper-bot"]["nicknames"])
system = self.store.load_bot_system("t", "helper-bot")
self.assertIn("helps a lot", system)
self.assertIn("Rules:", system)
self.assertIn("Behavior:", system)
self.assertEqual(self.store.resolve_bot("t", "helper-bot"), "helper-bot")
self.assertEqual(self.store.resolve_bot("t", "HELP"), "helper-bot")
self.assertEqual(self.store.resolve_bot("t", "main"), "main")
self.assertIsNone(self.store.resolve_bot("t", "nope"))
def test_create_bot_validation(self):
self.assertIn("invalid bot name", self.tools.dispatch("create_bot", json.dumps({"name": "Bad Name!"})))
self.assertIn("default bot", self.tools.dispatch("create_bot", json.dumps({"name": "main", "description": "x"})))
self.assertIn("at least one", self.tools.dispatch("create_bot", json.dumps({"name": "empty"})))
self.assertIn("nicknames must be", self.tools.dispatch("create_bot", json.dumps({"name": "x", "description": "y", "nicknames": "nope"})))
self.tools.dispatch("create_bot", json.dumps({"name": "coder", "description": "writes code", "nicknames": ["cd"]}))
self.assertIn("taken", self.tools.dispatch("create_bot", json.dumps({"name": "coder", "description": "again"})))
self.assertIn("taken", self.tools.dispatch("create_bot", json.dumps({"name": "other", "description": "y", "nicknames": ["cd"]})))
self.assertIn("invalid nickname", self.tools.dispatch("create_bot", json.dumps({"name": "ok", "description": "y", "nicknames": ["bad nick!"]})))
def test_switch_bot_resumes(self):
self.tools.dispatch("create_bot", json.dumps({"name": "coder", "description": "writes code", "nicknames": ["cd"]}))
self.agent.messages.append({"role": "user", "content": "main hello"})
self.agent.messages.append({"role": "assistant", "content": "main hi"})
with mock.patch("builtins.print"):
self.assertIn("switched", self.agent.switch_bot("coder"))
self.assertEqual(self.agent.bot, "coder")
self.assertEqual([item["role"] for item in self.agent.messages], ["system"])
self.assertIn("writes code", self.agent.system_message)
self.agent.messages.append({"role": "user", "content": "coder hello"})
with mock.patch("builtins.print"):
self.agent.switch_bot("main")
contents = [item.get("content", "") for item in self.agent.messages]
self.assertIn("main hello", contents)
self.assertNotIn("coder hello", contents)
with mock.patch("builtins.print"):
self.assertIn("switched", self.agent.switch_bot("cd"))
self.assertIn("coder hello", [item.get("content", "") for item in self.agent.messages])
def test_mention_routes_and_crossposts(self):
self.tools.dispatch("create_bot", json.dumps({"name": "coder", "description": "coder-brain-9"}))
seen = []
reply = {"role": "assistant", "content": "coded-it", "reasoning": "", "tool_calls": [], "backend": "x"}
def fake_complete(messages, tools, stream_sink=None):
seen.append(messages[0]["content"])
return dict(reply)
with mock.patch.object(self.agent.chat, "complete", side_effect=fake_complete):
answer = self.agent.run_turn("@coder write frob", capture=True)
self.assertEqual(answer, "coded-it")
self.assertIn("coder-brain-9", seen[0])
self.assertEqual(self.agent.bot, "main")
tail = self.agent.messages[-2:]
self.assertEqual(tail[0]["content"], "@coder write frob")
self.assertEqual(tail[1]["content"], "coded-it")
restored = self.store.load_session("t", "coder")
texts = [item.get("content", "") for item in restored]
self.assertIn("write frob", texts)
self.assertIn("coded-it", texts)
def test_mention_unknown_and_self(self):
reply = {"role": "assistant", "content": "hi", "reasoning": "", "tool_calls": [], "backend": "x"}
with mock.patch.object(self.agent.chat, "complete", return_value=dict(reply)):
self.assertIn("unknown bot", self.agent.run_turn("@ghost hi", capture=True))
self.assertEqual(self.agent.run_turn("@main hi", capture=True), "hi")
self.assertEqual(self.agent.messages[-2]["content"], "hi")
def test_remember_goes_to_current_bot(self):
self.tools.dispatch("create_bot", json.dumps({"name": "coder", "description": "v1"}))
with mock.patch("builtins.print"):
self.agent.switch_bot("coder")
self.agent.update_system("coder rules v2")
self.assertIn("v2", self.store.load_bot_system("t", "coder"))
self.assertNotIn("v2", self.store.load_system("t"))
def test_profile_switch_resets_bot(self):
self.tools.dispatch("create_bot", json.dumps({"name": "coder", "description": "v1"}))
with mock.patch("builtins.print"):
self.agent.switch_bot("coder")
self.agent.switch_profile("other")
self.assertEqual(self.agent.bot, "main")
def test_bots_repl(self):
self.tools.dispatch("create_bot", json.dumps({"name": "coder", "description": "v1", "nicknames": ["cd"]}))
with mock.patch("builtins.print") as shown:
tai.handle_command(self.agent, "/bots")
tai.handle_command(self.agent, "/bot coder")
tai.handle_command(self.agent, "/bot")
tai.handle_command(self.agent, "/bot nope")
printed = "\n".join(str(call.args[0]) for call in shown.call_args_list)
self.assertIn("* main", printed)
self.assertIn("coder aka cd", printed)
self.assertIn("switched to bot", printed)
self.assertIn("use /bot <name>", printed)
self.assertIn("unknown bot", printed)
self.assertEqual(self.agent.bot, "coder")
class FTSSearchTests(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, "fts-test-1"))
self.app = FakeSecretApp(self.store)
self.app.profile = "t"
self.app.config = self.config
self.tools = tai.Tools(self.app)
def tearDown(self):
self.store.close()
self.tmp.cleanup()
os.environ.pop("TAI_HOME", None)
def mem_id(self, text):
return re.search(r"mem:[0-9a-f]{16}", text).group(0)
def test_triggers_sync_index(self):
self.assertTrue(self.store.fts_ok)
record_id = self.mem_id(self.tools.dispatch("record_save", json.dumps({"title": "tangerine", "content": "alpha zonk"})))
self.assertEqual(len(self.store.fts_search("zonk", ("record",), "t", 5)), 1)
self.store.upsert_file_record("/tmp/x.txt", "beta zonk")
self.assertEqual(len(self.store.fts_search("zonk", ("record",), "t", 5)), 2)
self.store.delete_record(record_id, "t")
hits = self.store.fts_search("zonk", ("record",), "t", 5)
self.assertEqual(len(hits), 1)
self.assertNotEqual(hits[0]["item"], record_id)
self.store.log_event("t", "user", "message", "zonk event here")
event_hits = self.store.fts_search("zonk", ("event",), "t", 5)
self.assertEqual(len(event_hits), 1)
self.assertIn("[zonk]", event_hits[0]["snippet"])
found_events = self.store.search_events("t", "zonk event")
self.assertEqual(len(found_events), 1)
self.assertIn("zonk event here", found_events[0][3])
self.assertNotIn("tai1$", found_events[0][3])
path = os.path.join(self.tmp.name, "zonk.txt")
self.tools.dispatch("write_file", json.dumps({"path": path, "content": "x"}))
found = self.store.audit_search("zonk", "t", 5)
self.assertEqual(len(found), 1)
self.assertEqual(found[0]["path"], path)
def test_ranked_order_and_snippet(self):
self.store.add_record("note", "sparse", "quasar convenes", [], "t")
self.store.add_record("note", "dense", "quasar " * 20, [], "t")
hits = self.store.fts_search("quasar", ("record",), "t", 5)
self.assertEqual(len(hits), 2)
self.assertEqual(hits[0]["title"], "dense")
self.assertIn("[quasar]", hits[0]["snippet"])
def test_unified_search_mixed_and_expand(self):
first = self.mem_id(self.tools.dispatch("record_save", json.dumps({"title": "plan", "content": "harbor launch", "tags": ["harbor"]})))
second = self.mem_id(self.tools.dispatch("record_save", json.dumps({"title": "log", "content": "harbor diary", "tags": ["harbor"]})))
self.store.log_event("t", "user", "message", "harbor standup notes")
out = self.tools.dispatch("search", json.dumps({"query": "harbor"}))
self.assertIn("[record/", out)
self.assertIn(first, out)
self.assertIn("[event/", out)
narrowed = self.tools.dispatch("search", json.dumps({"query": "harbor", "kinds": ["event"]}))
self.assertIn("[event/", narrowed)
self.assertNotIn("[record/", narrowed)
expanded = self.tools.dispatch("search", json.dumps({"query": "harbor launch", "kinds": ["record"], "expand": True}))
self.assertIn(first, expanded)
self.assertIn("linked:", expanded)
self.assertIn("log", expanded)
self.assertIn(second, self.tools.dispatch("search", json.dumps({"query": "harbor", "kinds": ["record"], "expand": False})))
def test_record_search_fallback_partial_token(self):
record_id = self.mem_id(self.tools.dispatch("record_save", json.dumps({"title": "soup", "content": "alphabet soup serving", "tags": ["lunch"]})))
found = self.tools.dispatch("record_search", json.dumps({"query": "alphab", "tags": ["lunch"]}))
self.assertIn(record_id, found)
def test_recall_sealed_fallback(self):
other_home = os.path.join(self.tmp.name, "sealed")
os.makedirs(other_home)
class SealedArgs:
profile = "t"
yes = True
with mock.patch.dict(os.environ, {"TAI_HOME": other_home}):
config = tai.Config(SealedArgs())
sealed = tai.Store(config, tai.Seal(config.home, "sealed-pw-3"))
try:
sealed.log_event("t", "user", "message", "sealed recall marker words")
rows = sealed.search_events("t", "recall marker")
self.assertEqual(len(rows), 1)
self.assertIn("sealed recall marker", rows[0][3])
finally:
sealed.close()
def test_backfill_on_reopen(self):
self.store.add_record("note", "keep", "backfill beacon", [], "t")
self.store.db.execute("DELETE FROM fts_docs")
self.store.db.commit()
self.assertEqual(self.store.fts_search("beacon", ("record",), "t", 5), [])
self.store.close()
reopened = tai.Store(self.config, tai.Seal(self.config.home, "fts-test-1"))
try:
hits = reopened.fts_search("beacon", ("record",), "t", 5)
self.assertEqual(len(hits), 1)
finally:
reopened.close()
self.store = tai.Store(self.config, tai.Seal(self.config.home, "fts-test-1"))
def test_search_tool_validation(self):
self.assertIn("empty query", self.tools.dispatch("search", json.dumps({"query": ""})))
self.assertIn("kinds must be", self.tools.dispatch("search", json.dumps({"query": "x", "kinds": ["nope"]})))
self.assertIn("invalid limit", self.tools.dispatch("search", json.dumps({"query": "x", "limit": "z"})))
self.assertIn("no matches", self.tools.dispatch("search", json.dumps({"query": "zzz-nothing-here"})))
def test_audit_query(self):
path = os.path.join(self.tmp.name, "queryme.txt")
self.tools.dispatch("write_file", json.dumps({"path": path, "content": "x"}))
out = self.tools.dispatch("audit", json.dumps({"query": "queryme"}))
self.assertIn("queryme.txt", out)
def test_search_lazy_tag(self):
found = {schema["function"]["name"] for schema in tai.select_tools("find everything about the plan")}
self.assertIn("search", found)
def test_search_repl(self):
self.tools.dispatch("record_save", json.dumps({"title": "repl", "content": "repl beacon words"}))
self.app.tools = self.tools
with mock.patch("builtins.print") as shown:
tai.handle_command(self.app, "/search repl beacon")
tai.handle_command(self.app, "/search")
printed = "\n".join(str(call.args[0]) for call in shown.call_args_list)
self.assertIn("[record/", printed)
self.assertIn("use /search <query>", printed)
def test_mem_index_ranked_sealed(self):
self.assertIsNotNone(self.store.memdb)
self.store.log_event("t", "user", "message", "sparse comet sighting")
self.store.log_event("t", "user", "message", "comet " * 15)
hits = self.store.fts_search("comet", ("event",), "t", 5)
self.assertEqual(len(hits), 2)
self.assertLess(hits[0]["rank"], hits[1]["rank"])
self.assertIn("[comet]", hits[0]["snippet"])
self.assertNotIn("tai1$", hits[0]["snippet"])
stored = self.store.db.execute("SELECT text FROM events WHERE text LIKE 'tai1$%'").fetchall()
self.assertEqual(len(stored), 2)
def test_mem_index_cross_process_sync(self):
self.store.log_event("t", "user", "message", "first syncable event")
other = tai.Store(self.config, tai.Seal(self.config.home, "fts-test-1"))
try:
other.log_event("t", "user", "message", "second syncable event")
finally:
other.close()
hits = self.store.fts_search("syncable", ("event",), "t", 5)
self.assertEqual(len(hits), 2)
def test_mem_index_profile_scoped(self):
self.store.log_event("t", "user", "message", "scoped beacon words")
self.store.log_event("other", "user", "message", "scoped beacon words")
self.assertEqual(len(self.store.fts_search("beacon", ("event",), "t", 5)), 1)
self.assertEqual(len(self.store.fts_search("beacon", ("event",), "other", 5)), 1)
def test_mem_index_survives_rotation(self):
self.store.log_event("t", "user", "message", "rotation beacon words")
old = tai.Seal(self.config.home, "fts-test-1")
self.store.close()
tai.rotate_seal(self.config, old, "fts-test-2")
reopened = tai.Store(self.config, tai.Seal(self.config.home, "fts-test-2"))
try:
hits = reopened.fts_search("rotation", ("event",), "t", 5)
self.assertEqual(len(hits), 1)
self.assertIn("[rotation]", hits[0]["snippet"])
finally:
reopened.close()
self.store = tai.Store(self.config, tai.Seal(self.config.home, "fts-test-2"))
def test_run_search_sealed_events_ranked_and_partial(self):
self.store.log_event("t", "user", "message", "harbor standup notes")
out = self.tools.dispatch("search", json.dumps({"query": "harbor", "kinds": ["event"]}))
self.assertIn("[event/", out)
self.assertIn("[harbor]", out)
partial = self.tools.dispatch("search", json.dumps({"query": "standu", "kinds": ["event"]}))
self.assertIn("[event/", partial)
self.assertIn("standup", partial)
def test_rotate_multi_profile_secrets(self):
self.store.save_secret("api", "t-secret-1", None, None, "t")
self.store.save_secret("api", "o-secret-2", None, None, "other")
old = tai.Seal(self.config.home, "fts-test-1")
self.store.close()
tai.rotate_seal(self.config, old, "fts-test-2")
reopened = tai.Store(self.config, tai.Seal(self.config.home, "fts-test-2"))
try:
self.assertEqual(reopened.load_secret("api", "t"), "t-secret-1")
self.assertEqual(reopened.load_secret("api", "other"), "o-secret-2")
finally:
reopened.close()
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 DiffTests(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, "diff-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 test_highlight_python(self):
out = tai.highlight_python("def f(x=1): # hi")
self.assertIn(tai.Ansi.MAGENTA + "def" + tai.Ansi.RESET, out)
self.assertIn(tai.Ansi.CYAN + "f" + tai.Ansi.RESET, out)
self.assertIn(tai.FG_ORANGE + "1" + tai.Ansi.RESET, out)
self.assertIn(tai.Ansi.GRAY + "# hi" + tai.Ansi.RESET, out)
self.assertEqual(tai.strip_ansi(out), "def f(x=1): # hi")
def test_highlight_string_with_hash(self):
out = tai.highlight_python("x = 'a # b'")
self.assertIn(tai.Ansi.YELLOW + "'a # b'" + tai.Ansi.RESET, out)
self.assertNotIn(tai.Ansi.GRAY, out)
def test_render_diff_plain(self):
out = tai.render_diff("a.py", "x = 1\n", "x = 2\n", width=60, colors=False)
self.assertEqual(out, "│ ── a.py (+1 -1)\n 1 - x = 1\n 1 + x = 2")
def test_render_diff_colors(self):
out = tai.render_diff("a.py", "x = 1\n", "x = 2\n", width=60, colors=True)
self.assertIn(tai.BG_ADD, out)
self.assertIn(tai.BG_DEL, out)
self.assertIn(tai.FG_ORANGE + "2" + tai.Ansi.RESET, out)
for line in out.splitlines()[1:]:
self.assertEqual(len(tai.strip_ansi(line)), 60)
def test_render_diff_gap_and_cap(self):
old = "".join("line %d\n" % num for num in range(20))
new = old.replace("line 0\n", "line zero\n").replace("line 19\n", "line nineteen\n")
out = tai.render_diff("a.py", old, new, width=60, colors=False)
self.assertIn("···", out)
big = "".join("row %d\n" % num for num in range(200))
capped = tai.render_diff("a.py", "", big, width=60, colors=False)
self.assertIn("more lines hidden", capped)
self.assertEqual(len(capped.splitlines()), 1 + tai.DIFF_MAX_LINES + 1)
def test_render_diff_shapes(self):
self.assertEqual(tai.render_diff("x.py", "a\n", "a\n"), "")
created = tai.render_diff("x.py", None, "a\nb\n", width=60, colors=False)
self.assertIn("(new file, 2 lines)", created.splitlines()[0])
removed = tai.render_diff("x.py", "a\nb\n", None, width=60, colors=False)
self.assertIn("(deleted, 2 lines)", removed.splitlines()[0])
plain = tai.render_diff("notes.txt", "a\n", "b\n", width=60, colors=True)
self.assertIn(tai.BG_ADD, plain)
self.assertNotIn(tai.Ansi.MAGENTA, plain)
self.assertNotIn(tai.Ansi.YELLOW, plain)
def test_render_diff_docstring_fence(self):
out = tai.render_diff("a.py", "", '"""\ndoc body\n"""\nx = 1\n', width=60, colors=True)
self.assertIn(tai.Ansi.YELLOW + "doc body" + tai.Ansi.RESET, out)
def test_tools_note_preview(self):
path = os.path.join(self.tmp.name, "note.py")
self.assertTrue(self.tools.dispatch("write_file", json.dumps({"path": path, "content": "x = 1\n"})).startswith("wrote "))
preview = self.tools.diff_preview
self.assertEqual((preview["path"], preview["old"], preview["new"]), (path, None, "x = 1\n"))
self.tools.dispatch("read_file", json.dumps({"path": path}))
self.assertTrue(self.tools.dispatch("edit_file", json.dumps({"path": path, "find": "x = 1", "replace": "x = 2"})).startswith("edited "))
preview = self.tools.diff_preview
self.assertEqual((preview["old"], preview["new"]), ("x = 1\n", "x = 2\n"))
self.assertTrue(self.tools.dispatch("delete_file", json.dumps({"path": path})).startswith("deleted "))
preview = self.tools.diff_preview
self.assertEqual((preview["old"], preview["new"]), ("x = 2\n", None))
def test_show_result_renders_and_clears(self):
agent = tai.Agent(self.config, self.store, persist=False, quiet=False)
agent.tools.diff_preview = {"path": "a.py", "old": "x = 1\n", "new": "x = 2\n"}
with mock.patch("builtins.print") as shown:
agent.show_result("ok", 1)
agent.show_result("ok", 1)
printed = "\n".join(str(call.args[0]) for call in shown.call_args_list)
self.assertEqual(printed.count("── a.py"), 1)
self.assertEqual(printed.count("└─"), 2)
self.assertNotIn("\033", printed)
self.assertIsNone(agent.tools.diff_preview)
agent.quiet = True
agent.tools.diff_preview = {"path": "a.py", "old": "x = 1\n", "new": "x = 2\n"}
with mock.patch("builtins.print") as shown:
agent.show_result("ok", 1)
shown.assert_not_called()
class ResumeTests(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, "resume-test-1"))
def tearDown(self):
self.store.close()
self.tmp.cleanup()
os.environ.pop("TAI_HOME", None)
def text_reply(self, text="done"):
return {"role": "assistant", "content": text, "reasoning": "", "tool_calls": [], "backend": "x"}
def test_atomic_session_survives_failed_replace(self):
first = [{"role": "user", "content": "first"}]
self.store.save_session("t", first)
with mock.patch("os.replace", side_effect=OSError("disk full")):
self.store.save_session("t", [{"role": "user", "content": "second"}])
self.assertEqual([item["content"] for item in self.store.load_session("t")], ["first"])
def test_corrupt_session_quarantined(self):
path = self.store.session_path("t")
self.store.save_session("t", [{"role": "user", "content": "good"}])
with open(path, "w", encoding="utf-8") as handle:
handle.write("{not json")
self.assertEqual(self.store.load_session("t"), [])
leftovers = [entry for entry in os.listdir(self.config.profiles_dir) if ".corrupt-" in entry]
self.assertEqual(len(leftovers), 1)
with open(os.path.join(self.config.profiles_dir, leftovers[0]), encoding="utf-8") as handle:
self.assertEqual(handle.read(), "{not json")
def test_normal_turn_closes_state(self):
agent = tai.Agent(self.config, self.store, persist=True, quiet=True)
with mock.patch.object(agent.chat, "complete", return_value=self.text_reply()):
agent.run_turn("hello", capture=True)
state = self.store.load_turn_state("t", "main")
self.assertFalse(state["open"])
self.assertEqual(state["goal"], "hello")
def test_crash_resume_end_to_end(self):
agent = tai.Agent(self.config, self.store, persist=True, quiet=True)
tool_reply = {"role": "assistant", "content": "", "reasoning": "", "tool_calls": [{"id": "1", "name": "shell", "arguments": json.dumps({"command": "echo resume-mark"})}], "backend": "x"}
with mock.patch.object(agent.chat, "complete", side_effect=[dict(tool_reply), KeyboardInterrupt()]):
with self.assertRaises(KeyboardInterrupt):
agent.run_turn("do the thing", capture=True)
state = self.store.load_turn_state("t", "main")
self.assertTrue(state["open"])
self.assertEqual(state["goal"], "do the thing")
self.store.close()
store2 = tai.Store(self.config, tai.Seal(self.config.home, "resume-test-1"))
self.store = store2
agent2 = tai.Agent(self.config, store2, persist=True, quiet=True)
texts = [item.get("content") or "" for item in agent2.messages]
self.assertTrue(any("do the thing" in text for text in texts))
self.assertTrue(any("resume-mark" in text for text in texts))
with mock.patch.object(agent2.chat, "complete", return_value=self.text_reply("finished")):
with mock.patch("builtins.print"):
answer = tai.resume_turn(agent2)
self.assertEqual(answer, "finished")
self.assertFalse(store2.load_turn_state("t", "main")["open"])
def test_resume_nothing_open(self):
agent = tai.Agent(self.config, self.store, persist=True, quiet=True)
with mock.patch.object(agent.chat, "complete") as called:
with mock.patch("builtins.print") as shown:
self.assertEqual(tai.resume_turn(agent), "")
called.assert_not_called()
printed = "\n".join(str(call.args[0]) for call in shown.call_args_list)
self.assertIn("nothing to resume", printed)
def test_checkpoint_idle_writes_no_state(self):
agent = tai.Agent(self.config, self.store, persist=True, quiet=True)
agent.checkpoint()
self.assertEqual(self.store.load_turn_state("t", "main"), {})
self.assertFalse(os.path.exists(self.store.turn_state_path("t", "main")))
def test_parser_continue(self):
args = tai.build_parser().parse_args(["--continue"])
self.assertTrue(args.resume)
self.assertFalse(tai.build_parser().parse_args([]).resume)
class DenialTests(unittest.TestCase):
def setUp(self):
self.tmp = tempfile.TemporaryDirectory()
os.environ["TAI_HOME"] = self.tmp.name
class FakeArgs:
profile = "t"
yes = False
self.config = tai.Config(FakeArgs())
self.store = tai.Store(self.config, tai.Seal(self.config.home, ""))
def tearDown(self):
self.store.close()
self.tmp.cleanup()
os.environ.pop("TAI_HOME", None)
def call_reply(self, *calls):
return {"role": "assistant", "content": "", "reasoning": "", "tool_calls": list(calls), "backend": "x"}
def test_deny_continues_with_guidance(self):
agent = tai.Agent(self.config, self.store, persist=True, quiet=True)
first = self.call_reply({"id": "c1", "name": "shell", "arguments": json.dumps({"command": "ssh evil.example.com"})})
final = {"role": "assistant", "content": "done", "reasoning": "", "tool_calls": [], "backend": "x"}
with mock.patch.object(agent.chat, "complete", side_effect=[first, final]) as completer:
with mock.patch("sys.stdin") as fake_stdin:
fake_stdin.isatty.return_value = True
with mock.patch("builtins.input", side_effect=["n", "do X instead"]) as asker:
with mock.patch("builtins.print"):
result = agent.run_turn("try ssh", capture=True)
self.assertEqual(result, "done")
self.assertEqual(asker.call_count, 2)
second_messages = completer.call_args_list[1].args[0]
users = [item["content"] for item in second_messages if item["role"] == "user"]
self.assertIn("do X instead", users)
tools = [item for item in second_messages if item["role"] == "tool"]
self.assertEqual(tools[0]["content"], "denied by user")
def test_deny_empty_aborts(self):
agent = tai.Agent(self.config, self.store, persist=True, quiet=True)
first = self.call_reply(
{"id": "c1", "name": "shell", "arguments": json.dumps({"command": "ssh a.example.com"})},
{"id": "c2", "name": "shell", "arguments": json.dumps({"command": "ssh b.example.com"})},
)
with mock.patch.object(agent.chat, "complete", return_value=first) as completer:
with mock.patch("sys.stdin") as fake_stdin:
fake_stdin.isatty.return_value = True
with mock.patch("builtins.input", side_effect=["n", ""]):
with mock.patch("builtins.print"):
result = agent.run_turn("try ssh", capture=True)
self.assertEqual(result, "stopped by user")
self.assertEqual(completer.call_count, 1)
by_id = {}
for item in agent.messages:
for recorded in item.get("tool_calls") or []:
by_id[recorded["id"]] = None
for item in agent.messages:
if item.get("role") == "tool":
by_id[item["tool_call_id"]] = item["content"]
self.assertEqual(by_id, {"c1": "denied by user", "c2": "skipped: stopped after denial"})
def test_auto_never_asks(self):
agent = tai.Agent(self.config, self.store, persist=True, quiet=True, auto=True)
first = self.call_reply({"id": "c1", "name": "shell", "arguments": json.dumps({"command": "echo auto-ok"})})
final = {"role": "assistant", "content": "done", "reasoning": "", "tool_calls": [], "backend": "x"}
with mock.patch.object(agent.chat, "complete", side_effect=[first, final]):
with mock.patch("sys.stdin") as fake_stdin:
fake_stdin.isatty.return_value = True
with mock.patch("builtins.input", side_effect=AssertionError("must not ask")) as asker:
with mock.patch("builtins.print"):
result = agent.run_turn("try echo", capture=True)
self.assertEqual(result, "done")
self.assertEqual(asker.call_count, 0)
ran = [item for item in agent.messages if item.get("role") == "tool"]
self.assertIn("auto-ok", ran[0]["content"])
def test_dispatch_reraises_denied(self):
tools = tai.Tools(FakeApp())
def raiser(args):
raise tai.Denied("go left")
tools.handlers["boom"] = raiser
with self.assertRaises(tai.Denied) as caught:
tools.dispatch("boom", "{}")
self.assertEqual(caught.exception.guidance, "go left")
def test_secret_delete_passes_no_guidance(self):
agent = mock.Mock()
seen = []
def approver(command, guidance=True):
seen.append(guidance)
return False
agent.ask_approval = approver
agent.store = self.store
agent.profile = "t"
self.store.save_secret("wifi", "repl-value-9")
with mock.patch("builtins.print"):
tai.handle_command(agent, "/secret delete wifi")
self.assertEqual(seen, [False])
self.assertEqual(self.store.load_secret("wifi"), "repl-value-9")
def test_show_call_masks_secrets(self):
agent = tai.Agent(self.config, self.store, persist=False, quiet=False)
self.store.save_secret("api", "known-token-5")
with mock.patch("builtins.print") as printer:
agent.show_call({"name": "store_secret", "arguments": json.dumps({"name": "api", "value": "brand-new-1"})})
agent.show_call({"name": "shell", "arguments": json.dumps({"command": "curl known-token-5"})})
shown = "\n".join(str(call.args[0]) for call in printer.call_args_list if call.args)
self.assertNotIn("brand-new-1", shown)
self.assertIn("[hidden]", shown)
self.assertNotIn("known-token-5", shown)
self.assertIn("[redacted:api]", shown)
class FileGuardTests(unittest.TestCase):
def test_write_requires_read(self):
with tempfile.TemporaryDirectory() as tmp:
target = os.path.join(tmp, "notes.txt")
with open(target, "w", encoding="utf-8") as handle:
handle.write("original one")
tools = tai.Tools(FakeApp())
denied = tools.dispatch("write_file", json.dumps({"path": target, "content": "clobber"}))
self.assertIn("read it first", denied)
with open(target, encoding="utf-8") as handle:
self.assertEqual(handle.read(), "original one")
self.assertIn("original one", tools.dispatch("read_file", json.dumps({"path": target})))
self.assertIn("wrote", tools.dispatch("write_file", json.dumps({"path": target, "content": "updated two"})))
with open(target, encoding="utf-8") as handle:
self.assertEqual(handle.read(), "updated two")
def test_new_files_always_writable(self):
with tempfile.TemporaryDirectory() as tmp:
target = os.path.join(tmp, "sub", "fresh.txt")
tools = tai.Tools(FakeApp())
self.assertIn("wrote", tools.dispatch("write_file", json.dumps({"path": target, "content": "hello"})))
self.assertIn("edited", tools.dispatch("edit_file", json.dumps({"path": target, "find": "hello", "replace": "hi"})))
def test_edit_requires_read(self):
with tempfile.TemporaryDirectory() as tmp:
target = os.path.join(tmp, "code.py")
with open(target, "w", encoding="utf-8") as handle:
handle.write("print(1)")
tools = tai.Tools(FakeApp())
self.assertIn("read it first", tools.dispatch("edit_file", json.dumps({"path": target, "find": "1", "replace": "2"})))
tools.dispatch("read_file", json.dumps({"path": target}))
self.assertIn("edited", tools.dispatch("edit_file", json.dumps({"path": target, "find": "1", "replace": "2"})))
def test_guard_is_per_session(self):
with tempfile.TemporaryDirectory() as tmp:
target = os.path.join(tmp, "data.txt")
with open(target, "w", encoding="utf-8") as handle:
handle.write("v1")
first = tai.Tools(FakeApp())
first.dispatch("read_file", json.dumps({"path": target}))
second = tai.Tools(FakeApp())
self.assertIn("read it first", second.dispatch("write_file", json.dumps({"path": target, "content": "v2"})))
class FakeBot:
def __init__(self):
self.sent = []
self.actions = []
def call(self, method, params, timeout=70):
if method == "sendMessage":
self.sent.append(params["text"])
if method == "sendChatAction":
self.actions.append(params["action"])
return {"ok": True}
class FakeAgent:
def __init__(self):
self.reset = 0
self.turns = []
def reset_history(self):
self.reset += 1
def run_turn(self, text, capture=False):
self.turns.append((text, capture))
return "canned reply"
class TelegramTests(unittest.TestCase):
def test_send_chunks(self):
bot = FakeBot()
tai.telegram_send(bot, 7, "x" * 5000)
self.assertEqual([len(part) for part in bot.sent], [4000, 1000])
def test_start_and_new(self):
agent = FakeAgent()
bot = FakeBot()
tai.handle_telegram_update(agent, bot, {"message": {"chat": {"id": 1}, "text": "/start"}})
tai.handle_telegram_update(agent, bot, {"message": {"chat": {"id": 1}, "text": "/new"}})
self.assertEqual(agent.reset, 1)
self.assertEqual(len(bot.sent), 2)
self.assertEqual(agent.turns, [])
def test_text_turn(self):
agent = FakeAgent()
bot = FakeBot()
tai.handle_telegram_update(agent, bot, {"message": {"chat": {"id": 1}, "text": "hello"}})
self.assertEqual(agent.turns, [("hello", True)])
self.assertEqual(bot.sent, ["canned reply"])
self.assertEqual(bot.actions, ["typing"])
def test_token_loading(self):
with tempfile.TemporaryDirectory() as home:
with mock.patch.dict(os.environ, {}, clear=False):
os.environ.pop("TELEGRAM_BOT_TOKEN", None)
self.assertEqual(tai.load_telegram_token(home), "")
with open(os.path.join(home, "telegram.env"), "w", encoding="utf-8") as handle:
handle.write("TELEGRAM_BOT_TOKEN=file-token-1\n")
self.assertEqual(tai.load_telegram_token(home), "file-token-1")
os.environ["TELEGRAM_BOT_TOKEN"] = "env-token-2"
self.assertEqual(tai.load_telegram_token(home), "env-token-2")
os.environ.pop("TELEGRAM_BOT_TOKEN", None)
class InstallTests(unittest.TestCase):
def test_bashrc_upsert(self):
with tempfile.TemporaryDirectory() as tmp:
path = os.path.join(tmp, ".bashrc")
with open(path, "w", encoding="utf-8") as handle:
handle.write("export PATH=$PATH:/x\n")
self.assertTrue(tai.upsert_bashrc_block(path))
self.assertFalse(tai.upsert_bashrc_block(path))
with open(path, encoding="utf-8") as handle:
content = handle.read()
self.assertIn("export PATH=$PATH:/x", content)
self.assertEqual(content.count(tai.BASHRC_MARK_BEGIN), 1)
with open(path + ".bak-tai", encoding="utf-8") as handle:
self.assertNotIn("tai command-not-found", handle.read())
def test_bashrc_replaces_stale(self):
with tempfile.TemporaryDirectory() as tmp:
path = os.path.join(tmp, ".bashrc")
stale = tai.BASHRC_MARK_BEGIN + "\nold hook\n" + tai.BASHRC_MARK_END + "\n"
with open(path, "w", encoding="utf-8") as handle:
handle.write("alias x=y\n" + stale)
self.assertTrue(tai.upsert_bashrc_block(path))
with open(path, encoding="utf-8") as handle:
content = handle.read()
self.assertNotIn("old hook", content)
self.assertIn("alias x=y", content)
self.assertEqual(content.count(tai.BASHRC_MARK_BEGIN), 1)
def test_parser_prompt(self):
args = tai.build_parser().parse_args(["what", "is", "(2+3)?"])
self.assertEqual(args.prompt, ["what", "is", "(2+3)?"])
self.assertFalse(args.install)
args = tai.build_parser().parse_args(["--install-telegram"])
self.assertTrue(args.install_telegram)
self.assertEqual(args.prompt, [])
args = tai.build_parser().parse_args(["--profile", "work", "--yes"])
self.assertEqual(args.profile, "work")
self.assertTrue(args.yes)
class FakePrincipal:
def __init__(self):
self.depth = 0
self.profile = "t"
self.config = None
self.store = mock.Mock()
self.store.redact = lambda text, profile=None: text
self.runner_override = None
class OrchestrationTests(unittest.TestCase):
def setUp(self):
with tai.AGENTS_LOCK:
tai.AGENTS.clear()
tai.AGENTS_NEXT[0] = 1
def tearDown(self):
with tai.AGENTS_LOCK:
tai.AGENTS.clear()
tai.AGENTS_NEXT[0] = 1
def test_spawn_poll_instant(self):
agent_id = tai.spawn_agent("do it", "t", 60, None, None, 0, lambda task, profile, timeout: "stub-done")
self.assertEqual(agent_id, 1)
status, text = tai.poll_agent(agent_id, wait=5)
self.assertEqual(status, "done")
self.assertIn("stub-done", text)
def test_poll_running_and_missing(self):
def slow(task, profile, timeout):
time.sleep(2)
return "slow-done"
agent_id = tai.spawn_agent("slow", "t", 60, None, None, 0, slow)
status, text = tai.poll_agent(agent_id, wait=0)
self.assertEqual(status, "running")
self.assertIn("still running", text)
status, _text = tai.poll_agent(999, wait=0)
self.assertEqual(status, "missing")
status, text = tai.poll_agent(agent_id, wait=5)
self.assertEqual(status, "done")
def test_timeout_and_error_status(self):
slow_id = tai.spawn_agent("t", "t", 60, None, None, 0, lambda task, profile, timeout: "partial\n[time limit reached]")
self.assertEqual(tai.poll_agent(slow_id, wait=5)[0], "timeout")
def broken(task, profile, timeout):
raise RuntimeError("boom")
bad_id = tai.spawn_agent("t", "t", 60, None, None, 0, broken)
status, text = tai.poll_agent(bad_id, wait=5)
self.assertEqual(status, "error")
self.assertIn("boom", text)
def test_list_and_clear(self):
tai.spawn_agent("one", "t", 60, None, None, 0, lambda task, profile, timeout: "r1")
tai.spawn_agent("two", "t", 60, None, None, 0, lambda task, profile, timeout: time.sleep(2) or "r2")
self.assertEqual(tai.poll_agent(1, wait=5)[0], "done")
records = tai.list_agents()
self.assertEqual(len(records), 2)
self.assertEqual(tai.clear_agents(), 1)
self.assertEqual(len(tai.list_agents()), 1)
def test_fork_tool(self):
app = FakePrincipal()
app.runner_override = lambda task, profile, timeout: "forked-ok"
tools = tai.Tools(app)
started = tools.dispatch("fork", json.dumps({"task": "research x"}))
self.assertIn("agent 1 started", started)
result = tools.dispatch("poll", json.dumps({"id": 1, "wait": 5}))
self.assertIn("forked-ok", result)
def test_fork_depth_limit(self):
app = FakePrincipal()
app.depth = 2
tools = tai.Tools(app)
self.assertIn("depth limit", tools.dispatch("fork", json.dumps({"task": "x"})))
def test_reasoning_only_becomes_result(self):
with tempfile.TemporaryDirectory() as tmp:
old_home = os.environ.get("TAI_HOME")
os.environ["TAI_HOME"] = tmp
try:
class FakeArgs:
profile = "t"
yes = True
config = tai.Config(FakeArgs())
store = tai.Store(config, tai.Seal(config.home, ""))
agent = tai.Agent(config, store, persist=False, quiet=True)
reply = {"role": "assistant", "content": "", "reasoning": "thought out", "tool_calls": [], "backend": "x"}
with mock.patch.object(agent.chat, "complete", return_value=reply):
self.assertEqual(agent.run_turn("hi", capture=True), "thought out")
store.close()
finally:
if old_home is None:
os.environ.pop("TAI_HOME", None)
else:
os.environ["TAI_HOME"] = old_home
def test_deadline_shortcircuit(self):
with tempfile.TemporaryDirectory() as tmp:
old_home = os.environ.get("TAI_HOME")
os.environ["TAI_HOME"] = tmp
try:
class FakeArgs:
profile = "t"
yes = True
config = tai.Config(FakeArgs())
store = tai.Store(config, tai.Seal(config.home, ""))
agent = tai.Agent(config, store, persist=False, quiet=True)
agent.deadline = time.time() - 1
with mock.patch.object(agent.chat, "complete", side_effect=AssertionError("network used")):
result = agent.run_turn("hi", capture=True)
self.assertTrue(agent.timed_out)
self.assertIn("time limit reached", result)
store.close()
finally:
if old_home is None:
os.environ.pop("TAI_HOME", None)
else:
os.environ["TAI_HOME"] = old_home
class MarkdownTests(unittest.TestCase):
def test_passthrough_without_tty(self):
with mock.patch("sys.stdout") as fake_out:
fake_out.isatty.return_value = False
self.assertEqual(tai.render_markdown("**x**"), "**x**")
with mock.patch("sys.stdout") as fake_out:
fake_out.isatty.return_value = True
with mock.patch.dict(os.environ, {"NO_COLOR": "1"}):
self.assertEqual(tai.render_markdown("**x**"), "**x**")
def test_headings(self):
out = tai.render_markdown("# One\n\n## Two\n\n### Three", width=80, color=True)
self.assertNotIn("#", out)
self.assertIn(tai.Ansi.BOLD + tai.Ansi.CYAN + "One", out)
self.assertIn(tai.Ansi.BOLD + "Two", out)
self.assertIn("Three", tai.strip_ansi(out))
def test_inline_styles(self):
out = tai.render_markdown("**b** *i* `c` ~~s~~ _u_ __w__", width=80, color=True)
self.assertIn(tai.style_text("b", tai.Ansi.BOLD), out)
self.assertIn(tai.style_text("i", tai.Ansi.DIM), out)
self.assertIn(tai.style_text("c", tai.Ansi.CYAN), out)
self.assertIn(tai.style_text("s", tai.Ansi.STRIKE), out)
self.assertIn(tai.style_text("u", tai.Ansi.DIM), out)
self.assertIn(tai.style_text("w", tai.Ansi.BOLD), out)
self.assertNotIn("**", out)
def test_snake_case_survives(self):
out = tai.render_markdown("use my_var_name and `__init__` here", width=80, color=True)
self.assertIn("my_var_name", tai.strip_ansi(out))
self.assertIn("__init__", tai.strip_ansi(out))
self.assertNotIn(tai.Ansi.BOLD, out)
def test_link_and_image(self):
out = tai.render_markdown("[docs](https://x.example/d) and ![alt](pic.png)", width=80, color=True)
self.assertIn("docs", out)
self.assertIn("https://x.example/d", out)
self.assertIn("[image: alt]", tai.strip_ansi(out))
self.assertNotIn("[docs]", out)
def test_code_fence_verbatim(self):
out = tai.render_markdown("```python\nreturn \"**x**\"\n```", width=80, color=True)
self.assertIn('"**x**"', out)
self.assertNotIn(tai.Ansi.BOLD, out)
self.assertNotIn("```", out)
def test_unclosed_fence(self):
out = tai.render_markdown("```\ncode **x**", width=80, color=True)
self.assertIn("**x**", out)
self.assertNotIn(tai.Ansi.BOLD, out)
def test_table_alignment(self):
source = "| item | qty |\n|---|---:|\n| apple | 12 |\n| fig | 3 |"
out = tai.render_markdown(source, width=80, color=True)
plain = tai.strip_ansi(out)
self.assertIn("┌───────┬─────┐", plain)
self.assertIn("│ apple │ 12 │", plain)
self.assertIn("│ fig │ 3 │", plain)
self.assertEqual({len(line) for line in plain.splitlines()}, {15})
self.assertIn(tai.Ansi.BOLD, out)
def test_table_escaped_pipe(self):
out = tai.render_markdown("| a |\n|---|\n| x\\|y |", width=80, color=True)
plain = tai.strip_ansi(out)
self.assertIn("x|y", plain)
self.assertEqual(len(plain.splitlines()), 5)
def test_lists(self):
source = "- a\n - b\n- [x] done\n- [ ] open\n1. x\n1. y"
plain = tai.strip_ansi(tai.render_markdown(source, width=80, color=True))
self.assertIn("• a", plain)
self.assertIn(" ◦ b", plain)
self.assertIn("☑ done", plain)
self.assertIn("☐ open", plain)
self.assertIn("1. x", plain)
self.assertIn("2. y", plain)
def test_quote_and_rule(self):
out = tai.render_markdown("> wise\n> words\n\n---", width=40, color=True)
plain = tai.strip_ansi(out)
self.assertIn("│ wise words", plain)
self.assertIn("─" * 40, plain)
def test_wrap_keeps_styles(self):
out = tai.render_markdown("**word " + " ".join("w%d" % n for n in range(20)) + " end**", width=30, color=True)
lines = out.splitlines()
self.assertGreater(len(lines), 1)
self.assertNotIn("**", out)
for line in lines:
self.assertLessEqual(len(tai.strip_ansi(line)), 30)
self.assertIn(tai.Ansi.BOLD, out)
def test_escapes_and_breaks(self):
out = tai.render_markdown("a \\*b\\* c<br>d", width=80, color=True)
plain = tai.strip_ansi(out)
self.assertIn("*b*", plain)
self.assertNotIn(tai.Ansi.DIM, out)
self.assertEqual(plain.splitlines(), ["a *b* c", "d"])
def test_run_turn_renders_reply(self):
with tempfile.TemporaryDirectory() as tmp:
old_home = os.environ.get("TAI_HOME")
os.environ["TAI_HOME"] = tmp
try:
class FakeArgs:
profile = "t"
yes = True
config = tai.Config(FakeArgs())
store = tai.Store(config, tai.Seal(config.home, ""))
agent = tai.Agent(config, store)
reply = {"role": "assistant", "content": "# Title\n\nHello **bold**", "reasoning": "", "tool_calls": [], "backend": "x"}
with mock.patch.object(agent.chat, "complete", return_value=reply):
with mock.patch.dict(os.environ):
os.environ.pop("NO_COLOR", None)
with mock.patch("sys.stdout") as fake_out:
fake_out.isatty.return_value = True
with mock.patch("builtins.print") as printer:
with mock.patch.object(tai, "Spinner"):
result = agent.run_turn("hi")
self.assertEqual(result, "# Title\n\nHello **bold**")
printed = "\n".join(str(call.args[0]) for call in printer.call_args_list if call.args)
self.assertIn(tai.Ansi.BOLD, printed)
self.assertIn("Title", printed)
self.assertNotIn("# Title", printed)
store.close()
finally:
if old_home is None:
os.environ.pop("TAI_HOME", None)
else:
os.environ["TAI_HOME"] = old_home
if __name__ == "__main__":
unittest.main()