feat: add optional API key authentication to REST API server (#4442)
Co-authored-by: utkarsh240799 <utkarsh240799@users.noreply.github.com> Co-authored-by: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -14,7 +14,7 @@ The Mem0 REST API server exposes every OSS memory operation over HTTP. Run it al
|
||||
</Info>
|
||||
|
||||
<Warning>
|
||||
Add your own authentication and HTTPS before exposing the server to anything beyond your internal network. The default image does not include auth.
|
||||
Enable API key authentication (see below) and HTTPS before exposing the server to anything beyond your internal network.
|
||||
</Warning>
|
||||
|
||||
---
|
||||
@@ -22,6 +22,7 @@ The Mem0 REST API server exposes every OSS memory operation over HTTP. Run it al
|
||||
## Feature
|
||||
|
||||
- **CRUD endpoints:** Create, retrieve, search, update, delete, and reset memories by `user_id`, `agent_id`, or `run_id`.
|
||||
- **API key authentication:** Optionally secure all endpoints with a shared API key via the `X-API-Key` header.
|
||||
- **Status health check:** Access base routes to confirm the server is online.
|
||||
- **OpenAPI explorer:** Visit `/docs` for interactive testing and schema reference.
|
||||
|
||||
@@ -91,6 +92,41 @@ uvicorn main:app --reload
|
||||
|
||||
---
|
||||
|
||||
## Authentication
|
||||
|
||||
The server supports optional API key authentication. When the `ADMIN_API_KEY` environment variable is set, every endpoint requires a valid `X-API-Key` header. The `/` redirect, `/docs`, and `/openapi.json` routes remain open so you can always reach the interactive API explorer.
|
||||
|
||||
| `ADMIN_API_KEY` value | Behavior |
|
||||
|---|---|
|
||||
| Not set / empty | All endpoints are open (no auth) |
|
||||
| Any non-empty string | Requests must include `X-API-Key: <your-key>` |
|
||||
|
||||
### Enable authentication
|
||||
|
||||
Add the key to your `.env` file:
|
||||
|
||||
```bash
|
||||
ADMIN_API_KEY=your-secret-api-key
|
||||
```
|
||||
|
||||
Then include the header in every request:
|
||||
|
||||
```bash
|
||||
curl -X POST http://localhost:8000/memories \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "X-API-Key: your-secret-api-key" \
|
||||
-d '{
|
||||
"messages": [{"role": "user", "content": "I love pizza."}],
|
||||
"user_id": "alice"
|
||||
}'
|
||||
```
|
||||
|
||||
<Warning>
|
||||
The server logs a warning at startup when `ADMIN_API_KEY` is not set. Always set it in production.
|
||||
</Warning>
|
||||
|
||||
---
|
||||
|
||||
## See it in action
|
||||
|
||||
### Create and search memories via HTTP
|
||||
@@ -137,7 +173,7 @@ curl "http://localhost:8000/memories/search?user_id=alice&query=vegetable"
|
||||
|
||||
## Best practices
|
||||
|
||||
1. **Add authentication:** Protect endpoints with API gateways, proxies, or custom FastAPI middleware.
|
||||
1. **Enable authentication:** Set `ADMIN_API_KEY` to secure all endpoints, or use an API gateway for more advanced schemes.
|
||||
2. **Use HTTPS:** Terminate TLS at your load balancer or reverse proxy.
|
||||
3. **Monitor uptime:** Track request rates, latency, and error codes per endpoint.
|
||||
4. **Version configs:** Keep environment files and Docker Compose definitions in source control.
|
||||
|
||||
@@ -10,3 +10,6 @@ POSTGRES_DB=
|
||||
POSTGRES_USER=
|
||||
POSTGRES_PASSWORD=
|
||||
POSTGRES_COLLECTION_NAME=
|
||||
|
||||
# Optional: set to enable API key authentication on all endpoints
|
||||
ADMIN_API_KEY=
|
||||
|
||||
+55
-10
@@ -1,10 +1,12 @@
|
||||
import logging
|
||||
import os
|
||||
import secrets
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from dotenv import load_dotenv
|
||||
from fastapi import FastAPI, HTTPException
|
||||
from fastapi import Depends, FastAPI, HTTPException
|
||||
from fastapi.responses import JSONResponse, RedirectResponse
|
||||
from fastapi.security import APIKeyHeader
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from mem0 import Memory
|
||||
@@ -14,6 +16,22 @@ logging.basicConfig(level=logging.INFO, format="%(asctime)s - %(levelname)s - %(
|
||||
# Load environment variables
|
||||
load_dotenv()
|
||||
|
||||
ADMIN_API_KEY = os.environ.get("ADMIN_API_KEY", "")
|
||||
|
||||
MIN_KEY_LENGTH = 16
|
||||
|
||||
if not ADMIN_API_KEY:
|
||||
logging.warning(
|
||||
"ADMIN_API_KEY not set - API endpoints are UNSECURED! "
|
||||
"Set ADMIN_API_KEY environment variable for production use."
|
||||
)
|
||||
else:
|
||||
if len(ADMIN_API_KEY) < MIN_KEY_LENGTH:
|
||||
logging.warning(
|
||||
"ADMIN_API_KEY is shorter than %d characters - consider using a longer key for production.",
|
||||
MIN_KEY_LENGTH,
|
||||
)
|
||||
logging.info("API key authentication enabled")
|
||||
|
||||
POSTGRES_HOST = os.environ.get("POSTGRES_HOST", "postgres")
|
||||
POSTGRES_PORT = os.environ.get("POSTGRES_PORT", "5432")
|
||||
@@ -60,10 +78,35 @@ MEMORY_INSTANCE = Memory.from_config(DEFAULT_CONFIG)
|
||||
|
||||
app = FastAPI(
|
||||
title="Mem0 REST APIs",
|
||||
description="A REST API for managing and searching memories for your AI Agents and Apps.",
|
||||
description=(
|
||||
"A REST API for managing and searching memories for your AI Agents and Apps.\n\n"
|
||||
"## Authentication\n"
|
||||
"When the ADMIN_API_KEY environment variable is set, all endpoints require "
|
||||
"the `X-API-Key` header for authentication."
|
||||
),
|
||||
version="1.0.0",
|
||||
)
|
||||
|
||||
api_key_header = APIKeyHeader(name="X-API-Key", auto_error=False)
|
||||
|
||||
|
||||
async def verify_api_key(api_key: Optional[str] = Depends(api_key_header)):
|
||||
"""Validate the API key when ADMIN_API_KEY is configured. No-op otherwise."""
|
||||
if ADMIN_API_KEY:
|
||||
if api_key is None:
|
||||
raise HTTPException(
|
||||
status_code=401,
|
||||
detail="X-API-Key header is required.",
|
||||
headers={"WWW-Authenticate": "ApiKey"},
|
||||
)
|
||||
if not secrets.compare_digest(api_key, ADMIN_API_KEY):
|
||||
raise HTTPException(
|
||||
status_code=401,
|
||||
detail="Invalid API key.",
|
||||
headers={"WWW-Authenticate": "ApiKey"},
|
||||
)
|
||||
return api_key
|
||||
|
||||
|
||||
class Message(BaseModel):
|
||||
role: str = Field(..., description="Role of the message (user or assistant).")
|
||||
@@ -87,7 +130,7 @@ class SearchRequest(BaseModel):
|
||||
|
||||
|
||||
@app.post("/configure", summary="Configure Mem0")
|
||||
def set_config(config: Dict[str, Any]):
|
||||
def set_config(config: Dict[str, Any], _api_key: Optional[str] = Depends(verify_api_key)):
|
||||
"""Set memory configuration."""
|
||||
global MEMORY_INSTANCE
|
||||
MEMORY_INSTANCE = Memory.from_config(config)
|
||||
@@ -95,7 +138,7 @@ def set_config(config: Dict[str, Any]):
|
||||
|
||||
|
||||
@app.post("/memories", summary="Create memories")
|
||||
def add_memory(memory_create: MemoryCreate):
|
||||
def add_memory(memory_create: MemoryCreate, _api_key: Optional[str] = Depends(verify_api_key)):
|
||||
"""Store new memories."""
|
||||
if not any([memory_create.user_id, memory_create.agent_id, memory_create.run_id]):
|
||||
raise HTTPException(status_code=400, detail="At least one identifier (user_id, agent_id, run_id) is required.")
|
||||
@@ -114,6 +157,7 @@ def get_all_memories(
|
||||
user_id: Optional[str] = None,
|
||||
run_id: Optional[str] = None,
|
||||
agent_id: Optional[str] = None,
|
||||
_api_key: Optional[str] = Depends(verify_api_key),
|
||||
):
|
||||
"""Retrieve stored memories."""
|
||||
if not any([user_id, run_id, agent_id]):
|
||||
@@ -129,7 +173,7 @@ def get_all_memories(
|
||||
|
||||
|
||||
@app.get("/memories/{memory_id}", summary="Get a memory")
|
||||
def get_memory(memory_id: str):
|
||||
def get_memory(memory_id: str, _api_key: Optional[str] = Depends(verify_api_key)):
|
||||
"""Retrieve a specific memory by ID."""
|
||||
try:
|
||||
return MEMORY_INSTANCE.get(memory_id)
|
||||
@@ -139,7 +183,7 @@ def get_memory(memory_id: str):
|
||||
|
||||
|
||||
@app.post("/search", summary="Search memories")
|
||||
def search_memories(search_req: SearchRequest):
|
||||
def search_memories(search_req: SearchRequest, _api_key: Optional[str] = Depends(verify_api_key)):
|
||||
"""Search for memories based on a query."""
|
||||
try:
|
||||
params = {k: v for k, v in search_req.model_dump().items() if v is not None and k != "query"}
|
||||
@@ -150,7 +194,7 @@ def search_memories(search_req: SearchRequest):
|
||||
|
||||
|
||||
@app.put("/memories/{memory_id}", summary="Update a memory")
|
||||
def update_memory(memory_id: str, updated_memory: Dict[str, Any]):
|
||||
def update_memory(memory_id: str, updated_memory: Dict[str, Any], _api_key: Optional[str] = Depends(verify_api_key)):
|
||||
"""Update an existing memory with new content.
|
||||
|
||||
Args:
|
||||
@@ -168,7 +212,7 @@ def update_memory(memory_id: str, updated_memory: Dict[str, Any]):
|
||||
|
||||
|
||||
@app.get("/memories/{memory_id}/history", summary="Get memory history")
|
||||
def memory_history(memory_id: str):
|
||||
def memory_history(memory_id: str, _api_key: Optional[str] = Depends(verify_api_key)):
|
||||
"""Retrieve memory history."""
|
||||
try:
|
||||
return MEMORY_INSTANCE.history(memory_id=memory_id)
|
||||
@@ -178,7 +222,7 @@ def memory_history(memory_id: str):
|
||||
|
||||
|
||||
@app.delete("/memories/{memory_id}", summary="Delete a memory")
|
||||
def delete_memory(memory_id: str):
|
||||
def delete_memory(memory_id: str, _api_key: Optional[str] = Depends(verify_api_key)):
|
||||
"""Delete a specific memory by ID."""
|
||||
try:
|
||||
MEMORY_INSTANCE.delete(memory_id=memory_id)
|
||||
@@ -193,6 +237,7 @@ def delete_all_memories(
|
||||
user_id: Optional[str] = None,
|
||||
run_id: Optional[str] = None,
|
||||
agent_id: Optional[str] = None,
|
||||
_api_key: Optional[str] = Depends(verify_api_key),
|
||||
):
|
||||
"""Delete all memories for a given identifier."""
|
||||
if not any([user_id, run_id, agent_id]):
|
||||
@@ -209,7 +254,7 @@ def delete_all_memories(
|
||||
|
||||
|
||||
@app.post("/reset", summary="Reset all memories")
|
||||
def reset_memory():
|
||||
def reset_memory(_api_key: Optional[str] = Depends(verify_api_key)):
|
||||
"""Completely reset stored memories."""
|
||||
try:
|
||||
MEMORY_INSTANCE.reset()
|
||||
|
||||
@@ -0,0 +1,496 @@
|
||||
"""Comprehensive E2E tests for REST API server authentication.
|
||||
|
||||
Tests the actual server/main.py app through FastAPI's TestClient (full ASGI
|
||||
round-trip) covering:
|
||||
- Auth disabled mode (ADMIN_API_KEY unset)
|
||||
- Auth enabled mode (ADMIN_API_KEY set)
|
||||
- Edge cases: empty keys, near-miss keys, timing-safe comparison, header
|
||||
casing, response headers, startup logging, and full CRUD flows through auth.
|
||||
"""
|
||||
|
||||
import importlib
|
||||
import logging
|
||||
import os
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
pytest.importorskip("fastapi", reason="fastapi not installed")
|
||||
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Fixtures
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.fixture
|
||||
def _mock_memory():
|
||||
"""Patch Memory.from_config so the server imports without a real backend."""
|
||||
mock_instance = MagicMock()
|
||||
# Set up return values so CRUD endpoints return realistic responses
|
||||
mock_instance.get.return_value = {"id": "mem-1", "memory": "test memory", "user_id": "alice"}
|
||||
mock_instance.get_all.return_value = [
|
||||
{"id": "mem-1", "memory": "test memory", "user_id": "alice"},
|
||||
]
|
||||
mock_instance.add.return_value = {"results": [{"id": "mem-1", "event": "ADD", "memory": "test"}]}
|
||||
mock_instance.search.return_value = [{"id": "mem-1", "memory": "test", "score": 0.9}]
|
||||
mock_instance.update.return_value = {"message": "Memory updated"}
|
||||
mock_instance.history.return_value = [{"id": "mem-1", "old_memory": "a", "new_memory": "b"}]
|
||||
mock_instance.delete.return_value = None
|
||||
mock_instance.delete_all.return_value = {"message": "Memories deleted successfully!"}
|
||||
mock_instance.reset.return_value = None
|
||||
|
||||
with patch.dict(os.environ, {"OPENAI_API_KEY": "fake-key"}):
|
||||
with patch("mem0.Memory.from_config", return_value=mock_instance):
|
||||
yield mock_instance
|
||||
|
||||
|
||||
def _load_app(env_overrides: dict):
|
||||
"""Reload server/main.py with the given environment and return the FastAPI app."""
|
||||
import server.main as server_main
|
||||
|
||||
with patch.dict(os.environ, env_overrides, clear=False):
|
||||
importlib.reload(server_main)
|
||||
return server_main.app
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Auth disabled (ADMIN_API_KEY not set)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestAuthDisabled:
|
||||
"""All endpoints should be freely accessible when ADMIN_API_KEY is empty."""
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _setup(self, _mock_memory):
|
||||
self.app = _load_app({"ADMIN_API_KEY": ""})
|
||||
self.client = TestClient(self.app)
|
||||
self.mock = _mock_memory
|
||||
|
||||
def test_root_redirects_to_docs(self):
|
||||
resp = self.client.get("/", follow_redirects=False)
|
||||
assert resp.status_code == 307
|
||||
assert "/docs" in resp.headers["location"]
|
||||
|
||||
def test_get_memory_without_key(self):
|
||||
resp = self.client.get("/memories/mem-1")
|
||||
assert resp.status_code == 200
|
||||
assert resp.json()["id"] == "mem-1"
|
||||
|
||||
def test_get_all_memories_without_key(self):
|
||||
resp = self.client.get("/memories", params={"user_id": "alice"})
|
||||
assert resp.status_code == 200
|
||||
|
||||
def test_create_memory_without_key(self):
|
||||
resp = self.client.post("/memories", json={
|
||||
"messages": [{"role": "user", "content": "I like pizza"}],
|
||||
"user_id": "alice",
|
||||
})
|
||||
assert resp.status_code == 200
|
||||
|
||||
def test_search_without_key(self):
|
||||
resp = self.client.post("/search", json={"query": "pizza", "user_id": "alice"})
|
||||
assert resp.status_code == 200
|
||||
|
||||
def test_update_memory_without_key(self):
|
||||
resp = self.client.put("/memories/mem-1", json={"data": "updated"})
|
||||
assert resp.status_code == 200
|
||||
|
||||
def test_history_without_key(self):
|
||||
resp = self.client.get("/memories/mem-1/history")
|
||||
assert resp.status_code == 200
|
||||
|
||||
def test_delete_memory_without_key(self):
|
||||
resp = self.client.delete("/memories/mem-1")
|
||||
assert resp.status_code == 200
|
||||
|
||||
def test_delete_all_without_key(self):
|
||||
resp = self.client.delete("/memories", params={"user_id": "alice"})
|
||||
assert resp.status_code == 200
|
||||
|
||||
def test_reset_without_key(self):
|
||||
resp = self.client.post("/reset")
|
||||
assert resp.status_code == 200
|
||||
|
||||
def test_configure_without_key(self):
|
||||
self.mock.from_config = MagicMock()
|
||||
resp = self.client.post("/configure", json={"version": "v1.1"})
|
||||
assert resp.status_code == 200
|
||||
|
||||
def test_supplying_key_still_works_when_auth_disabled(self):
|
||||
"""A client that sends X-API-Key should not be penalized when auth is off."""
|
||||
resp = self.client.get(
|
||||
"/memories/mem-1", headers={"X-API-Key": "some-random-key"}
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"method,path",
|
||||
[
|
||||
("POST", "/configure"),
|
||||
("POST", "/memories"),
|
||||
("GET", "/memories"),
|
||||
("GET", "/memories/test-id"),
|
||||
("POST", "/search"),
|
||||
("PUT", "/memories/test-id"),
|
||||
("GET", "/memories/test-id/history"),
|
||||
("DELETE", "/memories/test-id"),
|
||||
("DELETE", "/memories"),
|
||||
("POST", "/reset"),
|
||||
],
|
||||
)
|
||||
def test_no_endpoint_returns_401_when_auth_disabled(self, method, path):
|
||||
resp = self.client.request(method, path)
|
||||
assert resp.status_code != 401, f"{method} {path} should not require auth"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Auth enabled (ADMIN_API_KEY set)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestAuthEnabled:
|
||||
"""All protected endpoints must enforce the API key."""
|
||||
|
||||
API_KEY = "test-secret-key-12345"
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _setup(self, _mock_memory):
|
||||
self.app = _load_app({"ADMIN_API_KEY": self.API_KEY})
|
||||
self.client = TestClient(self.app)
|
||||
self.mock = _mock_memory
|
||||
|
||||
# --- Rejection cases ---
|
||||
|
||||
def test_missing_key_returns_401(self):
|
||||
resp = self.client.get("/memories/mem-1")
|
||||
assert resp.status_code == 401
|
||||
|
||||
def test_missing_key_detail_mentions_header(self):
|
||||
resp = self.client.get("/memories/mem-1")
|
||||
assert "X-API-Key" in resp.json()["detail"]
|
||||
|
||||
def test_wrong_key_returns_401(self):
|
||||
resp = self.client.get("/memories/mem-1", headers={"X-API-Key": "wrong"})
|
||||
assert resp.status_code == 401
|
||||
|
||||
def test_wrong_key_detail_says_invalid(self):
|
||||
resp = self.client.get("/memories/mem-1", headers={"X-API-Key": "wrong"})
|
||||
assert "Invalid" in resp.json()["detail"]
|
||||
|
||||
def test_empty_string_key_returns_401(self):
|
||||
resp = self.client.get("/memories/mem-1", headers={"X-API-Key": ""})
|
||||
assert resp.status_code == 401
|
||||
|
||||
def test_401_includes_www_authenticate_header(self):
|
||||
resp = self.client.get("/memories/mem-1")
|
||||
assert resp.headers.get("www-authenticate") == "ApiKey"
|
||||
|
||||
def test_near_miss_key_rejected(self):
|
||||
"""Key that differs by one character should be rejected."""
|
||||
near_miss = self.API_KEY[:-1] + ("6" if self.API_KEY[-1] != "6" else "7")
|
||||
resp = self.client.get("/memories/mem-1", headers={"X-API-Key": near_miss})
|
||||
assert resp.status_code == 401
|
||||
|
||||
def test_key_with_extra_whitespace_rejected(self):
|
||||
resp = self.client.get("/memories/mem-1", headers={"X-API-Key": f" {self.API_KEY} "})
|
||||
assert resp.status_code == 401
|
||||
|
||||
def test_key_prefix_rejected(self):
|
||||
resp = self.client.get("/memories/mem-1", headers={"X-API-Key": self.API_KEY[:5]})
|
||||
assert resp.status_code == 401
|
||||
|
||||
def test_key_with_different_case_rejected(self):
|
||||
resp = self.client.get("/memories/mem-1", headers={"X-API-Key": self.API_KEY.upper()})
|
||||
assert resp.status_code == 401
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"method,path",
|
||||
[
|
||||
("POST", "/configure"),
|
||||
("POST", "/memories"),
|
||||
("GET", "/memories"),
|
||||
("GET", "/memories/test-id"),
|
||||
("POST", "/search"),
|
||||
("PUT", "/memories/test-id"),
|
||||
("GET", "/memories/test-id/history"),
|
||||
("DELETE", "/memories/test-id"),
|
||||
("DELETE", "/memories"),
|
||||
("POST", "/reset"),
|
||||
],
|
||||
)
|
||||
def test_all_endpoints_reject_without_key(self, method, path):
|
||||
resp = self.client.request(method, path)
|
||||
assert resp.status_code == 401, f"{method} {path} should require auth"
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"method,path",
|
||||
[
|
||||
("POST", "/configure"),
|
||||
("POST", "/memories"),
|
||||
("GET", "/memories"),
|
||||
("GET", "/memories/test-id"),
|
||||
("POST", "/search"),
|
||||
("PUT", "/memories/test-id"),
|
||||
("GET", "/memories/test-id/history"),
|
||||
("DELETE", "/memories/test-id"),
|
||||
("DELETE", "/memories"),
|
||||
("POST", "/reset"),
|
||||
],
|
||||
)
|
||||
def test_all_endpoints_reject_wrong_key(self, method, path):
|
||||
resp = self.client.request(method, path, headers={"X-API-Key": "wrong-key"})
|
||||
assert resp.status_code == 401, f"{method} {path} should reject wrong key"
|
||||
|
||||
# --- Acceptance cases ---
|
||||
|
||||
def test_root_does_not_require_key(self):
|
||||
resp = self.client.get("/", follow_redirects=False)
|
||||
assert resp.status_code == 307
|
||||
|
||||
def _authed(self, method, path, **kwargs):
|
||||
headers = kwargs.pop("headers", {})
|
||||
headers["X-API-Key"] = self.API_KEY
|
||||
return self.client.request(method, path, headers=headers, **kwargs)
|
||||
|
||||
def test_get_memory_with_key(self):
|
||||
resp = self._authed("GET", "/memories/mem-1")
|
||||
assert resp.status_code == 200
|
||||
assert resp.json()["id"] == "mem-1"
|
||||
|
||||
def test_get_all_memories_with_key(self):
|
||||
resp = self._authed("GET", "/memories", params={"user_id": "alice"})
|
||||
assert resp.status_code == 200
|
||||
|
||||
def test_create_memory_with_key(self):
|
||||
resp = self._authed("POST", "/memories", json={
|
||||
"messages": [{"role": "user", "content": "I like pizza"}],
|
||||
"user_id": "alice",
|
||||
})
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
assert "results" in data
|
||||
|
||||
def test_search_with_key(self):
|
||||
resp = self._authed("POST", "/search", json={"query": "pizza", "user_id": "alice"})
|
||||
assert resp.status_code == 200
|
||||
|
||||
def test_update_memory_with_key(self):
|
||||
resp = self._authed("PUT", "/memories/mem-1", json={"data": "updated"})
|
||||
assert resp.status_code == 200
|
||||
|
||||
def test_history_with_key(self):
|
||||
resp = self._authed("GET", "/memories/mem-1/history")
|
||||
assert resp.status_code == 200
|
||||
|
||||
def test_delete_memory_with_key(self):
|
||||
resp = self._authed("DELETE", "/memories/mem-1")
|
||||
assert resp.status_code == 200
|
||||
|
||||
def test_delete_all_with_key(self):
|
||||
resp = self._authed("DELETE", "/memories", params={"user_id": "alice"})
|
||||
assert resp.status_code == 200
|
||||
|
||||
def test_reset_with_key(self):
|
||||
resp = self._authed("POST", "/reset")
|
||||
assert resp.status_code == 200
|
||||
|
||||
def test_configure_with_key(self):
|
||||
resp = self._authed("POST", "/configure", json={"version": "v1.1"})
|
||||
assert resp.status_code == 200
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Full CRUD flow through auth
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestAuthenticatedCRUDFlow:
|
||||
"""Verify a complete create → read → search → update → history → delete
|
||||
cycle works end-to-end through the auth layer."""
|
||||
|
||||
API_KEY = "flow-test-key-99"
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _setup(self, _mock_memory):
|
||||
self.app = _load_app({"ADMIN_API_KEY": self.API_KEY})
|
||||
self.client = TestClient(self.app)
|
||||
self.mock = _mock_memory
|
||||
|
||||
def _authed(self, method, path, **kwargs):
|
||||
headers = kwargs.pop("headers", {})
|
||||
headers["X-API-Key"] = self.API_KEY
|
||||
return self.client.request(method, path, headers=headers, **kwargs)
|
||||
|
||||
def test_full_crud_cycle(self):
|
||||
# 1. Create
|
||||
resp = self._authed("POST", "/memories", json={
|
||||
"messages": [{"role": "user", "content": "I love fresh vegetable pizza"}],
|
||||
"user_id": "alice",
|
||||
})
|
||||
assert resp.status_code == 200
|
||||
self.mock.add.assert_called_once()
|
||||
|
||||
# 2. Read single
|
||||
resp = self._authed("GET", "/memories/mem-1")
|
||||
assert resp.status_code == 200
|
||||
self.mock.get.assert_called_once_with("mem-1")
|
||||
|
||||
# 3. Read all
|
||||
resp = self._authed("GET", "/memories", params={"user_id": "alice"})
|
||||
assert resp.status_code == 200
|
||||
self.mock.get_all.assert_called_once_with(user_id="alice")
|
||||
|
||||
# 4. Search
|
||||
resp = self._authed("POST", "/search", json={"query": "pizza", "user_id": "alice"})
|
||||
assert resp.status_code == 200
|
||||
self.mock.search.assert_called_once()
|
||||
|
||||
# 5. Update
|
||||
resp = self._authed("PUT", "/memories/mem-1", json={"data": "updated content"})
|
||||
assert resp.status_code == 200
|
||||
self.mock.update.assert_called_once()
|
||||
|
||||
# 6. History
|
||||
resp = self._authed("GET", "/memories/mem-1/history")
|
||||
assert resp.status_code == 200
|
||||
self.mock.history.assert_called_once_with(memory_id="mem-1")
|
||||
|
||||
# 7. Delete single
|
||||
resp = self._authed("DELETE", "/memories/mem-1")
|
||||
assert resp.status_code == 200
|
||||
self.mock.delete.assert_called_once_with(memory_id="mem-1")
|
||||
|
||||
# 8. Delete all
|
||||
resp = self._authed("DELETE", "/memories", params={"user_id": "alice"})
|
||||
assert resp.status_code == 200
|
||||
self.mock.delete_all.assert_called_once()
|
||||
|
||||
def test_crud_flow_blocked_without_auth(self):
|
||||
"""Same flow should fail at every step without the key."""
|
||||
endpoints = [
|
||||
("POST", "/memories", {"json": {
|
||||
"messages": [{"role": "user", "content": "test"}], "user_id": "alice"
|
||||
}}),
|
||||
("GET", "/memories/mem-1", {}),
|
||||
("GET", "/memories", {"params": {"user_id": "alice"}}),
|
||||
("POST", "/search", {"json": {"query": "pizza", "user_id": "alice"}}),
|
||||
("PUT", "/memories/mem-1", {"json": {"data": "x"}}),
|
||||
("GET", "/memories/mem-1/history", {}),
|
||||
("DELETE", "/memories/mem-1", {}),
|
||||
("DELETE", "/memories", {"params": {"user_id": "alice"}}),
|
||||
("POST", "/reset", {}),
|
||||
]
|
||||
for method, path, kwargs in endpoints:
|
||||
resp = self.client.request(method, path, **kwargs)
|
||||
assert resp.status_code == 401, f"Unauthenticated {method} {path} should be 401"
|
||||
# Verify the mock was NOT called (auth blocked before reaching handler)
|
||||
self.mock.add.assert_not_called()
|
||||
self.mock.get.assert_not_called()
|
||||
self.mock.search.assert_not_called()
|
||||
self.mock.update.assert_not_called()
|
||||
self.mock.history.assert_not_called()
|
||||
self.mock.delete.assert_not_called()
|
||||
self.mock.delete_all.assert_not_called()
|
||||
self.mock.reset.assert_not_called()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Edge cases
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestAuthEdgeCases:
|
||||
"""Boundary conditions and unusual inputs."""
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _setup(self, _mock_memory):
|
||||
self.mock = _mock_memory
|
||||
|
||||
def test_very_long_api_key(self):
|
||||
"""Server should handle a very long key without crashing."""
|
||||
long_key = "k" * 4096
|
||||
app = _load_app({"ADMIN_API_KEY": long_key})
|
||||
client = TestClient(app)
|
||||
resp = client.get("/memories/mem-1", headers={"X-API-Key": long_key})
|
||||
assert resp.status_code == 200
|
||||
|
||||
def test_special_characters_in_api_key(self):
|
||||
"""Keys with special ASCII characters should work."""
|
||||
special_key = "sk-!@#$%^&*()_+-=[]{}|;:',.<>?/~`"
|
||||
app = _load_app({"ADMIN_API_KEY": special_key})
|
||||
client = TestClient(app)
|
||||
|
||||
resp = client.get("/memories/mem-1", headers={"X-API-Key": special_key})
|
||||
assert resp.status_code == 200
|
||||
|
||||
resp = client.get("/memories/mem-1", headers={"X-API-Key": "wrong"})
|
||||
assert resp.status_code == 401
|
||||
|
||||
def test_key_env_var_not_present_at_all(self):
|
||||
"""When the env var is completely absent, auth should be disabled."""
|
||||
import server.main as server_main
|
||||
env = os.environ.copy()
|
||||
env.pop("ADMIN_API_KEY", None)
|
||||
with patch.dict(os.environ, env, clear=True):
|
||||
importlib.reload(server_main)
|
||||
client = TestClient(server_main.app)
|
||||
resp = client.get("/memories/mem-1")
|
||||
assert resp.status_code != 401
|
||||
|
||||
def test_switching_from_enabled_to_disabled(self):
|
||||
"""Simulates a server restart with auth toggled off."""
|
||||
# First: auth enabled
|
||||
app1 = _load_app({"ADMIN_API_KEY": "secret"})
|
||||
c1 = TestClient(app1)
|
||||
assert c1.get("/memories/mem-1").status_code == 401
|
||||
|
||||
# Then: auth disabled
|
||||
app2 = _load_app({"ADMIN_API_KEY": ""})
|
||||
c2 = TestClient(app2)
|
||||
assert c2.get("/memories/mem-1").status_code != 401
|
||||
|
||||
def test_openapi_schema_accessible_without_key(self):
|
||||
"""The /docs and /openapi.json endpoints should always be reachable."""
|
||||
app = _load_app({"ADMIN_API_KEY": "secret"})
|
||||
client = TestClient(app)
|
||||
|
||||
resp = client.get("/openapi.json")
|
||||
assert resp.status_code == 200
|
||||
schema = resp.json()
|
||||
assert "paths" in schema
|
||||
|
||||
resp = client.get("/docs")
|
||||
assert resp.status_code == 200
|
||||
|
||||
def test_openapi_schema_documents_auth(self):
|
||||
"""The OpenAPI schema should mention authentication."""
|
||||
app = _load_app({"ADMIN_API_KEY": "secret"})
|
||||
client = TestClient(app)
|
||||
schema = client.get("/openapi.json").json()
|
||||
assert "Authentication" in schema.get("info", {}).get("description", "")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Startup logging
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestStartupLogging:
|
||||
"""Verify the server emits the correct log messages at import time."""
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _setup(self, _mock_memory):
|
||||
pass
|
||||
|
||||
def test_warning_when_auth_disabled(self, caplog):
|
||||
with caplog.at_level(logging.WARNING):
|
||||
_load_app({"ADMIN_API_KEY": ""})
|
||||
assert any("UNSECURED" in r.message for r in caplog.records)
|
||||
|
||||
def test_info_when_auth_enabled(self, caplog):
|
||||
with caplog.at_level(logging.INFO):
|
||||
_load_app({"ADMIN_API_KEY": "a-long-enough-secret-key"})
|
||||
assert any("authentication enabled" in r.message for r in caplog.records)
|
||||
|
||||
def test_warning_when_key_too_short(self, caplog):
|
||||
with caplog.at_level(logging.WARNING):
|
||||
_load_app({"ADMIN_API_KEY": "short"})
|
||||
assert any("shorter than" in r.message for r in caplog.records)
|
||||
Reference in New Issue
Block a user