import json import os import unittest from agent_egress_guard import EgressDenied, Guard, guard_requests from agent_egress_guard.cli import main class FakeProvider: """Stands in for the page-type database.""" def __init__(self, table): self.table = table def lookup(self, host): h = host[4:] if host.startswith("www.") else host return self.table.get(h, {"found": False, "page_types": {}}) class FreeEdition(unittest.TestCase): def setUp(self): self.g = Guard() def d(self, url, method="GET"): return self.g.check(url, method).decision def test_bundled_policy(self): self.assertEqual(sorted(r["id"] for r in self.g.rules), ["login", "password_reset", "signup"]) self.assertEqual(sorted(self.g.hosts_exact), ["169.254.169.254", "metadata.google.internal"]) self.assertEqual(self.g.edition, "free rules") def test_reads_pass(self): self.assertEqual(self.d("https://docs.python.org/3/library/json.html"), "allow") self.assertEqual(self.d("https://example.com/pricing", "HEAD"), "allow") self.assertEqual(self.d("example.com"), "allow") def test_writes_denied_by_default(self): for m in ("POST", "PUT", "PATCH", "DELETE", "MKCOL", "PROPFIND"): v = self.g.check("https://example.com/api/items", m) self.assertEqual((v.decision, v.layer, v.rule), ("deny", "default", "unclassified_write"), m) def test_identity_rules(self): for u in ("https://example.com/login", "https://example.com/wp-login.php", "https://example.com/sign-in/", "https://example.com/register.php", "https://example.com/signup?plan=pro", "https://example.com/forgot-password", "https://example.com/users/sign_up"): v = self.g.check(u) self.assertEqual((v.decision, v.layer), ("deny", "rules"), u) self.assertEqual(self.d("https://example.com/registration-fees"), "allow") self.assertEqual(self.d("https://example.com/blog/how-to-login-safely-guide"), "allow") def test_metadata_hosts(self): v = self.g.check("http://169.254.169.254/latest/meta-data/iam/security-credentials/") self.assertEqual((v.decision, v.layer), ("deny", "high_value_hosts")) self.assertEqual(self.d("http://metadata.google.internal/computeMetadata/v1/"), "deny") def test_bypass_forms(self): for u in ("http://2852039166/latest/meta-data/", "http://0xA9FEA9FE/latest/", "http://0251.0376.0251.0376/", "http://169.254.43518/", "http://169.254.169.254./latest/", "http://[::ffff:a9fe:a9fe]/latest/", "http://[::169.254.169.254]/", "http://METADATA.GOOGLE.INTERNAL./computeMetadata/v1/"): v = self.g.check(u) self.assertEqual((v.decision, v.layer), ("deny", "high_value_hosts"), u) for u in ("https://example.com/%6Cogin", "https://example.com/%6c%6f%67%69%6e", "https://example.com/login;jsessionid=1", "https://example.com/Register;x=1/"): self.assertEqual(self.d(u), "deny", u) for u in ("https://example.com/%256Cogin", "https://example.com/search?q=login%20help", "https://cafe.be/menu", "https://example.com/my/cookie?continue=https%3A%2F%2Fexample.com%2Fmy%2Flogin%2F", "http://10.0.0.1/"): self.assertEqual(self.d(u), "allow", u) def test_page_type_database_layer(self): db = FakeProvider({"shop.example": {"found": True, "page_types": { "login": "https://shop.example/my-account", "pricing": "https://shop.example/plans", "documentation": "https://shop.example/register"}}}) g = Guard(page_types=db) self.assertEqual(g.edition, "free rules + page-type database") # a login page no URL rule can recognise self.assertEqual(Guard().check("https://shop.example/my-account").decision, "allow") v = g.check("https://shop.example/my-account") self.assertEqual((v.decision, v.layer, v.rule), ("deny", "page_type_db", "login")) # verified read page: reads pass, writes do not self.assertEqual(g.check("https://shop.example/plans").decision, "allow") self.assertEqual(g.check("https://shop.example/plans", "POST").decision, "deny") # a read label never overrides a matching rule v = g.check("https://shop.example/register") self.assertEqual((v.decision, v.rule, v.extra.get("db_label")), ("deny", "signup", "documentation")) def test_strict_mode(self): db = FakeProvider({"known.example": {"found": True, "page_types": {}}}) g = Guard(page_types=db, strict=True) self.assertEqual(g.check("https://known.example/docs").decision, "allow") v = g.check("https://unknown.example/docs") self.assertEqual((v.decision, v.rule), ("deny", "unknown_domain")) def test_requests_hook(self): class Req: def __init__(self, url, method): self.url, self.method = url, method class Session: def send(self, request, **kw): return "sent" s = guard_requests(Session()) self.assertEqual(s.send(Req("https://example.com/docs", "GET")), "sent") with self.assertRaises(EgressDenied): s.send(Req("https://example.com/login", "POST")) def test_replay_free(self): from io import StringIO from contextlib import redirect_stdout buf = StringIO() with redirect_stdout(buf): main(["replay", "--json"]) r = json.loads(buf.getvalue()) self.assertEqual((r["denied"], r["of"]), (15, 18)) def test_check_exit_codes(self): from contextlib import redirect_stdout from io import StringIO with redirect_stdout(StringIO()): self.assertEqual(main(["check", "https://example.com/docs"]), 0) self.assertEqual(main(["check", "https://example.com/login", "-X", "POST"]), 2) @unittest.skipUnless(os.environ.get("EGRESS_GUARD_FULL_RULES"), "licensed rules file not available") class FullEdition(unittest.TestCase): def test_replay_full(self): from contextlib import redirect_stdout from io import StringIO buf = StringIO() with redirect_stdout(buf): main(["replay", "--json", "--rules", os.environ["EGRESS_GUARD_FULL_RULES"], "--hosts", os.environ["EGRESS_GUARD_FULL_HOSTS"]]) r = json.loads(buf.getvalue()) self.assertEqual((r["denied"], r["of"]), (18, 18)) if __name__ == "__main__": unittest.main()