From cc4f88aa7666018da730c9450e16b7449770efa2 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 22:20:12 -0700 Subject: [PATCH] refactor(memory/mem0): inline single-use setup helpers, joined summary prints, blank-line squeeze --- plugins/memory/mem0/__init__.py | 48 ++++++--------- plugins/memory/mem0/_backend.py | 6 +- plugins/memory/mem0/_setup.py | 104 +++++++++----------------------- 3 files changed, 48 insertions(+), 110 deletions(-) diff --git a/plugins/memory/mem0/__init__.py b/plugins/memory/mem0/__init__.py index c78c1b9c95..d20660a356 100644 --- a/plugins/memory/mem0/__init__.py +++ b/plugins/memory/mem0/__init__.py @@ -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): diff --git a/plugins/memory/mem0/_backend.py b/plugins/memory/mem0/_backend.py index bb455e622e..c1abd8925a 100644 --- a/plugins/memory/mem0/_backend.py +++ b/plugins/memory/mem0/_backend.py @@ -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. diff --git a/plugins/memory/mem0/_setup.py b/plugins/memory/mem0/_setup.py index 3aabd2be7e..e2dcfce907 100644 --- a/plugins/memory/mem0/_setup.py +++ b/plugins/memory/mem0/_setup.py @@ -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()