refactor(memory/mem0): inline single-use setup helpers, joined summary prints, blank-line squeeze

This commit is contained in:
Teknium
2026-09-02 22:20:12 -07:00
parent 7ce8de140c
commit cc4f88aa76
3 changed files with 48 additions and 110 deletions

View File

@@ -40,13 +40,6 @@ def _is_client_error(exc: Exception) -> bool:
return type(exc).__name__ in _CLIENT_ERROR_TYPES or any(s in err_str for s in ("404", "not found", "valid uuid"))
def _truthy(value: Any, falsy_strings: tuple[str, ...] | None = None) -> bool:
"""Coerce a flag; strings match ``falsy_strings`` (deny-list) when given, else a truthy allow-list."""
if not isinstance(value, str):
return bool(value)
return value.lower() not in falsy_strings if falsy_strings else value.lower() in ("true", "1", "yes")
def _read_mem0_json(config_path: Path) -> dict:
"""Best-effort read of mem0.json; missing/corrupt file -> {}."""
if config_path.exists():
@@ -117,9 +110,7 @@ class Mem0MemoryProvider(MemoryProvider):
"""Merge-write config to $HERMES_HOME/mem0.json."""
from utils import atomic_json_write
config_path = Path(hermes_home) / "mem0.json"
existing = _read_mem0_json(config_path)
existing.update(values)
atomic_json_write(config_path, existing, mode=0o600)
atomic_json_write(config_path, {**_read_mem0_json(config_path), **values}, mode=0o600)
def get_config_schema(self):
api_key_required = _load_config().get("mode", "platform") != "oss"
@@ -137,9 +128,7 @@ class Mem0MemoryProvider(MemoryProvider):
def _oss_hint(self, template: str, default: str = "vector store") -> str:
"""OSS-only hint; ``{vs}`` is the configured vector-store provider. "" in other modes."""
if self._mode != "oss":
return ""
return template.format(vs=self._config.get("oss", {}).get("vector_store", {}).get("provider", default))
return template.format(vs=self._config.get("oss", {}).get("vector_store", {}).get("provider", default)) if self._mode == "oss" else ""
def _create_backend(self):
# Lazy-install the mem0 SDK before the backend imports it (honors security.allow_lazy_installs);
@@ -151,9 +140,7 @@ class Mem0MemoryProvider(MemoryProvider):
from . import _backend
if self._mode == "oss":
return _backend.OSSBackend(self._config.get("oss", {}))
if self._host:
return _backend.SelfHostedBackend(self._api_key, self._host)
return _backend.PlatformBackend(self._api_key)
return _backend.SelfHostedBackend(self._api_key, self._host) if self._host else _backend.PlatformBackend(self._api_key)
except Exception as e:
logger.error("Mem0 backend failed to initialize (%s mode): %s", self._mode, e)
self._init_error = str(e)
@@ -191,12 +178,12 @@ class Mem0MemoryProvider(MemoryProvider):
"""Background-path wrapper: run ``call`` under the breaker; on error log ``msg`` and return None."""
try:
result = call()
self._record_success()
return result
except Exception as e:
self._record_failure()
log(msg, e)
return None
self._record_success()
return result
def initialize(self, session_id: str, **kwargs) -> None:
self._config = cfg = _load_config()
@@ -206,7 +193,8 @@ class Mem0MemoryProvider(MemoryProvider):
configured = cfg.get("user_id")
self._user_id = (None if configured == _DEFAULT_USER_ID else configured) or kwargs.get("user_id") or _DEFAULT_USER_ID
# Persisted rerank preference: default for mem0_search when the model omits ``rerank``. Platform-only.
self._rerank_default = _truthy(cfg.get("rerank", False))
_rr = cfg.get("rerank", False)
self._rerank_default = _rr.lower() in ("true", "1", "yes") if isinstance(_rr, str) else bool(_rr)
self._channel = kwargs.get("platform") or "cli"
self._backend = self._create_backend()
if self._backend and not self._atexit_registered:
@@ -243,11 +231,6 @@ class Mem0MemoryProvider(MemoryProvider):
backend = self._backend
if not query or backend is None or self._is_breaker_open():
return
with self._prefetch_lock:
# Same query already answered or still in flight: don't restart it.
if self._prefetch_query == query and (self._prefetch_done or (self._prefetch_thread and self._prefetch_thread.is_alive())):
return
self._prefetch_query, self._prefetch_result, self._prefetch_done = query, "", False
def _run():
results = self._try(lambda: self._search(query, backend=backend), logger.debug, "Mem0 prefetch failed: %s")
@@ -257,9 +240,12 @@ class Mem0MemoryProvider(MemoryProvider):
if self._prefetch_query == query:
self._prefetch_result, self._prefetch_done = body, True
t = threading.Thread(target=_run, daemon=True, name="mem0-prefetch")
with self._prefetch_lock:
self._prefetch_thread = t
# Same query already answered or still in flight: don't restart it.
if self._prefetch_query == query and (self._prefetch_done or (self._prefetch_thread and self._prefetch_thread.is_alive())):
return
self._prefetch_query, self._prefetch_result, self._prefetch_done = query, "", False
self._prefetch_thread = t = threading.Thread(target=_run, daemon=True, name="mem0-prefetch")
t.start()
def prefetch(self, query: str, *, session_id: str = "") -> str:
@@ -271,8 +257,7 @@ class Mem0MemoryProvider(MemoryProvider):
thread = self._prefetch_thread if self._prefetch_query == query else None
if thread:
thread.join(timeout=_PREFETCH_WAIT_SECS)
# Slow backend: skip injection; mem0_search tool remains the backstop.
return self._consume_prefetch_result(query) or ""
return self._consume_prefetch_result(query) or "" # slow backend: skip injection; mem0_search remains the backstop
def sync_turn(self, user_content: str, assistant_content: str, *, session_id: str = "") -> None:
"""Send the turn to Mem0 for server-side fact extraction (non-blocking)."""
@@ -302,7 +287,8 @@ class Mem0MemoryProvider(MemoryProvider):
def _tool_search(self, args: dict) -> str:
top_k = max(1, min(int(args.get("top_k", 10)), 50))
rerank = _truthy(args.get("rerank", self._rerank_default), falsy_strings=("false", "0", "no"))
rerank_raw = args.get("rerank", self._rerank_default)
rerank = rerank_raw.lower() not in ("false", "0", "no") if isinstance(rerank_raw, str) else bool(rerank_raw)
results = self._search(args["query"], top_k, rerank)
if not results:
return json.dumps({"result": "No relevant memories found."})
@@ -336,8 +322,6 @@ class Mem0MemoryProvider(MemoryProvider):
return tool_error(f"Missing required parameter: {missing}")
try:
result = body(self, args)
self._record_success()
return result
except Exception as e:
client = _is_client_error(e)
if client and on_client_error == "not_found":
@@ -345,6 +329,8 @@ class Mem0MemoryProvider(MemoryProvider):
if not client or on_client_error == "count":
self._record_failure()
return tool_error(self._format_error(label, e))
self._record_success()
return result
def _shutdown_backend(self):
with suppress(Exception):

