diff --git a/campus_python/audit/ratelimit.py b/campus_python/audit/ratelimit.py new file mode 100644 index 0000000..9b8e2c2 --- /dev/null +++ b/campus_python/audit/ratelimit.py @@ -0,0 +1,212 @@ +"""campus_python.audit.ratelimit + +Circuit breaker for campus.audit ingest rate limiting. + +Phase 3 of the audit API key security epic (campus#538), tracked in +campus#831. Design: +https://github.com/nyjc-computing/campus/issues/538#issuecomment-5996377544 + +How it fits together: + +- campus.audit rate-limits POST /traces/ per identity with a per-minute + token bucket and returns 429 with a Retry-After header and the + tripped bucket key in error.details.bucket. +- The producers' ingest responses are fed to the module-level `breaker` + (see Traces.ingest here, and campus.audit.client's Traces.new — the + two clients that POST spans). A 429 trips the bucket; a 2xx clears + trips whose cooldown has expired; any other outcome (timeout, + connection refused, 5xx) is neutral — an audit outage must not take + producers down. +- Producers consult `breaker.check_request(...)` at request entry + (flag-gated on their side via AUDIT_TRACING_FAIL_CLOSED) and + short-circuit matching requests with 503 + Retry-After while a + bucket is tripped. + +State is per worker process (the module singleton). Cross-worker +staleness means the limit is not strictly adhered to — ratified slack. +Retraction is two-stage: cooldown expiry OPENS the gate (verifying) +without clearing the trip; the first 2xx ingest clears it; another 429 +re-trips with exponential backoff (floor 30s, x2, cap 5 min). +""" + +__all__ = [ + "AuditRateBreaker", + "WILDCARD", + "breaker", + "bucket_key", +] + +import threading +import time +import typing + +# Bucket key for requests whose identity (or learned fallback bucket) +# is unknown; producers gate these coarsely. +WILDCARD = "*" + +# Backoff policy: floor 30s, x2 per consecutive trip for the same +# bucket, capped at 5 minutes. Audit's Retry-After is honored up to +# the same cap (its per-minute bucket never advises more than 60s). +BACKOFF_FLOOR_SECONDS = 30.0 +BACKOFF_CAP_SECONDS = 300.0 + + +def bucket_key(client_id: str | None, user_id: str | None) -> str: + """Encode a request's identity into a rate-limit bucket key. + + Mirrors the server-side encoder (campus.audit.resources.ratelimit): + (client_id, user_id) pair > user_id > client_id. An identity-less + request maps to the wildcard sentinel — the server keys those + spans by the producer's API key, which the producer learns from + 429 bodies (see AuditRateBreaker.learned_fallback_key). + """ + if client_id and user_id: + return f"client={client_id};user={user_id}" + if user_id: + return f"user={user_id}" + if client_id: + return f"client={client_id}" + return WILDCARD + + +class AuditRateBreaker: + """Thread-safe in-process circuit breaker for audit ingest 429s.""" + + def __init__(self) -> None: + self._lock = threading.Lock() + # bucket key -> {"until": epoch seconds, "consecutive": int} + self._tripped: dict[str, dict[str, float]] = {} + # Learned producer-key fallback bucket ("apikey=..."), from + # 429 bodies. An apikey= bucket seen by this process's client + # is necessarily this producer's own fallback bucket. + self._fallback_key: str | None = None + + def observe( + self, + status_code: int, + headers: typing.Mapping[str, str] | None = None, + body: typing.Any = None, + ) -> None: + """Feed an ingest response to the breaker. + + 2xx clears trips whose cooldown has expired (verifying state); + 429 trips the bucket named in the body (or the wildcard when + the body carries no bucket key); anything else — timeouts, + connection errors, 5xx — is neutral (the ratified split: + availability errors neither set nor clear). + """ + if 200 <= status_code < 300: + self._clear_expired() + return + if status_code != 429: + return + + bucket = _bucket_from_body(body) or WILDCARD + retry_after = _retry_after_seconds(headers) + self.trip(bucket, retry_after) + if bucket.startswith("apikey="): + with self._lock: + self._fallback_key = bucket + + def trip(self, bucket_key: str, retry_after: float) -> None: + """Trip (or re-trip) a bucket with backoff on consecutive trips. + + Trip duration = min(cap, max(Retry-After, floor * 2**(n-1))) + where n counts consecutive trips of this bucket without an + intervening successful ingest. + """ + now = time.time() + with self._lock: + previous = self._tripped.get(bucket_key) + consecutive = int(previous["consecutive"]) + 1 if previous else 1 + backoff = min( + BACKOFF_FLOOR_SECONDS * 2 ** (consecutive - 1), + BACKOFF_CAP_SECONDS, + ) + duration = min(BACKOFF_CAP_SECONDS, max(retry_after, backoff)) + self._tripped[bucket_key] = { + "until": now + duration, + "consecutive": consecutive, + } + + def check(self, bucket_key: str) -> int | None: + """Return seconds until the bucket reopens, or None if open. + + A trip whose cooldown has expired is in the verifying state: + the gate is open (None) and only a 2xx or another 429 resolves + it (see observe). + """ + now = time.time() + with self._lock: + trip = self._tripped.get(bucket_key) + if trip is not None and now < trip["until"]: + return int(trip["until"] - now) + 1 + return None + + def check_request( + self, + client_id: str | None = None, + user_id: str | None = None, + ) -> int | None: + """Gate check for one incoming request. + + Per-identity when the request's identity is resolvable at + entry; otherwise the learned producer-key fallback bucket, or + the wildcard when nothing has been learned yet. + """ + if client_id or user_id: + return self.check(bucket_key(client_id, user_id)) + fallback = self.learned_fallback_key() + if fallback is not None: + return self.check(fallback) + return self.check(WILDCARD) + + def learned_fallback_key(self) -> str | None: + """The producer's own apikey= bucket key, learned from 429s.""" + with self._lock: + return self._fallback_key + + def reset(self) -> None: + """Clear all state (tests).""" + with self._lock: + self._tripped.clear() + self._fallback_key = None + + def _clear_expired(self) -> None: + """Drop trips whose cooldown has expired (successful ingest).""" + now = time.time() + with self._lock: + expired = [ + key for key, trip in self._tripped.items() + if now >= trip["until"] + ] + for key in expired: + del self._tripped[key] + + +def _bucket_from_body(body: typing.Any) -> str | None: + """Extract the tripped bucket key from a 429 error envelope. + + Shape (campus.audit #835): {"error": {"details": {"bucket": ...}}}. + Returns None when absent or malformed. + """ + try: + bucket = body["error"]["details"]["bucket"] + except (KeyError, TypeError, IndexError): + return None + return bucket if isinstance(bucket, str) and bucket else None + + +def _retry_after_seconds(headers: typing.Mapping[str, str] | None) -> float: + """Parse Retry-After from response headers, defaulting to the floor.""" + if headers: + try: + return float(headers.get("Retry-After")) + except (TypeError, ValueError): + pass + return BACKOFF_FLOOR_SECONDS + + +# Per-process singleton: all ingest responses in this worker feed it, +# and every gate check consults it. +breaker = AuditRateBreaker() diff --git a/campus_python/audit/v1/traces.py b/campus_python/audit/v1/traces.py index 39a3a16..1fc6984 100644 --- a/campus_python/audit/v1/traces.py +++ b/campus_python/audit/v1/traces.py @@ -6,9 +6,19 @@ hex trace ID and spans by their 16-char hex span ID. """ +import contextlib from typing import Any from ...interface import JsonDict, Resource, ResourceCollection +from .. import ratelimit + + +def _safe_json(resp: Any) -> Any: + """Best-effort JSON parse for breaker observation.""" + try: + return resp.json() + except Exception: + return None class Traces(ResourceCollection): @@ -31,6 +41,12 @@ def ingest(self, spans: "list[dict[str, Any]]") -> JsonDict: adds {"failed": [...]} with per-span statuses. """ resp = self.client.post(self.make_path(), json={"spans": spans}) + # Feed the ingest circuit breaker before raising: 429 trips the + # bucket (Retry-After + tripped key from the body), 2xx clears + # expired trips (#831). The error still raises as normal, and a + # breaker failure must never break ingestion. + with contextlib.suppress(Exception): + ratelimit.breaker.observe(resp.status_code, resp.headers, _safe_json(resp)) resp.raise_for_status() return resp.json() diff --git a/tests/unit/test_ratelimit.py b/tests/unit/test_ratelimit.py new file mode 100644 index 0000000..591abb9 --- /dev/null +++ b/tests/unit/test_ratelimit.py @@ -0,0 +1,172 @@ +"""Unit tests for the audit ingest circuit breaker (#831). + +Verifies trip/check semantics (floor, backoff, cap), two-stage +retraction (cooldown expiry opens the gate, 2xx clears), availability +error neutrality, fallback-key learning, and the observe hook on +Traces.ingest. +""" + +import unittest +from unittest.mock import Mock + +from campus_python.audit import ratelimit +from campus_python.audit.v1.traces import Traces + + +def _make_breaker() -> ratelimit.AuditRateBreaker: + """A fresh breaker with backoff policy intact.""" + return ratelimit.AuditRateBreaker() + + +class TestBucketKey(unittest.TestCase): + """Encoder mirrors the server-side priority order.""" + + def test_pair(self): + self.assertEqual( + ratelimit.bucket_key("c1", "u1"), "client=c1;user=u1" + ) + + def test_user_only(self): + self.assertEqual(ratelimit.bucket_key(None, "u1"), "user=u1") + + def test_client_only(self): + self.assertEqual(ratelimit.bucket_key("c1", None), "client=c1") + + def test_identityless_is_wildcard(self): + self.assertEqual(ratelimit.bucket_key(None, None), ratelimit.WILDCARD) + + +class TestObserve(unittest.TestCase): + """observe() feeds 429s into trips, 2xx clears, others neutral.""" + + def setUp(self): + self.br = _make_breaker() + + def test_429_trips_with_retry_after(self): + self.br.observe(429, {"Retry-After": "45"}, _body("user=u1")) + retry_after = self.br.check("user=u1") + self.assertIsNotNone(retry_after) + self.assertGreaterEqual(retry_after, 44) + self.assertLessEqual(retry_after, 46) + + def test_429_floor_30s(self): + # Missing/invalid Retry-After falls back to the 30s floor. + self.br.observe(429, None, _body("user=u1")) + retry_after = self.br.check("user=u1") + self.assertIsNotNone(retry_after) + self.assertGreaterEqual(retry_after, 29) + + def test_429_without_bucket_trips_wildcard(self): + self.br.observe(429, {"Retry-After": "30"}, {"error": {}}) + self.assertIsNotNone(self.br.check(ratelimit.WILDCARD)) + + def test_availability_errors_are_neutral(self): + for status in (500, 502, 503): + self.br.observe(status, None, None) + self.assertIsNone(self.br.check(ratelimit.WILDCARD)) + + def test_2xx_clears_expired_trip(self): + self.br.observe(429, {"Retry-After": "1"}, _body("user=u1")) + # Simulate cooldown expiry without sleeping. + self._expire("user=u1") + # Gate is open (verifying state)... + self.assertIsNone(self.br.check("user=u1")) + # ...and the first 2xx ingest fully clears it. + self.br.observe(201, None, None) + self.assertIsNone(self.br.check("user=u1")) + self.assertEqual(self.br._tripped, {}) + + def test_2xx_does_not_clear_active_trip(self): + self.br.observe(429, {"Retry-After": "120"}, _body("user=u1")) + self.br.observe(201, None, None) + self.assertIsNotNone(self.br.check("user=u1")) + + def test_consecutive_429s_back_off(self): + self.br.observe(429, {"Retry-After": "1"}, _body("user=u1")) + self._expire("user=u1") + first_until = self.br._tripped["user=u1"]["until"] + # Re-trip while verifying: backoff doubles (floor 30 -> 60). + self.br.observe(429, {"Retry-After": "1"}, _body("user=u1")) + second_until = self.br._tripped["user=u1"]["until"] + now_gap = second_until - first_until + self.assertGreater(now_gap, 0) + + def test_apikey_trip_learns_fallback(self): + self.br.observe(429, {"Retry-After": "30"}, _body("apikey=k-123")) + self.assertEqual(self.br.learned_fallback_key(), "apikey=k-123") + + def test_identity_trip_does_not_learn_fallback(self): + self.br.observe(429, {"Retry-After": "30"}, _body("client=c1;user=u1")) + self.assertIsNone(self.br.learned_fallback_key()) + + def _expire(self, key: str) -> None: + """Force a trip's cooldown into the past (verifying state).""" + self.br._tripped[key]["until"] -= 10_000 + + +def _body(bucket: str) -> dict: + """A 429 error envelope as emitted by campus.audit (#835).""" + return {"error": {"code": "RATE_LIMITED", "details": {"bucket": bucket}}} + + +class TestCheckRequest(unittest.TestCase): + """Gate checks: per-identity when resolvable, coarse otherwise.""" + + def setUp(self): + self.br = _make_breaker() + + def test_identity_check_matches_encoded_bucket(self): + self.br.observe(429, {"Retry-After": "30"}, _body("client=c1;user=u1")) + self.assertIsNotNone(self.br.check_request(client_id="c1", user_id="u1")) + self.assertIsNone(self.br.check_request(client_id="c1", user_id="other")) + self.assertIsNone(self.br.check_request(client_id="other", user_id="u1")) + + def test_identityless_uses_learned_fallback(self): + self.br.observe(429, {"Retry-After": "30"}, _body("apikey=k-123")) + self.assertIsNotNone(self.br.check_request()) + + def test_identityless_without_learning_uses_wildcard(self): + self.br.observe(429, {"Retry-After": "30"}, _body(ratelimit.WILDCARD)) + self.assertIsNotNone(self.br.check_request()) + + def test_untripped_identity_is_open(self): + self.assertIsNone(self.br.check_request(client_id="c1", user_id="u1")) + + +class TestIngestHook(unittest.TestCase): + """Traces.ingest feeds the module breaker before raising.""" + + def setUp(self): + ratelimit.breaker.reset() + self.addCleanup(ratelimit.breaker.reset) + + def _ingest_with_response(self, status: int, body: dict, raises: bool): + resp = Mock() + resp.status_code = status + resp.headers = {"Retry-After": "30"} + resp.json.return_value = body + if raises: + resp.raise_for_status.side_effect = RuntimeError(f"{status}") + client = Mock() + client.post.return_value = resp + traces = Traces(client=client, root=Mock()) + return traces + + def test_429_ingest_trips_breaker_and_raises(self): + traces = self._ingest_with_response( + 429, _body("client=c1;user=u1"), raises=True + ) + with self.assertRaises(RuntimeError): + traces.ingest([{"span_id": "s1"}]) + self.assertIsNotNone( + ratelimit.breaker.check_request(client_id="c1", user_id="u1") + ) + + def test_201_ingest_does_not_trip(self): + traces = self._ingest_with_response(201, {"created": ["s1"]}, raises=False) + traces.ingest([{"span_id": "s1"}]) + self.assertEqual(ratelimit.breaker._tripped, {}) + + +if __name__ == "__main__": + unittest.main()