feat: add official mem0 CLI (Python & TypeScript) (#4575)

This commit is contained in:
Saket Aryan
2026-03-28 05:03:01 +05:30
committed by GitHub
parent 88fd0e77d0
commit 3225e30859
58 changed files with 13438 additions and 0 deletions
+30
View File
@@ -0,0 +1,30 @@
.PHONY: install dev lint format test build clean publish publish-test
install:
pip install -e .
dev:
pip install -e ".[dev]"
lint:
ruff check .
ruff format --check .
format:
ruff check --fix .
ruff format .
test:
pytest
build: clean
hatch build
clean:
rm -rf dist/
publish: build
hatch publish
publish-test: build
hatch publish --repo test
+29
View File
@@ -0,0 +1,29 @@
# mem0 CLI
The official command-line interface for [mem0](https://mem0.ai) — the memory layer for AI agents.
## Installation
```bash
pip install mem0-cli
```
## Quick Start
```bash
# Set up your configuration
mem0 init
# Add a memory
mem0 add "I prefer dark mode and use vim keybindings" --user-id alice
# Search memories
mem0 search "What are Alice's preferences?" --user-id alice
# List all memories
mem0 list --user-id alice
```
## License
Apache-2.0
+76
View File
@@ -0,0 +1,76 @@
# Development
## Prerequisites
- Python **3.10+**
## Setup
All commands below should be run from the `python/` directory:
```bash
cd python
```
## Install local (editable) + run
```bash
python3 -m venv .venv
source .venv/bin/activate
python -m pip install -U pip
# Install in editable mode
pip install -e .
# Run
mem0 --help
mem0 version
```
> **After moving to the new directory structure:** If you previously had the CLI installed from the old repo root, you need to re-run `pip install -e .` from inside the `python/` directory to pick up the new location.
## Run without installing globally
This still installs the package into your active virtualenv (editable), but you can invoke it via module execution:
```bash
source .venv/bin/activate
pip install -e .
python -m mem0_cli --help
```
## Optional extras
### OSS integration extras
```bash
pip install -e ".[oss]"
```
### Dev tools (tests/lint)
```bash
pip install -e ".[dev]"
```
## Run tests
```bash
pip install -e ".[dev]"
# Run all tests
pytest
# Run a specific test file
pytest tests/test_cli_integration.py
# Run a single test
pytest -k test_help
```
## Lint
```bash
ruff check .
ruff format .
```
+77
View File
@@ -0,0 +1,77 @@
[build-system]
requires = ["hatchling"]
build-backend = "hatchling.build"
[project]
name = "mem0-cli"
version = "0.1.0"
description = "The official CLI for mem0 — the memory layer for AI agents"
readme = "README.md"
license = "Apache-2.0"
requires-python = ">=3.10"
authors = [
{ name = "mem0.ai", email = "founders@mem0.ai" },
]
keywords = ["mem0", "memory", "ai", "agents", "cli"]
classifiers = [
"Development Status :: 4 - Beta",
"Environment :: Console",
"Intended Audience :: Developers",
"License :: OSI Approved :: Apache Software License",
"Programming Language :: Python :: 3",
"Programming Language :: Python :: 3.10",
"Programming Language :: Python :: 3.11",
"Programming Language :: Python :: 3.12",
"Topic :: Software Development :: Libraries",
]
dependencies = [
"typer>=0.9.0",
"rich>=13.0.0",
"httpx>=0.24.0",
]
[project.optional-dependencies]
oss = ["mem0ai>=0.1.0"]
dev = [
"pytest>=7.0",
"pytest-asyncio>=0.21",
"ruff>=0.1.0",
]
[project.scripts]
mem0 = "mem0_cli.app:main"
[tool.hatch.build.targets.wheel]
packages = ["src/mem0_cli"]
[tool.hatch.build.targets.sdist]
include = ["src/mem0_cli"]
[tool.ruff]
target-version = "py310"
line-length = 100
[tool.ruff.lint]
select = [
"E", # pycodestyle errors
"F", # pyflakes
"I", # isort (import sorting)
"W", # pycodestyle warnings
"UP", # pyupgrade (modern Python syntax)
"B", # flake8-bugbear (common bugs)
"SIM", # flake8-simplify
"RUF", # ruff-specific rules
]
ignore = [
"E501", # line too long — handled by formatter
"B008", # function call in default arg — required by Typer's Option/Argument pattern
"SIM108", # ternary operator — sometimes less readable
]
[tool.ruff.lint.isort]
known-first-party = ["mem0_cli"]
[tool.ruff.format]
quote-style = "double"
indent-style = "space"
docstring-code-format = true
+3
View File
@@ -0,0 +1,3 @@
"""mem0 CLI — the command-line interface for the mem0 memory layer."""
__version__ = "0.1.0"
+5
View File
@@ -0,0 +1,5 @@
"""Allow running with `python -m mem0_cli`."""
from mem0_cli.app import main
main()
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,5 @@
"""Backend abstraction layer for mem0 CLI."""
from mem0_cli.backend.base import Backend, get_backend
__all__ = ["Backend", "get_backend"]
+113
View File
@@ -0,0 +1,113 @@
"""Abstract backend interface and factory."""
from __future__ import annotations
from abc import ABC, abstractmethod
from typing import Any
from mem0_cli.config import Mem0Config
class Backend(ABC):
"""Abstract interface for mem0 backends."""
@abstractmethod
def add(
self,
content: str | None = None,
messages: list[dict] | None = None,
*,
user_id: str | None = None,
agent_id: str | None = None,
app_id: str | None = None,
run_id: str | None = None,
metadata: dict | None = None,
immutable: bool = False,
infer: bool = True,
expires: str | None = None,
categories: list[str] | None = None,
enable_graph: bool = False,
) -> dict: ...
@abstractmethod
def search(
self,
query: str,
*,
user_id: str | None = None,
agent_id: str | None = None,
app_id: str | None = None,
run_id: str | None = None,
top_k: int = 10,
threshold: float = 0.3,
rerank: bool = False,
keyword: bool = False,
filters: dict | None = None,
fields: list[str] | None = None,
enable_graph: bool = False,
) -> list[dict]: ...
@abstractmethod
def get(self, memory_id: str) -> dict: ...
@abstractmethod
def list_memories(
self,
*,
user_id: str | None = None,
agent_id: str | None = None,
app_id: str | None = None,
run_id: str | None = None,
page: int = 1,
page_size: int = 100,
category: str | None = None,
after: str | None = None,
before: str | None = None,
enable_graph: bool = False,
) -> list[dict]: ...
@abstractmethod
def update(
self, memory_id: str, content: str | None = None, metadata: dict | None = None
) -> dict: ...
@abstractmethod
def delete(
self,
memory_id: str | None = None,
*,
all: bool = False,
user_id: str | None = None,
agent_id: str | None = None,
app_id: str | None = None,
run_id: str | None = None,
) -> dict: ...
@abstractmethod
def delete_entities(
self,
*,
user_id: str | None = None,
agent_id: str | None = None,
app_id: str | None = None,
run_id: str | None = None,
) -> dict: ...
@abstractmethod
def status(
self,
*,
user_id: str | None = None,
agent_id: str | None = None,
) -> dict[str, Any]: ...
@abstractmethod
def entities(self, entity_type: str) -> list[dict]: ...
def get_backend(config: Mem0Config) -> Backend:
"""Return the Platform backend."""
from mem0_cli.backend.platform import PlatformBackend
return PlatformBackend(config.platform)
+325
View File
@@ -0,0 +1,325 @@
"""Platform (SaaS) backend — communicates with api.mem0.ai."""
from __future__ import annotations
from typing import Any
import httpx
from mem0_cli.backend.base import Backend
from mem0_cli.config import PlatformConfig
class PlatformBackend(Backend):
"""Backend that talks to the mem0 Platform API."""
def __init__(self, config: PlatformConfig) -> None:
self.config = config
self.base_url = config.base_url.rstrip("/")
self._client = httpx.Client(
base_url=self.base_url,
headers={
"Authorization": f"Token {config.api_key}",
"Content-Type": "application/json",
},
timeout=30.0,
)
def _request(self, method: str, path: str, **kwargs: Any) -> Any:
resp = self._client.request(method, path, **kwargs)
if resp.status_code == 401:
raise AuthError("Authentication failed. Your API key may be invalid or expired.")
if resp.status_code == 404:
raise NotFoundError(f"Resource not found: {path}")
if resp.status_code == 400:
# Extract API error detail when available
try:
detail = resp.json().get("detail", resp.text)
except Exception:
detail = resp.text
raise APIError(f"Bad request to {path}: {detail}")
resp.raise_for_status()
if resp.status_code == 204:
return {}
return resp.json()
def add(
self,
content: str | None = None,
messages: list[dict] | None = None,
*,
user_id: str | None = None,
agent_id: str | None = None,
app_id: str | None = None,
run_id: str | None = None,
metadata: dict | None = None,
immutable: bool = False,
infer: bool = True,
expires: str | None = None,
categories: list[str] | None = None,
enable_graph: bool = False,
) -> dict:
payload: dict[str, Any] = {}
if messages:
payload["messages"] = messages
elif content:
payload["messages"] = [{"role": "user", "content": content}]
if user_id:
payload["user_id"] = user_id
if agent_id:
payload["agent_id"] = agent_id
if app_id:
payload["app_id"] = app_id
if run_id:
payload["run_id"] = run_id
if metadata:
payload["metadata"] = metadata
if immutable:
payload["immutable"] = True
if not infer:
payload["infer"] = False
if expires:
payload["expiration_date"] = expires
if categories:
payload["categories"] = categories
if enable_graph:
payload["enable_graph"] = True
return self._request("POST", "/v1/memories/", json=payload)
def _build_filters(
self,
*,
user_id: str | None = None,
agent_id: str | None = None,
app_id: str | None = None,
run_id: str | None = None,
extra_filters: dict | None = None,
) -> dict | None:
"""Build a filters dict for v2 API endpoints.
Entity IDs are ANDed (all provided IDs must match).
Extra filters (date ranges, categories) are also ANDed.
"""
# If caller passed a pre-built filter structure (e.g. --filter from CLI), use it directly
if extra_filters and ("AND" in extra_filters or "OR" in extra_filters):
return extra_filters
# Build AND conditions for entity IDs
and_conditions: list[dict[str, Any]] = []
if user_id:
and_conditions.append({"user_id": user_id})
if agent_id:
and_conditions.append({"agent_id": agent_id})
if app_id:
and_conditions.append({"app_id": app_id})
if run_id:
and_conditions.append({"run_id": run_id})
# Append any extra filters (dates, categories)
if extra_filters:
for k, v in extra_filters.items():
and_conditions.append({k: v})
if len(and_conditions) == 1:
return and_conditions[0]
elif and_conditions:
return {"AND": and_conditions}
else:
return None
def search(
self,
query: str,
*,
user_id: str | None = None,
agent_id: str | None = None,
app_id: str | None = None,
run_id: str | None = None,
top_k: int = 10,
threshold: float = 0.3,
rerank: bool = False,
keyword: bool = False,
filters: dict | None = None,
fields: list[str] | None = None,
enable_graph: bool = False,
) -> list[dict]:
payload: dict[str, Any] = {"query": query, "top_k": top_k, "threshold": threshold}
api_filters = self._build_filters(
user_id=user_id,
agent_id=agent_id,
app_id=app_id,
run_id=run_id,
extra_filters=filters,
)
if api_filters:
payload["filters"] = api_filters
if rerank:
payload["rerank"] = True
if keyword:
payload["keyword_search"] = True
if fields:
payload["fields"] = fields
if enable_graph:
payload["enable_graph"] = True
result = self._request("POST", "/v2/memories/search/", json=payload)
return (
result
if isinstance(result, list)
else result.get("results", result.get("memories", []))
)
def get(self, memory_id: str) -> dict:
return self._request("GET", f"/v1/memories/{memory_id}/")
def list_memories(
self,
*,
user_id: str | None = None,
agent_id: str | None = None,
app_id: str | None = None,
run_id: str | None = None,
page: int = 1,
page_size: int = 100,
category: str | None = None,
after: str | None = None,
before: str | None = None,
enable_graph: bool = False,
) -> list[dict]:
payload: dict[str, Any] = {}
params = {"page": str(page), "page_size": str(page_size)}
# Build filters for v2 API — entity IDs and date filters go inside "filters"
extra: dict[str, Any] = {}
if category:
extra["categories"] = {"contains": category}
if after:
extra["created_at"] = {**(extra.get("created_at", {})), "gte": after}
if before:
extra["created_at"] = {**(extra.get("created_at", {})), "lte": before}
api_filters = self._build_filters(
user_id=user_id,
agent_id=agent_id,
app_id=app_id,
run_id=run_id,
extra_filters=extra if extra else None,
)
if api_filters:
payload["filters"] = api_filters
if enable_graph:
payload["enable_graph"] = True
result = self._request("POST", "/v2/memories/", json=payload, params=params)
return (
result
if isinstance(result, list)
else result.get("results", result.get("memories", []))
)
def update(
self, memory_id: str, content: str | None = None, metadata: dict | None = None
) -> dict:
payload: dict[str, Any] = {}
if content:
payload["text"] = content
if metadata:
payload["metadata"] = metadata
return self._request("PUT", f"/v1/memories/{memory_id}/", json=payload)
def delete(
self,
memory_id: str | None = None,
*,
all: bool = False,
user_id: str | None = None,
agent_id: str | None = None,
app_id: str | None = None,
run_id: str | None = None,
) -> dict:
if all:
params: dict[str, str] = {}
if user_id:
params["user_id"] = user_id
if agent_id:
params["agent_id"] = agent_id
if app_id:
params["app_id"] = app_id
if run_id:
params["run_id"] = run_id
return self._request("DELETE", "/v1/memories/", params=params)
elif memory_id:
return self._request("DELETE", f"/v1/memories/{memory_id}/")
else:
raise ValueError("Either memory_id or --all is required")
def delete_entities(
self,
*,
user_id: str | None = None,
agent_id: str | None = None,
app_id: str | None = None,
run_id: str | None = None,
) -> dict:
params: dict[str, str] = {}
if user_id:
params["user_id"] = user_id
if agent_id:
params["agent_id"] = agent_id
if app_id:
params["app_id"] = app_id
if run_id:
params["run_id"] = run_id
if not params:
raise ValueError("At least one entity ID is required for delete_entities.")
return self._request("DELETE", "/v1/entities/", params=params)
def status(
self,
*,
user_id: str | None = None,
agent_id: str | None = None,
) -> dict[str, Any]:
"""Check connectivity by making a lightweight API call."""
try:
# If entity IDs are available, validate with a minimal memories list
if user_id or agent_id:
payload: dict[str, Any] = {}
params = {"page": "1", "page_size": "1"}
api_filters = self._build_filters(user_id=user_id, agent_id=agent_id)
if api_filters:
payload["filters"] = api_filters
self._request("POST", "/v2/memories/", json=payload, params=params)
else:
# No entity IDs — use entities endpoint to validate API key
self._request("GET", "/v1/entities/")
return {"connected": True, "backend": "platform", "base_url": self.base_url}
except Exception as e:
return {"connected": False, "backend": "platform", "error": str(e)}
def entities(self, entity_type: str) -> list[dict]:
result = self._request("GET", "/v1/entities/")
items = result if isinstance(result, list) else result.get("results", [])
# Filter by entity type client-side (API returns all types)
type_map = {"users": "user", "agents": "agent", "apps": "app", "runs": "run"}
target_type = type_map.get(entity_type)
if target_type:
items = [e for e in items if e.get("type", "").lower() == target_type]
return items
class AuthError(Exception):
pass
class NotFoundError(Exception):
pass
class APIError(Exception):
pass
+130
View File
@@ -0,0 +1,130 @@
"""Branding and ASCII art for mem0 CLI."""
import os
import sys
import time
from contextlib import contextmanager
from rich.console import Console
from rich.panel import Panel
from rich.status import Status
from rich.text import Text
# stderr console for spinners, errors, and timing messages
_err = Console(stderr=True)
LOGO = r"""
███╗ ███╗███████╗███╗ ███╗ ██████╗ ██████╗██╗ ██╗
████╗ ████║██╔════╝████╗ ████║██╔═████╗ ██╔════╝██║ ██║
██╔████╔██║█████╗ ██╔████╔██║██║██╔██║ ██║ ██║ ██║
██║╚██╔╝██║██╔══╝ ██║╚██╔╝██║████╔╝██║ ██║ ██║ ██║
██║ ╚═╝ ██║███████╗██║ ╚═╝ ██║╚██████╔╝ ╚██████╗███████╗██║
╚═╝ ╚═╝╚══════╝╚═╝ ╚═╝ ╚═════╝ ╚═════╝╚══════╝╚═╝
"""
LOGO_MINI = "◆ mem0"
TAGLINE = "The Memory Layer for AI Agents"
BRAND_COLOR = "#8b5cf6" # Purple
ACCENT_COLOR = "#a78bfa"
SUCCESS_COLOR = "#22c55e"
ERROR_COLOR = "#ef4444"
WARNING_COLOR = "#f59e0b"
DIM_COLOR = "#6b7280"
def _sym(fancy: str, plain: str) -> str:
"""Return *fancy* when stdout is a TTY with colour, else *plain*."""
if not sys.stdout.isatty() or os.environ.get("NO_COLOR") is not None:
return plain
return fancy
def print_banner(console: Console) -> None:
"""Print the mem0 welcome banner."""
logo_text = Text(LOGO, style=f"bold {BRAND_COLOR}")
tagline = Text(f" {TAGLINE}\n", style=f"{ACCENT_COLOR}")
content = Text()
content.append_text(logo_text)
content.append_text(tagline)
panel = Panel(
content,
border_style=BRAND_COLOR,
padding=(0, 2),
subtitle=f"[{DIM_COLOR}]Python SDK · v{_get_version()}[/]",
subtitle_align="right",
)
console.print(panel)
def print_success(console: Console, message: str) -> None:
sym = _sym("✓", "[ok]")
console.print(f"[{SUCCESS_COLOR}]{sym}[/] {message}")
def print_error(console: Console, message: str, hint: str | None = None) -> None:
sym = _sym("✗", "[error]")
console.print(f"[{ERROR_COLOR}]{sym} Error:[/] {message}")
if hint:
console.print(f" [{DIM_COLOR}]{hint}[/]")
def print_warning(console: Console, message: str) -> None:
sym = _sym("⚠", "[warn]")
console.print(f"[{WARNING_COLOR}]{sym}[/] {message}")
def print_info(console: Console, message: str) -> None:
sym = _sym("◆", "*")
console.print(f"[{BRAND_COLOR}]{sym}[/] {message}")
@contextmanager
def timed_status(console: Console, message: str):
"""Spinner with automatic timing. Yields a context object for setting the final message.
The spinner and timing output are sent to stderr (via ``_err``) so they
never contaminate machine-readable stdout. The *console* parameter is
kept for backward compatibility but is not used for spinner output.
"""
class _Ctx:
def __init__(self):
self.success_msg = ""
self.error_msg = ""
ctx = _Ctx()
start = time.perf_counter()
try:
with Status(f"[{DIM_COLOR}]{message}[/]", console=_err):
yield ctx
except Exception:
elapsed = time.perf_counter() - start
if ctx.error_msg:
print_error(_err, f"{ctx.error_msg} ({elapsed:.2f}s)")
raise
else:
elapsed = time.perf_counter() - start
if ctx.success_msg:
print_success(_err, f"{ctx.success_msg} ({elapsed:.2f}s)")
def print_scope(console: Console, **ids: str | None) -> None:
"""Show active entity scope if any IDs are set."""
parts = []
for key, val in ids.items():
if val:
label = key.replace("_", " ").replace("id", "ID").strip()
parts.append(f"{label}={val}")
if parts:
scope_str = ", ".join(parts)
console.print(f" [{DIM_COLOR}]Scope: {scope_str}[/]")
def _get_version() -> str:
from mem0_cli import __version__
return __version__
@@ -0,0 +1 @@
"""CLI command modules."""
@@ -0,0 +1,108 @@
"""Config management commands: show, set, get."""
from __future__ import annotations
from rich.console import Console
from rich.table import Table
from mem0_cli.branding import ACCENT_COLOR, BRAND_COLOR, DIM_COLOR, print_error, print_success
from mem0_cli.config import (
get_nested_value,
load_config,
redact_key,
save_config,
set_nested_value,
)
console = Console()
err_console = Console(stderr=True)
def cmd_config_show(*, output: str = "text") -> None:
"""Display current configuration (secrets redacted)."""
from mem0_cli.output import format_json_envelope
config = load_config()
if output == "json":
format_json_envelope(
console,
command="config show",
data={
"defaults": {
"user_id": config.defaults.user_id or None,
"agent_id": config.defaults.agent_id or None,
"app_id": config.defaults.app_id or None,
"run_id": config.defaults.run_id or None,
"enable_graph": config.defaults.enable_graph,
},
"platform": {
"api_key": redact_key(config.platform.api_key),
"base_url": config.platform.base_url,
},
},
)
return
console.print()
console.print(f" [{BRAND_COLOR}]◆ mem0 Configuration[/]\n")
table = Table(border_style=BRAND_COLOR, header_style=f"bold {ACCENT_COLOR}", padding=(0, 2))
table.add_column("Key", style="bold")
table.add_column("Value")
# Defaults
table.add_row(
"defaults.user_id",
config.defaults.user_id or f"[{DIM_COLOR}](not set)[/]",
)
table.add_row(
"defaults.agent_id",
config.defaults.agent_id or f"[{DIM_COLOR}](not set)[/]",
)
table.add_row(
"defaults.app_id",
config.defaults.app_id or f"[{DIM_COLOR}](not set)[/]",
)
table.add_row(
"defaults.run_id",
config.defaults.run_id or f"[{DIM_COLOR}](not set)[/]",
)
table.add_row(
"defaults.enable_graph",
str(config.defaults.enable_graph).lower(),
)
table.add_row("", "")
# Platform
table.add_row("[bold]platform.api_key[/]", redact_key(config.platform.api_key))
table.add_row("platform.base_url", config.platform.base_url)
console.print(table)
console.print()
def cmd_config_get(key: str) -> None:
"""Get a config value."""
config = load_config()
value = get_nested_value(config, key)
if value is None:
print_error(err_console, f"Unknown config key: {key}")
else:
# Redact secrets
if "api_key" in key or "key" in key.split(".")[-1:]:
console.print(redact_key(str(value)))
else:
console.print(str(value))
def cmd_config_set(key: str, value: str) -> None:
"""Set a config value."""
config = load_config()
if set_nested_value(config, key, value):
save_config(config)
display = redact_key(value) if "key" in key else value
print_success(console, f"{key} = {display}")
else:
print_error(err_console, f"Unknown config key: {key}")
@@ -0,0 +1,133 @@
"""Entity management commands."""
from __future__ import annotations
import time as _time
import typer
from rich.console import Console
from rich.table import Table
from mem0_cli.backend.base import Backend
from mem0_cli.branding import (
ACCENT_COLOR,
BRAND_COLOR,
DIM_COLOR,
print_error,
print_info,
print_success,
timed_status,
)
from mem0_cli.output import format_json
console = Console()
err_console = Console(stderr=True)
def cmd_entities_list(backend: Backend, entity_type: str, *, output: str) -> None:
"""List entities of a given type."""
valid_types = {"users", "agents", "apps", "runs"}
if entity_type not in valid_types:
print_error(err_console, f"Invalid entity type: {entity_type}. Use: {', '.join(valid_types)}")
raise typer.Exit(1)
_start = _time.perf_counter()
with timed_status(err_console, f"Fetching {entity_type}...") as _ts:
try:
results = backend.entities(entity_type)
except Exception as e:
print_error(err_console, str(e), hint="This feature may require the mem0 Platform.")
raise typer.Exit(1) from None
_elapsed = _time.perf_counter() - _start
if output == "json":
format_json(console, results)
return
if not results:
print_info(console, f"No {entity_type} found.")
return
table = Table(border_style=BRAND_COLOR, header_style=f"bold {ACCENT_COLOR}", padding=(0, 1))
table.add_column("Name / ID", style="bold")
table.add_column("Created", max_width=12)
for entity in results:
name = entity.get("name", entity.get("id", "—"))
created = str(entity.get("created_at", "—"))[:10]
table.add_row(str(name), created)
console.print()
console.print(table)
console.print(f" [{DIM_COLOR}]{len(results)} {entity_type} ({_elapsed:.2f}s)[/]")
console.print()
def cmd_entities_delete(
backend: Backend,
*,
user_id: str | None,
agent_id: str | None,
app_id: str | None,
run_id: str | None,
force: bool,
dry_run: bool = False,
output: str,
) -> None:
"""Delete an entity and all its memories (cascade delete)."""
if not any([user_id, agent_id, app_id, run_id]):
print_error(err_console, "Provide at least one of --user-id, --agent-id, --app-id, --run-id.")
raise typer.Exit(1)
if dry_run:
scope_parts = []
if user_id:
scope_parts.append(f"user={user_id}")
if agent_id:
scope_parts.append(f"agent={agent_id}")
if app_id:
scope_parts.append(f"app={app_id}")
if run_id:
scope_parts.append(f"run={run_id}")
scope = ", ".join(scope_parts)
print_info(console, f"Would delete entity {scope} and all its memories.")
print_info(console, "No changes made (dry run).")
return
if not force:
scope_parts = []
if user_id:
scope_parts.append(f"user={user_id}")
if agent_id:
scope_parts.append(f"agent={agent_id}")
if app_id:
scope_parts.append(f"app={app_id}")
if run_id:
scope_parts.append(f"run={run_id}")
scope = ", ".join(scope_parts)
confirm = typer.confirm(
f"\n \u26a0 Delete entity {scope} AND all its memories? This cannot be undone."
)
if not confirm:
print_info(console, "Cancelled.")
raise typer.Exit(0)
_start = _time.perf_counter()
with timed_status(err_console, "Deleting entity...") as _ts:
try:
result = backend.delete_entities(
user_id=user_id,
agent_id=agent_id,
app_id=app_id,
run_id=run_id,
)
except Exception as e:
print_error(err_console, str(e))
raise typer.Exit(1) from None
_elapsed = _time.perf_counter() - _start
if output == "json":
format_json(console, result)
elif output != "quiet":
print_success(console, f"Entity deleted with all memories ({_elapsed:.2f}s)")
@@ -0,0 +1,195 @@
"""mem0 init — interactive setup wizard."""
from __future__ import annotations
import sys
import typer
from rich.console import Console
from rich.prompt import Prompt
from mem0_cli.branding import (
BRAND_COLOR,
DIM_COLOR,
print_banner,
print_error,
print_info,
print_success,
)
from mem0_cli.config import Mem0Config, save_config
console = Console()
err_console = Console(stderr=True)
def _prompt_secret(label: str) -> str:
"""Prompt for a secret value, echoing '*' for each character typed."""
sys.stdout.write(label)
sys.stdout.flush()
chars: list[str] = []
if sys.platform == "win32":
import msvcrt
while True:
ch = msvcrt.getwch()
if ch in ("\r", "\n"):
sys.stdout.write("\n")
sys.stdout.flush()
break
if ch == "\x03":
raise KeyboardInterrupt
if ch in ("\x08", "\x7f"): # backspace
if chars:
chars.pop()
sys.stdout.write("\b \b")
sys.stdout.flush()
else:
chars.append(ch)
sys.stdout.write("*")
sys.stdout.flush()
else:
import termios
import tty
fd = sys.stdin.fileno()
old_settings = termios.tcgetattr(fd)
try:
tty.setraw(fd)
while True:
ch = sys.stdin.read(1)
if ch in ("\r", "\n"):
sys.stdout.write("\r\n")
sys.stdout.flush()
break
if ch == "\x03":
raise KeyboardInterrupt
if ch in ("\x7f", "\x08"): # backspace/delete
if chars:
chars.pop()
sys.stdout.write("\b \b")
sys.stdout.flush()
elif ch == "\x15": # Ctrl+U — clear line
sys.stdout.write("\b \b" * len(chars))
sys.stdout.flush()
chars = []
elif ch >= " ": # ignore other control characters
chars.append(ch)
sys.stdout.write("*")
sys.stdout.flush()
finally:
termios.tcsetattr(fd, termios.TCSADRAIN, old_settings)
return "".join(chars)
def run_init(*, api_key: str | None = None, user_id: str | None = None) -> None:
"""Interactive setup wizard for mem0 CLI.
When both *api_key* and *user_id* are supplied, all prompts are skipped
(non-interactive mode). When running in a non-TTY without the required
flags, an error message is printed.
"""
config = Mem0Config()
# Fully non-interactive when both flags provided
if api_key and user_id:
config.platform.api_key = api_key
config.defaults.user_id = user_id
_validate_platform(config)
save_config(config)
print_success(console, "Configuration saved to ~/.mem0/config.json")
return
# Non-TTY without full flags -> error
if not sys.stdin.isatty():
if not api_key or not user_id:
print_error(
err_console,
"Non-interactive terminal detected and required flags missing.",
hint="Run: mem0 init --api-key <key> --user-id <id>",
)
raise typer.Exit(1)
print_banner(console)
console.print()
print_info(console, "Welcome! Let's set up your mem0 CLI.\n")
# Use provided flags or prompt
if api_key:
config.platform.api_key = api_key
else:
_setup_platform(config)
if user_id:
config.defaults.user_id = user_id
else:
_setup_defaults(config)
_validate_platform(config)
save_config(config)
console.print()
print_success(console, "Configuration saved to ~/.mem0/config.json")
console.print()
console.print(f" [{DIM_COLOR}]Get started:[/]")
if config.defaults.user_id:
console.print(f' [{DIM_COLOR}] mem0 add "I prefer dark mode"[/]')
console.print(f' [{DIM_COLOR}] mem0 search "preferences"[/]')
else:
console.print(f' [{DIM_COLOR}] mem0 add "I prefer dark mode" --user-id alice[/]')
console.print(f' [{DIM_COLOR}] mem0 search "preferences" --user-id alice[/]')
console.print()
def _setup_platform(config: Mem0Config) -> None:
"""Platform setup flow."""
console.print()
console.print(f" [{DIM_COLOR}]Get your API key at https://app.mem0.ai/dashboard/api-keys[/]")
console.print()
console.print(f" [{BRAND_COLOR}]API Key[/]: ", end="")
api_key = _prompt_secret("")
if not api_key:
print_error(err_console, "API key is required.")
raise typer.Exit(1)
config.platform.api_key = api_key
def _setup_defaults(config: Mem0Config) -> None:
"""Collect default entity IDs."""
console.print()
print_info(console, "Set default entity IDs (press Enter to skip).\n")
user_id = Prompt.ask(
f" [{BRAND_COLOR}]Default User ID[/] [{DIM_COLOR}](recommended)[/]",
default="mem0-cli",
)
if user_id:
config.defaults.user_id = user_id
def _validate_platform(config: Mem0Config) -> None:
"""Validate platform connection after all inputs are collected."""
console.print()
print_info(console, "Validating connection...")
try:
from mem0_cli.backend.platform import PlatformBackend
backend = PlatformBackend(config.platform)
status = backend.status(
user_id=config.defaults.user_id or None,
agent_id=config.defaults.agent_id or None,
)
if status.get("connected"):
print_success(console, "Connected to mem0 Platform!")
else:
print_error(
err_console,
f"Could not connect: {status.get('error', 'Unknown error')}",
hint="Check your API key and try again.",
)
except Exception as e:
print_error(err_console, f"Connection test failed: {e}")
+469
View File
@@ -0,0 +1,469 @@
"""Memory CRUD commands: add, search, get, list, update, delete."""
from __future__ import annotations
import json
import sys
import time as _time
from pathlib import Path
import typer
from rich.console import Console
from mem0_cli.backend.base import Backend
from mem0_cli.branding import (
print_error,
print_info,
print_scope,
print_success,
timed_status,
)
from mem0_cli.output import (
format_add_result,
format_json,
format_memories_table,
format_memories_text,
format_single_memory,
print_result_summary,
)
console = Console()
err_console = Console(stderr=True)
def cmd_add(
backend: Backend,
text: str | None,
*,
user_id: str | None,
agent_id: str | None,
app_id: str | None,
run_id: str | None,
messages: str | None,
file: Path | None,
metadata: str | None,
immutable: bool,
no_infer: bool,
expires: str | None,
categories: str | None,
enable_graph: bool = False,
output: str = "text",
) -> None:
"""Add a memory."""
msgs = None
content = text
# Read from file
if file:
try:
raw = Path(file).read_text()
msgs = json.loads(raw)
except (FileNotFoundError, json.JSONDecodeError) as e:
print_error(err_console, f"Failed to read file: {e}")
raise typer.Exit(1) from None
# Parse messages JSON
elif messages:
try:
msgs = json.loads(messages)
except json.JSONDecodeError as e:
print_error(err_console, f"Invalid JSON in --messages: {e}")
raise typer.Exit(1) from None
# Read from stdin if no text and stdin is piped
elif not content and not sys.stdin.isatty():
content = sys.stdin.read().strip()
if not content and not msgs:
print_error(
err_console, "No content provided. Pass text, --messages, --file, or pipe via stdin."
)
raise typer.Exit(1)
meta = None
if metadata:
try:
meta = json.loads(metadata)
except json.JSONDecodeError:
print_error(err_console, "Invalid JSON in --metadata.")
raise typer.Exit(1) from None
cats = None
if categories:
try:
cats = json.loads(categories)
except json.JSONDecodeError:
cats = [c.strip() for c in categories.split(",")]
with timed_status(err_console, "Adding memory...") as ts:
try:
result = backend.add(
content=content,
messages=msgs,
user_id=user_id,
agent_id=agent_id,
app_id=app_id,
run_id=run_id,
metadata=meta,
immutable=immutable,
infer=not no_infer,
expires=expires,
categories=cats,
enable_graph=enable_graph,
)
except Exception as e:
ts.error_msg = str(e)
print_error(err_console, str(e))
raise typer.Exit(1) from None
if output == "quiet":
return
if output == "json":
format_add_result(console, result, output)
return
console.print()
print_scope(console, user_id=user_id, agent_id=agent_id, app_id=app_id, run_id=run_id)
# Count results
results = result if isinstance(result, list) else result.get("results", [result])
count = len(results) if results else 0
print_success(
console, f"Memory processed — {count} memor{'y' if count == 1 else 'ies'} extracted"
)
format_add_result(console, result, output)
def cmd_search(
backend: Backend,
query: str,
*,
user_id: str | None,
agent_id: str | None,
app_id: str | None,
run_id: str | None,
top_k: int,
threshold: float,
rerank: bool,
keyword: bool,
filter_json: str | None,
fields: str | None,
enable_graph: bool = False,
output: str = "text",
) -> None:
"""Search memories."""
filters = None
if filter_json:
try:
filters = json.loads(filter_json)
except json.JSONDecodeError:
print_error(err_console, "Invalid JSON in --filter.")
raise typer.Exit(1) from None
field_list = None
if fields:
field_list = [f.strip() for f in fields.split(",")]
_start = _time.perf_counter()
with timed_status(err_console, "Searching memories...") as _ts:
try:
results = backend.search(
query,
user_id=user_id,
agent_id=agent_id,
app_id=app_id,
run_id=run_id,
top_k=top_k,
threshold=threshold,
rerank=rerank,
keyword=keyword,
filters=filters,
fields=field_list,
enable_graph=enable_graph,
)
except Exception as e:
print_error(err_console, str(e))
raise typer.Exit(1) from None
_elapsed = _time.perf_counter() - _start
if output == "json":
format_json(console, results)
elif output == "table":
if results:
format_memories_table(console, results)
print_result_summary(
console, len(results), duration_secs=_elapsed, user_id=user_id, agent_id=agent_id
)
else:
console.print()
print_info(console, "No memories found matching your query.")
console.print()
else:
if results:
format_memories_text(console, results)
print_result_summary(
console, len(results), duration_secs=_elapsed, user_id=user_id, agent_id=agent_id
)
else:
console.print()
print_info(console, "No memories found matching your query.")
console.print()
def cmd_get(backend: Backend, memory_id: str, *, output: str) -> None:
"""Get a specific memory by ID."""
with timed_status(err_console, "Fetching memory...") as _ts:
try:
result = backend.get(memory_id)
except Exception as e:
print_error(err_console, str(e))
raise typer.Exit(1) from None
format_single_memory(console, result, output)
def cmd_list(
backend: Backend,
*,
user_id: str | None,
agent_id: str | None,
app_id: str | None,
run_id: str | None,
page: int,
page_size: int,
category: str | None,
after: str | None,
before: str | None,
enable_graph: bool = False,
output: str = "table",
) -> None:
"""List memories."""
_start = _time.perf_counter()
with timed_status(err_console, "Listing memories...") as _ts:
try:
results = backend.list_memories(
user_id=user_id,
agent_id=agent_id,
app_id=app_id,
run_id=run_id,
page=page,
page_size=page_size,
category=category,
after=after,
before=before,
enable_graph=enable_graph,
)
except Exception as e:
print_error(err_console, str(e))
raise typer.Exit(1) from None
_elapsed = _time.perf_counter() - _start
if output == "json":
format_json(console, results)
elif output == "table":
if results:
format_memories_table(console, results)
print_result_summary(
console,
len(results),
duration_secs=_elapsed,
page=page,
user_id=user_id,
agent_id=agent_id,
)
else:
console.print()
print_info(console, "No memories found.")
console.print()
else:
if results:
format_memories_text(console, results, title="memories")
print_result_summary(
console,
len(results),
duration_secs=_elapsed,
page=page,
user_id=user_id,
agent_id=agent_id,
)
else:
console.print()
print_info(console, "No memories found.")
console.print()
def cmd_update(
backend: Backend,
memory_id: str,
text: str | None,
*,
metadata: str | None,
output: str,
) -> None:
"""Update a memory."""
meta = None
if metadata:
try:
meta = json.loads(metadata)
except json.JSONDecodeError:
print_error(err_console, "Invalid JSON in --metadata.")
raise typer.Exit(1) from None
_start = _time.perf_counter()
with timed_status(err_console, "Updating memory...") as _ts:
try:
result = backend.update(memory_id, content=text, metadata=meta)
except Exception as e:
print_error(err_console, str(e))
raise typer.Exit(1) from None
_elapsed = _time.perf_counter() - _start
if output == "json":
format_json(console, result)
elif output != "quiet":
print_success(console, f"Memory {memory_id[:8]} updated ({_elapsed:.2f}s)")
def cmd_delete(
backend: Backend,
memory_id: str,
*,
dry_run: bool = False,
force: bool = False,
output: str,
) -> None:
"""Delete a single memory by ID."""
if dry_run:
# Fetch and display what would be deleted
try:
mem = backend.get(memory_id)
except Exception as e:
print_error(err_console, str(e))
raise typer.Exit(1) from None
format_single_memory(console, mem, output)
print_info(console, "No changes made (dry run).")
return
_start = _time.perf_counter()
with timed_status(err_console, "Deleting...") as _ts:
try:
result = backend.delete(memory_id=memory_id)
except Exception as e:
print_error(err_console, str(e))
raise typer.Exit(1) from None
_elapsed = _time.perf_counter() - _start
if output == "json":
format_json(console, result)
elif output != "quiet":
print_success(console, f"Memory {memory_id[:8]} deleted ({_elapsed:.2f}s)")
def cmd_delete_all(
backend: Backend,
*,
force: bool,
dry_run: bool = False,
all_: bool = False,
user_id: str | None,
agent_id: str | None,
app_id: str | None,
run_id: str | None,
output: str,
) -> None:
"""Delete all memories matching a scope."""
if all_:
# Project-wide wipe using wildcard entity IDs
if dry_run:
print_info(console, "Would delete ALL memories project-wide.")
print_info(console, "No changes made (dry run).")
return
if not force:
confirm = typer.confirm(
"\n ⚠ Delete ALL memories across the ENTIRE project? This cannot be undone."
)
if not confirm:
print_info(console, "Cancelled.")
raise typer.Exit(0)
_start = _time.perf_counter()
with timed_status(err_console, "Deleting all memories project-wide...") as _ts:
try:
result = backend.delete(
all=True,
user_id="*",
agent_id="*",
app_id="*",
run_id="*",
)
except Exception as e:
print_error(err_console, str(e))
raise typer.Exit(1) from None
_elapsed = _time.perf_counter() - _start
if output == "json":
format_json(console, result)
elif output != "quiet":
if isinstance(result, dict) and "message" in result:
print_info(console, "Deletion started. Memories will be removed in the background.")
else:
print_success(console, f"All project memories deleted ({_elapsed:.2f}s)")
return
if dry_run:
# List matching memories and show count
try:
results = backend.list_memories(
user_id=user_id,
agent_id=agent_id,
app_id=app_id,
run_id=run_id,
)
except Exception as e:
print_error(err_console, str(e))
raise typer.Exit(1) from None
count = len(results)
print_info(console, f"Would delete {count} memor{'y' if count == 1 else 'ies'}.")
print_info(console, "No changes made (dry run).")
return
if not force:
scope_parts = []
if user_id:
scope_parts.append(f"user={user_id}")
if agent_id:
scope_parts.append(f"agent={agent_id}")
if app_id:
scope_parts.append(f"app={app_id}")
if run_id:
scope_parts.append(f"run={run_id}")
scope = ", ".join(scope_parts) if scope_parts else "ALL entities"
confirm = typer.confirm(f"\n ⚠ Delete ALL memories for {scope}? This cannot be undone.")
if not confirm:
print_info(console, "Cancelled.")
raise typer.Exit(0)
_start = _time.perf_counter()
with timed_status(err_console, "Deleting all memories...") as _ts:
try:
result = backend.delete(
all=True,
user_id=user_id,
agent_id=agent_id,
app_id=app_id,
run_id=run_id,
)
except Exception as e:
print_error(err_console, str(e))
raise typer.Exit(1) from None
_elapsed = _time.perf_counter() - _start
if output == "json":
format_json(console, result)
elif output != "quiet":
if isinstance(result, dict) and "message" in result:
print_info(console, "Deletion started. Memories will be removed in the background.")
else:
print_success(console, f"All matching memories deleted ({_elapsed:.2f}s)")
+142
View File
@@ -0,0 +1,142 @@
"""Utility commands: status, version, import."""
from __future__ import annotations
import json
import time as _time
from pathlib import Path
import typer
from rich.console import Console
from rich.panel import Panel
from rich.progress import track
from mem0_cli import __version__
from mem0_cli.backend.base import Backend
from mem0_cli.branding import (
BRAND_COLOR,
DIM_COLOR,
ERROR_COLOR,
SUCCESS_COLOR,
print_error,
print_success,
timed_status,
)
from mem0_cli.config import load_config
console = Console()
err_console = Console(stderr=True)
def cmd_status(
backend: Backend,
*,
user_id: str | None = None,
agent_id: str | None = None,
output: str = "text",
) -> None:
"""Check connectivity and auth."""
from mem0_cli.output import format_json_envelope
_start = _time.perf_counter()
with timed_status(err_console, "Checking connection...") as _ts:
result = backend.status(user_id=user_id, agent_id=agent_id)
_elapsed = _time.perf_counter() - _start
if output == "json":
format_json_envelope(
console,
command="status",
data={
"connected": result.get("connected", False),
"backend": result.get("backend", "?"),
"base_url": result.get("base_url", ""),
"latency_ms": int(_elapsed * 1000),
},
duration_ms=int(_elapsed * 1000),
)
return
lines = []
if result.get("connected"):
lines.append(f" [{SUCCESS_COLOR}]●[/] Connected")
else:
lines.append(f" [{ERROR_COLOR}]●[/] Disconnected")
lines.append(f" [{DIM_COLOR}]Backend:[/] {result.get('backend', '?')}")
if result.get("base_url"):
lines.append(f" [{DIM_COLOR}]API URL:[/] {result['base_url']}")
if result.get("error"):
lines.append(f" [{ERROR_COLOR}]Error:[/] {result['error']}")
lines.append(f" [{DIM_COLOR}]Latency:[/] {_elapsed:.2f}s")
content = "\n".join(lines)
panel = Panel(
content,
title=f"[{BRAND_COLOR}]Connection Status[/]",
title_align="left",
border_style=BRAND_COLOR,
padding=(1, 1),
)
console.print()
console.print(panel)
console.print()
def cmd_version() -> None:
"""Show version."""
console.print(f" [{BRAND_COLOR}]◆ Mem0[/] CLI v{__version__}")
def cmd_import(
backend: Backend,
file_path: str,
*,
user_id: str | None,
agent_id: str | None,
output: str = "text",
) -> None:
"""Import memories from a JSON file."""
from mem0_cli.output import format_json_envelope
try:
data = json.loads(Path(file_path).read_text())
except (FileNotFoundError, json.JSONDecodeError) as e:
print_error(err_console, f"Failed to read file: {e}")
raise typer.Exit(1) from None
if not isinstance(data, list):
data = [data]
added = 0
failed = 0
_start = _time.perf_counter()
for item in track(data, description=f"[{DIM_COLOR}]Importing memories...[/]", console=err_console):
content = item.get("memory", item.get("text", item.get("content", "")))
if not content:
failed += 1
continue
try:
backend.add(
content=content,
user_id=user_id or item.get("user_id"),
agent_id=agent_id or item.get("agent_id"),
metadata=item.get("metadata"),
)
added += 1
except Exception:
failed += 1
_elapsed = _time.perf_counter() - _start
if output == "json":
format_json_envelope(
console,
command="import",
data={"added": added, "failed": failed, "duration_s": round(_elapsed, 2)},
duration_ms=int(_elapsed * 1000),
)
return
print_success(err_console, f"Imported {added} memories ({_elapsed:.2f}s)")
if failed:
print_error(err_console, f"{failed} memories failed to import.")
+176
View File
@@ -0,0 +1,176 @@
"""Configuration management for mem0 CLI.
Config precedence (highest to lowest):
1. CLI flags (--api-key, --base-url, etc.)
2. Environment variables (MEM0_API_KEY, etc.)
3. Config file (~/.mem0/config.json)
4. Defaults
"""
from __future__ import annotations
import json
import os
import stat
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any
CONFIG_DIR = Path.home() / ".mem0"
CONFIG_FILE = CONFIG_DIR / "config.json"
DEFAULT_BASE_URL = "https://api.mem0.ai"
CONFIG_VERSION = 1
@dataclass
class PlatformConfig:
api_key: str = ""
base_url: str = DEFAULT_BASE_URL
@dataclass
class DefaultsConfig:
user_id: str = ""
agent_id: str = ""
app_id: str = ""
run_id: str = ""
enable_graph: bool = False
@dataclass
class Mem0Config:
version: int = CONFIG_VERSION
defaults: DefaultsConfig = field(default_factory=DefaultsConfig)
platform: PlatformConfig = field(default_factory=PlatformConfig)
def ensure_config_dir() -> Path:
"""Create ~/.mem0 directory with secure permissions if it doesn't exist."""
CONFIG_DIR.mkdir(parents=True, exist_ok=True)
os.chmod(CONFIG_DIR, stat.S_IRWXU) # 0700
return CONFIG_DIR
def load_config() -> Mem0Config:
"""Load config from file, applying env var overrides."""
config = Mem0Config()
if CONFIG_FILE.exists():
with open(CONFIG_FILE) as f:
data = json.load(f)
config.version = data.get("version", CONFIG_VERSION)
plat = data.get("platform", {})
config.platform.api_key = plat.get("api_key", "")
config.platform.base_url = plat.get("base_url", DEFAULT_BASE_URL)
defaults = data.get("defaults", {})
config.defaults.user_id = defaults.get("user_id", "")
config.defaults.agent_id = defaults.get("agent_id", "")
config.defaults.app_id = defaults.get("app_id", "")
config.defaults.run_id = defaults.get("run_id", "")
config.defaults.enable_graph = defaults.get("enable_graph", False)
# Environment variable overrides
env_key = os.environ.get("MEM0_API_KEY")
if env_key:
config.platform.api_key = env_key
env_base = os.environ.get("MEM0_BASE_URL")
if env_base:
config.platform.base_url = env_base
env_user_id = os.environ.get("MEM0_USER_ID")
if env_user_id:
config.defaults.user_id = env_user_id
env_agent_id = os.environ.get("MEM0_AGENT_ID")
if env_agent_id:
config.defaults.agent_id = env_agent_id
env_app_id = os.environ.get("MEM0_APP_ID")
if env_app_id:
config.defaults.app_id = env_app_id
env_run_id = os.environ.get("MEM0_RUN_ID")
if env_run_id:
config.defaults.run_id = env_run_id
env_graph = os.environ.get("MEM0_ENABLE_GRAPH")
if env_graph:
config.defaults.enable_graph = env_graph.lower() in ("true", "1", "yes")
return config
def save_config(config: Mem0Config) -> None:
"""Write config to disk with secure permissions."""
ensure_config_dir()
data: dict[str, Any] = {
"version": config.version,
"defaults": {
"user_id": config.defaults.user_id,
"agent_id": config.defaults.agent_id,
"app_id": config.defaults.app_id,
"run_id": config.defaults.run_id,
"enable_graph": config.defaults.enable_graph,
},
"platform": {
"api_key": config.platform.api_key,
"base_url": config.platform.base_url,
},
}
with open(CONFIG_FILE, "w") as f:
json.dump(data, f, indent=2)
os.chmod(CONFIG_FILE, stat.S_IRUSR | stat.S_IWUSR) # 0600
def redact_key(key: str) -> str:
"""Redact an API key for display: m0-xxx...xxx"""
if not key:
return "(not set)"
if len(key) <= 8:
return key[:2] + "***"
return key[:4] + "..." + key[-4:]
def get_nested_value(config: Mem0Config, dotted_key: str) -> Any:
"""Get a config value by dotted path, e.g. 'platform.api_key'."""
parts = dotted_key.split(".")
obj: Any = config
for part in parts:
if hasattr(obj, part):
obj = getattr(obj, part)
else:
return None
return obj
def set_nested_value(config: Mem0Config, dotted_key: str, value: str) -> bool:
"""Set a config value by dotted path. Returns True on success."""
parts = dotted_key.split(".")
obj: Any = config
for part in parts[:-1]:
if hasattr(obj, part):
obj = getattr(obj, part)
else:
return False
final_key = parts[-1]
if not hasattr(obj, final_key):
return False
current = getattr(obj, final_key)
# Type coercion
if isinstance(current, bool):
value = value.lower() in ("true", "1", "yes") # type: ignore[assignment]
elif isinstance(current, int):
value = int(value) # type: ignore[assignment]
setattr(obj, final_key, value)
return True
+242
View File
@@ -0,0 +1,242 @@
"""Output formatting for mem0 CLI — text, JSON, table, quiet modes."""
from __future__ import annotations
import json
from datetime import datetime
from typing import Any
from rich.console import Console
from rich.panel import Panel
from rich.table import Table
from rich.text import Text
from mem0_cli.branding import ACCENT_COLOR, BRAND_COLOR, DIM_COLOR, SUCCESS_COLOR, _sym
def format_memories_text(console: Console, memories: list[dict], title: str = "memories") -> None:
"""Render memories in human-friendly text mode."""
count = len(memories)
console.print(f"\n[{BRAND_COLOR}]Found {count} {title}:[/]\n")
for i, mem in enumerate(memories, 1):
memory_text = mem.get("memory", mem.get("text", ""))
mem_id = mem.get("id", "")[:8]
score = mem.get("score")
created = _format_date(mem.get("created_at"))
category = mem.get("categories", [None])
if isinstance(category, list):
category = category[0] if category else None
line = Text()
line.append(f" {i}. ", style="bold")
line.append(memory_text, style="white")
console.print(line)
details = []
if score is not None:
details.append(f"Score: {score:.2f}")
if mem_id:
details.append(f"ID: {mem_id}")
if created:
details.append(f"Created: {created}")
if category:
details.append(f"Category: {category}")
if details:
detail_str = " · ".join(details)
console.print(f" [{DIM_COLOR}]{detail_str}[/]")
console.print()
def format_memories_table(console: Console, memories: list[dict]) -> None:
"""Render memories in a rich table."""
table = Table(
border_style=BRAND_COLOR,
header_style=f"bold {ACCENT_COLOR}",
row_styles=["", "dim"],
padding=(0, 1),
)
table.add_column("ID", style="dim", max_width=10)
table.add_column("Memory", max_width=50, no_wrap=False)
table.add_column("Category", max_width=14)
table.add_column("Created", max_width=12)
for mem in memories:
mem_id = mem.get("id", "")[:8]
memory_text = mem.get("memory", mem.get("text", ""))
if len(memory_text) > 60:
memory_text = memory_text[:57] + "..."
categories = mem.get("categories", [])
cat = categories[0] if isinstance(categories, list) and categories else "—"
created = _format_date(mem.get("created_at")) or "—"
table.add_row(mem_id, memory_text, cat, created)
console.print()
console.print(table)
console.print()
def format_json(console: Console, data: Any) -> None:
"""Output data as pretty-printed JSON."""
console.print_json(json.dumps(data, default=str))
def format_single_memory(console: Console, mem: dict, output: str = "text") -> None:
"""Format a single memory for display."""
if output == "json":
format_json(console, mem)
return
memory_text = mem.get("memory", mem.get("text", ""))
mem_id = mem.get("id", "")
lines = []
lines.append(f" [white bold]{memory_text}[/]")
lines.append("")
if mem_id:
lines.append(f" [{DIM_COLOR}]ID:[/] {mem_id}")
created = _format_date(mem.get("created_at"))
if created:
lines.append(f" [{DIM_COLOR}]Created:[/] {created}")
updated = _format_date(mem.get("updated_at"))
if updated:
lines.append(f" [{DIM_COLOR}]Updated:[/] {updated}")
meta = mem.get("metadata")
if meta:
lines.append(f" [{DIM_COLOR}]Metadata:[/] {json.dumps(meta)}")
categories = mem.get("categories")
if categories:
cat_str = ", ".join(categories) if isinstance(categories, list) else categories
lines.append(f" [{DIM_COLOR}]Categories:[/] {cat_str}")
content = "\n".join(lines)
panel = Panel(
content,
title=f"[{BRAND_COLOR}]Memory[/]",
title_align="left",
border_style=BRAND_COLOR,
padding=(1, 1),
)
console.print()
console.print(panel)
console.print()
def format_add_result(console: Console, result: dict | list, output: str = "text") -> None:
"""Format the result of an add operation."""
if output == "json":
format_json(console, result)
return
if output == "quiet":
return
# result from API is typically {"results": [...]}
results = result if isinstance(result, list) else result.get("results", [result])
if not results:
console.print(f" [{DIM_COLOR}]No memories extracted.[/]")
return
console.print()
for r in results:
# Detect async PENDING response from Platform API
if r.get("status") == "PENDING":
event_id = r.get("event_id", "")[:8]
icon = f"[{ACCENT_COLOR}]{_sym('⧗', '...')}[/]"
parts = [f" {icon} [{DIM_COLOR}]{'Queued':<10}[/]"]
parts.append("[white]Processing in background[/]")
if event_id:
parts.append(f"[{DIM_COLOR}](event {event_id})[/]")
console.print(" ".join(parts))
continue
event = r.get("event", "ADD")
memory = r.get("memory") or r.get("text") or r.get("content") or r.get("data") or ""
mem_id = (r.get("id") or r.get("memory_id") or "")[:8]
if event == "ADD":
icon = f"[{SUCCESS_COLOR}]+[/]"
label = "Added"
elif event == "UPDATE":
icon = f"[{ACCENT_COLOR}]~[/]"
label = "Updated"
elif event == "DELETE":
icon = "[red]-[/]"
label = "Deleted"
elif event == "NOOP":
icon = f"[{DIM_COLOR}]·[/]"
label = "No change"
else:
icon = f"[{DIM_COLOR}]?[/]"
label = event
# Build the display line
parts = [f" {icon} [{DIM_COLOR}]{label:<10}[/]"]
if memory:
parts.append(f"[white]{memory}[/]")
if mem_id:
parts.append(f"[{DIM_COLOR}]({mem_id})[/]")
console.print(" ".join(parts))
console.print()
def format_json_envelope(
console: Console,
*,
command: str,
data: Any,
duration_ms: int | None = None,
scope: dict | None = None,
count: int | None = None,
status: str = "success",
error: str | None = None,
) -> None:
"""Output structured JSON envelope for AI agent consumption."""
envelope: dict[str, Any] = {
"status": status,
"command": command,
}
if duration_ms is not None:
envelope["duration_ms"] = duration_ms
if scope is not None:
envelope["scope"] = scope
if count is not None:
envelope["count"] = count
if error:
envelope["error"] = error
envelope["data"] = data
console.print_json(json.dumps(envelope, default=str))
def print_result_summary(
console: Console,
count: int,
*,
duration_secs: float | None = None,
page: int | None = None,
**scope_ids: str | None,
) -> None:
"""Print a summary footer after result lists."""
parts = [f"{count} result{'s' if count != 1 else ''}"]
if page is not None:
parts.append(f"page {page}")
scope_parts = [f"{k.replace('_', ' ')}={v}" for k, v in scope_ids.items() if v]
if scope_parts:
parts.append(", ".join(scope_parts))
if duration_secs is not None:
parts.append(f"{duration_secs:.2f}s")
summary = " · ".join(parts)
console.print(f" [{DIM_COLOR}]{summary}[/]")
console.print()
def _format_date(dt_str: str | None) -> str | None:
if not dt_str:
return None
try:
dt = datetime.fromisoformat(dt_str.replace("Z", "+00:00"))
return dt.strftime("%Y-%m-%d")
except (ValueError, AttributeError):
return str(dt_str)[:10] if dt_str else None
View File
+109
View File
@@ -0,0 +1,109 @@
"""Shared fixtures for mem0 CLI tests."""
from __future__ import annotations
import os
from unittest.mock import MagicMock
import pytest
from mem0_cli.backend.base import Backend
from mem0_cli.config import Mem0Config
@pytest.fixture(autouse=True)
def isolate_config(tmp_path, monkeypatch):
"""Redirect config to a temp directory so tests don't touch real config."""
fake_config_dir = tmp_path / ".mem0"
fake_config_file = fake_config_dir / "config.json"
monkeypatch.setattr("mem0_cli.config.CONFIG_DIR", fake_config_dir)
monkeypatch.setattr("mem0_cli.config.CONFIG_FILE", fake_config_file)
# Also patch the commands that import config
monkeypatch.setattr("mem0_cli.commands.config_cmd.CONFIG_DIR", fake_config_dir, raising=False)
# Clear any MEM0 env vars
for key in list(os.environ.keys()):
if key.startswith("MEM0_"):
monkeypatch.delenv(key, raising=False)
return fake_config_dir
@pytest.fixture
def mock_backend():
"""Return a mock backend with all methods stubbed."""
backend = MagicMock(spec=Backend)
# Default return values
backend.add.return_value = {
"results": [
{
"id": "abc-123-def-456",
"memory": "User prefers dark mode",
"event": "ADD",
}
]
}
backend.search.return_value = [
{
"id": "abc-123-def-456",
"memory": "User prefers dark mode",
"score": 0.92,
"created_at": "2026-02-15T10:30:00Z",
"categories": ["preferences"],
},
{
"id": "ghi-789-jkl-012",
"memory": "User uses vim keybindings",
"score": 0.78,
"created_at": "2026-03-01T14:00:00Z",
"categories": ["tools"],
},
]
backend.get.return_value = {
"id": "abc-123-def-456",
"memory": "User prefers dark mode",
"created_at": "2026-02-15T10:30:00Z",
"updated_at": "2026-02-20T08:00:00Z",
"metadata": {"source": "onboarding"},
"categories": ["preferences"],
}
backend.list_memories.return_value = [
{
"id": "abc-123-def-456",
"memory": "User prefers dark mode",
"created_at": "2026-02-15T10:30:00Z",
"categories": ["preferences"],
},
{
"id": "ghi-789-jkl-012",
"memory": "User uses vim keybindings",
"created_at": "2026-03-01T14:00:00Z",
"categories": ["tools"],
},
]
backend.update.return_value = {"id": "abc-123-def-456", "memory": "Updated memory"}
backend.delete.return_value = {"status": "deleted"}
backend.status.return_value = {
"connected": True,
"backend": "platform",
"base_url": "https://api.mem0.ai",
}
backend.delete_entities.return_value = {"message": "Entity deleted"}
backend.entities.return_value = [
{"name": "alice", "count": 5},
{"name": "bob", "count": 3},
]
return backend
@pytest.fixture
def sample_config():
"""Return a sample config object."""
config = Mem0Config()
config.platform.api_key = "m0-test-key-12345678"
config.platform.base_url = "https://api.mem0.ai"
return config
+54
View File
@@ -0,0 +1,54 @@
"""Tests for branding and output helpers."""
from __future__ import annotations
from io import StringIO
from rich.console import Console
from mem0_cli.branding import print_banner, print_error, print_info, print_success, print_warning
def _make_console() -> tuple[Console, StringIO]:
buf = StringIO()
return Console(file=buf, force_terminal=False, no_color=True, width=80), buf
class TestBranding:
def test_print_banner(self):
console, buf = _make_console()
print_banner(console)
output = buf.getvalue()
# Banner contains the mem0 ASCII art and tagline
assert "Memory Layer" in output or "mem" in output.lower()
def test_print_success(self):
console, buf = _make_console()
print_success(console, "It worked!")
output = buf.getvalue()
assert "It worked!" in output
def test_print_error(self):
console, buf = _make_console()
print_error(console, "Something failed", hint="Try this fix")
output = buf.getvalue()
assert "Something failed" in output
assert "Try this fix" in output
def test_print_error_no_hint(self):
console, buf = _make_console()
print_error(console, "Failed")
output = buf.getvalue()
assert "Failed" in output
def test_print_warning(self):
console, buf = _make_console()
print_warning(console, "Watch out")
output = buf.getvalue()
assert "Watch out" in output
def test_print_info(self):
console, buf = _make_console()
print_info(console, "FYI")
output = buf.getvalue()
assert "FYI" in output
+238
View File
@@ -0,0 +1,238 @@
"""Integration tests — invoke CLI as subprocess to test end-to-end.
These tests launch the CLI as a real subprocess, so they must manage
environment isolation themselves (monkeypatch doesn't cross process
boundaries).
"""
from __future__ import annotations
import os
import subprocess
import sys
import pytest
def _run(
args: list[str],
env_override: dict | None = None,
home_dir: str | None = None,
) -> subprocess.CompletedProcess:
"""Run mem0 CLI command and capture output.
Args:
args: CLI arguments.
env_override: Extra env vars to set.
home_dir: If provided, set HOME to this path so the subprocess
reads config from ``<home_dir>/.mem0/config.json`` instead
of the user's real config. This is critical for tests that
depend on a clean (no API key) or custom config state.
"""
env = os.environ.copy()
# Strip all MEM0_ env vars so tests start clean
for key in list(env.keys()):
if key.startswith("MEM0_"):
del env[key]
if home_dir:
env["HOME"] = home_dir
if env_override:
env.update(env_override)
return subprocess.run(
[sys.executable, "-m", "mem0_cli", *args],
capture_output=True,
text=True,
env=env,
)
@pytest.fixture
def clean_home(tmp_path):
"""Return a temp directory to use as HOME, ensuring no ~/.mem0 exists."""
return str(tmp_path)
class TestCLIIntegration:
"""Tests that only inspect help text / version — no config needed."""
def test_help(self):
result = _run(["--help"])
assert result.returncode == 0
assert "mem0" in result.stdout
assert "add" in result.stdout
assert "search" in result.stdout
def test_version_flag(self):
result = _run(["--version"])
assert result.returncode == 0
assert "0.1.0" in result.stdout
def test_add_help(self):
result = _run(["add", "--help"])
assert result.returncode == 0
assert "user-id" in result.stdout
assert "messages" in result.stdout
def test_add_help_has_scope_panel(self):
"""Verify rich_help_panel grouping shows in help output."""
result = _run(["add", "--help"])
assert result.returncode == 0
assert "Scope" in result.stdout
def test_search_help(self):
result = _run(["search", "--help"])
assert result.returncode == 0
assert "top-k" in result.stdout
def test_list_help(self):
result = _run(["list", "--help"])
assert result.returncode == 0
assert "page-size" in result.stdout
def test_delete_help(self):
result = _run(["delete", "--help"])
assert result.returncode == 0
assert "--all" in result.stdout
assert "--entity" in result.stdout
assert "--project" in result.stdout
assert "--force" in result.stdout
assert "--dry-run" in result.stdout
def test_entity_list_help(self):
result = _run(["entity", "list", "--help"])
assert result.returncode == 0
assert "entity-type" in result.stdout.lower() or "entity_type" in result.stdout.lower()
def test_entity_delete_help(self):
result = _run(["entity", "delete", "--help"])
assert result.returncode == 0
assert "--user-id" in result.stdout
assert "--force" in result.stdout
def test_import_help(self):
result = _run(["import", "--help"])
assert result.returncode == 0
def test_no_args_shows_help(self):
"""no_args_is_help=True makes Typer print help and exit with code 2."""
result = _run([])
# Typer returns exit code 2 for "no command given" — this is standard
# Click/Typer behaviour and not an error.
assert result.returncode in (0, 2)
assert "Usage" in result.stdout
class TestCLIIsolated:
"""Tests that need a clean HOME to avoid reading the user's real config."""
def test_add_no_key_errors(self, clean_home):
"""Without an API key, `mem0 add` must fail with a helpful message."""
result = _run(
["add", "test", "--user-id", "alice"],
home_dir=clean_home,
)
assert result.returncode != 0
combined = result.stderr + result.stdout
assert "API key" in combined or "api" in combined.lower() or "Error" in combined
def test_search_no_key_errors(self, clean_home):
"""Without an API key, `mem0 search` must fail."""
result = _run(
["search", "preferences", "--user-id", "alice"],
home_dir=clean_home,
)
assert result.returncode != 0
combined = result.stderr + result.stdout
assert "API key" in combined or "Error" in combined
def test_list_no_key_errors(self, clean_home):
"""Without an API key, `mem0 list` must fail."""
result = _run(["list"], home_dir=clean_home)
assert result.returncode != 0
combined = result.stderr + result.stdout
assert "API key" in combined or "Error" in combined
def test_delete_no_id_no_all_errors(self, clean_home):
"""Delete without memory_id, --all, or --entity must fail."""
result = _run(
["delete", "--api-key", "m0-fake-key"],
home_dir=clean_home,
)
assert result.returncode != 0
combined = result.stderr + result.stdout
assert "memory ID" in combined.lower() or "--all" in combined or "--entity" in combined or "Error" in combined
def test_config_show_clean(self, clean_home):
"""config show with no config should still work."""
result = _run(["config", "show"], home_dir=clean_home)
assert result.returncode == 0
assert "backend" in result.stdout.lower() or "platform" in result.stdout.lower()
def test_config_set_and_get_roundtrip(self, clean_home):
"""config set then config get should return the set value."""
_run(
["config", "set", "defaults.user_id", "integration-test-user"],
home_dir=clean_home,
)
result = _run(
["config", "get", "defaults.user_id"],
home_dir=clean_home,
)
assert result.returncode == 0
assert "integration-test-user" in result.stdout
def test_import_nonexistent_file(self, clean_home):
"""Importing a nonexistent file should fail gracefully."""
result = _run(
["import", "/nonexistent/file.json", "--api-key", "m0-fake"],
home_dir=clean_home,
)
assert result.returncode != 0
combined = result.stderr + result.stdout
assert "Failed" in combined or "Error" in combined or "error" in combined
def test_add_no_content_errors(self, clean_home):
"""add with no text/messages/file should fail."""
result = _run(
["add", "--user-id", "alice", "--api-key", "m0-fake"],
home_dir=clean_home,
)
assert result.returncode != 0
combined = result.stderr + result.stdout
assert "No content" in combined or "Error" in combined
class TestCLINewFeatures:
"""Tests for MCP parity features: --graph, --limit, entities delete."""
def test_add_help_has_graph(self):
result = _run(["add", "--help"])
assert result.returncode == 0
assert "--graph" in result.stdout
def test_search_help_has_graph_and_limit(self):
result = _run(["search", "--help"])
assert result.returncode == 0
assert "--graph" in result.stdout
assert "--limit" in result.stdout
def test_list_help_has_graph(self):
result = _run(["list", "--help"])
assert result.returncode == 0
assert "--graph" in result.stdout
def test_delete_entity_via_delete_flag(self):
"""delete --entity should appear in help output."""
result = _run(["delete", "--help"])
assert result.returncode == 0
assert "--entity" in result.stdout
def test_entity_delete_has_scope_options(self):
"""entity delete should expose scope options."""
result = _run(["entity", "delete", "--help"])
assert result.returncode == 0
assert "--user-id" in result.stdout
assert "--force" in result.stdout
assert "--app-id" in result.stdout
assert "--run-id" in result.stdout
File diff suppressed because it is too large Load Diff
+240
View File
@@ -0,0 +1,240 @@
"""Tests for configuration management."""
from __future__ import annotations
import os
from mem0_cli.config import (
Mem0Config,
get_nested_value,
load_config,
redact_key,
save_config,
set_nested_value,
)
class TestRedactKey:
def test_empty_key(self):
assert redact_key("") == "(not set)"
def test_short_key(self):
assert redact_key("abc") == "ab***"
def test_normal_key(self):
result = redact_key("m0-abcdefgh12345678")
assert result == "m0-a...5678"
assert "abcdefgh" not in result
def test_exact_8_chars(self):
# 8 chars is <= 8, so it gets the short redaction
assert redact_key("12345678") == "12***"
class TestConfig:
def test_default_config(self):
config = Mem0Config()
assert config.platform.base_url == "https://api.mem0.ai"
assert config.platform.api_key == ""
def test_save_and_load(self, isolate_config):
config = Mem0Config()
config.platform.api_key = "m0-test-key"
save_config(config)
loaded = load_config()
assert loaded.platform.api_key == "m0-test-key"
def test_env_var_override(self, isolate_config, monkeypatch):
config = Mem0Config()
config.platform.api_key = "file-key"
save_config(config)
monkeypatch.setenv("MEM0_API_KEY", "env-key")
loaded = load_config()
assert loaded.platform.api_key == "env-key"
def test_load_nonexistent_config(self, isolate_config):
config = load_config()
assert config.platform.api_key == ""
def test_config_file_permissions(self, isolate_config):
config = Mem0Config()
config.platform.api_key = "secret"
save_config(config)
from mem0_cli.config import CONFIG_FILE
mode = os.stat(CONFIG_FILE).st_mode & 0o777
assert mode == 0o600
def test_defaults_save_and_load(self, isolate_config):
config = Mem0Config()
config.defaults.user_id = "alice"
config.defaults.agent_id = "support-bot"
config.defaults.app_id = "my-app"
config.defaults.run_id = "run-001"
save_config(config)
loaded = load_config()
assert loaded.defaults.user_id == "alice"
assert loaded.defaults.agent_id == "support-bot"
assert loaded.defaults.app_id == "my-app"
assert loaded.defaults.run_id == "run-001"
def test_defaults_env_var_override(self, isolate_config, monkeypatch):
config = Mem0Config()
config.defaults.user_id = "file-user"
save_config(config)
monkeypatch.setenv("MEM0_USER_ID", "env-user")
monkeypatch.setenv("MEM0_AGENT_ID", "env-agent")
loaded = load_config()
assert loaded.defaults.user_id == "env-user"
assert loaded.defaults.agent_id == "env-agent"
def test_backward_compat_no_defaults_key(self, isolate_config):
"""Old config files without 'defaults' key should load fine."""
import json
from mem0_cli.config import CONFIG_FILE, ensure_config_dir
ensure_config_dir()
# Write a config without the "defaults" key
data = {
"version": 1,
"platform": {"api_key": "m0-test", "base_url": "https://api.mem0.ai"},
}
with open(CONFIG_FILE, "w") as f:
json.dump(data, f)
loaded = load_config()
assert loaded.platform.api_key == "m0-test"
assert loaded.defaults.user_id == ""
assert loaded.defaults.agent_id == ""
def test_default_config_has_empty_defaults(self):
config = Mem0Config()
assert config.defaults.user_id == ""
assert config.defaults.agent_id == ""
assert config.defaults.app_id == ""
assert config.defaults.run_id == ""
assert config.defaults.enable_graph is False
def test_enable_graph_save_and_load(self, isolate_config):
config = Mem0Config()
config.defaults.enable_graph = True
save_config(config)
loaded = load_config()
assert loaded.defaults.enable_graph is True
def test_enable_graph_env_var_true(self, isolate_config, monkeypatch):
monkeypatch.setenv("MEM0_ENABLE_GRAPH", "true")
loaded = load_config()
assert loaded.defaults.enable_graph is True
def test_enable_graph_env_var_false(self, isolate_config, monkeypatch):
config = Mem0Config()
config.defaults.enable_graph = True
save_config(config)
monkeypatch.setenv("MEM0_ENABLE_GRAPH", "false")
loaded = load_config()
assert loaded.defaults.enable_graph is False
def test_backward_compat_no_enable_graph_key(self, isolate_config):
"""Old config files without 'enable_graph' key should default to False."""
import json
from mem0_cli.config import CONFIG_FILE, ensure_config_dir
ensure_config_dir()
data = {
"version": 1,
"defaults": {"user_id": "alice"},
"platform": {"api_key": "m0-test", "base_url": "https://api.mem0.ai"},
}
with open(CONFIG_FILE, "w") as f:
json.dump(data, f)
loaded = load_config()
assert loaded.defaults.enable_graph is False
assert loaded.defaults.user_id == "alice"
class TestNestedAccess:
def test_get_nested_value(self):
config = Mem0Config()
config.platform.api_key = "test-key"
assert get_nested_value(config, "platform.api_key") == "test-key"
def test_get_nonexistent_key(self):
config = Mem0Config()
assert get_nested_value(config, "nonexistent.key") is None
def test_set_nested_value(self):
config = Mem0Config()
assert set_nested_value(config, "platform.api_key", "new-key")
assert config.platform.api_key == "new-key"
def test_set_nonexistent_key(self):
config = Mem0Config()
assert set_nested_value(config, "nonexistent.key", "val") is False
def test_get_defaults_user_id(self):
config = Mem0Config()
config.defaults.user_id = "alice"
assert get_nested_value(config, "defaults.user_id") == "alice"
def test_set_defaults_user_id(self):
config = Mem0Config()
assert set_nested_value(config, "defaults.user_id", "bob")
assert config.defaults.user_id == "bob"
def test_set_defaults_enable_graph(self):
config = Mem0Config()
assert set_nested_value(config, "defaults.enable_graph", "true")
assert config.defaults.enable_graph is True
class TestResolveIds:
def test_cli_flag_overrides_default(self):
from mem0_cli.app import _resolve_ids
config = Mem0Config()
config.defaults.user_id = "default-user"
ids = _resolve_ids(
config,
user_id="cli-user",
agent_id=None,
)
assert ids["user_id"] == "cli-user"
def test_default_used_when_flag_is_none(self):
from mem0_cli.app import _resolve_ids
config = Mem0Config()
config.defaults.user_id = "default-user"
config.defaults.agent_id = "default-agent"
ids = _resolve_ids(config, user_id=None, agent_id=None)
assert ids["user_id"] == "default-user"
assert ids["agent_id"] == "default-agent"
def test_none_when_neither_set(self):
from mem0_cli.app import _resolve_ids
config = Mem0Config()
ids = _resolve_ids(config, user_id=None, agent_id=None)
assert ids["user_id"] is None
assert ids["agent_id"] is None
assert ids["app_id"] is None
assert ids["run_id"] is None
def test_empty_string_treated_as_unset(self):
from mem0_cli.app import _resolve_ids
config = Mem0Config()
config.defaults.user_id = ""
ids = _resolve_ids(config, user_id=None)
assert ids["user_id"] is None
+136
View File
@@ -0,0 +1,136 @@
"""Tests for output formatting."""
from __future__ import annotations
from io import StringIO
from rich.console import Console
from mem0_cli.output import (
format_add_result,
format_memories_table,
format_memories_text,
format_single_memory,
)
def _make_console() -> tuple[Console, StringIO]:
buf = StringIO()
return Console(file=buf, force_terminal=False, no_color=True, width=120, highlight=False), buf
SAMPLE_MEMORIES = [
{
"id": "abc-123-def-456",
"memory": "User prefers dark mode",
"score": 0.92,
"created_at": "2026-02-15T10:30:00Z",
"categories": ["preferences"],
},
{
"id": "ghi-789-jkl-012",
"memory": "User uses vim keybindings",
"score": 0.78,
"created_at": "2026-03-01T14:00:00Z",
"categories": ["tools"],
},
]
class TestTextFormat:
def test_format_memories_text(self):
console, buf = _make_console()
format_memories_text(console, SAMPLE_MEMORIES)
output = buf.getvalue()
assert "Found 2 memories" in output
assert "dark mode" in output
assert "vim keybindings" in output
assert "0.92" in output
def test_format_memories_text_empty(self):
console, buf = _make_console()
format_memories_text(console, [])
output = buf.getvalue()
assert "Found 0" in output
class TestTableFormat:
def test_format_memories_table(self):
console, buf = _make_console()
format_memories_table(console, SAMPLE_MEMORIES)
output = buf.getvalue()
assert "dark mode" in output
assert "abc-123-" in output
def test_format_memories_table_empty(self):
console, buf = _make_console()
format_memories_table(console, [])
output = buf.getvalue()
# Should still render (empty table)
assert "ID" in output
class TestSingleMemory:
def test_format_single_memory_text(self):
console, buf = _make_console()
mem = SAMPLE_MEMORIES[0]
format_single_memory(console, mem, "text")
output = buf.getvalue()
assert "dark mode" in output
assert "abc-123-def-456" in output
def test_format_single_memory_json(self):
console, buf = _make_console()
mem = SAMPLE_MEMORIES[0]
format_single_memory(console, mem, "json")
output = buf.getvalue()
assert '"memory"' in output
class TestAddResult:
def test_format_add_result_text(self):
console, buf = _make_console()
result = {
"results": [
{"id": "abc-123-def-456", "memory": "User prefers dark mode", "event": "ADD"},
]
}
format_add_result(console, result, "text")
output = buf.getvalue()
assert "dark mode" in output
assert "Added" in output
def test_format_add_result_update_event(self):
console, buf = _make_console()
result = {
"results": [
{"id": "abc-123", "memory": "Updated pref", "event": "UPDATE"},
]
}
format_add_result(console, result, "text")
output = buf.getvalue()
assert "Updated" in output
def test_format_add_result_noop(self):
console, buf = _make_console()
result = {
"results": [
{"id": "abc-123", "memory": "Same thing", "event": "NOOP"},
]
}
format_add_result(console, result, "text")
output = buf.getvalue()
assert "No change" in output
def test_format_add_result_quiet(self):
console, buf = _make_console()
result = {"results": [{"id": "abc-123", "memory": "Quiet", "event": "ADD"}]}
format_add_result(console, result, "quiet")
output = buf.getvalue()
assert output.strip() == ""
def test_format_add_result_empty(self):
console, buf = _make_console()
format_add_result(console, {"results": []}, "text")
output = buf.getvalue()
assert "No memories extracted" in output