diff --git a/docs/open-source/features/rest-api.mdx b/docs/open-source/features/rest-api.mdx index 3a77cb839..2d6f38930 100644 --- a/docs/open-source/features/rest-api.mdx +++ b/docs/open-source/features/rest-api.mdx @@ -14,7 +14,7 @@ The Mem0 REST API server exposes every OSS memory operation over HTTP. Run it al - 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. --- @@ -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: ` | + +### 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" + }' +``` + + + The server logs a warning at startup when `ADMIN_API_KEY` is not set. Always set it in production. + + +--- + ## 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. diff --git a/server/.env.example b/server/.env.example index 30c0a48a6..820221622 100644 --- a/server/.env.example +++ b/server/.env.example @@ -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= diff --git a/server/main.py b/server/main.py index 85c7cc7ea..f3667b7fe 100644 --- a/server/main.py +++ b/server/main.py @@ -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() diff --git a/tests/test_server_auth.py b/tests/test_server_auth.py new file mode 100644 index 000000000..d72a6789a --- /dev/null +++ b/tests/test_server_auth.py @@ -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)