Self-hosted dashboard and admin auth (#4837)

Co-authored-by: kartik-mem0 <kartik.labhshetwar@mem0.ai>
This commit is contained in:
Gabriel Stein
2026-04-23 06:36:36 -07:00
committed by GitHub
parent 15feaa8ac4
commit db8ac61713
252 changed files with 20275 additions and 331 deletions
+326 -106
View File
@@ -1,37 +1,102 @@
import asyncio
import logging
import os
import secrets
import time
from typing import Any, Dict, List, Optional
from dotenv import load_dotenv
from fastapi import Depends, FastAPI, HTTPException
from fastapi import Depends, FastAPI, HTTPException, Request
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import JSONResponse, RedirectResponse
from fastapi.security import APIKeyHeader
from pydantic import BaseModel, Field
from slowapi import _rate_limit_exceeded_handler
from slowapi.errors import RateLimitExceeded
from sqlalchemy import func, select
from mem0 import Memory
from auth import ADMIN_API_KEY, AUTH_DISABLED, JWT_SECRET, verify_auth
from errors import (
UpstreamError,
install_request_id_logging,
new_request_id,
request_id_var,
upstream_error,
upstream_error_handler,
)
from rate_limit import limiter
from db import SessionLocal
from models import RequestLog, User
import telemetry
from routers import auth as auth_router
from routers import api_keys as api_keys_router
from routers import entities as entities_router
from routers import requests as requests_router
from schemas import MessageResponse
from server_state import get_current_config, get_memory_instance, initialize_state, set_session_factory, update_config
logging.basicConfig(level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s")
# Load environment variables
load_dotenv()
ADMIN_API_KEY = os.environ.get("ADMIN_API_KEY", "")
install_request_id_logging()
logging.basicConfig(level=logging.INFO, format="%(asctime)s - %(levelname)s - [%(request_id)s] %(message)s")
MIN_KEY_LENGTH = 16
SENSITIVE_CONFIG_KEYS = {
"admin_api_key",
"api_key",
"authorization",
"jwt_secret",
"password",
"password_hash",
"secret",
"token",
}
SKIPPED_REQUEST_LOG_PATHS = {"/api/health", "/docs", "/redoc", "/openapi.json"}
SKIPPED_REQUEST_LOG_PREFIXES = ("/requests",)
BUNDLED_LLM_PROVIDERS = ("openai", "anthropic", "gemini")
BUNDLED_EMBEDDER_PROVIDERS = ("openai", "gemini")
def _warn_if_unconfigured() -> None:
"""Pre-auth deployments upgrading into this build will 401 everywhere until
an admin key or admin user exists. Surface the fix before the support tickets."""
try:
with SessionLocal() as session:
if session.scalar(select(func.count(User.id))) > 0:
return
except Exception:
return
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."
"\n%s\n"
" Auth is enabled by default and this server has no admin configured.\n"
" Protected endpoints will return 401 until you either:\n"
" 1. Set ADMIN_API_KEY=<long-random-value> (fastest, no client changes)\n"
" 2. Register an admin at http://<host>:3000/setup\n"
" 3. Set AUTH_DISABLED=true (local development only)\n"
" Docs: https://docs.mem0.ai/open-source/features/rest-api#authentication\n"
"%s",
"=" * 72,
"=" * 72,
)
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")
if not AUTH_DISABLED and not JWT_SECRET:
raise RuntimeError(
"JWT_SECRET is required. Set it in .env (generate with `openssl rand -base64 48`) "
"or set AUTH_DISABLED=true for local development only."
)
if AUTH_DISABLED:
logging.warning("AUTH_DISABLED is enabled. Protected endpoints are open for local development only.")
elif ADMIN_API_KEY and 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,
)
elif not ADMIN_API_KEY:
_warn_if_unconfigured()
telemetry.log_status()
POSTGRES_HOST = os.environ.get("POSTGRES_HOST", "postgres")
POSTGRES_PORT = os.environ.get("POSTGRES_PORT", "5432")
@@ -42,6 +107,8 @@ POSTGRES_COLLECTION_NAME = os.environ.get("POSTGRES_COLLECTION_NAME", "memories"
OPENAI_API_KEY = os.environ.get("OPENAI_API_KEY")
HISTORY_DB_PATH = os.environ.get("HISTORY_DB_PATH", "/app/history/history.db")
DEFAULT_LLM_MODEL = os.environ.get("MEM0_DEFAULT_LLM_MODEL", "gpt-4.1-nano-2025-04-14")
DEFAULT_EMBEDDER_MODEL = os.environ.get("MEM0_DEFAULT_EMBEDDER_MODEL", "text-embedding-3-small")
DEFAULT_CONFIG = {
"version": "v1.1",
@@ -56,44 +123,46 @@ DEFAULT_CONFIG = {
"collection_name": POSTGRES_COLLECTION_NAME,
},
},
"llm": {"provider": "openai", "config": {"api_key": OPENAI_API_KEY, "temperature": 0.2, "model": "gpt-4.1-nano-2025-04-14"}},
"embedder": {"provider": "openai", "config": {"api_key": OPENAI_API_KEY, "model": "text-embedding-3-small"}},
"llm": {
"provider": "openai",
"config": {"api_key": OPENAI_API_KEY, "temperature": 0.2, "model": DEFAULT_LLM_MODEL},
},
"embedder": {"provider": "openai", "config": {"api_key": OPENAI_API_KEY, "model": DEFAULT_EMBEDDER_MODEL}},
"history_db_path": HISTORY_DB_PATH,
}
MEMORY_INSTANCE = Memory.from_config(DEFAULT_CONFIG)
set_session_factory(SessionLocal)
initialize_state(DEFAULT_CONFIG)
app = FastAPI(
title="Mem0 REST APIs",
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."
"Supports Bearer JWT tokens, per-user API keys via `X-API-Key` header, "
"or the legacy `ADMIN_API_KEY` environment variable. Set `AUTH_DISABLED=true` for local development only."
),
version="1.0.0",
redirect_slashes=False,
)
app.state.limiter = limiter
app.add_exception_handler(RateLimitExceeded, _rate_limit_exceeded_handler)
app.add_exception_handler(UpstreamError, upstream_error_handler)
DASHBOARD_URL = os.environ.get("DASHBOARD_URL", "http://localhost:3000")
app.add_middleware(
CORSMiddleware,
allow_origins=[DASHBOARD_URL],
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
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
app.include_router(auth_router.router)
app.include_router(api_keys_router.router)
app.include_router(entities_router.router)
app.include_router(requests_router.router)
class Message(BaseModel):
@@ -127,27 +196,192 @@ class SearchRequest(BaseModel):
threshold: Optional[float] = Field(None, description="Minimum similarity score for results.")
class GenerateInstructionsRequest(BaseModel):
use_case: str = Field(..., description="Description of what the user will use Mem0 for.")
def _redact_config(value: Any, key: str | None = None) -> Any:
if isinstance(value, dict):
return {item_key: _redact_config(item_value, item_key) for item_key, item_value in value.items()}
if isinstance(value, list):
return [_redact_config(item_value, key) for item_value in value]
if key is not None and key.lower() in SENSITIVE_CONFIG_KEYS:
return "[redacted]" if value else value
return value
def _validate_bundled_providers(config: Dict[str, Any]) -> None:
llm = config.get("llm")
if isinstance(llm, dict) and (provider := llm.get("provider")) and provider not in BUNDLED_LLM_PROVIDERS:
raise HTTPException(
status_code=400,
detail=(
f"LLM provider '{provider}' is not bundled in this image. "
f"Bundled providers: {', '.join(BUNDLED_LLM_PROVIDERS)}. "
"To use another provider, install its Python package, rebuild the container, "
"and extend BUNDLED_LLM_PROVIDERS in server/main.py."
),
)
embedder = config.get("embedder")
if (
isinstance(embedder, dict)
and (provider := embedder.get("provider"))
and provider not in BUNDLED_EMBEDDER_PROVIDERS
):
raise HTTPException(
status_code=400,
detail=(
f"Embedder provider '{provider}' is not bundled in this image. "
f"Bundled providers: {', '.join(BUNDLED_EMBEDDER_PROVIDERS)}. "
"To use another provider, install its Python package, rebuild the container, "
"and extend BUNDLED_EMBEDDER_PROVIDERS in server/main.py."
),
)
def _should_log_request(request: Request) -> bool:
if request.method == "OPTIONS":
return False
path = request.url.path
if path in SKIPPED_REQUEST_LOG_PATHS:
return False
return not path.startswith(SKIPPED_REQUEST_LOG_PREFIXES)
def _persist_request_log(method: str, path: str, status_code: int, latency_ms: float, auth_type: str) -> None:
session = SessionLocal()
try:
session.add(
RequestLog(
method=method,
path=path,
status_code=status_code,
latency_ms=latency_ms,
auth_type=auth_type,
)
)
session.commit()
except Exception:
session.rollback()
logging.exception("Failed to persist request log")
finally:
session.close()
@app.middleware("http")
async def log_requests(request: Request, call_next):
request.state.auth_type = getattr(request.state, "auth_type", "none")
rid = new_request_id()
token = request_id_var.set(rid)
start = time.perf_counter()
status_code = 500
try:
response = await call_next(request)
status_code = response.status_code
response.headers["X-Request-ID"] = rid
return response
except Exception:
status_code = 500
raise
finally:
request_id_var.reset(token)
if _should_log_request(request):
asyncio.get_running_loop().run_in_executor(
None,
_persist_request_log,
request.method,
request.url.path,
status_code,
round((time.perf_counter() - start) * 1000, 2),
getattr(request.state, "auth_type", "none"),
)
@app.get("/configure", summary="Get current Mem0 configuration")
def get_config(_auth=Depends(verify_auth)):
return _redact_config(get_current_config())
@app.get("/configure/providers", summary="List bundled LLM and embedder providers")
def list_bundled_providers(_auth=Depends(verify_auth)):
return {"llm": list(BUNDLED_LLM_PROVIDERS), "embedder": list(BUNDLED_EMBEDDER_PROVIDERS)}
@app.post("/configure", summary="Configure Mem0")
def set_config(config: Dict[str, Any], _api_key: Optional[str] = Depends(verify_api_key)):
def set_config(config: Dict[str, Any], _auth=Depends(verify_auth)):
"""Set memory configuration."""
global MEMORY_INSTANCE
MEMORY_INSTANCE = Memory.from_config(config)
_validate_bundled_providers(config)
update_config(config)
return {"message": "Configuration set successfully"}
@app.post("/generate-instructions", summary="Generate custom instructions from a use case")
def generate_instructions(req: GenerateInstructionsRequest, _auth=Depends(verify_auth)):
"""Generate custom instructions and a contextual test message tailored to a use case."""
try:
llm = get_memory_instance().llm
prompt = (
"You are configuring a memory system. Given the use case below, produce two things:\n"
"1. INSTRUCTIONS: A short paragraph of custom instructions telling the memory extraction system "
"what kinds of facts, preferences, and context to prioritize. Be specific to the use case.\n"
"2. TEST_MESSAGE: A single realistic sentence a user in this use case would say, suitable for "
"testing that the memory system works.\n\n"
"Respond in exactly this format (no markdown, no extra text):\n"
"INSTRUCTIONS: <your instructions>\n"
f"TEST_MESSAGE: <your test message>\n\nUse case: {req.use_case}"
)
response = llm.generate_response([{"role": "user", "content": prompt}])
instructions = response
test_message = "I like to hike on weekends."
if "INSTRUCTIONS:" in response and "TEST_MESSAGE:" in response:
parts = response.split("TEST_MESSAGE:")
instructions = parts[0].replace("INSTRUCTIONS:", "").strip()
test_message = parts[1].strip()
return {"custom_instructions": instructions, "test_message": test_message}
except Exception:
raise upstream_error()
@app.post("/memories", summary="Create memories")
def add_memory(memory_create: MemoryCreate, _api_key: Optional[str] = Depends(verify_api_key)):
def add_memory(memory_create: MemoryCreate, _auth=Depends(verify_auth)):
"""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.")
params = {k: v for k, v in memory_create.model_dump().items() if v is not None and k != "messages"}
try:
response = MEMORY_INSTANCE.add(messages=[m.model_dump() for m in memory_create.messages], **params)
response = get_memory_instance().add(messages=[m.model_dump() for m in memory_create.messages], **params)
return JSONResponse(content=response)
except Exception as e:
logging.exception("Error in add_memory:") # This will log the full traceback
raise HTTPException(status_code=500, detail=str(e))
except Exception:
raise upstream_error()
ALL_MEMORIES_LIMIT = 1000
_RESERVED_PAYLOAD_KEYS = {"data", "user_id", "agent_id", "run_id", "hash", "created_at", "updated_at"}
def _serialize_memory(row: Any) -> Dict[str, Any]:
payload = getattr(row, "payload", None) or {}
return {
"id": getattr(row, "id", None),
"memory": payload.get("data"),
"user_id": payload.get("user_id"),
"agent_id": payload.get("agent_id"),
"run_id": payload.get("run_id"),
"hash": payload.get("hash"),
"metadata": {k: v for k, v in payload.items() if k not in _RESERVED_PAYLOAD_KEYS},
"created_at": payload.get("created_at"),
"updated_at": payload.get("updated_at"),
}
def _list_all_memories(limit: int = ALL_MEMORIES_LIMIT) -> Dict[str, Any]:
results = get_memory_instance().vector_store.list(top_k=limit)
rows = results[0] if results and isinstance(results, list) and isinstance(results[0], list) else results or []
return {"results": [_serialize_memory(row) for row in rows]}
@app.get("/memories", summary="Get memories")
@@ -155,87 +389,75 @@ 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),
_auth=Depends(verify_auth),
):
"""Retrieve stored memories."""
if not any([user_id, run_id, agent_id]):
raise HTTPException(status_code=400, detail="At least one identifier is required.")
"""Retrieve stored memories. Lists all memories when no identifier is provided."""
try:
if not any([user_id, run_id, agent_id]):
return _list_all_memories()
params = {
k: v for k, v in {"user_id": user_id, "run_id": run_id, "agent_id": agent_id}.items() if v is not None
}
return MEMORY_INSTANCE.get_all(**params)
except Exception as e:
logging.exception("Error in get_all_memories:")
raise HTTPException(status_code=500, detail=str(e))
return get_memory_instance().get_all(**params)
except Exception:
raise upstream_error()
@app.get("/memories/{memory_id}", summary="Get a memory")
def get_memory(memory_id: str, _api_key: Optional[str] = Depends(verify_api_key)):
def get_memory(memory_id: str, _auth=Depends(verify_auth)):
"""Retrieve a specific memory by ID."""
try:
return MEMORY_INSTANCE.get(memory_id)
except Exception as e:
logging.exception("Error in get_memory:")
raise HTTPException(status_code=500, detail=str(e))
return get_memory_instance().get(memory_id)
except Exception:
raise upstream_error()
@app.post("/search", summary="Search memories")
def search_memories(search_req: SearchRequest, _api_key: Optional[str] = Depends(verify_api_key)):
def search_memories(search_req: SearchRequest, _auth=Depends(verify_auth)):
"""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"}
return MEMORY_INSTANCE.search(query=search_req.query, **params)
except Exception as e:
logging.exception("Error in search_memories:")
raise HTTPException(status_code=500, detail=str(e))
return get_memory_instance().search(query=search_req.query, **params)
except Exception:
raise upstream_error()
@app.put("/memories/{memory_id}", summary="Update a memory")
def update_memory(memory_id: str, updated_memory: MemoryUpdate, _api_key: Optional[str] = Depends(verify_api_key)):
"""Update an existing memory with new content.
Args:
memory_id (str): ID of the memory to update
updated_memory (MemoryUpdate): New content and optional metadata to update the memory with
Returns:
dict: Success message indicating the memory was updated
"""
def update_memory(memory_id: str, updated_memory: MemoryUpdate, _auth=Depends(verify_auth)):
"""Update an existing memory."""
try:
return MEMORY_INSTANCE.update(memory_id=memory_id, data=updated_memory.text, metadata=updated_memory.metadata)
except Exception as e:
logging.exception("Error in update_memory:")
raise HTTPException(status_code=500, detail=str(e))
return get_memory_instance().update(
memory_id=memory_id, data=updated_memory.text, metadata=updated_memory.metadata
)
except Exception:
raise upstream_error()
@app.get("/memories/{memory_id}/history", summary="Get memory history")
def memory_history(memory_id: str, _api_key: Optional[str] = Depends(verify_api_key)):
def memory_history(memory_id: str, _auth=Depends(verify_auth)):
"""Retrieve memory history."""
try:
return MEMORY_INSTANCE.history(memory_id=memory_id)
except Exception as e:
logging.exception("Error in memory_history:")
raise HTTPException(status_code=500, detail=str(e))
return get_memory_instance().history(memory_id=memory_id)
except Exception:
raise upstream_error()
@app.delete("/memories/{memory_id}", summary="Delete a memory")
def delete_memory(memory_id: str, _api_key: Optional[str] = Depends(verify_api_key)):
@app.delete("/memories/{memory_id}", summary="Delete a memory", response_model=MessageResponse)
def delete_memory(memory_id: str, _auth=Depends(verify_auth)):
"""Delete a specific memory by ID."""
try:
MEMORY_INSTANCE.delete(memory_id=memory_id)
return {"message": "Memory deleted successfully"}
except Exception as e:
logging.exception("Error in delete_memory:")
raise HTTPException(status_code=500, detail=str(e))
get_memory_instance().delete(memory_id=memory_id)
return MessageResponse(message="Memory deleted successfully")
except Exception:
raise upstream_error()
@app.delete("/memories", summary="Delete all memories")
@app.delete("/memories", summary="Delete all memories", response_model=MessageResponse)
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),
_auth=Depends(verify_auth),
):
"""Delete all memories for a given identifier."""
if not any([user_id, run_id, agent_id]):
@@ -244,22 +466,20 @@ def delete_all_memories(
params = {
k: v for k, v in {"user_id": user_id, "run_id": run_id, "agent_id": agent_id}.items() if v is not None
}
MEMORY_INSTANCE.delete_all(**params)
return {"message": "All relevant memories deleted"}
except Exception as e:
logging.exception("Error in delete_all_memories:")
raise HTTPException(status_code=500, detail=str(e))
get_memory_instance().delete_all(**params)
return MessageResponse(message="All relevant memories deleted")
except Exception:
raise upstream_error()
@app.post("/reset", summary="Reset all memories")
def reset_memory(_api_key: Optional[str] = Depends(verify_api_key)):
def reset_memory(_auth=Depends(verify_auth)):
"""Completely reset stored memories."""
try:
MEMORY_INSTANCE.reset()
get_memory_instance().reset()
return {"message": "All memories reset"}
except Exception as e:
logging.exception("Error in reset_memory:")
raise HTTPException(status_code=500, detail=str(e))
except Exception:
raise upstream_error()
@app.get("/", summary="Redirect to the OpenAPI documentation", include_in_schema=False)