refactor(memory/mem0): contextlib.suppress for swallow-all blocks, derived _FLAG_KEYS, packed initialize
This commit is contained in:
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user