View File

@@ -13,9 +13,7 @@ def _add_kwargs(user_id: str, agent_id: str, infer: bool, metadata: dict | None)
def _unwrap_results(response: Any) -> list:
"""Normalize API response — extract results list from dict or pass through."""
if isinstance(response, dict):
return response.get("results", [])
return response if isinstance(response, list) else []
return response.get("results", []) if isinstance(response, dict) else response if isinstance(response, list) else []
class Mem0Backend(ABC):
@@ -105,7 +103,6 @@ def _register_direct_openai_provider() -> None:
"""Register Hermes' OpenAI-only Mem0 LLM provider once per factory."""
from mem0.configs.llms.openai import OpenAIConfig
from mem0.utils.factory import LlmFactory
provider_map = getattr(LlmFactory, "provider_to_class", None)
register_provider = getattr(LlmFactory, "register_provider", None)
if not isinstance(provider_map, dict) or not callable(register_provider):
@@ -143,7 +140,6 @@ class OSSBackend(Mem0Backend):
vs_config["embedding_model_dims"] = dims
self._recreate_collection_if_dims_changed(vector_store.get("provider", "qdrant"), vs_config, dims)
vector_store["config"] = vs_config
config = {"vector_store": vector_store, "llm": _provider_block("llm", LLM_PROVIDERS), "embedder": _provider_block("embedder", EMBEDDER_PROVIDERS), "version": "v1.1"}
if str(config["llm"].get("provider") or "").strip().lower() == "openai":
# mem0 validates LlmConfig.provider before its factory lookup: build the supported OpenAI config, then swap the provider.

