diff --git a/app/routes/v1.py b/app/routes/v1.py index c6024ac..557106b 100644 --- a/app/routes/v1.py +++ b/app/routes/v1.py @@ -61,8 +61,14 @@ def _caller(request: Request) -> tuple[str, str, str]: return subject, org_id, user_id -async def _rate_limited(request: Request, subject: str, request_id: str): - limit = Config.V1_RATE_LIMIT_PER_MINUTE +async def _rate_limited(request: Request, subject: str, org_id: str, request_id: str): + """limit is the caller's organization's plan-specific override + (published by omnibioai-billing's gateway_quota_sync_service.py) + when one is set, falling back to the configured global default -- + `is not None`, not `or`, since a plan-specific limit of exactly 0 + is a real (if unusual) value, not "unset".""" + org_limit = await store.rate_limit_for_org(org_id) if org_id else None + limit = org_limit if org_limit is not None else Config.V1_RATE_LIMIT_PER_MINUTE allowed, remaining, reset = await store.hit_rate_limit(subject, limit) headers = { "X-RateLimit-Limit": str(limit), @@ -124,7 +130,7 @@ async def _handle_billable_literature_call( except UnsupportedRequestError as exc: return _error(400, "unsupported_request", exc.message, request_id, detail={"field": exc.field}) - headers, limited = await _rate_limited(request, subject, request_id) + headers, limited = await _rate_limited(request, subject, org_id, request_id) if limited: return limited @@ -226,8 +232,8 @@ async def literature_studies(request: Request): """Free: the queryable studies/domains, from omnibioai-rag GET /v1/studies. Rate-limited like every /v1 call, never billed.""" request_id = getattr(request.state, "trace_id", "") - subject, _, _ = _caller(request) - headers, limited = await _rate_limited(request, subject, request_id) + subject, org_id, _ = _caller(request) + headers, limited = await _rate_limited(request, subject, org_id, request_id) if limited: return limited status, response = await _forward(request, "GET", "v1/studies") @@ -245,8 +251,8 @@ async def literature_domains(request: Request): call, reshaped to the public contract. Rate-limited like every /v1 call, never billed.""" request_id = getattr(request.state, "trace_id", "") - subject, _, _ = _caller(request) - headers, limited = await _rate_limited(request, subject, request_id) + subject, org_id, _ = _caller(request) + headers, limited = await _rate_limited(request, subject, org_id, request_id) if limited: return limited status, response = await _forward(request, "GET", "v1/studies") @@ -282,7 +288,7 @@ async def literature_usage(request: Request): return _error(403, "organization_required", "This API is billed to an organization; your account has none.", request_id) - headers, limited = await _rate_limited(request, subject, request_id) + headers, limited = await _rate_limited(request, subject, org_id, request_id) if limited: return limited @@ -333,8 +339,8 @@ async def literature_models(request: Request): passes back as this endpoint's own `model` request field). """ request_id = getattr(request.state, "trace_id", "") - subject, _, _ = _caller(request) - headers, limited = await _rate_limited(request, subject, request_id) + subject, org_id, _ = _caller(request) + headers, limited = await _rate_limited(request, subject, org_id, request_id) if limited: return limited models = [ diff --git a/app/services/v1_store.py b/app/services/v1_store.py index 22fdb18..8c8a6e5 100644 --- a/app/services/v1_store.py +++ b/app/services/v1_store.py @@ -122,6 +122,33 @@ async def hit_rate_limit(self, subject: str, limit: int) -> tuple[bool, int, int return True, limit, reset return count <= limit, max(0, limit - count), reset + # ---------------- rate limit override (maintained by omnibioai-billing) + def _org_rate_limit_key(self, org_id: str) -> str: + # Nested under the same "quota:" prefix omnibioai-billing's + # gateway_quota_sync_service.py already writes allowance keys + # to -- deliberately not a new "ratelimit:" top-level prefix, + # so this reuses that service's existing Redis ACL grant + # (~gateway:v1:quota:*) exactly as-is. See that module's own + # docstring for the full reasoning. + return f"{_PREFIX}quota:{org_id}:ratelimit" + + async def rate_limit_for_org(self, org_id: str) -> Optional[int]: + """The org's plan-specific requests-per-minute override, or + None when no plan override is set (every plan this platform + currently seeds leaves it unset -- a real number is a product + decision, not an engineering one) or on a Redis error. Callers + fall back to the configured global default in either case.""" + try: + raw = await self.redis.get(self._org_rate_limit_key(org_id)) + except Exception: + return None + if raw is None: + return None + try: + return int(raw) + except ValueError: + return None + # ---------------- quota (maintained by omnibioai-billing) ------------- def _quota_key(self, org_id: str, resource: str) -> str: return f"{_PREFIX}quota:{org_id}:{resource}" diff --git a/tests/test_v1_literature.py b/tests/test_v1_literature.py index ee0bfb0..37aeb31 100644 --- a/tests/test_v1_literature.py +++ b/tests/test_v1_literature.py @@ -222,6 +222,65 @@ def test_rate_limit(client, redis, upstream, monkeypatch): assert len(_usage(redis)) == 2 +def test_rate_limit_uses_the_organizations_plan_specific_override(client, redis, upstream, monkeypatch): + """omnibioai-billing publishes this org's plan-specific override + under gateway:v1:quota:{org}:ratelimit (see + gateway_quota_sync_service.py) -- when set, it wins over the + configured global default.""" + monkeypatch.setattr(Config, "V1_RATE_LIMIT_PER_MINUTE", 100) + redis.kv["gateway:v1:quota:42:ratelimit"] = 1 + + first = _post(client) + assert first.status_code == 200 + assert first.headers["X-RateLimit-Limit"] == "1" + second = _post(client) + assert second.status_code == 429 + + +def test_rate_limit_falls_back_to_the_global_default_when_no_override_is_set(client, redis, upstream, monkeypatch): + monkeypatch.setattr(Config, "V1_RATE_LIMIT_PER_MINUTE", 100) + # No gateway:v1:quota:42:ratelimit key at all. + + resp = _post(client) + assert resp.headers["X-RateLimit-Limit"] == "100" + + +def test_rate_limit_override_of_zero_is_honored_not_treated_as_unset(client, redis, upstream, monkeypatch): + """`is not None`, not a truthiness/`or` check -- a plan-specific + limit of exactly 0 is a real (if unusual) value.""" + monkeypatch.setattr(Config, "V1_RATE_LIMIT_PER_MINUTE", 100) + redis.kv["gateway:v1:quota:42:ratelimit"] = 0 + + resp = _post(client) + assert resp.status_code == 429 + assert resp.headers["X-RateLimit-Limit"] == "0" + + +def test_rate_limit_override_fails_open_on_redis_error(client, upstream, monkeypatch): + """A broken Redis for the rate-limit *lookup* must still let the + request through at the global default -- the same fail-open + posture every other V1Store method takes on a Redis error.""" + monkeypatch.setattr(Config, "V1_RATE_LIMIT_PER_MINUTE", 100) + + class BrokenGetRedis: + async def get(self, key): + raise ConnectionError("redis down") + + async def incr(self, key): + return 1 + + async def expire(self, key, ttl): + return True + + with patch.object(v1.store, "redis", BrokenGetRedis()), \ + patch.object(v1.store, "usage_redis", BrokenGetRedis()), \ + patch.object(v1.store, "outbox", Outbox(":memory:")): + resp = _post(client) + + assert resp.status_code == 200 + assert resp.headers["X-RateLimit-Limit"] == "100" + + def test_quota_exhausted_returns_402_without_calling_upstream(client, redis, upstream): redis.kv["gateway:v1:quota:42:literature.answer"] = 0 resp = _post(client, idem="q-1") @@ -511,6 +570,17 @@ def test_store_edge_cases(tmp_path): assert asyncio.run(store.idempotency_begin("s", "k", "f")) == {"state": "in_progress"} +def test_rate_limit_for_org_edge_cases(tmp_path): + import asyncio + with patch("app.services.v1_store.aioredis.from_url", return_value=FakeRedis()): + store = V1Store("redis://x", "redis://y", outbox_path=str(tmp_path / "outbox.db")) + assert asyncio.run(store.rate_limit_for_org("1")) is None # no key set at all + store.redis.kv["gateway:v1:quota:1:ratelimit"] = "not-a-number" + assert asyncio.run(store.rate_limit_for_org("1")) is None + store.redis.kv["gateway:v1:quota:1:ratelimit"] = "30" + assert asyncio.run(store.rate_limit_for_org("1")) == 30 + + def test_memory_store_semantics(): import asyncio