Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
26 changes: 16 additions & 10 deletions app/routes/v1.py
Original file line number Diff line number Diff line change
Expand Up @@ -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),
Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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")
Expand All @@ -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")
Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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 = [
Expand Down
27 changes: 27 additions & 0 deletions app/services/v1_store.py
Original file line number Diff line number Diff line change
Expand Up @@ -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}"
Expand Down
70 changes: 70 additions & 0 deletions tests/test_v1_literature.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand Down Expand Up @@ -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

Expand Down
Loading