"""Unit tests for granthi-link: login derivation, identity binding rules, link/repos flows against a stub HTTP server that plays both Zitadel userinfo and the Gitea API, plus startup/transport hardening.""" import http.client import json import os import sys import tempfile import threading import unittest import urllib.request from unittest import mock from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "server")) import granthi_link # noqa: E402 class StubUpstream(BaseHTTPRequestHandler): """Plays Zitadel (/oidc/v1/userinfo) and Gitea (everything else). state["users"]: dict login -> email (existing Gitea users) state["hide_once"]: logins whose next GET 404s (simulates a concurrent create racing between the existence check and the create call) """ state = None # dict injected per-test def _json(self, status, obj): payload = json.dumps(obj).encode() self.send_response(status) self.send_header("Content-Type", "application/json") self.send_header("Content-Length", str(len(payload))) self.end_headers() self.wfile.write(payload) def do_GET(self): st = self.state if self.path == "/oidc/v1/userinfo": auth = self.headers.get("Authorization", "") if auth == "Bearer good-token": return self._json(200, {"sub": "123", "preferred_username": "Alice.Smith@org.example", "email": "alice@example.com", "name": "Alice"}) return self._json(401, {"error": "invalid token"}) if self.path.startswith("/api/v1/users/") and not self.path.endswith("/tokens"): login = self.path.rsplit("/", 1)[1] if login in st["hide_once"]: st["hide_once"].discard(login) return self._json(404, {"message": "not found"}) if login in st["users"]: return self._json(200, {"login": login, "email": st["users"][login]}) return self._json(404, {"message": "not found"}) self._json(404, {}) def do_POST(self): st = self.state length = int(self.headers.get("Content-Length", 0)) body = json.loads(self.rfile.read(length) or b"{}") if self.path == "/api/v1/admin/users": if body["username"] in st["users"]: return self._json(409, {"message": "user already exists"}) st["users"][body["username"]] = body["email"] st["created"].append(body) return self._json(201, {"login": body["username"]}) if self.path.startswith("/api/v1/users/") and self.path.endswith("/tokens"): # must arrive with basic auth + Sudo (the verified 1.27 mechanism) st["token_reqs"].append({ "auth": self.headers.get("Authorization", ""), "sudo": self.headers.get("Sudo", ""), "body": body}) if not self.headers.get("Authorization", "").startswith("Basic "): return self._json(401, {"message": "auth required"}) return self._json(201, {"sha1": "MINTED", "name": body["name"]}) if self.path == "/api/v1/user/repos": if body["name"] in st["repos"]: return self._json(409, {"message": "exists"}) st["repos"].add(body["name"]) return self._json(201, {"name": body["name"], "private": body.get("private"), "full_name": f"alice/{body['name']}"}) self._json(404, {}) def log_message(self, *a): pass class ServiceTestBase(unittest.TestCase): def setUp(self): StubUpstream.state = {"users": {}, "created": [], "repos": set(), "token_reqs": [], "hide_once": set()} self.upstream = ThreadingHTTPServer(("127.0.0.1", 0), StubUpstream) threading.Thread(target=self.upstream.serve_forever, daemon=True).start() self.addCleanup(self.upstream.shutdown) base = f"http://127.0.0.1:{self.upstream.server_address[1]}" self.state_dir = tempfile.mkdtemp(prefix="granthi-link-state-") self.state_path = os.path.join(self.state_dir, "state.json") self.svc = granthi_link.LinkService({ "gitea_base": base, "public_gitea_base": "http://public.example:3041", "zitadel_userinfo": f"{base}/oidc/v1/userinfo", "admin_token": "ADMTOK", "admin_login": "root", "admin_password": "rootpw", "test_mode": False, "state_path": self.state_path, }) def enable_test_mode(self): self.svc.cfg["test_mode"] = True patcher = mock.patch.dict( os.environ, {granthi_link.TEST_MODE_ENV: "1"}) patcher.start() self.addCleanup(patcher.stop) def stub_link(self, sub, username, email=None, verified=None, device="d"): ui = {"sub": sub, "preferred_username": username} if email is not None: ui["email"] = email if verified is not None: ui["email_verified"] = verified return self.svc.link({"test_userinfo": ui, "device_name": device}) def read_state(self): with open(self.state_path) as f: return json.load(f) class TestDeriveLogin(unittest.TestCase): def test_strips_domain_and_sanitizes(self): self.assertEqual( granthi_link.LinkService.derive_login( {"preferred_username": "Alice.Smith@org.example"}), "alice.smith") def test_email_fallback(self): self.assertEqual( granthi_link.LinkService.derive_login({"email": "Bob+x@e.com"}), "bob-x") def test_empty_returns_none(self): self.assertIsNone(granthi_link.LinkService.derive_login({})) class TestLink(ServiceTestBase): def test_link_creates_user_and_mints_token(self): status, resp = self.svc.link({"zitadel_access_token": "good-token", "device_name": "mac studio"}) self.assertEqual(status, 200) self.assertEqual(resp["login"], "alice.smith") self.assertEqual(resp["token"], "MINTED") self.assertEqual(resp["gitea_base"], "http://public.example:3041") st = StubUpstream.state self.assertEqual(len(st["created"]), 1) created = st["created"][0] self.assertFalse(created["must_change_password"]) self.assertEqual(created["visibility"], "private") self.assertGreaterEqual(len(created["password"]), 30) req = st["token_reqs"][0] self.assertTrue(req["auth"].startswith("Basic ")) self.assertEqual(req["sudo"], "alice.smith") self.assertEqual(sorted(req["body"]["scopes"]), ["write:repository", "write:user"]) def test_link_records_identity_mapping(self): self.svc.link({"zitadel_access_token": "good-token", "device_name": "d"}) state = self.read_state() rec = state["identities"]["123"] self.assertEqual(rec["login"], "alice.smith") self.assertTrue(rec["created_by_service"]) mode = os.stat(self.state_path).st_mode & 0o777 self.assertEqual(mode, 0o600) def test_link_bad_token_401(self): status, resp = self.svc.link({"zitadel_access_token": "bad", "device_name": "d"}) self.assertEqual(status, 401) def test_link_missing_token_400(self): status, _ = self.svc.link({"device_name": "d"}) self.assertEqual(status, 400) def test_link_missing_sub_422(self): self.enable_test_mode() status, _ = self.svc.link({"test_userinfo": {"preferred_username": "nosub"}, "device_name": "d"}) self.assertEqual(status, 422) class TestIdentityBinding(ServiceTestBase): """Finding 1: account takeover by login collision.""" def setUp(self): super().setUp() self.enable_test_mode() def test_repeat_link_same_sub_reuses_mapping(self): status, resp = self.stub_link("s1", "alice") self.assertEqual(status, 200) # same sub again -- even with a different preferred_username the # mapping wins and no second user is created status, resp = self.stub_link("s1", "totally-different") self.assertEqual(status, 200) self.assertEqual(resp["login"], "alice") self.assertEqual(len(StubUpstream.state["created"]), 1) def test_colliding_username_different_sub_409_no_token(self): status, _ = self.stub_link("s1", "alice") self.assertEqual(status, 200) minted_before = len(StubUpstream.state["token_reqs"]) status, resp = self.stub_link("s2", "alice") # attacker self.assertEqual(status, 409) self.assertIn("not linked to this identity", resp["error"]) # no token minted for the refused identity self.assertEqual(len(StubUpstream.state["token_reqs"]), minted_before) self.assertNotIn("s2", self.read_state()["identities"]) def test_existing_user_binds_on_verified_email_match(self): StubUpstream.state["users"]["bob"] = "bob@example.com" status, resp = self.stub_link("s9", "bob", email="bob@example.com", verified=True) self.assertEqual(status, 200) self.assertEqual(resp["login"], "bob") rec = self.read_state()["identities"]["s9"] self.assertFalse(rec["created_by_service"]) def test_existing_user_unverified_email_409(self): StubUpstream.state["users"]["bob"] = "bob@example.com" status, _ = self.stub_link("s9", "bob", email="bob@example.com", verified=False) self.assertEqual(status, 409) self.assertEqual(StubUpstream.state["token_reqs"], []) def test_existing_user_wrong_email_409(self): StubUpstream.state["users"]["bob"] = "bob@example.com" status, _ = self.stub_link("s9", "bob", email="evil@example.com", verified=True) self.assertEqual(status, 409) self.assertEqual(StubUpstream.state["token_reqs"], []) def test_deleted_service_created_login_is_recreated(self): self.stub_link("s1", "alice") del StubUpstream.state["users"]["alice"] # user deleted in Gitea status, resp = self.stub_link("s1", "alice") self.assertEqual(status, 200) self.assertEqual(resp["login"], "alice") self.assertIn("alice", StubUpstream.state["users"]) def test_deleted_adopted_login_is_refused(self): StubUpstream.state["users"]["bob"] = "bob@example.com" status, _ = self.stub_link("s9", "bob", email="bob@example.com", verified=True) self.assertEqual(status, 200) del StubUpstream.state["users"]["bob"] status, resp = self.stub_link("s9", "bob", email="bob@example.com", verified=True) self.assertEqual(status, 409) self.assertIn("not created by this service", resp["error"]) def test_corrupt_state_fails_closed(self): with open(self.state_path, "w") as f: f.write("{ not json") status, resp = self.svc.link({"test_userinfo": {"sub": "s1", "preferred_username": "alice"}, "device_name": "d"}) self.assertEqual(status, 500) self.assertEqual(StubUpstream.state["token_reqs"], []) class TestConcurrentCreateRace(ServiceTestBase): """Finding 7: Gitea 409 on user create is handled idempotently.""" def setUp(self): super().setUp() self.enable_test_mode() def test_409_on_create_refetches_and_continues(self): # user exists (created by a concurrent request with OUR email) but # the first existence check misses it email = "alice@example.com" StubUpstream.state["users"]["alice"] = email StubUpstream.state["hide_once"].add("alice") status, resp = self.stub_link("s1", "alice", email=email) self.assertEqual(status, 200) self.assertEqual(resp["login"], "alice") self.assertEqual(self.read_state()["identities"]["s1"]["login"], "alice") def test_409_on_create_with_foreign_email_refused(self): StubUpstream.state["users"]["alice"] = "someoneelse@example.com" StubUpstream.state["hide_once"].add("alice") status, resp = self.stub_link("s1", "alice", email="alice@example.com") self.assertEqual(status, 409) self.assertEqual(StubUpstream.state["token_reqs"], []) class TestTestModeGate(ServiceTestBase): """Finding 2: test_mode requires the env gate.""" def test_config_flag_alone_is_ignored(self): self.svc.cfg["test_mode"] = True env = {k: v for k, v in os.environ.items() if k != granthi_link.TEST_MODE_ENV} with mock.patch.dict(os.environ, env, clear=True), \ self.assertLogs("granthi-link", level="ERROR"): status, _ = self.svc.link({"test_userinfo": { "sub": "1", "preferred_username": "x"}}) self.assertEqual(status, 400) # falls through to token-required def test_env_gate_wrong_value_is_ignored(self): self.svc.cfg["test_mode"] = True with mock.patch.dict(os.environ, {granthi_link.TEST_MODE_ENV: "true"}): status, _ = self.svc.link({"test_userinfo": { "sub": "1", "preferred_username": "x"}}) self.assertEqual(status, 400) def test_enabled_with_config_and_env(self): self.enable_test_mode() status, resp = self.stub_link("1", "evetest") self.assertEqual(status, 200) self.assertEqual(resp["login"], "evetest") def test_env_alone_without_config_flag_disabled(self): with mock.patch.dict(os.environ, {granthi_link.TEST_MODE_ENV: "1"}): status, _ = self.svc.link({"test_userinfo": { "sub": "1", "preferred_username": "x"}}) self.assertEqual(status, 400) class TestConfigPerms(unittest.TestCase): """Finding 3: refuse startup on permissive or foreign-owned config.""" def setUp(self): self.tmp = tempfile.NamedTemporaryFile(delete=False, suffix=".json") self.tmp.write(b"{}") self.tmp.close() self.addCleanup(os.unlink, self.tmp.name) def test_0600_ok(self): os.chmod(self.tmp.name, 0o600) self.assertIsNone(granthi_link.check_config_perms(self.tmp.name)) def test_0400_ok(self): os.chmod(self.tmp.name, 0o400) self.assertIsNone(granthi_link.check_config_perms(self.tmp.name)) def test_0644_refused(self): os.chmod(self.tmp.name, 0o644) err = granthi_link.check_config_perms(self.tmp.name) self.assertIn("refusing to start", err) self.assertIn("0o644", err) def test_0640_refused(self): os.chmod(self.tmp.name, 0o640) self.assertIsNotNone(granthi_link.check_config_perms(self.tmp.name)) def test_foreign_owner_refused(self): os.chmod(self.tmp.name, 0o600) not_me = os.geteuid() + 1 err = granthi_link.check_config_perms(self.tmp.name, euid=not_me) self.assertIn("owned by uid", err) def test_main_exits_2_on_permissive_config(self): """The refusal must actually stop startup (exit nonzero), not just return a string -- verified through main().""" with open(self.tmp.name, "w") as f: json.dump({"gitea_base": "http://x", "admin_token": "t", "admin_login": "r", "admin_password": "p"}, f) os.chmod(self.tmp.name, 0o644) with mock.patch.object(granthi_link.sys, "argv", ["granthi_link.py", self.tmp.name]), \ mock.patch.object(granthi_link, "serve") as served, \ self.assertLogs("granthi-link", level="ERROR"): with self.assertRaises(SystemExit) as cm: granthi_link.main() self.assertEqual(cm.exception.code, 2) served.assert_not_called() # never reached serve() class HandlerTestBase(ServiceTestBase): def setUp(self): super().setUp() granthi_link.Handler.service = self.svc self.srv = ThreadingHTTPServer(("127.0.0.1", 0), granthi_link.Handler) threading.Thread(target=self.srv.serve_forever, daemon=True).start() self.addCleanup(self.srv.shutdown) self.port = self.srv.server_address[1] def raw_post(self, path, body_bytes=None, headers=None): conn = http.client.HTTPConnection("127.0.0.1", self.port, timeout=10) self.addCleanup(conn.close) conn.putrequest("POST", path) for k, v in (headers or {}).items(): conn.putheader(k, v) conn.endheaders() if body_bytes: conn.send(body_bytes) resp = conn.getresponse() return resp.status, json.loads(resp.read() or b"{}") class TestBodyLimits(HandlerTestBase): """Finding 6: bounded reads, Content-Length required on POST.""" def test_oversized_content_length_413(self): status, resp = self.raw_post( "/v1/link", headers={"Content-Length": str(granthi_link.MAX_BODY_BYTES + 1)}) self.assertEqual(status, 413) self.assertIn("too large", resp["error"]) def test_oversized_real_body_413(self): """Send an actual over-limit payload on the wire (not just the header), so the response is proven, not the predicate alone.""" body = json.dumps( {"pad": "x" * (granthi_link.MAX_BODY_BYTES + 2000)}).encode() self.assertGreater(len(body), granthi_link.MAX_BODY_BYTES) status, resp = self.raw_post( "/v1/link", body_bytes=body, headers={"Content-Length": str(len(body)), "Content-Type": "application/json"}) self.assertEqual(status, 413) def test_missing_content_length_411(self): status, _ = self.raw_post("/v1/link") self.assertEqual(status, 411) def test_invalid_content_length_400(self): status, _ = self.raw_post("/v1/link", headers={"Content-Length": "banana"}) self.assertEqual(status, 400) def test_normal_post_still_works(self): body = json.dumps({"device_name": "d"}).encode() status, resp = self.raw_post( "/v1/link", body_bytes=body, headers={"Content-Length": str(len(body)), "Content-Type": "application/json"}) self.assertEqual(status, 400) # missing token, but parsed fine self.assertIn("zitadel_access_token", resp["error"]) def test_at_limit_accepted(self): pad = "x" * (granthi_link.MAX_BODY_BYTES - 30) body = json.dumps({"pad": pad}).encode() self.assertLessEqual(len(body), granthi_link.MAX_BODY_BYTES) status, _ = self.raw_post( "/v1/link", body_bytes=body, headers={"Content-Length": str(len(body))}) self.assertEqual(status, 400) # parsed; fails on missing token class TestRepos(ServiceTestBase): def test_repo_create_returns_public_clone_url(self): status, resp = self.svc.repos({"token": "USERTOK", "name": "notes", "private": True}) self.assertEqual(status, 200) self.assertEqual(resp["clone_url"], "http://public.example:3041/alice/notes.git") def test_repo_conflict_409(self): self.svc.repos({"token": "T", "name": "notes"}) status, _ = self.svc.repos({"token": "T", "name": "notes"}) self.assertEqual(status, 409) def test_missing_fields_400(self): status, _ = self.svc.repos({"name": "x"}) self.assertEqual(status, 400) class TestHealth(HandlerTestBase): def test_health_endpoint(self): with urllib.request.urlopen( f"http://127.0.0.1:{self.port}/health") as r: body = json.loads(r.read()) self.assertEqual(body["status"], "ok") self.assertEqual(body["service"], "granthi-link") if __name__ == "__main__": unittest.main() class TestRateLimiterUnit(unittest.TestCase): """The limiter in isolation, on a fake clock -- no sleeping in tests.""" def setUp(self): self.now = 1000.0 self.rl = granthi_link.RateLimiter( rules={"/v1/link": (3, 60)}, clock=lambda: self.now) def test_allows_up_to_the_limit_then_denies(self): for i in range(3): allowed, _ = self.rl.check("/v1/link", "1.1.1.1") self.assertTrue(allowed, f"request {i} should pass") allowed, retry = self.rl.check("/v1/link", "1.1.1.1") self.assertFalse(allowed) self.assertGreater(retry, 0) self.assertLessEqual(retry, 61) def test_window_slides(self): for _ in range(3): self.rl.check("/v1/link", "1.1.1.1") self.assertFalse(self.rl.check("/v1/link", "1.1.1.1")[0]) self.now += 61 self.assertTrue(self.rl.check("/v1/link", "1.1.1.1")[0]) def test_denied_requests_do_not_extend_the_window(self): """A client that keeps hammering must not push its own window forward and lock itself out forever.""" for _ in range(3): self.rl.check("/v1/link", "1.1.1.1") for _ in range(20): # hammer while denied self.now += 1 self.assertFalse(self.rl.check("/v1/link", "1.1.1.1")[0]) self.now = 1000.0 + 61 # just past the ORIGINAL window self.assertTrue(self.rl.check("/v1/link", "1.1.1.1")[0]) def test_clients_are_isolated(self): for _ in range(3): self.rl.check("/v1/link", "1.1.1.1") self.assertFalse(self.rl.check("/v1/link", "1.1.1.1")[0]) self.assertTrue(self.rl.check("/v1/link", "2.2.2.2")[0]) def test_routes_are_isolated_and_unknown_routes_pass(self): for _ in range(3): self.rl.check("/v1/link", "1.1.1.1") self.assertFalse(self.rl.check("/v1/link", "1.1.1.1")[0]) self.assertTrue(self.rl.check("/v1/repos", "1.1.1.1")[0]) for _ in range(50): self.assertTrue(self.rl.check("/health", "1.1.1.1")[0]) def test_zero_limit_disables_the_endpoint(self): rl = granthi_link.RateLimiter(rules={"/v1/link": (0, 60)}, clock=lambda: self.now) self.assertFalse(rl.check("/v1/link", "1.1.1.1")[0]) def test_key_store_stays_bounded(self): rl = granthi_link.RateLimiter(rules={"/v1/link": (5, 60)}, max_keys=50, clock=lambda: self.now) for i in range(500): self.now += 0.001 rl.check("/v1/link", f"10.0.{i // 256}.{i % 256}") self.assertLessEqual(len(rl._hits), 50 + 1) def test_expired_keys_are_reclaimed(self): rl = granthi_link.RateLimiter(rules={"/v1/link": (5, 60)}, max_keys=10, clock=lambda: self.now) for i in range(10): rl.check("/v1/link", f"10.0.0.{i}") self.now += 120 # everything expires for i in range(10, 25): rl.check("/v1/link", f"10.0.0.{i}") self.assertLessEqual(len(rl._hits), 11) def test_concurrent_checks_never_exceed_the_limit(self): """The lock has to actually hold under threads: 40 racing callers against a limit of 10 must yield exactly 10 allows.""" rl = granthi_link.RateLimiter(rules={"/v1/link": (10, 60)}) results, lock = [], threading.Lock() def hit(): a, _ = rl.check("/v1/link", "9.9.9.9") with lock: results.append(a) ts = [threading.Thread(target=hit) for _ in range(40)] for t in ts: t.start() for t in ts: t.join() self.assertEqual(sum(results), 10) class _FakeHandler: def __init__(self, peer, xff=None): self.client_address = (peer, 12345) self.headers = {} if xff is None else {"X-Forwarded-For": xff} if xff is not None: self.headers = type("H", (), {"get": lambda s, k, d="": xff if k == "X-Forwarded-For" else d})() class TestClientIp(unittest.TestCase): PROXY = ("10.0.0.0/8",) def test_socket_peer_by_default(self): h = _FakeHandler("5.5.5.5", xff="1.2.3.4") self.assertEqual(granthi_link.client_ip(h, False, self.PROXY), "5.5.5.5") def test_trusted_proxy_uses_the_last_hop_not_the_client_supplied_first(self): """A caller can PREPEND anything to X-Forwarded-For; a trusted proxy appends the peer it really saw. Only the last entry is trustworthy.""" h = _FakeHandler("10.0.0.1", xff="1.2.3.4, 203.0.113.9") self.assertEqual(granthi_link.client_ip(h, True, self.PROXY), "203.0.113.9") def test_untrusted_peer_cannot_choose_its_own_key(self): """The origin also listens on the tailnet. Anyone reaching it directly must not be able to pick (and rotate) their rate-limit key just by sending a header.""" h = _FakeHandler("203.0.113.50", xff="9.9.9.9") self.assertEqual(granthi_link.client_ip(h, True, self.PROXY), "203.0.113.50") def test_non_ip_last_hop_falls_back_to_peer(self): for junk in ("not-an-ip", "', OR 1=1", "x" * 500, "10.0.0.1:8080"): with self.subTest(junk=junk): h = _FakeHandler("10.0.0.1", xff=f"1.2.3.4, {junk}") self.assertEqual( granthi_link.client_ip(h, True, self.PROXY), "10.0.0.1") def test_trusted_proxy_falls_back_when_header_absent(self): h = _FakeHandler("10.0.0.1") self.assertEqual(granthi_link.client_ip(h, True, self.PROXY), "10.0.0.1") def test_malformed_trusted_proxy_entry_does_not_grant_trust(self): h = _FakeHandler("10.0.0.1", xff="9.9.9.9") self.assertEqual( granthi_link.client_ip(h, True, ("not-a-cidr",)), "10.0.0.1") class TestRateLimitConfig(ServiceTestBase): def _svc(self, rl): cfg = dict(self.svc.cfg) cfg["rate_limit"] = rl return granthi_link.LinkService(cfg) def test_enabled_by_default_when_key_absent(self): self.assertIsNotNone(self.svc.limiter) def test_explicit_disable_is_honored(self): self.assertIsNone(self._svc({"enabled": False}).limiter) def test_custom_rule_overrides_default(self): svc = self._svc({"rules": {"/v1/link": [99, 120]}}) self.assertEqual(svc.limiter.rules["/v1/link"], (99, 120)) def test_malformed_rule_refuses_startup_rather_than_meaning_unlimited(self): for bad in ({"rules": {"/v1/link": [5]}}, {"rules": {"/v1/link": "5/hour"}}, {"rules": {"/v1/link": [5, 0]}}, {"rules": {"/v1/link": [5, -1]}}, {"rules": {"/v1/link": ["5", "60"]}}): with self.subTest(cfg=bad): with self.assertRaises(SystemExit): self._svc(bad) class TestRateLimitOverHttp(HandlerTestBase): """Proven on the wire: real 429, real Retry-After.""" def setUp(self): super().setUp() self.svc.limiter = granthi_link.RateLimiter(rules={"/v1/link": (2, 60)}) def test_third_request_gets_429_with_retry_after(self): body = json.dumps({"zitadel_access_token": "x"}).encode() hdrs = {"Content-Length": str(len(body)), "Content-Type": "application/json"} for _ in range(2): st, _ = self.raw_post("/v1/link", body, hdrs) self.assertNotEqual(st, 429) conn = http.client.HTTPConnection("127.0.0.1", self.port, timeout=10) self.addCleanup(conn.close) conn.request("POST", "/v1/link", body=body, headers=hdrs) resp = conn.getresponse() self.assertEqual(resp.status, 429) self.assertTrue(resp.getheader("Retry-After")) self.assertIn("rate limit", json.loads(resp.read())["error"]) def test_health_is_never_rate_limited(self): for _ in range(30): conn = http.client.HTTPConnection("127.0.0.1", self.port, timeout=10) conn.request("GET", "/health") self.assertEqual(conn.getresponse().status, 200) conn.close() class TestRateLimitHardening(unittest.TestCase): """The four codex [P2] findings, each pinned by a test.""" def test_capacity_fails_closed_instead_of_resetting_a_live_window(self): """Evicting a live window would let an attacker who can mint many distinct keys clear their OWN limit on demand. New keys are refused instead while every window is still live.""" now = [1000.0] rl = granthi_link.RateLimiter(rules={"/v1/link": (1, 3600)}, max_keys=5, clock=lambda: now[0]) for i in range(5): self.assertTrue(rl.check("/v1/link", f"10.0.0.{i}")[0]) victim_hits = list(rl._hits[("/v1/link", "10.0.0.0")]) for i in range(100, 140): # identity flood allowed, retry = rl.check("/v1/link", f"10.0.1.{i}") self.assertFalse(allowed) self.assertGreater(retry, 0) # the earlier client's window survived the flood untouched self.assertEqual(rl._hits[("/v1/link", "10.0.0.0")], victim_hits) self.assertLessEqual(len(rl._hits), 5) def test_capacity_recovers_once_windows_expire(self): now = [1000.0] rl = granthi_link.RateLimiter(rules={"/v1/link": (1, 60)}, max_keys=3, clock=lambda: now[0]) for i in range(3): rl.check("/v1/link", f"10.0.0.{i}") self.assertFalse(rl.check("/v1/link", "10.0.9.9")[0]) now[0] += 61 self.assertTrue(rl.check("/v1/link", "10.0.9.9")[0]) def test_hits_stay_chronological_under_thread_contention(self): """retry_after uses hits[0] and reclamation uses v[-1]; both assume the list is ordered. Reading the clock outside the lock let racing threads append out of order.""" rl = granthi_link.RateLimiter(rules={"/v1/link": (500, 3600)}) ts = [threading.Thread(target=lambda: rl.check("/v1/link", "7.7.7.7")) for _ in range(200)] for t in ts: t.start() for t in ts: t.join() hits = rl._hits[("/v1/link", "7.7.7.7")] self.assertEqual(len(hits), 200) self.assertEqual(hits, sorted(hits), "timestamps out of order") class TestRateLimitConfigTypes(ServiceTestBase): def _svc(self, rl): cfg = dict(self.svc.cfg) cfg["rate_limit"] = rl return granthi_link.LinkService(cfg) def test_non_boolean_enabled_refuses_startup(self): """`"enabled": null` or `0` must not quietly mean unlimited.""" for bad in (None, 0, "", "false", "no", []): with self.subTest(enabled=bad): with self.assertRaises(SystemExit): self._svc({"enabled": bad}) def test_string_false_does_not_enable_forwarded_trust(self): """Every non-empty string is truthy -- "false" used to mean True.""" with self.assertRaises(SystemExit): self._svc({"trust_forwarded_for": "false", "trusted_proxies": ["10.0.0.0/8"]}) def test_rate_limit_must_be_an_object(self): for bad in ("yes", 5, ["/v1/link"]): with self.subTest(rl=bad): with self.assertRaises(SystemExit): self._svc(bad) def test_null_rate_limit_means_defaults_not_disabled(self): svc = self._svc(None) self.assertIsNotNone(svc.limiter) def test_forwarded_trust_without_trusted_proxies_refuses_startup(self): """Trusting the header from ANY peer lets callers choose their own rate-limit key -- that must not be reachable by omission.""" with self.assertRaises(SystemExit): self._svc({"trust_forwarded_for": True}) with self.assertRaises(SystemExit): self._svc({"trust_forwarded_for": True, "trusted_proxies": []}) def test_forwarded_trust_with_proxies_is_accepted(self): svc = self._svc({"trust_forwarded_for": True, "trusted_proxies": ["10.0.0.0/8", "127.0.0.1/32"]}) self.assertTrue(svc.trust_forwarded_for) self.assertEqual(len(svc.trusted_proxies), 2)