diff --git a/README.md b/README.md index 3ccfbf9..d6e415d 100644 --- a/README.md +++ b/README.md @@ -237,8 +237,8 @@ This brings up Kafka, Neo4j, Redis, Postgres, applies graph migrations, writes t | `GET` | `/metrics` | Prometheus exposition (request + query latency histograms) | Set `OTEL_EXPORTER_OTLP_ENDPOINT` to enable distributed tracing (OpenTelemetry OTLP HTTP). -| `POST` | `/query` | Decision search (Neo4j full-text + optional Qdrant merge when `CORTEX_SEMANTIC_ENABLED=true`) | -| `POST` | `/inject` | Ranked context for agents | +| `POST` | `/query` | Decision search (Neo4j full-text + optional Qdrant merge when `CORTEX_SEMANTIC_ENABLED=true`); rate limited per IP (`CORTEX_RATE_LIMIT_QUERY`, default `30/minute`) | +| `POST` | `/inject` | Ranked context for agents; rate limited per IP (`CORTEX_RATE_LIMIT_INJECT`, default `60/minute`) | | `GET` | `/contradictions/pending` | Pending contradiction review items (`workspace_id` query param; `X-Cortex-Roles` for RBAC) | | `GET` | `/decisions/by-system/{system_id}` | Recent decisions affecting a service | | `GET` | `/decisions/{id}/chain` | SUPERSEDES / trigger lineage | diff --git a/api/main.py b/api/main.py index 59d1461..87c7358 100644 --- a/api/main.py +++ b/api/main.py @@ -22,6 +22,8 @@ from fastapi import FastAPI, HTTPException, Request, Response, status from fastapi.middleware.cors import CORSMiddleware from fastapi.responses import JSONResponse +from slowapi import _rate_limit_exceeded_handler +from slowapi.errors import RateLimitExceeded from api.contradictions import router as contradictions_router from api.gdpr import router as gdpr_router @@ -29,6 +31,7 @@ from api.decisions import router as decisions_router from api.deps import RolesDep, memory, set_memory_service from api.memory import MemoryService +from api.rate_limit import inject_rate_limit, limiter, query_rate_limit from api.remember import router as remember_router from api.schemas import ( DecisionResult, @@ -76,6 +79,9 @@ async def lifespan(app: FastAPI) -> AsyncIterator[None]: lifespan=lifespan, ) +app.state.limiter = limiter +app.add_exception_handler(RateLimitExceeded, _rate_limit_exceeded_handler) + # Browsers reject `Access-Control-Allow-Origin: *` together with credentials. # Only enable credentials when an explicit origin allowlist is configured. _cors_origins_raw = os.environ.get("CORS_ORIGINS", "").strip() @@ -197,8 +203,11 @@ async def metrics() -> Response: summary="Query organizational memory", tags=["memory"], ) +@limiter.limit(query_rate_limit) async def query( - request: QueryRequest, + request: Request, + response: Response, + payload: QueryRequest, roles: RolesDep, ) -> QueryResponse: """Search organizational decisions by natural language query. @@ -207,30 +216,31 @@ async def query( Decision: D-004 — Active context injection, not passive retrieval. RBAC: workspace_id scoped — cross-workspace results never returned. + Rate limited per IP (CORTEX_RATE_LIMIT_QUERY, default 30/minute). """ t0 = time.time() log.info( "query.received", - query=request.query[:100], - workspace_id=request.workspace_id, - limit=request.limit, + query=payload.query[:100], + workspace_id=payload.workspace_id, + limit=payload.limit, ) try: records = await memory().query_decisions( - query=request.query, - workspace_id=request.workspace_id, - limit=request.limit, - min_importance=request.min_importance, - min_trust=request.min_trust, - event_types=request.event_types, + query=payload.query, + workspace_id=payload.workspace_id, + limit=payload.limit, + min_importance=payload.min_importance, + min_trust=payload.min_trust, + event_types=payload.event_types, caller_roles=roles, ) results = [DecisionResult(**record) for record in records] except Exception as exc: record_query(status="error", duration_s=time.time() - t0) - log.error("query.failed", error=str(exc), query=request.query[:100]) + log.error("query.failed", error=str(exc), query=payload.query[:100]) raise HTTPException( status_code=status.HTTP_503_SERVICE_UNAVAILABLE, detail="Graph query failed — Neo4j may be unavailable", @@ -242,12 +252,12 @@ async def query( "query.complete", result_count=len(results), latency_ms=latency_ms, - workspace_id=request.workspace_id, + workspace_id=payload.workspace_id, ) return QueryResponse( - query=request.query, - workspace_id=request.workspace_id, + query=payload.query, + workspace_id=payload.workspace_id, results=results, total=len(results), latency_ms=latency_ms, @@ -260,8 +270,11 @@ async def query( summary="Active context injection for AI agents", tags=["memory"], ) +@limiter.limit(inject_rate_limit) async def inject( - request: InjectRequest, + request: Request, + response: Response, + payload: InjectRequest, roles: RolesDep, ) -> InjectResponse: """Inject relevant organizational memory into an AI agent's context window. @@ -270,25 +283,26 @@ async def inject( Returns the most relevant decisions ranked by importance × trust × recency. Ranks injectable decisions by importance × trust (see scoring.trust_scorer). + Rate limited per IP (CORTEX_RATE_LIMIT_INJECT, default 60/minute). """ t0 = time.time() log.info( "inject.received", - agent_id=request.agent_id, - workspace_id=request.workspace_id, - context_length=len(request.context), + agent_id=payload.agent_id, + workspace_id=payload.workspace_id, + context_length=len(payload.context), ) try: injected = await memory().inject_decisions( - context=request.context, - workspace_id=request.workspace_id, + context=payload.context, + workspace_id=payload.workspace_id, caller_roles=roles, - limit=min(request.max_tokens // 400, 10), + limit=min(payload.max_tokens // 400, 10), ) except Exception as exc: - log.error("inject.failed", error=str(exc), agent_id=request.agent_id) + log.error("inject.failed", error=str(exc), agent_id=payload.agent_id) raise HTTPException( status_code=status.HTTP_503_SERVICE_UNAVAILABLE, detail="Context injection failed", @@ -302,8 +316,8 @@ async def inject( token_estimate = sum(len(item.content.split()) for item in decisions) * 4 // 3 return InjectResponse( - agent_id=request.agent_id, - workspace_id=request.workspace_id, + agent_id=payload.agent_id, + workspace_id=payload.workspace_id, injected_decisions=decisions, context_summary=summary, token_estimate=token_estimate, diff --git a/api/rate_limit.py b/api/rate_limit.py new file mode 100644 index 0000000..4bf5875 --- /dev/null +++ b/api/rate_limit.py @@ -0,0 +1,25 @@ +"""Rate limiting for expensive Cortex graph endpoints.""" + +from __future__ import annotations + +import os + +from slowapi import Limiter +from slowapi.util import get_remote_address + + +def _limit(env_name: str, default: str) -> str: + return os.environ.get(env_name, default).strip() or default + + +def query_rate_limit(*_args, **_kwargs) -> str: + """Per-IP limit for POST /query (override with CORTEX_RATE_LIMIT_QUERY).""" + return _limit("CORTEX_RATE_LIMIT_QUERY", "30/minute") + + +def inject_rate_limit(*_args, **_kwargs) -> str: + """Per-IP limit for POST /inject (override with CORTEX_RATE_LIMIT_INJECT).""" + return _limit("CORTEX_RATE_LIMIT_INJECT", "60/minute") + + +limiter = Limiter(key_func=get_remote_address, headers_enabled=True) diff --git a/pyproject.toml b/pyproject.toml index 4c9c183..43d4e52 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -45,6 +45,8 @@ dependencies = [ # Auth "python-jose[cryptography]>=3.3.0", "passlib[bcrypt]>=1.7.4", + # Rate limiting + "slowapi>=0.1.9", # HTTP client "httpx>=0.27.0", # Utilities diff --git a/tests/api/test_main.py b/tests/api/test_main.py index 0c3464c..9bbe7cd 100644 --- a/tests/api/test_main.py +++ b/tests/api/test_main.py @@ -66,6 +66,31 @@ def test_query_endpoint_uses_memory_service() -> None: assert payload["results"][0]["made_by"] == ["alice@company.com"] +def test_query_rate_limit_returns_429(monkeypatch) -> None: + monkeypatch.setenv("CORTEX_RATE_LIMIT_QUERY", "2/minute") + # Re-bind limit callable so decorator sees the env override on next check. + from api import rate_limit as rl + + client = TestClient(app) + mock_memory = AsyncMock( + query_decisions=AsyncMock(return_value=[]), + ) + with patch("api.main.memory", return_value=mock_memory): + # Clear any prior hits for this client IP in the in-memory storage. + rl.limiter.reset() + for _ in range(2): + ok = client.post( + "/query", + json={"query": "rate limit probe", "workspace_id": "ws-1"}, + ) + assert ok.status_code == 200 + limited = client.post( + "/query", + json={"query": "rate limit probe", "workspace_id": "ws-1"}, + ) + assert limited.status_code == 429 + + def test_contradictions_pending() -> None: """Endpoint should return RBAC-filtered rows from the shared memory service.""" client = TestClient(app) diff --git a/uv.lock b/uv.lock index 5903200..6ca6dcd 100644 --- a/uv.lock +++ b/uv.lock @@ -790,6 +790,7 @@ dependencies = [ { name = "qdrant-client" }, { name = "redis" }, { name = "sentence-transformers" }, + { name = "slowapi" }, { name = "spacy" }, { name = "sqlalchemy", extra = ["asyncio"] }, { name = "structlog" }, @@ -843,6 +844,7 @@ requires-dist = [ { name = "redis", specifier = ">=5.0.4" }, { name = "ruff", marker = "extra == 'dev'", specifier = ">=0.4.7" }, { name = "sentence-transformers", specifier = ">=3.0.1" }, + { name = "slowapi", specifier = ">=0.1.9" }, { name = "spacy", specifier = ">=3.7.4" }, { name = "sqlalchemy", extras = ["asyncio"], specifier = ">=2.0.30" }, { name = "structlog", specifier = ">=24.2.0" }, @@ -1165,6 +1167,18 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/2c/fe/df2365e7ed0ccf08901bbaba48f882d876e16ee79c455e7c736a13d7b42b/databricks_sdk-0.107.0-py3-none-any.whl", hash = "sha256:77678ef08c05c276ad0827ad8feb36df240ed621dd544d3747aec91bd614bac0", size = 887698, upload-time = "2026-05-11T12:44:03.914Z" }, ] +[[package]] +name = "deprecated" +version = "1.3.1" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "wrapt" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/49/85/12f0a49a7c4ffb70572b6c2ef13c90c88fd190debda93b23f026b25f9634/deprecated-1.3.1.tar.gz", hash = "sha256:b1b50e0ff0c1fddaa5708a2c6b0a6588bb09b892825ab2b214ac9ea9d92a5223", size = 2932523, upload-time = "2025-10-30T08:19:02.757Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/84/d0/205d54408c08b13550c733c4b85429e7ead111c7f0014309637425520a9a/deprecated-1.3.1-py2.py3-none-any.whl", hash = "sha256:597bfef186b6f60181535a29fbe44865ce137a5079f295b479886c82729d5f3f", size = 11298, upload-time = "2025-10-30T08:19:00.758Z" }, +] + [[package]] name = "distlib" version = "0.4.0" @@ -2283,6 +2297,20 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/ce/62/b40b382fa0c66fee1478073eb8db352a4a6beda4a1adccf1df911d8c289c/librt-0.11.0-cp314-cp314t-win_arm64.whl", hash = "sha256:dee008f20b542e3cd162ba338a7f9ec0f6d23d395f66fe8aeeec3c9d067ea253", size = 102572, upload-time = "2026-05-10T18:17:06.809Z" }, ] +[[package]] +name = "limits" +version = "5.8.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "deprecated" }, + { name = "packaging" }, + { name = "typing-extensions" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/71/69/826a5d1f45426c68d8f6539f8d275c0e4fcaa57f0c017ec3100986558a41/limits-5.8.0.tar.gz", hash = "sha256:c9e0d74aed837e8f6f50d1fcebcf5fd8130957287206bc3799adaee5092655da", size = 226104, upload-time = "2026-02-05T07:17:35.859Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/b9/98/cb5ca20618d205a09d5bec7591fbc4130369c7e6308d9a676a28ff3ab22c/limits-5.8.0-py3-none-any.whl", hash = "sha256:ae1b008a43eb43073c3c579398bd4eb4c795de60952532dc24720ab45e1ac6b8", size = 60954, upload-time = "2026-02-05T07:17:34.425Z" }, +] + [[package]] name = "mako" version = "1.3.12" @@ -4570,6 +4598,18 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/e7/0e/3ae19fa941522cd98e119762e7181d371c8dba0b2d72bfaf9522692e329c/skops-0.14.0-py3-none-any.whl", hash = "sha256:60a5db78a9db46ccee2139a0ba13ab5afb1c96f4749b382e75a371291bbe3e36", size = 132198, upload-time = "2026-04-20T18:23:54.018Z" }, ] +[[package]] +name = "slowapi" +version = "0.1.10" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "limits" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/b9/52/24527cf25a8b508926aff53350b0136561dfe86c7125f61526653666e1b2/slowapi-0.1.10.tar.gz", hash = "sha256:d320d5bc04d9f171a77fb16700faf3036d85b00f420f22924c8a225f95bd14f9", size = 13841, upload-time = "2026-06-13T11:59:31.571Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/c4/8b/1d359f38706b4097d9a943bf8bd22599f537de4cbaff1e622d3e3936e164/slowapi-0.1.10-py3-none-any.whl", hash = "sha256:3acb61561dc9d687e3d3669362ff6a439de9ba44e2fed3a9c165da26b4b83e28", size = 14921, upload-time = "2026-06-13T11:59:30.485Z" }, +] + [[package]] name = "smart-open" version = "7.6.1"