refactor(memory/mem0): contextlib.suppress for swallow-all blocks, derived _FLAG_KEYS, packed initialize

This commit is contained in:
Teknium
2026-09-02 21:58:20 -07:00
parent ce45335720
commit 7ce8de140c
3 changed files with 21 additions and 48 deletions

View File

@@ -15,6 +15,7 @@ import logging
import os
import threading
import time
from contextlib import suppress
from pathlib import Path
from typing import Any, Dict, List
@@ -49,10 +50,8 @@ def _truthy(value: Any, falsy_strings: tuple[str, ...] | None = None) -> bool:
def _read_mem0_json(config_path: Path) -> dict:
"""Best-effort read of mem0.json; missing/corrupt file -> {}."""
if config_path.exists():
try:
with suppress(Exception):
return json.loads(config_path.read_text(encoding="utf-8"))
except Exception:
pass
return {}
@@ -62,8 +61,7 @@ def _load_config() -> dict:
like ``api_key`` that the user set in ``.env``."""
from hermes_constants import get_hermes_home
config = {"mode": os.environ.get("MEM0_MODE", "platform"), "api_key": get_secret("MEM0_API_KEY", ""), "host": os.environ.get("MEM0_HOST", ""), "agent_id": os.environ.get("MEM0_AGENT_ID", "hermes"), "oss": {}}
# Only carry user_id when explicitly configured so initialize() can fall back to the gateway-native id.
if os.environ.get("MEM0_USER_ID"):
if os.environ.get("MEM0_USER_ID"): # only when explicitly configured, so initialize() can fall back to the gateway-native id
config["user_id"] = os.environ["MEM0_USER_ID"]
file_cfg = _read_mem0_json(get_hermes_home() / "mem0.json")
config.update({k: v for k, v in file_cfg.items() if v is not None and v != ""})
@@ -146,11 +144,9 @@ class Mem0MemoryProvider(MemoryProvider):
def _create_backend(self):
# Lazy-install the mem0 SDK before the backend imports it (honors security.allow_lazy_installs);
# on failure the backend import raises the canonical error, captured below.
try:
with suppress(Exception):
from tools.lazy_deps import ensure as _lazy_ensure
_lazy_ensure("memory.mem0", prompt=False)
except Exception:
pass
try:
from . import _backend
if self._mode == "oss":
@@ -203,17 +199,14 @@ class Mem0MemoryProvider(MemoryProvider):
return None
def initialize(self, session_id: str, **kwargs) -> None:
self._config = _load_config()
self._mode = self._config.get("mode", "platform")
self._api_key = self._config.get("api_key", "")
self._host = self._config.get("host", "")
self._config = cfg = _load_config()
self._mode, self._api_key, self._host, self._agent_id = cfg.get("mode", "platform"), cfg.get("api_key", ""), cfg.get("host", ""), cfg.get("agent_id", "hermes")
# user_id precedence: operator-configured (env/mem0.json) > gateway-native id (kwargs) > _DEFAULT_USER_ID.
# The literal placeholder counts as unset so wizard users still get gateway-native ids.
configured = self._config.get("user_id")
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
self._agent_id = self._config.get("agent_id", "hermes")
# Persisted rerank preference: default for mem0_search when the model omits ``rerank``. Platform-only.
self._rerank_default = _truthy(self._config.get("rerank", False))
self._rerank_default = _truthy(cfg.get("rerank", False))
self._channel = kwargs.get("platform") or "cli"
self._backend = self._create_backend()
if self._backend and not self._atexit_registered:
@@ -354,12 +347,10 @@ class Mem0MemoryProvider(MemoryProvider):
return tool_error(self._format_error(label, e))
def _shutdown_backend(self):
try:
with suppress(Exception):
if self._backend:
self._backend.close()
self._backend = None
except Exception:
pass
def shutdown(self) -> None:
for t in (self._prefetch_thread, self._sync_thread):

View File

@@ -3,7 +3,7 @@
from __future__ import annotations
from abc import ABC, abstractmethod
from contextlib import closing
from contextlib import closing, suppress
from typing import Any
@@ -93,10 +93,8 @@ class SelfHostedBackend(Mem0Backend):
self._json("DELETE", f"/memories/{memory_id}")
def close(self) -> None:
try:
with suppress(Exception):
self._client.close()
except Exception:
pass
_DIRECT_OPENAI_PROVIDER = "hermes_openai"
@@ -164,7 +162,7 @@ class OSSBackend(Mem0Backend):
def _recreate_collection_if_dims_changed(provider: str, vs_config: dict, expected_dims: int) -> None:
"""Delete stale vector collection when embedding dimensions change."""
collection_name = vs_config.get("collection_name", "mem0")
try:
with suppress(Exception):
if provider == "qdrant":
from qdrant_client import QdrantClient
path, url = vs_config.get("path"), vs_config.get("url")
@@ -195,8 +193,6 @@ class OSSBackend(Mem0Backend):
row = cur.fetchone()
if row and row[0] > 0 and row[0] != expected_dims:
cur.execute(pgsql.SQL("DROP TABLE IF EXISTS {}").format(pgsql.Identifier(collection_name)))
except Exception:
pass
def search(self, query: str, *, filters: dict, top_k: int = 10, rerank: bool = False) -> list[dict]:
return _unwrap_results(self._memory.search(query, filters=filters, top_k=top_k))
@@ -211,17 +207,13 @@ class OSSBackend(Mem0Backend):
self._memory.delete(memory_id)
def close(self):
try:
with suppress(Exception):
telemetry = getattr(self._memory, "telemetry", None)
if telemetry and hasattr(telemetry, "posthog"):
try:
with suppress(Exception):
telemetry.posthog.shutdown()
except Exception:
pass
vs = getattr(self._memory, "vector_store", None)
# Memory, then its vector store, then the store's raw client; the first failure aborts the chain.
for obj in filter(None, (self._memory, vs, getattr(vs, "client", None))):
if hasattr(obj, "close"):
obj.close()
except Exception:
pass

View File

@@ -4,6 +4,7 @@ from __future__ import annotations
import getpass
import json
from contextlib import suppress
import os
import shutil
import socket
@@ -83,13 +84,11 @@ def _print_dry_run(summary: str, env_writes: dict, check=None) -> None:
print(" [dry-run] No files written.\n")
_FLAG_KEYS = (
"mode", "api_key", "host", "oss_llm", "oss_llm_key", "oss_llm_model", "oss_llm_url", "oss_embedder", "oss_embedder_key", "oss_embedder_model", "oss_embedder_url",
"oss_vector", "oss_vector_path", "oss_vector_url", "oss_vector_host", "oss_vector_port", "oss_vector_user", "oss_vector_password", "oss_vector_dbname", "user_id",
)
_FLAG_DEFAULTS = {"oss_llm": "openai", "oss_embedder": "openai", "oss_vector": "qdrant"}
# --oss-vector-<key> flags accepted per vector store (also the pgvector key order).
_VECTOR_FLAG_KEYS = {"qdrant": ("path", "url"), "pgvector": ("host", "port", "user", "password", "dbname")}
_FLAG_KEYS = ("mode", "api_key", "host", *(f"oss_{s}{k}" for s in ("llm", "embedder") for k in ("", "_key", "_model", "_url")),
"oss_vector", *(f"oss_vector_{k}" for ks in _VECTOR_FLAG_KEYS.values() for k in ks), "user_id")
_FLAG_DEFAULTS = {"oss_llm": "openai", "oss_embedder": "openai", "oss_vector": "qdrant"}
def parse_flags(argv: list[str] | None = None) -> dict[str, str]:
@@ -306,8 +305,7 @@ def _ensure_pgvector(host: str = "localhost", port: int = 5432) -> dict | None:
if not shutil.which("docker"):
print(" Docker not found. Install Docker to auto-start pgvector,\n or run PostgreSQL with pgvector manually.")
return None
# Restart our own container if it exists but is stopped.
try:
with suppress(Exception): # restart our own container if it exists but is stopped
result = _docker("inspect", _PGVECTOR_CONTAINER, "--format", "{{.State.Status}}", timeout=10, text=True, encoding='utf-8', errors='replace')
if result.returncode == 0 and "exited" in result.stdout:
print(f" Found stopped container '{_PGVECTOR_CONTAINER}', restarting...")
@@ -315,8 +313,6 @@ def _ensure_pgvector(host: str = "localhost", port: int = 5432) -> dict | None:
if _pg_ready(host, port, 15):
print(" ✓ PostgreSQL container restarted")
return None
except Exception:
pass
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.")
@@ -538,24 +534,18 @@ def _run_connectivity_checks(oss_config: dict) -> None:
_MODE_HANDLERS = {"oss": _setup_oss, "selfhosted": _setup_selfhosted, "self-hosted": _setup_selfhosted, "platform": _setup_platform}
# Interactive picker order: Platform, Self-hosted server, Open Source.
_MODE_ITEMS = [
("Platform", "Mem0 Cloud API (lightweight, just needs an API key)"),
("Self-hosted server", "Connect to an existing self-hosted Mem0 server (Docker/FastAPI)"),
("Open Source", "Run Mem0 locally (self-hosted LLM + vector store)"),
]
_MODE_ITEMS = [("Platform", "Mem0 Cloud API (lightweight, just needs an API key)"), ("Self-hosted server", "Connect to an existing self-hosted Mem0 server (Docker/FastAPI)"), ("Open Source", "Run Mem0 locally (self-hosted LLM + vector store)")]
_MODE_PICKER = (_setup_platform, _setup_selfhosted, _setup_oss)
def post_setup(hermes_home: str, config: dict) -> None:
"""Entry point for `hermes memory setup`: routes on --mode (platform / selfhosted / oss), else shows a picker.
OSS is non-interactive only when the mode came from the flag."""
try: # mem0ai must meet the minimum version from plugin.yaml
with suppress(Exception): # mem0ai must meet the minimum version from plugin.yaml
import mem0
installed_ver = getattr(mem0, "__version__", None)
if installed_ver and tuple(int(x) for x in installed_ver.split(".")[:3]) < (2, 0, 7):
print(f"\n ⚠ mem0ai {installed_ver} installed but >=2.0.7 required.\n Run: uv pip install --python {sys.executable} 'mem0ai>=2.0.7'")
except Exception:
pass
flags = parse_flags(sys.argv[1:])
handler = _MODE_HANDLERS.get(flags["mode"])
flags["_mode_from_flag"] = handler is not None