Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 8 additions & 1 deletion agentmain.py
Original file line number Diff line number Diff line change
Expand Up @@ -216,7 +216,14 @@ def run(self):
parser.add_argument('--nolog', action='store_true')
parser.add_argument('--no-user-tools', action='store_true')
args, _unknown = parser.parse_known_args()
_extra_args = dict(zip([k.lstrip('-') for k in _unknown[::2]], _unknown[1::2])) if _unknown else {}
reflect_enabled = args.reflect is not None
if args.reflect == '':
parser.error("--reflect requires a non-empty script path")
if _unknown and not reflect_enabled:
parser.error(f"unrecognized arguments: {' '.join(_unknown)}")
if reflect_enabled and (len(_unknown) % 2 or any(len(k) <= 2 or not k.startswith('--') or k.startswith('---') for k in _unknown[::2]) or any(v.startswith('--') for v in _unknown[1::2])):
parser.error("reflect extra arguments must be --key value pairs")
_extra_args = dict(zip([k[2:] for k in _unknown[::2]], _unknown[1::2])) if _unknown else {}

if (args.func or args.task) and not args.nobg:
import subprocess, platform
Expand Down
84 changes: 84 additions & 0 deletions tests/test_agentmain_cli_args.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,84 @@
import os
import subprocess
import sys
import tempfile
import unittest
from pathlib import Path


ROOT = Path(__file__).resolve().parents[1]
AGENTMAIN = ROOT / "agentmain.py"
_RUNNER = (
"import runpy, sys\n"
"target = sys.argv[1]\n"
"sys.argv = [target, *sys.argv[2:]]\n"
"runpy.run_path(target, run_name='__main__')\n"
)


def run_agentmain(*args):
with tempfile.TemporaryDirectory() as tmp:
state = Path(tmp)
(state / "mykey.py").write_text("", encoding="utf-8")
plugins = state / "plugins"
plugins.mkdir()
(plugins / "__init__.py").write_text("", encoding="utf-8")
(plugins / "hooks.py").write_text(
"def trigger(event, ctx): return ctx\n"
"def discover_and_load(plugin_dir=None): return None\n",
encoding="utf-8",
)
env = os.environ.copy()
env["GA_LANG"] = "en"
env["PYTHONNOUSERSITE"] = "1"
env["PYTHONPATH"] = os.pathsep.join((str(state), str(ROOT)))
return subprocess.run(
[sys.executable, "-c", _RUNNER, str(AGENTMAIN), *args],
cwd=state,
stdin=subprocess.DEVNULL,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
timeout=20,
env=env,
)


class AgentMainCliArgumentTests(unittest.TestCase):
def test_unknown_args_without_reflect_fail_fast(self):
result = run_agentmain("--goal", "dummy-goal.json")
self.assertEqual(result.returncode, 2)
self.assertIn("unrecognized arguments: --goal dummy-goal.json", result.stderr)
self.assertNotIn("EOFError", result.stderr)

def test_reflect_extras_require_complete_key_value_pairs(self):
with tempfile.TemporaryDirectory() as tmp:
script = Path(tmp) / "reflect_exit.py"
script.write_text("def check(): return '/exit'\n", encoding="utf-8")
for extras in (("--name",), ("--name", "--other"), ("---name", "hive-master")):
with self.subTest(extras=extras):
result = run_agentmain("--reflect", str(script), *extras)
self.assertEqual(result.returncode, 2)
self.assertIn("reflect extra arguments must be --key value pairs", result.stderr)

def test_reflect_rejects_empty_script_path_explicitly(self):
result = run_agentmain("--reflect", "", "--name", "hive-master")
self.assertEqual(result.returncode, 2)
self.assertIn("--reflect requires a non-empty script path", result.stderr)

def test_reflect_key_value_extras_remain_supported(self):
with tempfile.TemporaryDirectory() as tmp:
script = Path(tmp) / "reflect_exit.py"
script.write_text(
"def init(args): print('INIT_NAME=' + str(args.get('name')))\n"
"def check(): return '/exit'\n",
encoding="utf-8",
)
result = run_agentmain("--reflect", str(script), "--name", "hive-master")
self.assertEqual(result.returncode, 0, result.stderr)
self.assertIn("INIT_NAME=hive-master", result.stdout)
self.assertIn("[Reflect] loaded", result.stdout)


if __name__ == "__main__":
unittest.main()