All four were real bypass or fail-open paths on an endpoint about to be publicly exposed: - X-Forwarded-For was trusted from ANY peer. The origin also listens on the tailnet, so anyone reaching it directly could pick -- and rotate -- their own rate-limit key by sending a header. Now honored only when the socket peer is in a configured trusted_proxies list, and the last hop must parse as a real IP. trust_forwarded_for without trusted_proxies REFUSES startup. - Capacity eviction was fail-open and exploitable: an attacker able to mint many distinct keys could evict their own live window and start fresh. Now reclaims only EXPIRED windows and refuses the new key when all are live. Fail closed -- /v1/link is invite-only, so hitting the cap is an attack. - The clock was read outside the lock, so racing threads could append out of order; both retry_after (hits[0]) and reclamation (v[-1]) assume the list is chronological. Moved inside. - Config types were unvalidated: `"enabled": null` or `0` silently disabled limiting, and the string "false" enabled XFF trust (non-empty strings are truthy). Booleans must now be real JSON booleans; rate_limit must be an object. Codex confirmed no path-variant bypass (dispatch is exact-match) and no keep-alive/pipelining bypass (rejects set close_connection). Tests 87 -> 99: capacity fail-closed with the victim's window proven untouched through a 40-key flood, 200-thread chronological-order check, untrusted-peer spoof, junk XFF, and every config-type trap. Co-Authored-By: Claude Opus 5 <[email protected]> Claude-Session: https://claude.ai/code/session_01LTARYHX7GPepi3CH3tp5pg
773 lines
32 KiB
Python
773 lines
32 KiB
Python
"""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":
|
|
"[email protected]", "email":
|
|
"[email protected]", "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": "[email protected]"}),
|
|
"alice.smith")
|
|
|
|
def test_email_fallback(self):
|
|
self.assertEqual(
|
|
granthi_link.LinkService.derive_login({"email": "[email protected]"}),
|
|
"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"] = "[email protected]"
|
|
status, resp = self.stub_link("s9", "bob", email="[email protected]",
|
|
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"] = "[email protected]"
|
|
status, _ = self.stub_link("s9", "bob", email="[email protected]",
|
|
verified=False)
|
|
self.assertEqual(status, 409)
|
|
self.assertEqual(StubUpstream.state["token_reqs"], [])
|
|
|
|
def test_existing_user_wrong_email_409(self):
|
|
StubUpstream.state["users"]["bob"] = "[email protected]"
|
|
status, _ = self.stub_link("s9", "bob", email="[email protected]",
|
|
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"] = "[email protected]"
|
|
status, _ = self.stub_link("s9", "bob", email="[email protected]",
|
|
verified=True)
|
|
self.assertEqual(status, 200)
|
|
del StubUpstream.state["users"]["bob"]
|
|
status, resp = self.stub_link("s9", "bob", email="[email protected]",
|
|
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 = "[email protected]"
|
|
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"] = "[email protected]"
|
|
StubUpstream.state["hide_once"].add("alice")
|
|
status, resp = self.stub_link("s1", "alice",
|
|
email="[email protected]")
|
|
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)
|