3160 lines
152 KiB
Python
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 ", 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()
|