Self-hosted dashboard and admin auth (#4837)
Co-authored-by: kartik-mem0 <kartik.labhshetwar@mem0.ai>
This commit is contained in:
+326
-106
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user