View File

@@ -54,12 +54,9 @@ def _prompt_api_key(label: str, env_var: str, hermes_home: str) -> str:
"""Prompt for API key, showing masked existing value if found."""
existing = os.environ.get(env_var, "")
env_path = Path(hermes_home) / ".env"
if not existing and env_path.exists():
# utf-8-sig: a Notepad BOM on line 1 would otherwise defeat the key match.
for line in env_path.read_text(encoding="utf-8-sig", errors="replace").splitlines():
if line.startswith(f"{env_var}="):
existing = line.split("=", 1)[1].strip()
break
if not existing and env_path.exists(): # utf-8-sig: a Notepad BOM on line 1 would otherwise defeat the key match
lines = env_path.read_text(encoding="utf-8-sig", errors="replace").splitlines()
existing = next((line.split("=", 1)[1].strip() for line in lines if line.startswith(f"{env_var}=")), "")
hint = f" (current: {_masked(existing)}, blank to keep)" if existing else ""
return getpass.getpass(f" {label} API key{hint}: ").strip()
@@ -93,8 +90,7 @@ _FLAG_DEFAULTS = {"oss_llm": "openai", "oss_embedder": "openai", "oss_vector": "
def parse_flags(argv: list[str] | None = None) -> dict[str, str]:
args = argv if argv is not None else sys.argv[1:]
flags: dict[str, Any] = {k: _FLAG_DEFAULTS.get(k, "") for k in _FLAG_KEYS}
flags["dry_run"] = False
flags: dict[str, Any] = {**{k: _FLAG_DEFAULTS.get(k, "") for k in _FLAG_KEYS}, "dry_run": False}
flag_map = {"--" + k.replace("_", "-"): k for k in _FLAG_KEYS}
i = 0
while i < len(args):
@@ -127,7 +123,6 @@ def build_oss_config(flags: dict[str, str]) -> tuple[dict, dict[str, str]]:
dims = KNOWN_DIMS.get(embedder_config["model"])
if dims:
embedder_config["embedding_dims"] = dims
vector_id = flags.get("oss_vector", "qdrant")
vector_config = dict(VECTOR_PROVIDERS[vector_id]["default_config"])
for key in _VECTOR_FLAG_KEYS.get(vector_id, ()):
@@ -135,7 +130,6 @@ def build_oss_config(flags: dict[str, str]) -> tuple[dict, dict[str, str]]:
vector_config[key] = int(val) if key == "port" else val
if "url" in vector_config:
vector_config.pop("path", None) # a remote Qdrant URL replaces local storage
oss_config = {"llm": {"provider": llm_id, "config": llm_config}, "embedder": {"provider": embedder_id, "config": embedder_config}, "vector_store": {"provider": vector_id, "config": vector_config}}
# An embedder sharing the LLM's provider reuses the LLM key when no embedder key was given.
llm_key = flags.get("oss_llm_key") if llm_def.get("needs_key") else ""
@@ -170,14 +164,8 @@ def _persist_provider_config(hermes_home: str, config: dict, provider_config: di
_write_env(Path(hermes_home) / ".env", env_writes)
if server:
_check_selfhosted_server(server)
print(f"\n Memory provider: {label}")
if server:
print(f" Server: {server}")
print(" Activation saved to config.yaml")
print(" Provider config saved")
if env_writes:
print(f" {key_line}")
print("\n Start a new session to activate.\n")
print("\n".join(["", f" Memory provider: {label}", *([f" Server: {server}"] if server else []), " Activation saved to config.yaml", " Provider config saved",
*([f" {key_line}"] if env_writes else []), "", " Start a new session to activate.", ""]))
def _setup_platform(hermes_home: str, config: dict, flags: dict[str, str]) -> None:
@@ -191,14 +179,12 @@ def _setup_platform(hermes_home: str, config: dict, flags: dict[str, str]) -> No
choices = ["true", "false"]
current = str(provider_config.get("rerank", "false") or "").lower()
provider_config["rerank"] = choices[_curses_select(" Enable reranking for recall", [(c, "") for c in choices], default=choices.index(current) if current in choices else 0)]
if flags.get("dry_run"):
_print_dry_run(str(provider_config), env_writes)
return
provider_config["mode"] = "platform"
# Routing checks ``host`` before platform, so clear a stale self-hosted host. "" rather than
# pop(): save_config merges into the existing mem0.json, so a popped key would survive.
provider_config["host"] = ""
provider_config.update(mode="platform", host="")
# _load_config() also seeds ``host`` from MEM0_HOST (.env); the file clear can't help there, so warn.
if os.environ.get("MEM0_HOST", "").strip():
print(f"\n ⚠ MEM0_HOST is set in your environment ({os.environ['MEM0_HOST']}). It overrides platform mode — remove it from ~/.hermes/.env (or unset it) or Hermes will keep routing to the self-hosted server.")
@@ -229,7 +215,6 @@ def _setup_selfhosted(hermes_home: str, config: dict, flags: dict[str, str]) ->
env_writes = _api_key_writes(flags, "Server API key", fresh_label="Server API key (blank if AUTH_DISABLED)")
user_id = flags.get("user_id") or _prompt("User identifier", default=provider_config.get("user_id") or "hermes-user")
agent_id = _prompt("Agent identifier", default=provider_config.get("agent_id") or "hermes")
if flags.get("dry_run"):
_print_dry_run(f"host={host}, user_id={user_id}, agent_id={agent_id}", env_writes, lambda: _check_selfhosted_server(host))
return
@@ -240,19 +225,14 @@ def _setup_selfhosted(hermes_home: str, config: dict, flags: dict[str, str]) ->
def _print_oss_summary(oss_config: dict, env_writes: dict, dry_run: bool = False) -> None:
llm, emb = oss_config["llm"], oss_config["embedder"]
w = 0 if dry_run else 9 # final summary column-aligns the labels
print("\n [dry-run] OSS config would be:" if dry_run else "\n ✓ Mem0 configured (OSS mode)")
print(f" {'LLM:':<{w}} {llm['provider']} ({llm['config'].get('model', '')})")
print(f" {'Embedder:':<{w}} {emb['provider']} ({emb['config'].get('model', '')})")
print(f" {'Vector:':<{w}} {oss_config['vector_store']['provider']}")
lines = ["", " [dry-run] OSS config would be:" if dry_run else " ✓ Mem0 configured (OSS mode)",
f" {'LLM:':<{w}} {llm['provider']} ({llm['config'].get('model', '')})", f" {'Embedder:':<{w}} {emb['provider']} ({emb['config'].get('model', '')})",
f" {'Vector:':<{w}} {oss_config['vector_store']['provider']}"]
if dry_run:
if env_writes:
print(f" Env vars: {', '.join(env_writes.keys())}")
return
if env_writes:
print(" API keys saved to .env")
print(" Config saved to mem0.json")
print(" Provider set in config.yaml")
print("\n Start a new session to activate.\n")
lines += [f" Env vars: {', '.join(env_writes.keys())}"] if env_writes else []
else:
lines += [*([" API keys saved to .env"] if env_writes else []), " Config saved to mem0.json", " Provider set in config.yaml", "", " Start a new session to activate.", ""]
print("\n".join(lines))
def _finish_oss(hermes_home: str, config: dict, oss_config: dict, env_writes: dict[str, str], user_id: str, agent_id: str, pgvector_config: dict | None = None) -> None:
@@ -313,13 +293,9 @@ def _ensure_pgvector(host: str = "localhost", port: int = 5432) -> dict | None:
if _pg_ready(host, port, 15):
print(" ✓ PostgreSQL container restarted")
return None
if input(" Start pgvector via Docker? [Y/n]: ").strip().lower() in ("", "y", "yes"):
return _start_pgvector_docker(host, port)
print(" Skipping Docker setup. Make sure PostgreSQL with pgvector is running.")
return None
def _start_pgvector_docker(host: str, port: int) -> dict | None:
if input(" Start pgvector via Docker? [Y/n]: ").strip().lower() not in ("", "y", "yes"):
print(" Skipping Docker setup. Make sure PostgreSQL with pgvector is running.")
return None
try:
print(f" Pulling {_PGVECTOR_IMAGE}...")
_docker("pull", _PGVECTOR_IMAGE, timeout=120)
@@ -357,7 +333,11 @@ def _ensure_ollama(models: list[str]) -> bool:
print(" Warning: Ollama not reachable. Models cannot be pulled.")
return False
for model in models:
if any(model in n or model.split(":")[0] in n for n in _ollama_models(_OLLAMA_URL)):
try:
names = [m.get("name", "") for m in json.loads(_http_get(_OLLAMA_URL, "/api/tags", 5).read()).get("models", [])]
except Exception:
names = []
if any(model in n or model.split(":")[0] in n for n in names):
print(f" ✓ Model '{model}' available")
continue
print(f" Pulling '{model}'... (this may take a few minutes)")
@@ -369,13 +349,6 @@ def _ensure_ollama(models: list[str]) -> bool:
return True
def _ollama_models(url: str) -> list[str]:
try:
return [m.get("name", "") for m in json.loads(_http_get(url, "/api/tags", 5).read()).get("models", [])]
except Exception:
return []
def _ensure_pgvector_extension(pg_config: dict) -> None:
try:
import psycopg2
@@ -403,18 +376,13 @@ def _wait_for_port(host: str, port: int, timeout: int = 15) -> None:
# Picker descriptions: LLM/embedder show model (+ URL); vector stores by provider id (default: the id itself).
def _provider_description(v: dict) -> str:
model, url = v.get("default_model", ""), v.get("default_url")
return f"{model} ({url})" if url else model
_VECTOR_DESCRIPTIONS = {"qdrant": lambda cfg: cfg.get("path", "local storage"), "pgvector": lambda cfg: f"{cfg.get('host', 'localhost')}:{cfg.get('port', 5432)}"}
def _configure_model_provider(kind: str, registry: dict, hermes_home: str, env_writes: dict[str, str], llm: tuple[str, dict] | None = None) -> tuple[str, dict, str, str | None]:
"""Pick an LLM/embedder provider, collect its key, and (for Ollama) model + URL -> (id, definition, model, url).
For the embedder (``llm`` given), a provider shared with the LLM reuses the LLM key instead of prompting again."""
items = [(v["label"], _provider_description(v)) for v in registry.values()]
items = [(v["label"], f"{v.get('default_model', '')} ({v['default_url']})" if v.get("default_url") else v.get("default_model", "")) for v in registry.values()]
pid = list(registry)[_curses_select(f"{kind} Provider", items, 0)]
pdef = registry[pid]
model, url = pdef["default_model"], pdef.get("default_url")
@@ -430,31 +398,23 @@ def _configure_model_provider(kind: str, registry: dict, hermes_home: str, env_w
return pid, pdef, model, url
def _prompt_pgvector_config() -> dict:
"""Native PostgreSQL — prompt for connection details (user first, matching the historical order)."""
pg = {k: _input(f"PostgreSQL {label}", d) for k, label, d in (
("user", "user", os.getenv("USER", "postgres")), ("host", "host", "localhost"), ("port", "port", "5432"), ("dbname", "database", "postgres"))}
pg_password = getpass.getpass(" PostgreSQL password (blank if none): ").strip()
return {**pg, "port": int(pg["port"]), **({"password": pg_password} if pg_password else {})}
def _setup_oss_interactive(hermes_home: str, config: dict) -> None:
env_writes: dict[str, str] = {}
llm_id, llm_def, llm_model, llm_url = _configure_model_provider("LLM", LLM_PROVIDERS, hermes_home, env_writes)
embedder_id, _, embedder_model, embedder_url = _configure_model_provider("Embedder", EMBEDDER_PROVIDERS, hermes_home, env_writes, llm=(llm_id, llm_def))
vector_items = [(v["label"], _VECTOR_DESCRIPTIONS.get(pid, lambda cfg: pid)(v.get("default_config", {}))) for pid, v in VECTOR_PROVIDERS.items()]
vector_id = list(VECTOR_PROVIDERS)[_curses_select("Vector Store", vector_items, 0)]
# Auto-setup: ensure Ollama is running and models are pulled; ensure pgvector is reachable (offer Docker if not).
ollama_models = [m for pid, m in ((llm_id, llm_model), (embedder_id, embedder_model)) if pid == "ollama"]
if ollama_models:
_ensure_ollama(ollama_models)
pgvector_config = None
if vector_id == "pgvector":
pgvector_config = _ensure_pgvector() or _prompt_pgvector_config()
pgvector_config = _ensure_pgvector() if vector_id == "pgvector" else None
if vector_id == "pgvector" and not pgvector_config: # native PostgreSQL: prompt for connection details (user first, historical order)
pg = {k: _input(f"PostgreSQL {label}", d) for k, label, d in (("user", "user", os.getenv("USER", "postgres")), ("host", "host", "localhost"), ("port", "port", "5432"), ("dbname", "database", "postgres"))}
pg_password = getpass.getpass(" PostgreSQL password (blank if none): ").strip()
pgvector_config = {**pg, "port": int(pg["port"]), **({"password": pg_password} if pg_password else {})}
user_id = _input("User ID", os.getenv("USER", "hermes-user"))
agent_id = _input("Agent ID", "hermes")
flags = {
"oss_llm": llm_id, "oss_llm_model": llm_model, "oss_llm_url": llm_url or "",
"oss_llm_key": env_writes.get(llm_def["env_var"], "") if llm_def.get("env_var") else "",
@@ -476,12 +436,8 @@ def _install_provider_deps(llm_id: str, embedder_id: str, vector_id: str) -> Non
outcome = install_specs([dep], timeout=60)
except Exception:
outcome = None
if outcome is not None and outcome.ok:
print(f" ✓ Installed {dep}")
elif outcome is not None and outcome.blocked:
print(f" Warning: cannot install {dep}: {outcome.reason}")
else:
print(f" Warning: Could not install {dep}. Install manually: uv pip install {dep}")
print(f" ✓ Installed {dep}" if outcome is not None and outcome.ok else f" Warning: cannot install {dep}: {outcome.reason}" if outcome is not None and outcome.blocked
else f" Warning: Could not install {dep}. Install manually: uv pip install {dep}")
if deps:
import importlib
importlib.invalidate_caches()