diff --git a/plugins/memory/byterover/__init__.py b/plugins/memory/byterover/__init__.py index 55bc1e75f6..6d81225bf5 100644 --- a/plugins/memory/byterover/__init__.py +++ b/plugins/memory/byterover/__init__.py @@ -1,14 +1,9 @@ """ByteRover memory plugin — MemoryProvider interface. -Persistent memory via the ByteRover CLI (``brv``): hierarchical context tree with -tiered retrieval (fuzzy text → LLM-driven search), local-first with optional cloud sync. -Original PR #3499 by hieuntg81, adapted to the MemoryProvider ABC. - -Requires the ``brv`` CLI (npm install -g byterover-cli, or byterover.dev/install.sh). - -Config: BRV_API_KEY env var (cloud features; optional for local), and in config.yaml -``memory.byterover.auto_extract: false`` to disable automatic brv curate hooks. -Working directory: $HERMES_HOME/byterover/ (profile-scoped context tree). +Persistent memory via the ByteRover CLI (``brv``): hierarchical context tree with tiered retrieval +(fuzzy text → LLM-driven search), local-first with optional cloud sync (BRV_API_KEY). Requires the +``brv`` CLI (npm install -g byterover-cli, or byterover.dev/install.sh). Working directory is +$HERMES_HOME/byterover/ (profile-scoped); ``memory.byterover.auto_extract: false`` disables curate hooks. """ from __future__ import annotations @@ -20,7 +15,7 @@ import shutil import subprocess import threading from pathlib import Path -from typing import Any, Callable, Dict, List, Optional +from typing import Any, Dict, List, Optional from agent.memory_provider import MemoryProvider from tools.registry import tool_error @@ -46,22 +41,25 @@ def _load_plugin_config() -> Dict[str, Any]: """Read ``memory.byterover``; fall back to legacy ``memory.provider_config`` (early docs used it).""" try: from hermes_cli.config import load_config - memory_config = load_config().get("memory", {}) - if isinstance(memory_config, dict): - for key in ("byterover", "provider_config"): - block = memory_config.get(key, {}) - if isinstance(block, dict) and (block or key == "provider_config"): - return dict(block) except Exception: - pass + return {} + for key in ("byterover", "provider_config") if isinstance(memory_config, dict) else (): + block = memory_config.get(key, {}) + if isinstance(block, dict) and (block or key == "provider_config"): + return dict(block) return {} -# ── brv binary resolution (cached, thread-safe) ───────────────────────────── +def _get_brv_cwd() -> Path: + """Profile-scoped working directory for the brv context tree.""" + from hermes_constants import get_hermes_home + return get_hermes_home() / "byterover" + +# brv binary resolution (cached, thread-safe): None = unresolved, "" = resolved-missing _brv_path_lock = threading.Lock() -_cached_brv_path: Optional[str] = None # None = unresolved, "" = resolved-missing +_cached_brv_path: Optional[str] = None def _resolve_brv_path() -> Optional[str]: @@ -84,7 +82,6 @@ def _run_brv(args: List[str], timeout: int = _QUERY_TIMEOUT, cwd: str = None) -> brv_path = _resolve_brv_path() if not brv_path: return {"success": False, "error": "brv CLI not found. Install: npm install -g byterover-cli"} - effective_cwd = cwd or str(_get_brv_cwd()) Path(effective_cwd).mkdir(parents=True, exist_ok=True) env = {**os.environ, "PATH": str(Path(brv_path).parent) + os.pathsep + os.environ.get("PATH", "")} @@ -93,10 +90,6 @@ def _run_brv(args: List[str], timeout: int = _QUERY_TIMEOUT, cwd: str = None) -> [brv_path] + args, capture_output=True, text=True, encoding='utf-8', errors='replace', timeout=timeout, cwd=effective_cwd, env=env, stdin=subprocess.DEVNULL, ) - stdout, stderr = result.stdout.strip(), result.stderr.strip() - if result.returncode == 0: - return {"success": True, "output": stdout} - return {"success": False, "error": stderr or stdout or f"brv exited {result.returncode}"} except subprocess.TimeoutExpired: return {"success": False, "error": f"brv timed out after {timeout}s"} except FileNotFoundError: @@ -105,48 +98,35 @@ def _run_brv(args: List[str], timeout: int = _QUERY_TIMEOUT, cwd: str = None) -> return {"success": False, "error": "brv CLI not found"} except Exception as e: return {"success": False, "error": str(e)} + stdout, stderr = result.stdout.strip(), result.stderr.strip() + if result.returncode == 0: + return {"success": True, "output": stdout} + return {"success": False, "error": stderr or stdout or f"brv exited {result.returncode}"} -def _get_brv_cwd() -> Path: - """Profile-scoped working directory for the brv context tree.""" - from hermes_constants import get_hermes_home - return get_hermes_home() / "byterover" - - -# ── Tool schemas ───────────────────────────────────────────────────────────── - def _schema(name: str, description: str, arg: str = "", arg_desc: str = "") -> dict: props = {arg: {"type": "string", "description": arg_desc}} if arg else {} return {"name": name, "description": description, "parameters": {"type": "object", "properties": props, "required": [arg] if arg else []}} QUERY_SCHEMA = _schema( - "brv_query", - "Search ByteRover's persistent knowledge tree for relevant context. Returns memories, project knowledge, " + "brv_query", "Search ByteRover's persistent knowledge tree for relevant context. Returns memories, project knowledge, " "architectural decisions, and patterns from previous sessions. Use for any question where past context would help.", - "query", "What to search for.", -) + "query", "What to search for.") CURATE_SCHEMA = _schema( - "brv_curate", - "Store important information in ByteRover's persistent knowledge tree. Use for architectural decisions, bug fixes, " + "brv_curate", "Store important information in ByteRover's persistent knowledge tree. Use for architectural decisions, bug fixes, " "user preferences, project patterns — anything worth remembering across sessions. ByteRover's LLM automatically " - "categorizes and organizes the memory.", - "content", "The information to remember.", -) + "categorizes and organizes the memory.", "content", "The information to remember.") STATUS_SCHEMA = _schema("brv_status", "Check ByteRover status — CLI version, context tree stats, cloud sync state.") -# ── MemoryProvider implementation ──────────────────────────────────────────── - class ByteRoverMemoryProvider(MemoryProvider): """ByteRover persistent memory via the brv CLI.""" def __init__(self, config: Optional[Dict[str, Any]] = None): self._config = dict(config) if config is not None else _load_plugin_config() self._auto_extract = _coerce_bool(self._config.get("auto_extract"), True) - self._cwd = "" - self._session_id = "" - self._turn_count = 0 + self._cwd, self._session_id, self._turn_count = "", "", 0 self._sync_thread: Optional[threading.Thread] = None @property @@ -166,20 +146,14 @@ class ByteRoverMemoryProvider(MemoryProvider): ] def initialize(self, session_id: str, **kwargs) -> None: - self._cwd = str(_get_brv_cwd()) - self._session_id = session_id - self._turn_count = 0 + self._cwd, self._session_id, self._turn_count = str(_get_brv_cwd()), session_id, 0 Path(self._cwd).mkdir(parents=True, exist_ok=True) def system_prompt_block(self) -> str: if not _resolve_brv_path(): return "" - return ( - "# ByteRover Memory\n" - "Active. Persistent knowledge tree with hierarchical context.\n" - "Use brv_query to search past knowledge, brv_curate to store " - "important facts, brv_status to check state." - ) + return ("# ByteRover Memory\nActive. Persistent knowledge tree with hierarchical context.\n" + "Use brv_query to search past knowledge, brv_curate to store important facts, brv_status to check state.") def _query(self, query: str) -> dict: return _run_brv(["query", "--", query.strip()[:5000]], timeout=_QUERY_TIMEOUT, cwd=self._cwd) @@ -222,78 +196,59 @@ class ByteRoverMemoryProvider(MemoryProvider): self._turn_count += 1 if not self._auto_extract_enabled("sync_turn") or len(user_content.strip()) < _MIN_QUERY_LEN: return - # Wait for the previous sync so curates don't pile up. - if self._sync_thread and self._sync_thread.is_alive(): + if self._sync_thread and self._sync_thread.is_alive(): # wait for the previous sync so curates don't pile up self._sync_thread.join(timeout=5.0) - combined = f"User: {user_content[:2000]}\nAssistant: {assistant_content[:2000]}" - self._sync_thread = self._curate_in_background(combined, name="brv-sync", what="sync") + self._sync_thread = self._curate_in_background(f"User: {user_content[:2000]}\nAssistant: {assistant_content[:2000]}", + name="brv-sync", what="sync") def on_memory_write(self, action: str, target: str, content: str) -> None: """Mirror built-in memory writes to ByteRover.""" - if not self._auto_extract_enabled("memory mirror") or action not in {"add", "replace"} or not content: - return - label = "User profile" if target == "user" else "Agent memory" - self._curate_in_background(f"[{label}] {content}", name="brv-memwrite", what="memory mirror") + if self._auto_extract_enabled("memory mirror") and action in {"add", "replace"} and content: + label = "User profile" if target == "user" else "Agent memory" + self._curate_in_background(f"[{label}] {content}", name="brv-memwrite", what="memory mirror") def on_pre_compress(self, messages: List[Dict[str, Any]]) -> str: """Extract insights from the last 10 user/assistant messages before compression discards them.""" if not self._auto_extract_enabled("pre-compression flush") or not messages: return "" - parts = [] - for msg in messages[-10:]: - role, content = msg.get("role", ""), msg.get("content", "") - if isinstance(content, str) and content.strip() and role in {"user", "assistant"}: - parts.append(f"{role}: {content[:500]}") + parts = [f"{msg.get('role', '')}: {msg.get('content', '')[:500]}" for msg in messages[-10:] + if msg.get("role", "") in {"user", "assistant"} and isinstance(msg.get("content", ""), str) and msg.get("content", "").strip()] if parts: - self._curate_in_background( - "[Pre-compression context]\n" + "\n".join(parts), name="brv-flush", what="pre-compression flush", - on_done=f"ByteRover pre-compression flush: {len(parts)} messages", - ) + self._curate_in_background("[Pre-compression context]\n" + "\n".join(parts), name="brv-flush", what="pre-compression flush", + on_done=f"ByteRover pre-compression flush: {len(parts)} messages") return "" def get_tool_schemas(self) -> List[Dict[str, Any]]: return [QUERY_SCHEMA, CURATE_SCHEMA, STATUS_SCHEMA] def handle_tool_call(self, tool_name: str, args: dict, **kwargs) -> str: - handler = self._TOOLS.get(tool_name) - if handler is None: + if tool_name not in _TOOLS: return tool_error(f"Unknown tool: {tool_name}") - return handler(self, args) + arg, fail_msg, run, on_ok = _TOOLS[tool_name] + value = args.get(arg, "") if arg else None + if arg and not value: + return tool_error(f"{arg} is required") + result = run(self, value) + return json.dumps(on_ok(result.get("output", ""))) if result["success"] else tool_error(result.get("error", fail_msg)) def shutdown(self) -> None: if self._sync_thread and self._sync_thread.is_alive(): self._sync_thread.join(timeout=10.0) - # Tool implementations - @staticmethod - def _tool_result(result: dict, fail_msg: str, on_ok: Callable[[str], dict]) -> str: - if not result["success"]: - return tool_error(result.get("error", fail_msg)) - return json.dumps(on_ok(result.get("output", ""))) +def _format_query_output(output: str) -> dict: + output = output.strip() + if len(output) < _MIN_OUTPUT_LEN: + return {"result": "No relevant memories found."} + return {"result": output[:8000] + "\n\n[... truncated]" if len(output) > 8000 else output} - def _tool_query(self, args: dict) -> str: - query = args.get("query", "") - if not query: - return tool_error("query is required") - def fmt(output: str) -> dict: - output = output.strip() - if len(output) < _MIN_OUTPUT_LEN: - return {"result": "No relevant memories found."} - return {"result": output[:8000] + "\n\n[... truncated]" if len(output) > 8000 else output} - return self._tool_result(self._query(query), "Query failed", fmt) - - def _tool_curate(self, args: dict) -> str: - content = args.get("content", "") - if not content: - return tool_error("content is required") - return self._tool_result(self._curate(content), "Curate failed", lambda _: {"result": "Memory curated successfully."}) - - def _tool_status(self, args: dict) -> str: - return self._tool_result(_run_brv(["status"], timeout=15, cwd=self._cwd), "Status check failed", lambda out: {"status": out}) - - _TOOLS = {"brv_query": _tool_query, "brv_curate": _tool_curate, "brv_status": _tool_status} +# tool name -> (required arg or None, failure message, run(provider, arg_value) -> brv result, on_ok(output) -> JSON payload) +_TOOLS = { + "brv_query": ("query", "Query failed", lambda p, q: p._query(q), _format_query_output), + "brv_curate": ("content", "Curate failed", lambda p, c: p._curate(c), lambda _: {"result": "Memory curated successfully."}), + "brv_status": (None, "Status check failed", lambda p, _: _run_brv(["status"], timeout=15, cwd=p._cwd), lambda out: {"status": out}), +} def register(ctx) -> None: diff --git a/plugins/memory/holographic/__init__.py b/plugins/memory/holographic/__init__.py index 08f799005f..8ff1d0ce9a 100644 --- a/plugins/memory/holographic/__init__.py +++ b/plugins/memory/holographic/__init__.py @@ -1,19 +1,8 @@ -"""hermes-memory-store — holographic memory plugin using MemoryProvider interface. - -Registers as a MemoryProvider plugin, giving the agent structured fact storage -with entity resolution, trust scoring, and HRR-based compositional retrieval. - -Original plugin by dusterbloom (PR #2351), adapted to the MemoryProvider ABC. - -Config in $HERMES_HOME/config.yaml (profile-scoped): - plugins: - hermes-memory-store: - db_path: $HERMES_HOME/memory_store.db # omit to use the default - auto_extract: false - default_trust: 0.5 - min_trust_threshold: 0.3 - temporal_decay_half_life: 0 -""" +"""hermes-memory-store — holographic memory plugin (MemoryProvider): structured fact storage with entity +resolution, trust scoring, and HRR-based compositional retrieval. Original plugin by dusterbloom (PR #2351). +Config in $HERMES_HOME/config.yaml under plugins.hermes-memory-store: db_path ($HERMES_HOME/memory_store.db), +auto_extract (false), default_trust (0.5), min_trust_threshold (0.3), temporal_decay_half_life (0), +hrr_dim (1024), hrr_weight (0.3).""" from __future__ import annotations @@ -36,26 +25,18 @@ logger = logging.getLogger(__name__) FACT_STORE_SCHEMA = { "name": "fact_store", "description": ( - "Deep structured memory with algebraic reasoning. " - "Use alongside the memory tool — memory for always-on context, " - "fact_store for deep recall and compositional queries.\n\n" - "ACTIONS (simple → powerful):\n" - "• add — Store a fact the user would expect you to remember.\n" - "• search — Keyword lookup ('editor config', 'deploy process').\n" - "• probe — Entity recall: ALL facts about a person/thing.\n" - "• related — What connects to an entity? Structural adjacency.\n" + "Deep structured memory with algebraic reasoning. Use alongside the memory tool — memory for always-on " + "context, fact_store for deep recall and compositional queries.\n\nACTIONS (simple → powerful):\n" + "• add — Store a fact the user would expect you to remember.\n• search — Keyword lookup ('editor config', 'deploy process').\n" + "• probe — Entity recall: ALL facts about a person/thing.\n• related — What connects to an entity? Structural adjacency.\n" "• reason — Compositional: facts connected to MULTIPLE entities simultaneously.\n" - "• contradict — Memory hygiene: find facts making conflicting claims.\n" - "• update/remove/list — CRUD operations.\n\n" + "• contradict — Memory hygiene: find facts making conflicting claims.\n• update/remove/list — CRUD operations.\n\n" "IMPORTANT: Before answering questions about the user, ALWAYS probe or reason first." ), "parameters": { "type": "object", "properties": { - "action": { - "type": "string", - "enum": ["add", "search", "probe", "related", "reason", "contradict", "update", "remove", "list"], - }, + "action": {"type": "string", "enum": ["add", "search", "probe", "related", "reason", "contradict", "update", "remove", "list"]}, "content": {"type": "string", "description": "Fact content (required for 'add')."}, "query": {"type": "string", "description": "Search query (required for 'search')."}, "entity": {"type": "string", "description": "Entity name for 'probe'/'related'."}, @@ -73,44 +54,46 @@ FACT_STORE_SCHEMA = { FACT_FEEDBACK_SCHEMA = { "name": "fact_feedback", - "description": ( - "Rate a fact after using it. Mark 'helpful' if accurate, 'unhelpful' if outdated. " - "This trains the memory — good facts rise, bad facts sink." - ), + "description": ("Rate a fact after using it. Mark 'helpful' if accurate, 'unhelpful' if outdated. " + "This trains the memory — good facts rise, bad facts sink."), "parameters": { "type": "object", - "properties": { - "action": {"type": "string", "enum": ["helpful", "unhelpful"]}, - "fact_id": {"type": "integer", "description": "The fact ID to rate."}, - }, + "properties": {"action": {"type": "string", "enum": ["helpful", "unhelpful"]}, + "fact_id": {"type": "integer", "description": "The fact ID to rate."}}, "required": ["action", "fact_id"], }, } -# Auto-extraction patterns (on_session_end): user preferences -> user_pref, decisions -> project. -_PREF_PATTERNS = [ - re.compile(r'\bI\s+(?:prefer|like|love|use|want|need)\s+(.+)', re.IGNORECASE), - re.compile(r'\bmy\s+(?:favorite|preferred|default)\s+\w+\s+is\s+(.+)', re.IGNORECASE), - re.compile(r'\bI\s+(?:always|never|usually)\s+(.+)', re.IGNORECASE), -] -_DECISION_PATTERNS = [ - re.compile(r'\bwe\s+(?:decided|agreed|chose)\s+(?:to\s+)?(.+)', re.IGNORECASE), - re.compile(r'\bthe\s+project\s+(?:uses|needs|requires)\s+(.+)', re.IGNORECASE), -] +# Auto-extraction (on_session_end): (patterns, category) — user preferences -> user_pref, decisions -> project. +_EXTRACT_CATEGORIES = ( + ([re.compile(r'\bI\s+(?:prefer|like|love|use|want|need)\s+(.+)', re.IGNORECASE), + re.compile(r'\bmy\s+(?:favorite|preferred|default)\s+\w+\s+is\s+(.+)', re.IGNORECASE), + re.compile(r'\bI\s+(?:always|never|usually)\s+(.+)', re.IGNORECASE)], "user_pref"), + ([re.compile(r'\bwe\s+(?:decided|agreed|chose)\s+(?:to\s+)?(.+)', re.IGNORECASE), + re.compile(r'\bthe\s+project\s+(?:uses|needs|requires)\s+(.+)', re.IGNORECASE)], "project"), +) def _load_plugin_config() -> dict: try: - # Canonical loader: honors the managed-scope overlay + ${VAR} expansion. - from hermes_cli.config import load_config_readonly - all_config = load_config_readonly() - return cfg_get(all_config, "plugins", "hermes-memory-store", default={}) or {} + from hermes_cli.config import load_config_readonly # canonical: managed-scope overlay + ${VAR} expansion + return cfg_get(load_config_readonly(), "plugins", "hermes-memory-store", default={}) or {} except Exception: return {} -def _results(results: list) -> str: - return json.dumps({"results": results, "count": len(results)}) +def _results(items: list, key: str = "results") -> str: + return json.dumps({key: items, "count": len(items)}) + + +def _limit(args: dict) -> int: + return int(args.get("limit", 10)) + + +def _tool_handler(actions: dict): + """(self, args) handler dispatching on args["action"] over ``actions``; unknown action -> tool_error.""" + return lambda self, args: (actions[args["action"]](self, args) if args["action"] in actions + else tool_error(f"Unknown action: {args['action']}")) class HolographicMemoryProvider(MemoryProvider): @@ -118,8 +101,7 @@ class HolographicMemoryProvider(MemoryProvider): def __init__(self, config: dict | None = None): self._config = config or _load_plugin_config() - self._store = None - self._retriever = None + self._store = self._retriever = None self._min_trust = float(self._config.get("min_trust_threshold", 0.3)) @property @@ -134,9 +116,7 @@ class HolographicMemoryProvider(MemoryProvider): config_path = Path(hermes_home) / "config.yaml" try: import yaml - # Raw read for the write-back round-trip: merged defaults must not - # be persisted into the user's file. - from hermes_cli.config import read_user_config_raw + from hermes_cli.config import read_user_config_raw # raw read: merged defaults must not be persisted existing = read_user_config_raw(config_path) existing.setdefault("plugins", {})["hermes-memory-store"] = values with open(config_path, "w", encoding="utf-8") as f: @@ -146,9 +126,8 @@ class HolographicMemoryProvider(MemoryProvider): def get_config_schema(self): from hermes_constants import display_hermes_home - _default_db = f"{display_hermes_home()}/memory_store.db" return [ - {"key": "db_path", "description": "SQLite database path", "default": _default_db}, + {"key": "db_path", "description": "SQLite database path", "default": f"{display_hermes_home()}/memory_store.db"}, {"key": "auto_extract", "description": "Auto-extract facts at session end", "default": "false", "choices": ["true", "false"]}, {"key": "default_trust", "description": "Default trust score for new facts", "default": "0.5"}, {"key": "hrr_dim", "description": "HRR vector dimensions", "default": "1024"}, @@ -158,20 +137,12 @@ class HolographicMemoryProvider(MemoryProvider): from hermes_constants import get_hermes_home _hermes_home = str(get_hermes_home()) db_path = self._config.get("db_path", _hermes_home + "/memory_store.db") - # Expand $HERMES_HOME so configured paths resolve to the active profile's directory. - if isinstance(db_path, str): + if isinstance(db_path, str): # expand $HERMES_HOME so paths resolve to the active profile db_path = db_path.replace("$HERMES_HOME", _hermes_home).replace("${HERMES_HOME}", _hermes_home) hrr_dim = int(self._config.get("hrr_dim", 1024)) - - self._store = MemoryStore( - db_path=db_path, default_trust=float(self._config.get("default_trust", 0.5)), hrr_dim=hrr_dim, - ) - self._retriever = FactRetriever( - store=self._store, - temporal_decay_half_life=int(self._config.get("temporal_decay_half_life", 0)), - hrr_weight=float(self._config.get("hrr_weight", 0.3)), - hrr_dim=hrr_dim, - ) + self._store = MemoryStore(db_path=db_path, default_trust=float(self._config.get("default_trust", 0.5)), hrr_dim=hrr_dim) + self._retriever = FactRetriever(store=self._store, hrr_dim=hrr_dim, hrr_weight=float(self._config.get("hrr_weight", 0.3)), + temporal_decay_half_life=int(self._config.get("temporal_decay_half_life", 0))) self._session_id = session_id def system_prompt_block(self) -> str: @@ -181,54 +152,39 @@ class HolographicMemoryProvider(MemoryProvider): total = self._store._conn.execute("SELECT COUNT(*) FROM facts").fetchone()[0] except Exception: total = 0 - if total == 0: - return ( - "# Holographic Memory\n" - "Active. Empty fact store — proactively add facts the user would expect you to remember.\n" + body = ("Active. Empty fact store — proactively add facts the user would expect you to remember.\n" "Use fact_store(action='add') to store durable structured facts about people, projects, preferences, decisions.\n" - "Use fact_feedback to rate facts after using them (trains trust scores)." - ) - return ( - f"# Holographic Memory\n" - f"Active. {total} facts stored with entity resolution and trust scoring.\n" - f"Use fact_store to search, probe entities, reason across entities, or add facts.\n" - f"Use fact_feedback to rate facts after using them (trains trust scores)." - ) + if total == 0 else + f"Active. {total} facts stored with entity resolution and trust scoring.\n" + "Use fact_store to search, probe entities, reason across entities, or add facts.\n") + return "# Holographic Memory\n" + body + "Use fact_feedback to rate facts after using them (trains trust scores)." def prefetch(self, query: str, *, session_id: str = "") -> str: if not self._retriever or not query: return "" try: results = self._retriever.search(query, min_trust=self._min_trust, limit=5) - if not results: - return "" lines = [f"- [{r.get('trust_score', r.get('trust', 0)):.1f}] {r.get('content', '')}" for r in results] - return "## Holographic Memory\n" + "\n".join(lines) + return "## Holographic Memory\n" + "\n".join(lines) if results else "" except Exception as e: logger.debug("Holographic prefetch failed: %s", e) return "" - def sync_turn(self, user_content: str, assistant_content: str, *, session_id: str = "") -> None: - # Facts are stored explicitly via tools; on_session_end handles auto-extraction if configured. - pass - def get_tool_schemas(self) -> List[Dict[str, Any]]: return [FACT_STORE_SCHEMA, FACT_FEEDBACK_SCHEMA] def handle_tool_call(self, tool_name: str, args: Dict[str, Any], **kwargs) -> str: - handler = self._TOOL_HANDLERS.get(tool_name) - if handler is None: + if tool_name not in self._TOOL_HANDLERS: return tool_error(f"Unknown tool: {tool_name}") try: - return handler(self, args) + return self._TOOL_HANDLERS[tool_name](self, args) except KeyError as exc: return tool_error(f"Missing required argument: {exc}") except Exception as exc: return tool_error(str(exc)) def on_session_end(self, messages: List[Dict[str, Any]]) -> None: - # is_truthy_value: the config schema declares auto_extract as a string - # enum ("false"/"true"); plain truthiness would treat "false" as enabled. + # is_truthy_value: auto_extract is a string enum ("false"/"true"); plain truthiness would treat "false" as on. if is_truthy_value(self._config.get("auto_extract", False)) and self._store and messages: self._auto_extract_facts(messages) @@ -236,134 +192,73 @@ class HolographicMemoryProvider(MemoryProvider): """Mirror built-in memory writes as facts.""" if action == "add" and self._store and content: try: - category = "user_pref" if target == "user" else "general" - self._store.add_fact(content, category=category) + self._store.add_fact(content, category="user_pref" if target == "user" else "general") except Exception as e: logger.debug("Holographic memory_write mirror failed: %s", e) def shutdown(self) -> None: - # Release the shared SQLite connection on the caller's thread: leaving - # it to GC keeps the connection (and its write lock) alive on a - # long-running gateway. close() is idempotent and refcount-guarded. + # Close on the caller's thread: leaving the shared connection (+ write lock) to GC keeps it alive on a gateway. if self._store is not None: try: self._store.close() except Exception as e: logger.debug("Holographic shutdown close() failed: %s", e) - self._store = None - self._retriever = None + self._store = self._retriever = None - # -- Tool handlers ------------------------------------------------------- - # KeyError from args[...] / Exception are turned into tool_error by handle_tool_call. + # Tool handlers (self, args) -> str. KeyError from args[...] / Exception -> tool_error in handle_tool_call; + # argument coercion order (and therefore which error surfaces first) mirrors the underlying call order. - def _handle_fact_store(self, args: dict) -> str: - action = args["action"] - handler = self._FACT_STORE_ACTIONS.get(action) - if handler is None: - return tool_error(f"Unknown action: {action}") - return handler(self, args) + def _entity_query(self, method: str, a: dict) -> str: + """'probe' / 'related': single-entity retriever queries.""" + return _results(getattr(self._retriever, method)(a["entity"], category=a.get("category"), limit=_limit(a))) - def _act_add(self, args: dict) -> str: - fact_id = self._store.add_fact(args["content"], category=args.get("category", "general"), tags=args.get("tags", "")) - return json.dumps({"fact_id": fact_id, "status": "added"}) - - def _act_search(self, args: dict) -> str: - return _results(self._retriever.search( - args["query"], category=args.get("category"), - min_trust=float(args.get("min_trust", self._min_trust)), limit=int(args.get("limit", 10)), - )) - - def _act_entity(self, args: dict, method: str) -> str: - """Shared body of 'probe' and 'related' (single-entity retriever queries).""" - return _results(getattr(self._retriever, method)( - args["entity"], category=args.get("category"), limit=int(args.get("limit", 10)), - )) - - def _act_reason(self, args: dict) -> str: - entities = args.get("entities", []) - if not entities: - return tool_error("reason requires 'entities' list") - return _results(self._retriever.reason(entities, category=args.get("category"), limit=int(args.get("limit", 10)))) - - def _act_contradict(self, args: dict) -> str: - return _results(self._retriever.contradict(category=args.get("category"), limit=int(args.get("limit", 10)))) - - def _act_update(self, args: dict) -> str: - updated = self._store.update_fact( - int(args["fact_id"]), content=args.get("content"), - trust_delta=float(args["trust_delta"]) if "trust_delta" in args else None, - tags=args.get("tags"), category=args.get("category"), - ) - return json.dumps({"updated": updated}) - - def _act_remove(self, args: dict) -> str: - return json.dumps({"removed": self._store.remove_fact(int(args["fact_id"]))}) - - def _act_list(self, args: dict) -> str: - facts = self._store.list_facts( - category=args.get("category"), min_trust=float(args.get("min_trust", 0.0)), limit=int(args.get("limit", 10)), - ) - return json.dumps({"facts": facts, "count": len(facts)}) - - def _handle_fact_feedback(self, args: dict) -> str: - return json.dumps(self._store.record_feedback(int(args["fact_id"]), helpful=args["action"] == "helpful")) - - _FACT_STORE_ACTIONS = { - "add": _act_add, "search": _act_search, - "probe": lambda self, args: self._act_entity(args, "probe"), - "related": lambda self, args: self._act_entity(args, "related"), - "reason": _act_reason, "contradict": _act_contradict, - "update": _act_update, "remove": _act_remove, "list": _act_list, + _TOOL_HANDLERS = { + "fact_store": _tool_handler({ + "add": lambda self, a: json.dumps({"fact_id": self._store.add_fact( + a["content"], category=a.get("category", "general"), tags=a.get("tags", "")), "status": "added"}), + "search": lambda self, a: _results(self._retriever.search( + a["query"], category=a.get("category"), min_trust=float(a.get("min_trust", self._min_trust)), limit=_limit(a))), + "probe": lambda self, a: self._entity_query("probe", a), + "related": lambda self, a: self._entity_query("related", a), + "reason": lambda self, a: _results(self._retriever.reason(a["entities"], category=a.get("category"), limit=_limit(a))) + if a.get("entities") else tool_error("reason requires 'entities' list"), + "contradict": lambda self, a: _results(self._retriever.contradict(category=a.get("category"), limit=_limit(a))), + "update": lambda self, a: json.dumps({"updated": self._store.update_fact( + int(a["fact_id"]), content=a.get("content"), trust_delta=float(a["trust_delta"]) if "trust_delta" in a else None, + tags=a.get("tags"), category=a.get("category"))}), + "remove": lambda self, a: json.dumps({"removed": self._store.remove_fact(int(a["fact_id"]))}), + "list": lambda self, a: _results(self._store.list_facts( + category=a.get("category"), min_trust=float(a.get("min_trust", 0.0)), limit=_limit(a)), key="facts"), + }), + "fact_feedback": lambda self, a: json.dumps(self._store.record_feedback(int(a["fact_id"]), helpful=a["action"] == "helpful")), } - _TOOL_HANDLERS = {"fact_store": _handle_fact_store, "fact_feedback": _handle_fact_feedback} - - # -- Auto-extraction (on_session_end) ------------------------------------ def _auto_extract_facts(self, messages: list) -> None: - # Local import: the compressor module is heavier than this plugin and - # only needed when auto_extract is on. - from agent.context_compressor import ( - _MERGED_PRIOR_CONTEXT_HEADER, - _MERGED_SUMMARY_DELIMITER, - is_compaction_summary_message, - ) - + # Compaction handoff summaries arrive as role="user" and match the decision patterns; never store the + # compactor's own output as a fact. A merge-into-tail row holds genuine prior user text BEFORE + # _MERGED_SUMMARY_DELIMITER (after the header) and the summary AFTER it — harvest only that segment. + from agent.context_compressor import _MERGED_PRIOR_CONTEXT_HEADER, _MERGED_SUMMARY_DELIMITER, is_compaction_summary_message # heavy; lazy extracted = 0 for msg in messages: - if msg.get("role") != "user": - continue - content = msg.get("content", "") - # Compaction handoff summaries arrive as role="user" and reliably - # match the decision patterns; skip them so the compactor's own - # output is never stored as a durable fact. A merge-into-tail row - # holds genuine prior user text BEFORE _MERGED_SUMMARY_DELIMITER - # (prefixed with the header) and the summary AFTER it — harvest - # only the pre-delimiter segment. - if isinstance(content, str) and _MERGED_SUMMARY_DELIMITER in content: - pre = content.split(_MERGED_SUMMARY_DELIMITER, 1)[0] - if pre.startswith(_MERGED_PRIOR_CONTEXT_HEADER): - pre = pre[len(_MERGED_PRIOR_CONTEXT_HEADER):] - if pre.strip(): - content = pre.strip() - elif is_compaction_summary_message(msg): - continue - elif is_compaction_summary_message(msg): + content = msg.get("content", "") if msg.get("role") == "user" else None + pre = content.split(_MERGED_SUMMARY_DELIMITER, 1)[0].removeprefix(_MERGED_PRIOR_CONTEXT_HEADER).strip() \ + if isinstance(content, str) and _MERGED_SUMMARY_DELIMITER in content else "" + if pre: + content = pre + elif content is None or is_compaction_summary_message(msg): continue if not isinstance(content, str) or len(content) < 10: continue - - for patterns, category in ((_PREF_PATTERNS, "user_pref"), (_DECISION_PATTERNS, "project")): + for patterns, category in _EXTRACT_CATEGORIES: if any(p.search(content) for p in patterns): try: self._store.add_fact(content[:400], category=category) extracted += 1 except Exception: pass - if extracted: logger.info("Auto-extracted %d facts from conversation", extracted) - def register(ctx) -> None: """Register the holographic memory provider with the plugin system.""" ctx.register_memory_provider(HolographicMemoryProvider(config=_load_plugin_config())) diff --git a/plugins/memory/holographic/holographic.py b/plugins/memory/holographic/holographic.py index c6acf0fbef..05176df2ec 100644 --- a/plugins/memory/holographic/holographic.py +++ b/plugins/memory/holographic/holographic.py @@ -1,21 +1,13 @@ -"""Holographic Reduced Representations (HRR) with phase encoding. - -Each concept is a vector of angles in [0, 2π). Operations: - bind — circular convolution (phase addition) — associates two concepts - unbind — circular correlation (phase subtraction) — retrieves a bound value - bundle — superposition (circular mean) — merges multiple concepts - -Phase encoding avoids the magnitude collapse of complex-number HRRs and maps -cleanly to cosine similarity. Atoms derive deterministically from SHA-256 so -representations are identical across processes, machines, and Python versions. - -References: Plate (1995) HRRs; Gayler (2004) Vector Symbolic Architectures. -""" +"""Holographic Reduced Representations (HRR) with phase encoding. Each concept is a vector of angles in [0, 2π): +bind = circular convolution (phase addition), unbind = circular correlation (phase subtraction), bundle = +superposition (circular mean). Phase encoding avoids the magnitude collapse of complex-number HRRs and maps cleanly +to cosine similarity; atoms derive deterministically from SHA-256 so representations are identical across processes, +machines, and Python versions. References: Plate (1995) HRRs; Gayler (2004) Vector Symbolic Architectures.""" import hashlib import logging -import struct import math +import struct try: import numpy as np @@ -28,6 +20,8 @@ logger = logging.getLogger(__name__) _TWO_PI = 2.0 * math.pi _FLOAT32_BLOB_PREFIX = b"HRR1" +_F32, _F64 = 4, 8 # itemsizes of np.float32 / np.float64 +ROLE_CONTENT, ROLE_ENTITY = "__hrr_role_content__", "__hrr_role_entity__" # role atoms used by encode_fact def _require_numpy() -> None: @@ -37,14 +31,10 @@ def _require_numpy() -> None: def encode_atom(word: str, dim: int = 1024) -> "np.ndarray": """Deterministic phase vector: SHA-256 counter blocks of f"{word}:{i}" -> uint16 -> [0, 2π). - - hashlib rather than numpy RNG so atoms are reproducible across platforms. - """ + hashlib rather than numpy RNG so atoms are reproducible across platforms.""" _require_numpy() - uint16_values: list[int] = [] - for i in range(math.ceil(dim / 16)): # 32-byte digest = 16 uint16 values - digest = hashlib.sha256(f"{word}:{i}".encode()).digest() - uint16_values.extend(struct.unpack("<16H", digest)) + uint16_values = [v for i in range(math.ceil(dim / 16)) # 32-byte digest = 16 uint16 values + for v in struct.unpack("<16H", hashlib.sha256(f"{word}:{i}".encode()).digest())] return np.array(uint16_values[:dim], dtype=np.float64) * (_TWO_PI / 65536.0) @@ -63,8 +53,7 @@ def unbind(memory: "np.ndarray", key: "np.ndarray") -> "np.ndarray": def bundle(*vectors: "np.ndarray") -> "np.ndarray": """Superposition via circular mean; holds O(sqrt(dim)) items before similarity degrades.""" _require_numpy() - complex_sum = np.sum([np.exp(1j * v) for v in vectors], axis=0) - return np.angle(complex_sum) % _TWO_PI + return np.angle(np.sum([np.exp(1j * v) for v in vectors], axis=0)) % _TWO_PI def similarity(a: "np.ndarray", b: "np.ndarray") -> float: @@ -77,89 +66,60 @@ def encode_text(text: str, dim: int = 1024) -> "np.ndarray": """Bag-of-words bundle of token atoms; empty text -> encode_atom("__hrr_empty__").""" _require_numpy() tokens = [t for t in (tok.strip(".,!?;:\"'()[]{}") for tok in text.lower().split()) if t] - if not tokens: - return encode_atom("__hrr_empty__", dim) - return bundle(*[encode_atom(token, dim) for token in tokens]) + return bundle(*[encode_atom(token, dim) for token in tokens]) if tokens else encode_atom("__hrr_empty__", dim) def encode_fact(content: str, entities: list[str], dim: int = 1024) -> "np.ndarray": """bundle(bind(text, ROLE_CONTENT), bind(entity_i, ROLE_ENTITY)...), so unbind(fact, bind(entity, ROLE_ENTITY)) ≈ content_vector.""" _require_numpy() - role_content = encode_atom("__hrr_role_content__", dim) - role_entity = encode_atom("__hrr_role_entity__", dim) - components = [bind(encode_text(content, dim), role_content)] - for entity in entities: - components.append(bind(encode_atom(entity.lower(), dim), role_entity)) - return bundle(*components) + role_content, role_entity = encode_atom(ROLE_CONTENT, dim), encode_atom(ROLE_ENTITY, dim) + return bundle(bind(encode_text(content, dim), role_content), + *[bind(encode_atom(entity.lower(), dim), role_entity) for entity in entities]) def phases_to_bytes(phases: "np.ndarray", dim: int | None = None) -> bytes: - """Serialize as prefixed float32 (half the size of legacy float64 blobs). - - At dim=1 the prefixed float32 blob and a raw float64 blob are both 8 bytes, - so we write legacy float64 there to keep ``bytes_to_phases`` unambiguous. - """ + """Serialize as prefixed float32 (half the size of legacy float64 blobs). At dim=1 the prefixed float32 blob + and a raw float64 blob are both 8 bytes, so write legacy float64 there to keep ``bytes_to_phases`` unambiguous.""" _require_numpy() - if dim is None: - dim = int(phases.shape[0]) - if len(_FLOAT32_BLOB_PREFIX) + dim * np.dtype(np.float32).itemsize == dim * np.dtype(np.float64).itemsize: + dim = int(phases.shape[0]) if dim is None else dim + if len(_FLOAT32_BLOB_PREFIX) + dim * _F32 == dim * _F64: return np.asarray(phases, dtype=np.float64).tobytes() return _FLOAT32_BLOB_PREFIX + np.asarray(phases, dtype=np.float32).tobytes() def bytes_to_phases(data: bytes, dim: int | None = None) -> "np.ndarray": - """Deserialize prefixed float32 or legacy raw float64 blobs (always returns float64). - - With ``dim`` given, a prefixed blob whose size equals the float64 size - (dim=1) is read as legacy float64: ``phases_to_bytes`` never writes a - prefixed blob at that size, so such a blob must be legacy data. - """ + """Deserialize prefixed float32 or legacy raw float64 blobs (always returns float64). With ``dim`` given, a + prefixed blob whose size equals the float64 size (dim=1) is read as legacy float64: ``phases_to_bytes`` never + writes a prefixed blob at that size.""" _require_numpy() - f32, f64 = np.dtype(np.float32).itemsize, np.dtype(np.float64).itemsize plen = len(_FLOAT32_BLOB_PREFIX) prefixed = data.startswith(_FLOAT32_BLOB_PREFIX) - + f32 = lambda payload: np.frombuffer(payload, dtype=np.float32).astype(np.float64) # noqa: E731 + f64 = lambda payload: np.frombuffer(payload, dtype=np.float64).copy() # noqa: E731 if dim is None: - if prefixed: - payload = data[plen:] - if len(payload) % f32 != 0: - raise ValueError(f"HRR float32 vector blob has invalid payload byte length: {len(payload)}") - return np.frombuffer(payload, dtype=np.float32).astype(np.float64) - if len(data) % f64 != 0: - raise ValueError(f"HRR legacy vector blob has invalid byte length: {len(data)}") - return np.frombuffer(data, dtype=np.float64).copy() - - float32_blob_bytes = plen + dim * f32 - float64_bytes = dim * f64 + payload, size, what = (data[plen:], _F32, "float32 vector blob has invalid payload") if prefixed else (data, _F64, "legacy vector blob has invalid") + if len(payload) % size != 0: + raise ValueError(f"HRR {what} byte length: {len(payload)}") + return f32(payload) if prefixed else f64(payload) + float32_blob_bytes, float64_bytes = plen + dim * _F32, dim * _F64 collides = float32_blob_bytes == float64_bytes if not collides and prefixed and len(data) == float32_blob_bytes: - return np.frombuffer(data[plen:], dtype=np.float32).astype(np.float64) + return f32(data[plen:]) if len(data) == float64_bytes: - return np.frombuffer(data, dtype=np.float64).copy() - if prefixed: - expected = (f"{float64_bytes} (legacy float64)" if collides - else f"{float32_blob_bytes} (prefixed float32) or {float64_bytes} (legacy float64)") - raise ValueError( - f"HRR vector blob has {len(data)} bytes ({len(data) - plen} payload bytes after " - f"the float32 prefix); expected {expected} for dim={dim}" - ) - raise ValueError( - f"HRR legacy vector blob has {len(data)} bytes; expected " - f"{float64_bytes} (float64) for dim={dim}" - ) + return f64(data) + if not prefixed: + raise ValueError(f"HRR legacy vector blob has {len(data)} bytes; expected {float64_bytes} (float64) for dim={dim}") + expected = f"{float64_bytes} (legacy float64)" if collides else f"{float32_blob_bytes} (prefixed float32) or {float64_bytes} (legacy float64)" + raise ValueError(f"HRR vector blob has {len(data)} bytes ({len(data) - plen} payload bytes after the float32 prefix); " + f"expected {expected} for dim={dim}") def snr_estimate(dim: int, n_items: int) -> float: """SNR = sqrt(dim / n_items) (inf when empty); warns below 2.0 (n_items > dim/4).""" _require_numpy() - if n_items <= 0: - return float("inf") - snr = math.sqrt(dim / n_items) + snr = math.sqrt(dim / n_items) if n_items > 0 else float("inf") if snr < 2.0: - logger.warning( - "HRR storage near capacity: SNR=%.2f (dim=%d, n_items=%d). " - "Retrieval accuracy may degrade. Consider increasing dim or reducing stored items.", - snr, dim, n_items, - ) + logger.warning("HRR storage near capacity: SNR=%.2f (dim=%d, n_items=%d). " + "Retrieval accuracy may degrade. Consider increasing dim or reducing stored items.", snr, dim, n_items) return snr diff --git a/plugins/memory/holographic/retrieval.py b/plugins/memory/holographic/retrieval.py index a8a3fa86f1..1383e942bc 100644 --- a/plugins/memory/holographic/retrieval.py +++ b/plugins/memory/holographic/retrieval.py @@ -1,8 +1,5 @@ -"""Hybrid keyword/BM25 retrieval for the memory store. - -Ported from KIK memory_agent.py — combines FTS5 full-text search with -Jaccard similarity reranking and trust-weighted scoring. -""" +"""Hybrid keyword/BM25 retrieval for the memory store: FTS5 candidates reranked with +Jaccard similarity and HRR vector similarity, trust-weighted (ported from KIK memory_agent.py).""" from __future__ import annotations @@ -14,264 +11,173 @@ from typing import TYPE_CHECKING if TYPE_CHECKING: from .store import MemoryStore -try: - from . import holographic as hrr -except ImportError: - import holographic as hrr # type: ignore[no-redef] +from . import holographic as hrr -_FACT_COLUMNS = ( - "fact_id, content, category, tags, trust_score, " - "retrieval_count, helpful_count, created_at, updated_at" -) +_FACT_COLUMNS = "fact_id, content, category, tags, trust_score, retrieval_count, helpful_count, created_at, updated_at" +_ROLE_ENTITY, _ROLE_CONTENT = hrr.ROLE_ENTITY, hrr.ROLE_CONTENT +_PUNCT = ".,;:!?\"'()[]{}#@<>" +_FTS_OPERATORS = str.maketrans("", "", '"()*^:-+') +# Stopwords dropped before FTS5 OR-expansion: short English function words that +# carry no retrieval signal and force false-negative AND matches. +_FTS_STOPWORDS = frozenset(""" + a about above after again all am an and any are as at be because been before being between both but by can could + did do does doing don down during each few for from further had has have having he her here hers herself him himself + his how i if in into is it its itself just me more most my myself no nor not now of off on once only or other our + ours ourselves out over own same she should so some such than that the their theirs them themselves then there these + they this those through to too under until up very was we were what when where which while who whom why will with + would you your yours yourself yourselves""".split()) + + +def _shift(sim: float) -> float: + """Cosine similarity [-1, 1] -> [0, 1].""" + return (sim + 1.0) / 2.0 class FactRetriever: """Multi-strategy fact retrieval with trust-weighted scoring.""" - def __init__( - self, - store: MemoryStore, - temporal_decay_half_life: int = 0, # days, 0 = disabled - fts_weight: float = 0.4, jaccard_weight: float = 0.3, hrr_weight: float = 0.3, - hrr_dim: int = 1024, - ): - self.store = store - self.half_life = temporal_decay_half_life - self.hrr_dim = hrr_dim - - # Auto-redistribute weights if numpy unavailable - if hrr_weight > 0 and not hrr._HAS_NUMPY: + def __init__(self, store: MemoryStore, temporal_decay_half_life: int = 0, # days, 0 = disabled + fts_weight: float = 0.4, jaccard_weight: float = 0.3, hrr_weight: float = 0.3, hrr_dim: int = 1024): + self.store, self.half_life, self.hrr_dim = store, temporal_decay_half_life, hrr_dim + if hrr_weight > 0 and not hrr._HAS_NUMPY: # redistribute weights without numpy fts_weight, jaccard_weight, hrr_weight = 0.6, 0.4, 0.0 self.fts_weight, self.jaccard_weight, self.hrr_weight = fts_weight, jaccard_weight, hrr_weight + def _atom(self, word: str): + return hrr.encode_atom(word, self.hrr_dim) + + def _phases(self, blob: bytes): + return hrr.bytes_to_phases(blob, dim=self.hrr_dim) + def search(self, query: str, category: str | None = None, min_trust: float = 0.3, limit: int = 10) -> list[dict]: - """Hybrid search: FTS5 candidates (limit*3) → Jaccard + HRR rerank → trust - weighting → optional temporal decay 0.5^(age_days / half_life). - - Returns fact dicts with a 'score' field, sorted by score desc. - """ + """FTS5 candidates (limit*3) → Jaccard + HRR rerank → trust weighting → optional temporal decay + 0.5^(age_days / half_life). Returns fact dicts with 'score', sorted desc.""" candidates = self._fts_candidates(query, category, min_trust, limit * 3) - if not candidates: - return [] - query_tokens = self._tokenize(query) - # Query vector is loop-invariant; encode lazily on the first candidate - # that carries an HRR vector so migrated stores whose hrr_vector was - # never backfilled don't pay for an encode nothing uses. + # Query vector is loop-invariant; encode lazily on the first candidate that carries an HRR vector + # so stores whose hrr_vector was never backfilled don't pay for it. query_vec = None for fact in candidates: - all_tokens = self._tokenize(fact["content"]) | self._tokenize(fact.get("tags", "")) - jaccard = self._jaccard_similarity(query_tokens, all_tokens) + jaccard = self._jaccard_similarity(query_tokens, self._tokenize(fact["content"]) | self._tokenize(fact.get("tags", ""))) hrr_sim = 0.5 # neutral if self.hrr_weight > 0 and fact.get("hrr_vector"): - fact_vec = hrr.bytes_to_phases(fact["hrr_vector"], dim=self.hrr_dim) + fact_vec = self._phases(fact["hrr_vector"]) if query_vec is None: query_vec = hrr.encode_text(query, self.hrr_dim) - hrr_sim = (hrr.similarity(query_vec, fact_vec) + 1.0) / 2.0 # shift to [0,1] - relevance = (self.fts_weight * fact.get("fts_rank", 0.0) - + self.jaccard_weight * jaccard - + self.hrr_weight * hrr_sim) + hrr_sim = _shift(hrr.similarity(query_vec, fact_vec)) + relevance = self.fts_weight * fact.get("fts_rank", 0.0) + self.jaccard_weight * jaccard + self.hrr_weight * hrr_sim fact["score"] = relevance * fact["trust_score"] if self.half_life > 0: fact["score"] *= self._temporal_decay(fact.get("updated_at") or fact.get("created_at")) - - candidates.sort(key=lambda x: x["score"], reverse=True) - results = candidates[:limit] + results = sorted(candidates, key=lambda x: x["score"], reverse=True)[:limit] for fact in results: fact.pop("hrr_vector", None) # callers expect JSON-serializable dicts return results + def _vector_query(self, fallback: str, category: str | None, limit: int, sim_fn: Callable) -> list[dict]: + """Rank every fact vector (optionally per category) by sim_fn; FTS5 fallback when no vectors exist.""" + rows = self._vector_rows(category) + return self._rank_by_vector(rows, sim_fn, limit) if rows else self.search(fallback, category=category, limit=limit) + def probe(self, entity: str, category: str | None = None, limit: int = 10) -> list[dict]: - """Compositional entity query: unbind bind(entity, ROLE_ENTITY) from the - category bank (or each fact vector) to find facts where the entity plays - a structural role. Not keyword search. Falls back to FTS5 without numpy. - """ + """Compositional entity query: unbind bind(entity, ROLE_ENTITY) from the category bank (or each fact vector) + to find facts where the entity plays a structural role. Not keyword search. Falls back to FTS5 without numpy.""" if not hrr._HAS_NUMPY: return self.search(entity, category=category, limit=limit) - - role_entity = hrr.encode_atom("__hrr_role_entity__", self.hrr_dim) - entity_vec = hrr.encode_atom(entity.lower(), self.hrr_dim) - probe_key = hrr.bind(entity_vec, role_entity) - - # Try the category-specific bank first, then individual fact vectors - if category: - bank_row = self.store._conn.execute( - "SELECT vector FROM memory_banks WHERE bank_name = ?", - (f"cat:{category}",), - ).fetchone() + probe_key = hrr.bind(self._atom(entity.lower()), self._atom(_ROLE_ENTITY)) + if category: # category bank first, then individual fact vectors + bank_row = self.store._conn.execute("SELECT vector FROM memory_banks WHERE bank_name = ?", (f"cat:{category}",)).fetchone() if bank_row: - extracted = hrr.unbind(hrr.bytes_to_phases(bank_row["vector"], dim=self.hrr_dim), probe_key) - return self._rank_by_vector( - self._vector_rows(category), lambda _f, fact_vec: hrr.similarity(extracted, fact_vec), limit, - ) - - rows = self._vector_rows(category) - if not rows: - return self.search(entity, category=category, limit=limit) - - # role_content is loop-invariant — encode once, not per row. - role_content = hrr.encode_atom("__hrr_role_content__", self.hrr_dim) - - def _sim(fact: dict, fact_vec) -> float: - # Does unbinding the probe key leave the fact's content signal? - residual = hrr.unbind(fact_vec, probe_key) - content_vec = hrr.bind(hrr.encode_text(fact["content"], self.hrr_dim), role_content) - return hrr.similarity(residual, content_vec) - - return self._rank_by_vector(rows, _sim, limit) + extracted = hrr.unbind(self._phases(bank_row["vector"]), probe_key) + return self._rank_by_vector(self._vector_rows(category), lambda _f, fact_vec: hrr.similarity(extracted, fact_vec), limit) + role_content = self._atom(_ROLE_CONTENT) # loop-invariant: encode once, not per row + # Does unbinding the probe key leave the fact's content signal? + return self._vector_query(entity, category, limit, lambda fact, fact_vec: hrr.similarity( + hrr.unbind(fact_vec, probe_key), hrr.bind(hrr.encode_text(fact["content"], self.hrr_dim), role_content))) def related(self, entity: str, category: str | None = None, limit: int = 10) -> list[dict]: - """Facts structurally connected to an entity (shared context), not just - facts *about* it as in probe. Falls back to FTS5 without numpy. - """ + """Facts structurally connected to an entity (shared context), not just facts *about* it as in probe. + Falls back to FTS5 without numpy.""" if not hrr._HAS_NUMPY: return self.search(entity, category=category, limit=limit) - - # Bare atom, not role-bound — we want ANY structural match - entity_vec = hrr.encode_atom(entity.lower(), self.hrr_dim) - - rows = self._vector_rows(category) - if not rows: - return self.search(entity, category=category, limit=limit) - - # Both role atoms are loop-invariant — encode once, not per row. - role_entity = hrr.encode_atom("__hrr_role_entity__", self.hrr_dim) - role_content = hrr.encode_atom("__hrr_role_content__", self.hrr_dim) - - def _sim(fact: dict, fact_vec) -> float: - # A residual similar to ANY role vector means the entity plays a - # structural role in the fact; take the max over both roles. - residual = hrr.unbind(fact_vec, entity_vec) - return max(hrr.similarity(residual, role_entity), hrr.similarity(residual, role_content)) - - return self._rank_by_vector(rows, _sim, limit) + entity_vec = self._atom(entity.lower()) # bare atom, not role-bound: ANY structural match + roles = (self._atom(_ROLE_ENTITY), self._atom(_ROLE_CONTENT)) # loop-invariant: encode once + # A residual similar to ANY role vector means the entity plays a structural role in the fact. + return self._vector_query(entity, category, limit, lambda _f, fact_vec: max( + hrr.similarity(hrr.unbind(fact_vec, entity_vec), role) for role in roles)) def reason(self, entities: list[str], category: str | None = None, limit: int = 10) -> list[dict]: - """Multi-entity compositional query (vector-space JOIN): facts where ALL - entities play structural roles. Falls back to FTS5 without numpy. - """ + """Multi-entity compositional query (vector-space JOIN): facts where ALL entities play structural roles. + Falls back to FTS5 without numpy.""" if not hrr._HAS_NUMPY or not entities: return self.search(" ".join(entities), category=category, limit=limit) - - role_entity = hrr.encode_atom("__hrr_role_entity__", self.hrr_dim) - probe_keys = [ - hrr.bind(hrr.encode_atom(entity.lower(), self.hrr_dim), role_entity) - for entity in entities - ] - - rows = self._vector_rows(category) - if not rows: - return self.search(" ".join(entities), category=category, limit=limit) - - role_content = hrr.encode_atom("__hrr_role_content__", self.hrr_dim) - - def _sim(fact: dict, fact_vec) -> float: - # AND semantics via min: high only if EVERY entity is structurally present. - return min( - hrr.similarity(hrr.unbind(fact_vec, key), role_content) for key in probe_keys - ) - - return self._rank_by_vector(rows, _sim, limit) + role_entity, role_content = self._atom(_ROLE_ENTITY), self._atom(_ROLE_CONTENT) + probe_keys = [hrr.bind(self._atom(entity.lower()), role_entity) for entity in entities] + # AND semantics via min: high only if EVERY entity is structurally present. + return self._vector_query(" ".join(entities), category, limit, lambda _f, fact_vec: min( + hrr.similarity(hrr.unbind(fact_vec, key), role_content) for key in probe_keys)) def contradict(self, category: str | None = None, threshold: float = 0.3, limit: int = 10) -> list[dict]: - """Memory hygiene: pairs of facts that share entities (same subject) but - have low content-vector similarity (different claims). Empty without numpy. - """ + """Pairs of facts sharing entities (same subject) with low content-vector similarity (different claims). Empty without numpy.""" if not hrr._HAS_NUMPY: return [] - - rows = self._vector_rows( - category, - columns="fact_id, content, category, tags, trust_score, created_at, updated_at, hrr_vector", - ) + rows = self._vector_rows(category, columns="fact_id, content, category, tags, trust_score, created_at, updated_at, hrr_vector") if len(rows) < 2: return [] - # O(n²) guard: ~125K comparisons at 500 facts is acceptable; above that - # only compare the most recently updated facts. - if len(rows) > 500: + if len(rows) > 500: # O(n²) guard: only compare the most recently updated facts rows = sorted(rows, key=lambda r: r["updated_at"] or r["created_at"], reverse=True)[:500] - - facts = [dict(r) for r in rows] - for fact in facts: + facts = [] # (public dict, lower-cased entity names, phase vector) + for row in rows: + fact = dict(row) entity_rows = self.store._conn.execute( "SELECT e.name FROM entities e JOIN fact_entities fe ON fe.entity_id = e.entity_id WHERE fe.fact_id = ?", (fact["fact_id"],), ).fetchall() - fact["_entities"] = {r["name"].lower() for r in entity_rows} - fact["_vec"] = hrr.bytes_to_phases(fact.pop("hrr_vector"), dim=self.hrr_dim) - - def _public(fact: dict) -> dict: - return {k: v for k, v in fact.items() if k not in ("_entities", "_vec")} - + facts.append((fact, {r["name"].lower() for r in entity_rows}, self._phases(fact.pop("hrr_vector")))) contradictions = [] - for i, f1 in enumerate(facts): - for f2 in facts[i + 1:]: - ents1, ents2 = f1["_entities"], f2["_entities"] + for i, (f1, ents1, vec1) in enumerate(facts): + for f2, ents2, vec2 in facts[i + 1:]: if not ents1 or not ents2: continue entity_overlap = len(ents1 & ents2) / len(ents1 | ents2) if entity_overlap < 0.3: continue # not enough shared subject to be contradictory - content_sim = hrr.similarity(f1["_vec"], f2["_vec"]) - # High entity overlap + low content similarity = contradiction - contradiction_score = entity_overlap * (1.0 - (content_sim + 1.0) / 2.0) + content_sim = hrr.similarity(vec1, vec2) + contradiction_score = entity_overlap * (1.0 - _shift(content_sim)) # high overlap + low similarity if contradiction_score >= threshold: contradictions.append({ - "fact_a": _public(f1), - "fact_b": _public(f2), + "fact_a": f1, "fact_b": f2, "entity_overlap": round(entity_overlap, 3), "content_similarity": round(content_sim, 3), "contradiction_score": round(contradiction_score, 3), "shared_entities": sorted(ents1 & ents2), }) - - contradictions.sort(key=lambda x: x["contradiction_score"], reverse=True) - return contradictions[:limit] - - # -- Vector scoring helpers ----------------------------------------------- + return sorted(contradictions, key=lambda x: x["contradiction_score"], reverse=True)[:limit] def _vector_rows(self, category: str | None, columns: str = _FACT_COLUMNS + ", hrr_vector") -> list: """All facts that carry an HRR vector, optionally filtered by category.""" - where = "WHERE hrr_vector IS NOT NULL" - params: list = [] - if category: - where += " AND category = ?" - params.append(category) - return self.store._conn.execute(f"SELECT {columns} FROM facts {where}", params).fetchall() + where = "WHERE hrr_vector IS NOT NULL" + (" AND category = ?" if category else "") + return self.store._conn.execute(f"SELECT {columns} FROM facts {where}", [category] if category else []).fetchall() def _rank_by_vector(self, rows: list, sim_fn: Callable[[dict, object], float], limit: int) -> list[dict]: """Score each row as (sim + 1) / 2 * trust_score (sim shifted to [0, 1]), sorted desc.""" - scored = [] - for row in rows: - fact = dict(row) - fact_vec = hrr.bytes_to_phases(fact.pop("hrr_vector"), dim=self.hrr_dim) - fact["score"] = (sim_fn(fact, fact_vec) + 1.0) / 2.0 * fact["trust_score"] - scored.append(fact) - scored.sort(key=lambda x: x["score"], reverse=True) - return scored[:limit] - - # -- FTS / lexical helpers ------------------------------------------------ + scored = [dict(row) for row in rows] + for fact in scored: + fact["score"] = _shift(sim_fn(fact, self._phases(fact.pop("hrr_vector")))) * fact["trust_score"] + return sorted(scored, key=lambda x: x["score"], reverse=True)[:limit] def _fts_candidates(self, query: str, category: str | None, min_trust: float, limit: int) -> list[dict]: """Raw FTS5 MATCH candidates with rank normalized to [0, 1] as 'fts_rank'.""" category_clause = "AND f.category = ? " if category else "" params = [self._sanitize_fts_query(query)] + ([category] if category else []) + [min_trust, limit] - sql = ( - "SELECT f.*, facts_fts.rank as fts_rank_raw FROM facts_fts " - "JOIN facts f ON f.fact_id = facts_fts.rowid " - f"WHERE facts_fts MATCH ? {category_clause}AND f.trust_score >= ? " - "ORDER BY facts_fts.rank LIMIT ?" - ) + sql = ("SELECT f.*, facts_fts.rank as fts_rank_raw FROM facts_fts JOIN facts f ON f.fact_id = facts_fts.rowid " + f"WHERE facts_fts MATCH ? {category_clause}AND f.trust_score >= ? ORDER BY facts_fts.rank LIMIT ?") try: - rows = self.store._conn.execute(sql, params).fetchall() + results = [dict(row) for row in self.store._conn.execute(sql, params).fetchall()] except Exception: return [] # FTS5 MATCH can fail on malformed queries - if not rows: - return [] - - results = [dict(row) for row in rows] - # FTS5 rank is negative (lower = better); normalize |rank| / max to [0, 1] - max_rank = max(max(abs(f["fts_rank_raw"]) for f in results), 1e-6) # avoid div by zero + # FTS5 rank is negative (lower = better); normalize |rank| / max to [0, 1] (1e-6 floor avoids div by zero) + max_rank = max([abs(f["fts_rank_raw"]) for f in results] + [1e-6]) for fact in results: fact["fts_rank"] = abs(fact.pop("fts_rank_raw")) / max_rank return results @@ -279,40 +185,17 @@ class FactRetriever: @staticmethod def _tokenize(text: str) -> set[str]: """Lowercase whitespace tokens with surrounding punctuation stripped (no stemming).""" - if not text: - return set() - return {c for c in (w.strip(".,;:!?\"'()[]{}#@<>") for w in text.lower().split()) if c} + return {c for c in (w.strip(_PUNCT) for w in text.lower().split()) if c} if text else set() - # Stopwords dropped before FTS5 OR-expansion: short English function words - # that carry no retrieval signal and force false-negative AND matches. - _FTS_STOPWORDS = frozenset(""" - a about above after again all am an and any are as at be because been before being - between both but by can could did do does doing don down during each few for from - further had has have having he her here hers herself him himself his how i if in - into is it its itself just me more most my myself no nor not now of off on once - only or other our ours ourselves out over own same she should so some such than that - the their theirs them themselves then there these they this those through to too under - until up very was we were what when where which while who whom why will with would - you your yours yourself yourselves - """.split()) - - @classmethod - def _sanitize_fts_query(cls, query: str) -> str: - """Natural-language query -> FTS5-safe OR expression of quoted tokens. - - FTS5 AND-joins a multi-word MATCH by default, which tanks recall on prose. - Drops stopwords and <2-char tokens, strips FTS5 operator chars, and - phrase-quotes each survivor. If nothing survives, returns the raw query - (caller gets zero results rather than a SQL error). - """ + @staticmethod + def _sanitize_fts_query(query: str) -> str: + """Natural-language query -> FTS5-safe OR expression of quoted tokens. FTS5 AND-joins a multi-word + MATCH by default, which tanks recall on prose: drop stopwords and <2-char tokens, strip FTS5 operator + chars, phrase-quote each survivor. If nothing survives, return the raw query (zero results, not a SQL error).""" if not query: return "" - strip_special = str.maketrans("", "", '"()*^:-+') - tokens = [ - f'"{cleaned}"' - for cleaned in (raw.strip(".,;:!?\"'()[]{}#@<>").translate(strip_special) for raw in query.lower().split()) - if len(cleaned) >= 2 and cleaned not in cls._FTS_STOPWORDS - ] + tokens = [f'"{c}"' for c in (raw.strip(_PUNCT).translate(_FTS_OPERATORS) for raw in query.lower().split()) + if len(c) >= 2 and c not in _FTS_STOPWORDS] return " OR ".join(tokens) if tokens else query @staticmethod @@ -325,12 +208,8 @@ class FactRetriever: if not self.half_life or not timestamp_str: return 1.0 try: - ts = timestamp_str - if isinstance(ts, str): - ts = datetime.fromisoformat(ts.replace("Z", "+00:00")) - if ts.tzinfo is None: - ts = ts.replace(tzinfo=timezone.utc) - age_days = (datetime.now(timezone.utc) - ts).total_seconds() / 86400 + ts = datetime.fromisoformat(timestamp_str.replace("Z", "+00:00")) if isinstance(timestamp_str, str) else timestamp_str + age_days = (datetime.now(timezone.utc) - (ts if ts.tzinfo else ts.replace(tzinfo=timezone.utc))).total_seconds() / 86400 return 1.0 if age_days < 0 else math.pow(0.5, age_days / self.half_life) except (ValueError, TypeError): return 1.0 diff --git a/plugins/memory/holographic/store.py b/plugins/memory/holographic/store.py index 6e91c0791c..8b8819b742 100644 --- a/plugins/memory/holographic/store.py +++ b/plugins/memory/holographic/store.py @@ -6,10 +6,7 @@ import sqlite3 import threading from pathlib import Path -try: - from . import holographic as hrr -except ImportError: - import holographic as hrr # type: ignore[no-redef] +from . import holographic as hrr _SCHEMA = """ CREATE TABLE IF NOT EXISTS facts ( @@ -73,19 +70,16 @@ CREATE TABLE IF NOT EXISTS memory_banks ( ); """ -# Trust adjustment constants -_HELPFUL_DELTA = 0.05 -_UNHELPFUL_DELTA = -0.10 +_HELPFUL_DELTA, _UNHELPFUL_DELTA = 0.05, -0.10 -# Entity extraction patterns, applied in order: capitalized multi-word phrases -# ("John Doe"), double-quoted terms, single-quoted terms, then "X aka Y" (both sides). -_RE_SINGLE_ENTITY = ( - re.compile(r'\b([A-Z][a-z]+(?:\s+[A-Z][a-z]+)+)\b'), - re.compile(r'"([^"]+)"'), - re.compile(r"'([^']+)'"), -) +# Entity extraction patterns, applied in order: capitalized multi-word phrases ("John Doe"), double-quoted terms, +# single-quoted terms, then "X aka Y" (both sides). +_RE_SINGLE_ENTITY = (re.compile(r'\b([A-Z][a-z]+(?:\s+[A-Z][a-z]+)+)\b'), re.compile(r'"([^"]+)"'), re.compile(r"'([^']+)'")) _RE_AKA = re.compile(r'(\w+(?:\s+\w+)*)\s+(?:aka|also known as)\s+(\w+(?:\s+\w+)*)', re.IGNORECASE) _ENTITY_NAMES_SQL = "SELECT e.name FROM entities e JOIN fact_entities fe ON fe.entity_id = e.entity_id WHERE fe.fact_id = ?" +# Entity lookup order: exact name, then aliases (comma-separated; wrapped in commas for whole-alias matching). +_ENTITY_LOOKUPS = ("SELECT entity_id FROM entities WHERE name LIKE ?", + "SELECT entity_id FROM entities WHERE ',' || aliases || ',' LIKE '%,' || ? || ',%'") def _clamp_trust(value: float) -> float: @@ -93,14 +87,13 @@ def _clamp_trust(value: float) -> float: class MemoryStore: - """SQLite-backed fact store with entity resolution and trust scoring.""" + """SQLite-backed fact store with entity resolution and trust scoring. + + Process-wide shared connection registry: SQLite allows one writer at a time and several providers + coexist per process (main agent + every delegate_task subagent), so all instances for the same database + share ONE connection and ONE re-entrant lock — writes are fully serialized and "database is locked" is + impossible. Refcounted: closing one instance never tears the connection out from under a sibling.""" - # Process-wide shared connection registry. SQLite allows one writer at a - # time, and several providers coexist per process (main agent + every - # delegate_task subagent). All instances for the same database share ONE - # connection and ONE re-entrant lock, so writes are fully serialized and - # "database is locked" contention is impossible. Refcounted: closing one - # instance never tears the connection out from under a live sibling. _shared: dict = {} _shared_guard = threading.Lock() @@ -110,125 +103,88 @@ class MemoryStore: db_path = str(get_hermes_home() / "memory_store.db") self.db_path = Path(db_path).expanduser() self.db_path.parent.mkdir(parents=True, exist_ok=True) - self.default_trust = _clamp_trust(default_trust) - self.hrr_dim = hrr_dim - self._hrr_available = hrr._HAS_NUMPY - - # resolve() so symlinked/relative paths to the same file share ONE - # connection instead of reintroducing multi-writer contention. - try: + self.default_trust, self.hrr_dim, self._hrr_available = _clamp_trust(default_trust), hrr_dim, hrr._HAS_NUMPY + try: # resolve() so symlinked/relative paths to the same file share ONE connection self._key = str(self.db_path.resolve()) except OSError: self._key = str(self.db_path) with MemoryStore._shared_guard: entry = MemoryStore._shared.get(self._key) if entry is None: - # Autocommit: a write that raises mid-method can never leave a - # dangling transaction (and its write lock) open. The explicit - # commit() calls below are harmless no-ops. + # Autocommit: a write that raises mid-method can't leave a dangling transaction (and its + # write lock) open; the explicit commit() calls in _write are then harmless no-ops. conn = sqlite3.connect(self._key, check_same_thread=False, timeout=10.0, isolation_level=None) conn.row_factory = sqlite3.Row - entry = MemoryStore._shared[self._key] = { - "conn": conn, "lock": threading.RLock(), "refs": 0, "ready": False, - } + entry = MemoryStore._shared[self._key] = {"conn": conn, "lock": threading.RLock(), "refs": 0, "ready": False} entry["refs"] += 1 self._entry, self._conn, self._lock = entry, entry["conn"], entry["lock"] - - # Initialise the schema once per shared connection. - with self._lock: - if not self._entry["ready"]: + with self._lock: # schema initialised once per shared connection + if not entry["ready"]: self._init_db() - self._entry["ready"] = True + entry["ready"] = True def _init_db(self) -> None: - """Create tables/indexes/triggers, enable WAL (via the shared fallback helper so - NFS/SMB/FUSE HERMES_HOME degrades gracefully), and add hrr_vector to pre-HRR databases.""" + """Create schema, enable WAL via the shared fallback helper (NFS/SMB/FUSE degrade gracefully), add hrr_vector to pre-HRR DBs.""" from hermes_state import apply_wal_with_fallback apply_wal_with_fallback(self._conn, db_label="memory_store.db (holographic)") self._conn.executescript(_SCHEMA) - columns = {row[1] for row in self._conn.execute("PRAGMA table_info(facts)").fetchall()} - if "hrr_vector" not in columns: + if "hrr_vector" not in {row[1] for row in self._conn.execute("PRAGMA table_info(facts)").fetchall()}: self._conn.execute("ALTER TABLE facts ADD COLUMN hrr_vector BLOB") self._conn.commit() - def add_fact(self, content: str, category: str = "general", tags: str = "") -> int: - """Insert a fact and return its fact_id. + def _one(self, sql: str, params=()): + return self._conn.execute(sql, params).fetchone() - Deduplicates by content (UNIQUE constraint): on duplicate, returns the - existing fact_id without modifying the row. Links extracted entities. - """ + def _write(self, sql: str, params=()) -> sqlite3.Cursor: + cur = self._conn.execute(sql, params) + self._conn.commit() + return cur + + def add_fact(self, content: str, category: str = "general", tags: str = "") -> int: + """Insert a fact and return its fact_id; on duplicate content (UNIQUE) return the existing fact_id untouched. + Links extracted entities and rebuilds the category bank.""" with self._lock: content = content.strip() if not content: raise ValueError("content must not be empty") - try: - cur = self._conn.execute( - "INSERT INTO facts (content, category, tags, trust_score) VALUES (?, ?, ?, ?)", - (content, category, tags, self.default_trust), - ) - self._conn.commit() - fact_id: int = cur.lastrowid # type: ignore[assignment] + fact_id: int = self._write("INSERT INTO facts (content, category, tags, trust_score) VALUES (?, ?, ?, ?)", + (content, category, tags, self.default_trust)).lastrowid # type: ignore[assignment] except sqlite3.IntegrityError: - row = self._conn.execute("SELECT fact_id FROM facts WHERE content = ?", (content,)).fetchone() - return int(row["fact_id"]) - + return int(self._one("SELECT fact_id FROM facts WHERE content = ?", (content,))["fact_id"]) self._link_entities(fact_id, content) self._compute_hrr_vector(fact_id, content) self._rebuild_bank(category) return fact_id - def update_fact( - self, - fact_id: int, - content: str | None = None, - trust_delta: float | None = None, - tags: str | None = None, - category: str | None = None, - ) -> bool: + def update_fact(self, fact_id: int, content: str | None = None, trust_delta: float | None = None, + tags: str | None = None, category: str | None = None) -> bool: """Partially update a fact (trust clamped to [0, 1]). Returns True if the row existed.""" with self._lock: - row = self._conn.execute( - "SELECT fact_id, trust_score FROM facts WHERE fact_id = ?", (fact_id,) - ).fetchone() + row = self._one("SELECT fact_id, trust_score FROM facts WHERE fact_id = ?", (fact_id,)) if row is None: return False - - changes = [(col, val) for col, val in ( - ("content", content.strip() if content is not None else None), - ("tags", tags), - ("category", category), - ("trust_score", _clamp_trust(row["trust_score"] + trust_delta) if trust_delta is not None else None), - ) if val is not None] - assignments = ", ".join(["updated_at = CURRENT_TIMESTAMP"] + [f"{col} = ?" for col, _ in changes]) - self._conn.execute(f"UPDATE facts SET {assignments} WHERE fact_id = ?", [val for _, val in changes] + [fact_id]) - self._conn.commit() - - if content is not None: - # Content changed: re-extract entities and recompute the HRR vector. - self._conn.execute("DELETE FROM fact_entities WHERE fact_id = ?", (fact_id,)) + changes = {col: val for col, val in { + "content": content.strip() if content is not None else None, "tags": tags, "category": category, + "trust_score": _clamp_trust(row["trust_score"] + trust_delta) if trust_delta is not None else None, + }.items() if val is not None} + assignments = ", ".join(["updated_at = CURRENT_TIMESTAMP"] + [f"{col} = ?" for col in changes]) + self._write(f"UPDATE facts SET {assignments} WHERE fact_id = ?", [*changes.values(), fact_id]) + if content is not None: # re-extract entities and recompute the HRR vector + self._write("DELETE FROM fact_entities WHERE fact_id = ?", (fact_id,)) self._link_entities(fact_id, content) - self._conn.commit() self._compute_hrr_vector(fact_id, content) - cat = category or self._conn.execute( - "SELECT category FROM facts WHERE fact_id = ?", (fact_id,) - ).fetchone()["category"] - self._rebuild_bank(cat) - + self._rebuild_bank(category or self._one("SELECT category FROM facts WHERE fact_id = ?", (fact_id,))["category"]) return True def remove_fact(self, fact_id: int) -> bool: """Delete a fact and its entity links. Returns True if the row existed.""" with self._lock: - row = self._conn.execute( - "SELECT fact_id, category FROM facts WHERE fact_id = ?", (fact_id,) - ).fetchone() + row = self._one("SELECT fact_id, category FROM facts WHERE fact_id = ?", (fact_id,)) if row is None: return False - self._conn.execute("DELETE FROM fact_entities WHERE fact_id = ?", (fact_id,)) - self._conn.execute("DELETE FROM facts WHERE fact_id = ?", (fact_id,)) - self._conn.commit() + self._write("DELETE FROM facts WHERE fact_id = ?", (fact_id,)) self._rebuild_bank(row["category"]) return True @@ -237,131 +193,83 @@ class MemoryStore: with self._lock: category_clause = "AND category = ? " if category is not None else "" params = [min_trust] + ([category] if category is not None else []) + [limit] - sql = ( - "SELECT fact_id, content, category, tags, trust_score, retrieval_count, helpful_count, " - f"created_at, updated_at FROM facts WHERE trust_score >= ? {category_clause}" - "ORDER BY trust_score DESC LIMIT ?" - ) + sql = ("SELECT fact_id, content, category, tags, trust_score, retrieval_count, helpful_count, " + f"created_at, updated_at FROM facts WHERE trust_score >= ? {category_clause}" + "ORDER BY trust_score DESC LIMIT ?") return [dict(r) for r in self._conn.execute(sql, params).fetchall()] def record_feedback(self, fact_id: int, helpful: bool) -> dict: """Adjust trust asymmetrically: helpful -> +0.05 and helpful_count += 1; unhelpful -> -0.10. - - Returns {fact_id, old_trust, new_trust, helpful_count}. Raises KeyError if fact_id is unknown. - """ + Returns {fact_id, old_trust, new_trust, helpful_count}. Raises KeyError if fact_id is unknown.""" with self._lock: - row = self._conn.execute( - "SELECT fact_id, trust_score, helpful_count FROM facts WHERE fact_id = ?", - (fact_id,), - ).fetchone() + row = self._one("SELECT fact_id, trust_score, helpful_count FROM facts WHERE fact_id = ?", (fact_id,)) if row is None: raise KeyError(f"fact_id {fact_id} not found") - old_trust: float = row["trust_score"] new_trust = _clamp_trust(old_trust + (_HELPFUL_DELTA if helpful else _UNHELPFUL_DELTA)) - helpful_increment = 1 if helpful else 0 - self._conn.execute( - "UPDATE facts SET trust_score = ?, helpful_count = helpful_count + ?, " - "updated_at = CURRENT_TIMESTAMP WHERE fact_id = ?", - (new_trust, helpful_increment, fact_id), - ) - self._conn.commit() - return {"fact_id": fact_id, "old_trust": old_trust, "new_trust": new_trust, - "helpful_count": row["helpful_count"] + helpful_increment} - - # -- Entity / HRR helpers ------------------------------------------------- + increment = 1 if helpful else 0 + self._write("UPDATE facts SET trust_score = ?, helpful_count = helpful_count + ?, " + "updated_at = CURRENT_TIMESTAMP WHERE fact_id = ?", (new_trust, increment, fact_id)) + return {"fact_id": fact_id, "old_trust": old_trust, "new_trust": new_trust, "helpful_count": row["helpful_count"] + increment} def _extract_entities(self, text: str) -> list[str]: """Regex entity candidates (see the pattern table), deduplicated case-insensitively in first-seen order.""" raw = [m.group(1) for pattern in _RE_SINGLE_ENTITY for m in pattern.finditer(text)] for m in _RE_AKA.finditer(text): raw += [m.group(1), m.group(2)] - seen: set[str] = set() - candidates: list[str] = [] - for name in (n.strip() for n in raw): - if name and name.lower() not in seen: - seen.add(name.lower()) - candidates.append(name) - return candidates + uniq: dict[str, str] = {} # lower-cased key -> first-seen spelling, insertion-ordered + for name in filter(None, (n.strip() for n in raw)): + uniq.setdefault(name.lower(), name) + return list(uniq.values()) def _link_entities(self, fact_id: int, content: str) -> None: """Extract entities from content, resolve/create them, and link each to the fact.""" for name in self._extract_entities(content): - entity_id = self._resolve_entity(name) - self._conn.execute( - "INSERT OR IGNORE INTO fact_entities (fact_id, entity_id) VALUES (?, ?)", - (fact_id, entity_id), - ) - self._conn.commit() + self._write("INSERT OR IGNORE INTO fact_entities (fact_id, entity_id) VALUES (?, ?)", + (fact_id, self._resolve_entity(name))) def _resolve_entity(self, name: str) -> int: """Return the entity_id for a case-insensitive name or alias match, creating the entity if absent.""" - row = self._conn.execute("SELECT entity_id FROM entities WHERE name LIKE ?", (name,)).fetchone() - if row is not None: - return int(row["entity_id"]) - - # Aliases are comma-separated; wrap both sides in commas for whole-alias matching. - alias_row = self._conn.execute( - "SELECT entity_id FROM entities WHERE ',' || aliases || ',' LIKE '%,' || ? || ',%'", (name,) - ).fetchone() - if alias_row is not None: - return int(alias_row["entity_id"]) - - cur = self._conn.execute("INSERT INTO entities (name) VALUES (?)", (name,)) - self._conn.commit() - return int(cur.lastrowid) # type: ignore[return-value] + for sql in _ENTITY_LOOKUPS: + row = self._one(sql, (name,)) + if row is not None: + return int(row["entity_id"]) + return int(self._write("INSERT INTO entities (name) VALUES (?)", (name,)).lastrowid) # type: ignore[arg-type] def _compute_hrr_vector(self, fact_id: int, content: str) -> None: """Compute and store the HRR vector for a fact (linked entities as roles). No-op without numpy.""" if not self._hrr_available: return - - rows = self._conn.execute(_ENTITY_NAMES_SQL, (fact_id,)).fetchall() - vector = hrr.encode_fact(content, [row["name"] for row in rows], self.hrr_dim) - self._conn.execute("UPDATE facts SET hrr_vector = ? WHERE fact_id = ?", (hrr.phases_to_bytes(vector), fact_id)) - self._conn.commit() + entities = [row["name"] for row in self._conn.execute(_ENTITY_NAMES_SQL, (fact_id,)).fetchall()] + blob = hrr.phases_to_bytes(hrr.encode_fact(content, entities, self.hrr_dim)) + self._write("UPDATE facts SET hrr_vector = ? WHERE fact_id = ?", (blob, fact_id)) def _rebuild_bank(self, category: str) -> None: """Full rebuild of a category's memory bank from all its fact vectors.""" if not self._hrr_available: return - bank_name = f"cat:{category}" - rows = self._conn.execute( - "SELECT hrr_vector FROM facts WHERE category = ? AND hrr_vector IS NOT NULL", (category,), - ).fetchall() + rows = self._conn.execute("SELECT hrr_vector FROM facts WHERE category = ? AND hrr_vector IS NOT NULL", (category,)).fetchall() if not rows: - self._conn.execute("DELETE FROM memory_banks WHERE bank_name = ?", (bank_name,)) - self._conn.commit() + self._write("DELETE FROM memory_banks WHERE bank_name = ?", (bank_name,)) return - bank_vector = hrr.bundle(*[hrr.bytes_to_phases(row["hrr_vector"], dim=self.hrr_dim) for row in rows]) hrr.snr_estimate(self.hrr_dim, len(rows)) # warns when near capacity - self._conn.execute( - "INSERT INTO memory_banks (bank_name, vector, dim, fact_count, updated_at) " - "VALUES (?, ?, ?, ?, CURRENT_TIMESTAMP) ON CONFLICT(bank_name) DO UPDATE SET " - "vector = excluded.vector, dim = excluded.dim, fact_count = excluded.fact_count, " - "updated_at = excluded.updated_at", - (bank_name, hrr.phases_to_bytes(bank_vector), self.hrr_dim, len(rows)), - ) - self._conn.commit() - - # -- Lifecycle ------------------------------------------------------------ + self._write("INSERT INTO memory_banks (bank_name, vector, dim, fact_count, updated_at) " + "VALUES (?, ?, ?, ?, CURRENT_TIMESTAMP) ON CONFLICT(bank_name) DO UPDATE SET " + "vector = excluded.vector, dim = excluded.dim, fact_count = excluded.fact_count, " + "updated_at = excluded.updated_at", (bank_name, hrr.phases_to_bytes(bank_vector), self.hrr_dim, len(rows))) @classmethod def release_all_under(cls, directory: "str | Path") -> int: - """Force-close every shared connection whose database lives under ``directory``. - - close() is refcount-driven, so a live holder (e.g. an agent's provider) keeps - a profile's SQLite handle open; on Windows that makes rmtree of the profile - fail while any handle is open. The directory is going away, so later use by a - stale holder is expected to fail. Returns how many connections were closed. - """ + """Force-close every shared connection whose database lives under ``directory``; returns the count. + close() is refcount-driven, so a live holder (e.g. an agent's provider) keeps a profile's SQLite handle + open, which on Windows makes rmtree of the profile fail. The directory is going away, so later use by a + stale holder is expected to fail.""" root = os.path.normcase(str(Path(directory).expanduser().resolve())) + os.sep with cls._shared_guard: - doomed = [key for key in cls._shared if os.path.normcase(key).startswith(root)] - for key in doomed: - entry = cls._shared.pop(key) + doomed = [cls._shared.pop(key) for key in list(cls._shared) if os.path.normcase(key).startswith(root)] + for entry in doomed: try: with entry["lock"]: entry["conn"].close() @@ -380,9 +288,8 @@ class MemoryStore: try: entry["conn"].close() finally: - # Pop only OUR entry: after release_all_under() a same-path - # store may have registered a FRESH entry under this key, - # and a stale holder's late close() must not evict it. + # Pop only OUR entry: after release_all_under() a same-path store may have + # registered a FRESH entry under this key; a stale late close() must not evict it. if MemoryStore._shared.get(self._key) is entry: MemoryStore._shared.pop(self._key, None) self._entry = None diff --git a/plugins/memory/mem0/__init__.py b/plugins/memory/mem0/__init__.py index 9bd5a9ca49..d20660a356 100644 --- a/plugins/memory/mem0/__init__.py +++ b/plugins/memory/mem0/__init__.py @@ -1,12 +1,10 @@ """Mem0 memory plugin — MemoryProvider interface. -Server-side LLM fact extraction, semantic search and deduplication via the Mem0 -Platform API (cloud), a self-hosted Mem0 server (MEM0_HOST, HTTP), or OSS Memory. -Secrets live in $HERMES_HOME/.env (MEM0_API_KEY, MEM0_HOST); behavioral settings -in $HERMES_HOME/mem0.json via `hermes memory setup`: mode ("platform"|"oss"), host, -user_id (canonical id across every gateway so one human gets one merged store; -unset → gateway-native id), agent_id. MEM0_MODE/MEM0_USER_ID/MEM0_AGENT_ID env -vars remain a backward-compatible fallback. +Server-side fact extraction and semantic search via the Mem0 Platform API (cloud), a +self-hosted Mem0 server (MEM0_HOST, HTTP), or OSS Memory. Secrets live in $HERMES_HOME/.env +(MEM0_API_KEY, MEM0_HOST); settings in $HERMES_HOME/mem0.json via `hermes memory setup`: +mode ("platform"|"oss"), host, user_id (canonical id across gateways; unset → gateway-native +id), agent_id. MEM0_* env vars remain a fallback. """ from __future__ import annotations @@ -17,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 @@ -44,10 +43,8 @@ def _is_client_error(exc: Exception) -> 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 {} @@ -56,94 +53,44 @@ def _load_config() -> dict: Layering avoids a silent failure when the JSON file exists but lacks fields 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"): + 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": {}} + 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 != ""}) return config -# --------------------------------------------------------------------------- -# Tool schemas -# --------------------------------------------------------------------------- - -def _schema(name: str, description: str, properties: dict, required: list[str]) -> dict: - return {"name": name, "description": description, "parameters": {"type": "object", "properties": properties, "required": required}} +def _schema(name: str, description: str, properties: dict[str, tuple[str, str]], required: list[str]) -> dict: + props = {k: {"type": t, "description": d} for k, (t, d) in properties.items()} + return {"name": name, "description": description, "parameters": {"type": "object", "properties": props, "required": required}} -def _param(type_: str, description: str) -> dict: - return {"type": type_, "description": description} +TOOL_SCHEMAS = [ + _schema("mem0_search", "Search the user's memories by meaning; returns facts ranked by relevance. Use this before answering any question that may depend on what you know about the user (preferences, facts, history, people, projects, past decisions). For multi-part or multi-hop questions, call it several times — vary the wording and run follow-up searches on what earlier results reveal; one search is rarely enough.", + {"query": ("string", "What to search for."), "top_k": ("integer", "Max results (default: 10, max: 50)."), "rerank": ("boolean", "Rerank results for relevance (default: false, platform mode only).")}, ["query"]), + _schema("mem0_add", "Store a durable fact about the user, verbatim (no LLM extraction). Call this the moment the user states a lasting preference, correction, decision, or personal detail worth recalling on future turns — don't wait to be asked to remember. Skip transient chit-chat and facts you've already stored.", + {"content": ("string", "The fact to store.")}, ["content"]), + _schema("mem0_update", "Replace the text of an existing memory by its ID (take the ID from a mem0_search result). Use when a stored fact has changed or was wrong — correct it in place instead of adding a duplicate.", + {"memory_id": ("string", "Memory UUID to update."), "text": ("string", "New text content.")}, ["memory_id", "text"]), + _schema("mem0_delete", "Delete a memory by its ID (take the ID from a mem0_search result). Use when a stored fact is obsolete or the user asks you to forget it; prefer mem0_update if the fact merely changed.", + {"memory_id": ("string", "Memory UUID to delete.")}, ["memory_id"]), +] - -SEARCH_SCHEMA = _schema( - "mem0_search", - "Search the user's memories by meaning; returns facts ranked by " - "relevance. Use this before answering any question that may depend on " - "what you know about the user (preferences, facts, history, people, " - "projects, past decisions). For multi-part or multi-hop questions, " - "call it several times — vary the wording and run follow-up searches " - "on what earlier results reveal; one search is rarely enough.", - { - "query": _param("string", "What to search for."), - "top_k": _param("integer", "Max results (default: 10, max: 50)."), - "rerank": _param("boolean", "Rerank results for relevance (default: false, platform mode only)."), - }, - ["query"], +_PROMPT_BODY = ( + "You have persistent memory of this user from past conversations. You should call mem0_search before answering anything that could depend on prior context (the user's preferences, facts, history, people, projects, or earlier decisions) — do not rely on the chat window alone, and do not assume you have no memory.\n" + "For multi-part or multi-hop questions, run several searches with different wording/angles and follow-up searches on what the first results surface; one search is rarely enough. Keep searching until you have every fact the question needs before you answer.\n" + "Tools: mem0_search to find memories, mem0_add to store facts, mem0_update and mem0_delete to manage by ID." ) -ADD_SCHEMA = _schema( - "mem0_add", - "Store a durable fact about the user, verbatim (no LLM extraction). " - "Call this the moment the user states a lasting preference, correction, " - "decision, or personal detail worth recalling on future turns — don't " - "wait to be asked to remember. Skip transient chit-chat and facts you've " - "already stored.", - {"content": _param("string", "The fact to store.")}, - ["content"], -) - -UPDATE_SCHEMA = _schema( - "mem0_update", - "Replace the text of an existing memory by its ID (take the ID from a " - "mem0_search result). Use when a stored fact has changed " - "or was wrong — correct it in place instead of adding a duplicate.", - {"memory_id": _param("string", "Memory UUID to update."), "text": _param("string", "New text content.")}, - ["memory_id", "text"], -) - -DELETE_SCHEMA = _schema( - "mem0_delete", - "Delete a memory by its ID (take the ID from a mem0_search " - "result). Use when a stored fact is obsolete or the user asks you to " - "forget it; prefer mem0_update if the fact merely changed.", - {"memory_id": _param("string", "Memory UUID to delete.")}, - ["memory_id"], -) - - -# --------------------------------------------------------------------------- -# MemoryProvider implementation -# --------------------------------------------------------------------------- class Mem0MemoryProvider(MemoryProvider): """Mem0 memory with server-side extraction and semantic search (platform, self-hosted or OSS).""" def __init__(self): - self._config = self._backend = None - self._mode, self._api_key, self._host = "platform", "", "" - self._user_id, self._agent_id = _DEFAULT_USER_ID, "hermes" - self._rerank_default = False - self._channel = "cli" # gateway channel name (cli/telegram/discord/...) - self._sync_thread = self._prefetch_thread = None + self._config = self._backend = self._sync_thread = self._prefetch_thread = None + self._mode, self._api_key, self._host, self._user_id, self._agent_id = "platform", "", "", _DEFAULT_USER_ID, "hermes" + self._rerank_default, self._channel = False, "cli" # channel = gateway name (cli/telegram/discord/...) self._prefetch_query = self._prefetch_result = "" self._prefetch_done = self._atexit_registered = False self._consecutive_failures, self._breaker_open_until = 0, 0.0 # circuit breaker state @@ -157,17 +104,13 @@ class Mem0MemoryProvider(MemoryProvider): cfg = _load_config() if cfg.get("mode", "platform") == "oss": return bool(cfg.get("oss", {}).get("vector_store")) - # Platform needs an api_key; self-hosted needs a host (api_key optional - # when the server runs with AUTH_DISABLED). - return bool(cfg.get("api_key") or cfg.get("host")) + return bool(cfg.get("api_key") or cfg.get("host")) # platform needs a key; self-hosted a host (key optional with AUTH_DISABLED) def save_config(self, values, hermes_home): """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" @@ -183,49 +126,39 @@ class Mem0MemoryProvider(MemoryProvider): from ._setup import post_setup post_setup(hermes_home, config) - def _vs_provider(self, default: str) -> str: - """Configured OSS vector-store provider name (for error hints).""" - return self._config.get("oss", {}).get("vector_store", {}).get("provider", default) + def _oss_hint(self, template: str, default: str = "vector store") -> str: + """OSS-only hint; ``{vs}`` is the configured vector-store provider. "" in other modes.""" + 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 either backend imports it. ensure() honors - # security.allow_lazy_installs and redirects sealed Docker venvs to the durable - # target; on failure the backend import raises the canonical error, captured below. - try: + # 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. + 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": - from ._backend import OSSBackend - return OSSBackend(self._config.get("oss", {})) - if self._host: - from ._backend import SelfHostedBackend - return SelfHostedBackend(self._api_key, self._host) - from ._backend import PlatformBackend - return PlatformBackend(self._api_key) + return _backend.OSSBackend(self._config.get("oss", {})) + 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) return None def _is_breaker_open(self) -> bool: - """Return True if the circuit breaker is tripped (too many failures).""" + """True while the breaker is tripped; an expired cooldown resets the failure count.""" with self._breaker_lock: - if self._consecutive_failures < _BREAKER_THRESHOLD: - return False - if time.monotonic() >= self._breaker_open_until: + if self._consecutive_failures >= _BREAKER_THRESHOLD and time.monotonic() < self._breaker_open_until: + return True + if self._consecutive_failures >= _BREAKER_THRESHOLD: self._consecutive_failures = 0 - return False - return True + return False def _format_error(self, prefix: str, exc: Exception) -> str: msg = f"{prefix}: {exc}" - if self._mode == "oss": - err_str = str(exc).lower() - if "connection" in err_str or "refused" in err_str or "timeout" in err_str: - msg += f" (check that {self._vs_provider('vector store')} is running)" + if any(s in str(exc).lower() for s in ("connection", "refused", "timeout")): + msg += self._oss_hint(" (check that {vs} is running)") return msg def _record_success(self): @@ -235,31 +168,32 @@ class Mem0MemoryProvider(MemoryProvider): def _record_failure(self): with self._breaker_lock: self._consecutive_failures = count = self._consecutive_failures + 1 - tripped = count >= _BREAKER_THRESHOLD - if tripped: + if count >= _BREAKER_THRESHOLD: self._breaker_open_until = time.monotonic() + _BREAKER_COOLDOWN_SECS - if tripped: - hint = f" Check that your {self._vs_provider('unknown')} vector store is running and reachable." if self._mode == "oss" else "" - logger.warning( - "Mem0 circuit breaker tripped after %d consecutive failures. " - "Pausing API calls for %ds.%s", - count, _BREAKER_COOLDOWN_SECS, hint, - ) + if count >= _BREAKER_THRESHOLD: + hint = self._oss_hint(" Check that your {vs} vector store is running and reachable.", "unknown") + logger.warning("Mem0 circuit breaker tripped after %d consecutive failures. Pausing API calls for %ds.%s", count, _BREAKER_COOLDOWN_SECS, hint) + + def _try(self, call, log, msg: str): + """Background-path wrapper: run ``call`` under the breaker; on error log ``msg`` and return None.""" + try: + result = call() + 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 = _load_config() - self._mode = self._config.get("mode", "platform") - self._api_key = self._config.get("api_key", "") - self._host = self._config.get("host", "") - # user_id precedence: operator-configured (env/mem0.json) > gateway-native id - # from kwargs > _DEFAULT_USER_ID. The literal placeholder counts as unset so - # wizard users still get gateway-native ids instead of being bucketed together. - configured = self._config.get("user_id") + 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 = 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 doesn't - # pass ``rerank``; per-call args win. Platform-only; other backends ignore it. - _rr = self._config.get("rerank", False) + # Persisted rerank preference: default for mem0_search when the model omits ``rerank``. Platform-only. + _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() @@ -267,40 +201,26 @@ class Mem0MemoryProvider(MemoryProvider): atexit.register(self._shutdown_backend) self._atexit_registered = True - def _read_filters(self) -> Dict[str, Any]: - # Scoped to user_id only — by design — so recall surfaces memories from any - # gateway/agent under this principal; writes attach agent_id and metadata.channel - # (dashboard per-channel filtering) so narrower views remain possible at query time. - return {"user_id": self._user_id} + def _search(self, query: str, top_k: int = 10, rerank: bool = False, backend=None) -> list: + # Scoped to user_id only — by design — so recall surfaces memories from any gateway/agent under this + # principal; writes attach agent_id and metadata.channel so narrower views remain possible at query time. + return (backend or self._backend).search(query, filters={"user_id": self._user_id}, top_k=top_k, rerank=rerank) - def _write_metadata(self) -> Dict[str, Any]: - return {"channel": self._channel} if self._channel else {} + def _add(self, messages: list, infer: bool): + metadata = {"channel": self._channel} if self._channel else {} + return self._backend.add(messages, user_id=self._user_id, agent_id=self._agent_id, infer=infer, metadata=metadata) def system_prompt_block(self) -> str: - # Mirror _create_backend precedence (oss > host > platform) so the label names - # the backend that actually runs. Rerank is a Mem0 Platform feature only. + # Mirror _create_backend precedence (oss > host > platform). Rerank is a Mem0 Platform feature only. mode_label = "OSS (self-hosted)" if self._mode == "oss" else "self-hosted (HTTP API)" if self._host else "platform (cloud API)" rerank_note = " Rerank is available on search." if (self._mode == "platform" and not self._host) else "" - return ( - "# Mem0 Memory\n" - f"Active. Mode: {mode_label}. User: {self._user_id}.\n" - "You have persistent memory of this user from past conversations. " - "You should call mem0_search before answering anything that could depend " - "on prior context (the user's preferences, facts, history, people, " - "projects, or earlier decisions) — do not rely on the chat window " - "alone, and do not assume you have no memory.\n" - "For multi-part or multi-hop questions, run several searches with " - "different wording/angles and follow-up searches on what the first " - "results surface; one search is rarely enough. Keep searching until " - "you have every fact the question needs before you answer.\n" - "Tools: mem0_search to find memories, mem0_add to store facts, " - f"mem0_update and mem0_delete to manage by ID.{rerank_note}" - ) + return f"# Mem0 Memory\nActive. Mode: {mode_label}. User: {self._user_id}.\n{_PROMPT_BODY}{rerank_note}" def on_turn_start(self, turn_number: int, message: str, **kwargs) -> None: self._start_prefetch(message) def _consume_prefetch_result(self, query: str) -> str | None: + """Pop the finished prefetch body for ``query`` (None if absent or still running).""" with self._prefetch_lock: if self._prefetch_query != query or not self._prefetch_done: return None @@ -311,47 +231,33 @@ 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(): - body = "" - try: - results = backend.search(query, filters=self._read_filters(), top_k=10, rerank=False) - lines = [r.get("memory", "") for r in (results or []) if r.get("memory")] - if lines: - body = "## Mem0 Memory\n" + "\n".join(f"- {l}" for l in lines) - self._record_success() - except Exception as e: - self._record_failure() - logger.debug("Mem0 prefetch failed: %s", e) + results = self._try(lambda: self._search(query, backend=backend), logger.debug, "Mem0 prefetch failed: %s") + lines = [r.get("memory", "") for r in (results or []) if r.get("memory")] + body = "## Mem0 Memory\n" + "\n".join(f"- {l}" for l in lines) if lines else "" with self._prefetch_lock: if self._prefetch_query == query: - self._prefetch_result = body - self._prefetch_done = True + 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: """Recall memories for the CURRENT question with a short hot-path wait.""" - cached = self._consume_prefetch_result(query) - if cached is not None: + if (cached := self._consume_prefetch_result(query)) is not None: return cached self._start_prefetch(query) with self._prefetch_lock: 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).""" @@ -359,117 +265,78 @@ class Mem0MemoryProvider(MemoryProvider): return def _sync(): - backend = self._backend - if backend is None: - return - try: + if self._backend is not None: messages = [{"role": "user", "content": user_content}, {"role": "assistant", "content": assistant_content}] - backend.add(messages, user_id=self._user_id, agent_id=self._agent_id, infer=True, metadata=self._write_metadata()) - self._record_success() - except Exception as e: - self._record_failure() - logger.warning("Mem0 sync failed: %s", e) + self._try(lambda: self._add(messages, infer=True), logger.warning, "Mem0 sync failed: %s") with self._sync_lock: - if self._sync_thread and self._sync_thread.is_alive(): - self._sync_thread.join(timeout=5.0) - # If still alive after timeout, skip to avoid duplicate ingestion. - if self._sync_thread and self._sync_thread.is_alive(): - return + prev = self._sync_thread + if prev and prev.is_alive(): + prev.join(timeout=5.0) + if prev.is_alive(): # still busy after the wait: skip to avoid duplicate ingestion + return self._sync_thread = threading.Thread(target=_sync, daemon=True, name="mem0-sync") self._sync_thread.start() def get_tool_schemas(self) -> List[Dict[str, Any]]: - return [SEARCH_SCHEMA, ADD_SCHEMA, UPDATE_SCHEMA, DELETE_SCHEMA] + return list(TOOL_SCHEMAS) - # -- tool handlers ------------------------------------------------------- + # -- tool handlers: (required params, error label, body, client-error policy) --- + # Client errors (bad ID / not found) never trip the breaker, except for mem0_add + # where they count as failures; update/delete answer them with "Memory not found". def _tool_search(self, args: dict) -> str: - query = args.get("query", "") - if not query: - return tool_error("Missing required parameter: query") - try: - top_k = max(1, min(int(args.get("top_k", 10)), 50)) - 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._backend.search(query, filters=self._read_filters(), top_k=top_k, rerank=rerank) - self._record_success() - if not results: - return json.dumps({"result": "No relevant memories found."}) - items = [{"id": r.get("id"), "memory": r.get("memory", ""), "score": r.get("score", 0)} for r in results] - return json.dumps({"results": items, "count": len(items)}) - except Exception as e: - if not _is_client_error(e): - self._record_failure() - return tool_error(self._format_error("Search failed", e)) + top_k = max(1, min(int(args.get("top_k", 10)), 50)) + 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."}) + items = [{"id": r.get("id"), "memory": r.get("memory", ""), "score": r.get("score", 0)} for r in results] + return json.dumps({"results": items, "count": len(items)}) def _tool_add(self, args: dict) -> str: - content = args.get("content", "") - if not content: - return tool_error("Missing required parameter: content") - try: - result = self._backend.add( - [{"role": "user", "content": content}], - user_id=self._user_id, agent_id=self._agent_id, infer=False, metadata=self._write_metadata(), - ) - self._record_success() - event_id = result.get("event_id") if isinstance(result, dict) else None - # Cloud add is async (server-side extraction); OSS and self-hosted store synchronously. - msg = "Fact stored." if (self._mode == "oss" or self._host) else "Fact queued for storage." - return json.dumps({"result": msg, "event_id": event_id}) - except Exception as e: - self._record_failure() - return tool_error(self._format_error("Failed to store", e)) - - def _tool_by_id(self, args: dict, required: tuple[str, ...], label: str, method: str) -> str: - """Shared update/delete shape: required-param check, then backend.(*values).""" - values = [args.get(k, "") for k in required] - for k, v in zip(required, values): - if not v: - return tool_error(f"Missing required parameter: {k}") - try: - result = getattr(self._backend, method)(*values) - self._record_success() - return json.dumps(result) - except Exception as e: - if _is_client_error(e): - return tool_error(f"Memory not found: {values[0]}") - self._record_failure() - return tool_error(self._format_error(label, e)) - - def _tool_update(self, args: dict) -> str: - return self._tool_by_id(args, ("memory_id", "text"), "Update failed", "update") - - def _tool_delete(self, args: dict) -> str: - return self._tool_by_id(args, ("memory_id",), "Delete failed", "delete") + result = self._add([{"role": "user", "content": args["content"]}], infer=False) + event_id = result.get("event_id") if isinstance(result, dict) else None + # Cloud add is async (server-side extraction); OSS and self-hosted store synchronously. + msg = "Fact stored." if (self._mode == "oss" or self._host) else "Fact queued for storage." + return json.dumps({"result": msg, "event_id": event_id}) _TOOL_HANDLERS = { - "mem0_search": _tool_search, - "mem0_add": _tool_add, - "mem0_update": _tool_update, - "mem0_delete": _tool_delete, + "mem0_search": (("query",), "Search failed", _tool_search, "skip"), + "mem0_add": (("content",), "Failed to store", _tool_add, "count"), + "mem0_update": (("memory_id", "text"), "Update failed", lambda self, a: json.dumps(self._backend.update(a["memory_id"], a["text"])), "not_found"), + "mem0_delete": (("memory_id",), "Delete failed", lambda self, a: json.dumps(self._backend.delete(a["memory_id"])), "not_found"), } def handle_tool_call(self, tool_name: str, args: dict, **kwargs) -> str: if self._backend is None: err = getattr(self, "_init_error", "unknown error") - hint = f" Check that {self._vs_provider('vector store')} is running and reachable." if self._mode == "oss" else "" - return json.dumps({"error": f"Mem0 backend not initialized: {err}.{hint}"}) + return json.dumps({"error": f"Mem0 backend not initialized: {err}.{self._oss_hint(' Check that {vs} is running and reachable.')}"}) if self._is_breaker_open(): - hint = f" Check that your {self._vs_provider('vector store')} is running." if self._mode == "oss" else "" - return json.dumps({"error": f"Mem0 temporarily unavailable (multiple consecutive failures). Will retry automatically.{hint}"}) - handler = self._TOOL_HANDLERS.get(tool_name) - if handler is None: + return json.dumps({"error": f"Mem0 temporarily unavailable (multiple consecutive failures). Will retry automatically.{self._oss_hint(' Check that your {vs} is running.')}"}) + if tool_name not in self._TOOL_HANDLERS: return tool_error(f"Unknown tool: {tool_name}") - return handler(self, args) + required, label, body, on_client_error = self._TOOL_HANDLERS[tool_name] + if missing := next((k for k in required if not args.get(k, "")), None): + return tool_error(f"Missing required parameter: {missing}") + try: + result = body(self, args) + except Exception as e: + client = _is_client_error(e) + if client and on_client_error == "not_found": + return tool_error(f"Memory not found: {args['memory_id']}") + 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): - 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): diff --git a/plugins/memory/mem0/_backend.py b/plugins/memory/mem0/_backend.py index 930214958a..c1abd8925a 100644 --- a/plugins/memory/mem0/_backend.py +++ b/plugins/memory/mem0/_backend.py @@ -3,39 +3,29 @@ from __future__ import annotations from abc import ABC, abstractmethod +from contextlib import closing, suppress from typing import Any def _add_kwargs(user_id: str, agent_id: str, infer: bool, metadata: dict | None) -> dict[str, Any]: - kwargs: dict[str, Any] = {"user_id": user_id, "agent_id": agent_id, "infer": infer} - if metadata: - kwargs["metadata"] = metadata - return kwargs + return {"user_id": user_id, "agent_id": agent_id, "infer": infer, **({"metadata": metadata} if metadata else {})} 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): """Unified interface over Platform (MemoryClient), self-hosted (HTTP) and OSS (Memory) backends. - - update()/delete() are template methods: subclasses implement the raw - ``_update``/``_delete`` calls and the base wraps the uniform result dict. - """ + update()/delete() are template methods: subclasses implement raw ``_update``/``_delete``.""" @abstractmethod def search(self, query: str, *, filters: dict, top_k: int = 10, rerank: bool = False) -> list[dict]: ... - @abstractmethod def add(self, messages: list, *, user_id: str, agent_id: str, infer: bool = False, metadata: dict | None = None) -> dict: ... - @abstractmethod def _update(self, memory_id: str, text: str) -> None: ... - @abstractmethod def _delete(self, memory_id: str) -> None: ... @@ -73,24 +63,14 @@ class PlatformBackend(Mem0Backend): class SelfHostedBackend(Mem0Backend): """Direct HTTP backend for a self-hosted Mem0 server (the FastAPI ``server/``). - - mem0.MemoryClient is hardwired to the cloud API (``Authorization: Token`` auth, - ``GET /v1/ping/`` in ``__init__``) so it can't be reused here; this speaks the - server's real contract: ``X-API-Key`` auth and the ``/memories`` / ``/search`` routes. - """ + mem0.MemoryClient is hardwired to the cloud API (``Authorization: Token``, ``GET /v1/ping/`` in ``__init__``), + so this speaks the server's real contract: ``X-API-Key`` auth and the ``/memories`` / ``/search`` routes.""" def __init__(self, api_key: str, host: str, transport=None): import httpx - - headers = {"Content-Type": "application/json"} - if api_key: - headers["X-API-Key"] = api_key # omitted only for AUTH_DISABLED servers - # Connect-level retries keep a single dropped SYN from counting toward the - # provider failure breaker. ``transport`` is injectable for tests. - self._client = httpx.Client( - base_url=host.rstrip("/"), headers=headers, timeout=30.0, - transport=transport or httpx.HTTPTransport(retries=2), - ) + headers = {"Content-Type": "application/json", **({"X-API-Key": api_key} if api_key else {})} # key omitted only for AUTH_DISABLED servers + # Connect-level retries keep one dropped SYN from counting toward the breaker. ``transport`` is injectable for tests. + self._client = httpx.Client(base_url=host.rstrip("/"), headers=headers, timeout=30.0, transport=transport or httpx.HTTPTransport(retries=2)) def _json(self, method: str, path: str, **kwargs) -> Any: resp = self._client.request(method, path, **kwargs) @@ -98,15 +78,11 @@ class SelfHostedBackend(Mem0Backend): return resp.json() if resp.content else {} def search(self, query: str, *, filters: dict, top_k: int = 10, rerank: bool = False) -> list[dict]: - # rerank is platform-only; the self-hosted /search ignores it. - body: dict[str, Any] = {"query": query, "top_k": top_k} - if filters: - body["filters"] = filters # user_id belongs in filters (top-level is deprecated) - return _unwrap_results(self._json("POST", "/search", json=body)) + # rerank is platform-only; the self-hosted /search ignores it. user_id belongs in filters (top-level is deprecated). + return _unwrap_results(self._json("POST", "/search", json={"query": query, "top_k": top_k, **({"filters": filters} if filters else {})})) def add(self, messages: list, *, user_id: str, agent_id: str, infer: bool = False, metadata: dict | None = None) -> dict: - body: dict[str, Any] = {"messages": messages, **_add_kwargs(user_id, agent_id, infer, metadata)} - return self._json("POST", "/memories", json=body) + return self._json("POST", "/memories", json={"messages": messages, **_add_kwargs(user_id, agent_id, infer, metadata)}) def _update(self, memory_id: str, text: str) -> None: self._json("PUT", f"/memories/{memory_id}", json={"text": text}) @@ -115,10 +91,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" @@ -129,14 +103,10 @@ 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): - raise RuntimeError( - "mem0 LlmFactory does not support the provider registration required " - "for the Hermes OpenAI OSS backend" - ) + raise RuntimeError("mem0 LlmFactory does not support the provider registration required for the Hermes OpenAI OSS backend") if provider_map.get(_DIRECT_OPENAI_PROVIDER) != (_DIRECT_OPENAI_CLASS_PATH, OpenAIConfig): register_provider(_DIRECT_OPENAI_PROVIDER, _DIRECT_OPENAI_CLASS_PATH, OpenAIConfig) @@ -149,16 +119,14 @@ class OSSBackend(Mem0Backend): from mem0 import Memory from ._oss_providers import EMBEDDER_PROVIDERS, KNOWN_DIMS, LLM_PROVIDERS - def _provider_block(name: str) -> dict: + def _provider_block(name: str, registry: dict) -> dict: + """Copy of oss_config[name] with the legacy ``api_base`` key mapped to the provider's canonical base-URL key.""" block = dict(oss_config[name]) - provider = str(block.get("provider") or "").strip().lower() provider_config = dict(block.get("config", {})) legacy_base = provider_config.pop("api_base", None) - if legacy_base: - registry = LLM_PROVIDERS if name == "llm" else EMBEDDER_PROVIDERS - canonical_key = registry.get(provider, {}).get("base_url_key") - if canonical_key: - provider_config.setdefault(canonical_key, legacy_base) + canonical_key = registry.get(str(block.get("provider") or "").strip().lower(), {}).get("base_url_key") + if legacy_base and canonical_key: + provider_config.setdefault(canonical_key, legacy_base) block["config"] = provider_config return block @@ -166,33 +134,22 @@ class OSSBackend(Mem0Backend): vs_config = dict(vector_store.get("config", {})) if "path" in vs_config: vs_config["path"] = os.path.expanduser(vs_config["path"]) - embedder_config = oss_config.get("embedder", {}).get("config", {}) dims = embedder_config.get("embedding_dims") or KNOWN_DIMS.get(embedder_config.get("model", "")) if dims: 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"), - "embedder": _provider_block("embedder"), - "version": "v1.1", - } + 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 first, then swap the provider on the validated object. + # mem0 validates LlmConfig.provider before its factory lookup: build the supported OpenAI config, then swap the provider. _register_direct_openai_provider() from mem0.configs.base import MemoryConfig memory_config = MemoryConfig(**config) try: memory_config.llm.provider = _DIRECT_OPENAI_PROVIDER except (AttributeError, TypeError) as exc: - raise RuntimeError( - "mem0 MemoryConfig does not expose a mutable llm.provider " - "for the Hermes OpenAI OSS backend" - ) from exc + raise RuntimeError("mem0 MemoryConfig does not expose a mutable llm.provider for the Hermes OpenAI OSS backend") from exc self._memory = Memory(memory_config) else: self._memory = Memory.from_config(config) @@ -201,7 +158,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") @@ -211,41 +168,27 @@ class OSSBackend(Mem0Backend): client = QdrantClient(url=url, api_key=vs_config.get("api_key")) else: return - try: + with closing(client): if not client.collection_exists(collection_name): return vectors = client.get_collection(collection_name).config.params.vectors - # Named-vector collections expose a dict; unnamed expose an object with .size. + # Named-vector collections expose a dict; unnamed expose an object with .size. if isinstance(vectors, dict): vectors = next(iter(vectors.values()), None) current_dims = getattr(vectors, "size", None) if current_dims is not None and current_dims != expected_dims: client.delete_collection(collection_name) - finally: - client.close() elif provider == "pgvector": import psycopg2 from psycopg2 import sql as pgsql conn_params = {k: vs_config[k] for k in ("host", "port", "user", "password", "dbname", "sslmode") if vs_config.get(k)} - conn = psycopg2.connect(**conn_params) - conn.autocommit = True - try: - cur = conn.cursor() - try: - cur.execute( - "SELECT atttypmod FROM pg_attribute " - "WHERE attrelid = %s::regclass AND attname = 'vector'", - (collection_name,), - ) + with closing(psycopg2.connect(**conn_params)) as conn: + conn.autocommit = True + with closing(conn.cursor()) as cur: + cur.execute("SELECT atttypmod FROM pg_attribute WHERE attrelid = %s::regclass AND attname = 'vector'", (collection_name,)) 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))) - finally: - cur.close() - finally: - conn.close() - 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)) @@ -260,20 +203,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 - if hasattr(self._memory, "close"): - self._memory.close() vs = getattr(self._memory, "vector_store", None) - if vs and hasattr(vs, "close"): - vs.close() - client = getattr(vs, "client", None) - if client and hasattr(client, "close"): - client.close() - except Exception: - pass + # 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() diff --git a/plugins/memory/mem0/_openai_llm.py b/plugins/memory/mem0/_openai_llm.py index 64d1fe6312..a5e19ee446 100644 --- a/plugins/memory/mem0/_openai_llm.py +++ b/plugins/memory/mem0/_openai_llm.py @@ -11,6 +11,10 @@ from mem0.configs.llms.openai import OpenAIConfig from mem0.llms.base import LLMBase from mem0.llms.openai import OpenAILLM +# BaseLlmConfig fields copied into OpenAIConfig; the last two may be absent on older mem0. +_COPIED_FIELDS = ("model", "temperature", "api_key", "max_tokens", "top_p", "top_k", "enable_vision", "vision_details", "http_client_proxies") +_OPTIONAL_FIELDS = ("reasoning_effort", "is_reasoning_model") + class DirectOpenAILLM(OpenAILLM): """Use OpenAI credentials and requests regardless of router environment.""" @@ -21,15 +25,9 @@ class DirectOpenAILLM(OpenAILLM): elif isinstance(config, dict): config = OpenAIConfig(**config) elif isinstance(config, BaseLlmConfig) and not isinstance(config, OpenAIConfig): - config = OpenAIConfig( - model=config.model, temperature=config.temperature, api_key=config.api_key, - max_tokens=config.max_tokens, top_p=config.top_p, top_k=config.top_k, - enable_vision=config.enable_vision, vision_details=config.vision_details, - reasoning_effort=getattr(config, "reasoning_effort", None), - http_client_proxies=config.http_client_proxies, - is_reasoning_model=getattr(config, "is_reasoning_model", None), - ) - + fields = {k: getattr(config, k) for k in _COPIED_FIELDS} + fields.update({k: getattr(config, k, None) for k in _OPTIONAL_FIELDS}) + config = OpenAIConfig(**fields) if not config.model: config.model = "gpt-5-mini" # Configs predating the setup marker: keep the default model reasoning-safe @@ -39,21 +37,16 @@ class DirectOpenAILLM(OpenAILLM): # Bypass OpenAILLM.__init__ (it picks OpenRouter when OPENROUTER_API_KEY is # set); LLMBase still owns validation and supported-parameter filtering. LLMBase.__init__(self, config) - api_key = self.config.api_key or os.getenv("OPENAI_API_KEY") if not api_key: raise ValueError("OpenAI API key is required for the Hermes Mem0 OSS provider") from openai import OpenAI self.client = OpenAI(api_key=api_key, base_url=self.config.openai_base_url or os.getenv("OPENAI_BASE_URL") or "https://api.openai.com/v1") - def generate_response( - self, messages: List[Dict[str, str]], response_format=None, - tools: Optional[List[Dict]] = None, tool_choice: str = "auto", **kwargs, - ): + def generate_response(self, messages: List[Dict[str, str]], response_format=None, tools: Optional[List[Dict]] = None, tool_choice: str = "auto", **kwargs): params = self._get_supported_params(messages=messages, **kwargs) params.update({"model": self.config.model, "messages": messages}) - # No OpenRouter-only fields; ``store`` is opt-in so OpenAI-compatible - # endpoints never receive unknown fields. + # No OpenRouter-only fields; ``store`` is opt-in so OpenAI-compatible endpoints never receive unknown fields. if self.config.store is not None: params["store"] = self.config.store if response_format: diff --git a/plugins/memory/mem0/_oss_providers.py b/plugins/memory/mem0/_oss_providers.py index 69ead2009b..7fb87d998a 100644 --- a/plugins/memory/mem0/_oss_providers.py +++ b/plugins/memory/mem0/_oss_providers.py @@ -6,33 +6,17 @@ import os from typing import Any LLM_PROVIDERS: dict[str, dict[str, Any]] = { - "openai": { - "label": "OpenAI", "needs_key": True, "env_var": "OPENAI_API_KEY", - "default_model": "gpt-5-mini", "base_url_key": "openai_base_url", - }, - "ollama": { - "label": "Ollama (local)", "needs_key": False, "default_model": "llama3.1:8b", - "default_url": "http://localhost:11434", "base_url_key": "ollama_base_url", "pip_dep": "ollama", - }, + "openai": {"label": "OpenAI", "needs_key": True, "env_var": "OPENAI_API_KEY", "default_model": "gpt-5-mini", "base_url_key": "openai_base_url"}, + "ollama": {"label": "Ollama (local)", "needs_key": False, "default_model": "llama3.1:8b", "default_url": "http://localhost:11434", "base_url_key": "ollama_base_url", "pip_dep": "ollama"}, } EMBEDDER_PROVIDERS: dict[str, dict[str, Any]] = { - "openai": { - "label": "OpenAI", "needs_key": True, "env_var": "OPENAI_API_KEY", - "default_model": "text-embedding-3-small", "base_url_key": "openai_base_url", "dims": 1536, - }, - "ollama": { - "label": "Ollama (local)", "needs_key": False, "default_model": "nomic-embed-text", - "default_url": "http://localhost:11434", "base_url_key": "ollama_base_url", "dims": 768, "pip_dep": "ollama", - }, + "openai": {"label": "OpenAI", "needs_key": True, "env_var": "OPENAI_API_KEY", "default_model": "text-embedding-3-small", "base_url_key": "openai_base_url", "dims": 1536}, + "ollama": {"label": "Ollama (local)", "needs_key": False, "default_model": "nomic-embed-text", "default_url": "http://localhost:11434", "base_url_key": "ollama_base_url", "dims": 768, "pip_dep": "ollama"}, } VECTOR_PROVIDERS: dict[str, dict[str, Any]] = { - "qdrant": { - "label": "Qdrant", - "default_config": {"path": os.path.expanduser("~/.hermes/mem0_qdrant")}, - "pip_dep": "qdrant-client", - }, + "qdrant": {"label": "Qdrant", "default_config": {"path": os.path.expanduser("~/.hermes/mem0_qdrant")}, "pip_dep": "qdrant-client"}, "pgvector": { "label": "PGVector", "default_config": {"host": "localhost", "port": 5432, "user": os.getenv("USER", "postgres"), "dbname": "postgres"}, @@ -40,9 +24,7 @@ VECTOR_PROVIDERS: dict[str, dict[str, Any]] = { }, } -KNOWN_DIMS: dict[str, int] = { - "text-embedding-3-small": 1536, "text-embedding-3-large": 3072, "text-embedding-ada-002": 1536, "nomic-embed-text": 768, -} +KNOWN_DIMS: dict[str, int] = {"text-embedding-3-small": 1536, "text-embedding-3-large": 3072, "text-embedding-ada-002": 1536, "nomic-embed-text": 768} SECTION_REGISTRIES = (("llm", LLM_PROVIDERS), ("embedder", EMBEDDER_PROVIDERS), ("vector_store", VECTOR_PROVIDERS)) @@ -56,7 +38,6 @@ def validate_oss_config(oss_config: dict) -> list[str]: errors.append(f"Missing required section: {section}") elif block.get("provider", "") not in registry: errors.append(f"Unknown {section} provider '{block.get('provider', '')}'. Valid: {', '.join(registry.keys())}") - vs = oss_config.get("vector_store", {}) if vs.get("provider") == "pgvector" and not vs.get("config", {}).get("user"): errors.append("PGVector requires 'user' in vector_store.config") diff --git a/plugins/memory/mem0/_setup.py b/plugins/memory/mem0/_setup.py index 550cb7239b..e2dcfce907 100644 --- a/plugins/memory/mem0/_setup.py +++ b/plugins/memory/mem0/_setup.py @@ -4,6 +4,7 @@ from __future__ import annotations import getpass import json +from contextlib import suppress import os import shutil import socket @@ -20,9 +21,11 @@ from hermes_constants import get_hermes_home # noqa: F401 — patched by tests from . import _read_mem0_json from ._oss_providers import EMBEDDER_PROVIDERS, KNOWN_DIMS, LLM_PROVIDERS, SECTION_REGISTRIES, VECTOR_PROVIDERS, validate_oss_config +_OLLAMA_URL = "http://localhost:11434" +_PGVECTOR_CONTAINER, _PGVECTOR_IMAGE, _PGVECTOR_PASSWORD = "hermes-pgvector", "pgvector/pgvector:pg17", "hermes" + def _curses_select(title: str, items: list[tuple[str, str]], default: int = 0) -> int: - """Interactive single-select with arrow keys.""" from hermes_cli.curses_ui import curses_radiolist return curses_radiolist(title, [f"{label} {desc}" if desc else label for label, desc in items], selected=default, cancel_returns=default) @@ -36,7 +39,6 @@ def _prompt(label: str, default: str | None = None, secret: bool = False) -> str def _input(label: str, default: str) -> str: - """input() with a bracketed default shown and applied on blank.""" return input(f" {label} [{default}]: ").strip() or default @@ -52,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() @@ -67,12 +66,9 @@ def _api_key_writes(flags: dict, label: str, *, url: str | None = None, fresh_la if flags.get("api_key"): return {"MEM0_API_KEY": flags["api_key"]} existing = os.environ.get("MEM0_API_KEY", "") - if existing: - val = _prompt(f"{label} (current: {_masked(existing)}, blank to keep)", secret=True) - else: - if url: - print(f" Get yours at {url}") - val = _prompt(fresh_label or label, secret=True) + if url and not existing: + print(f" Get yours at {url}") + val = _prompt(f"{label} (current: {_masked(existing)}, blank to keep)" if existing else fresh_label or label, secret=True) return {"MEM0_API_KEY": val} if val else {} @@ -85,46 +81,25 @@ def _print_dry_run(summary: str, env_writes: dict, check=None) -> None: print(" [dry-run] No files written.\n") -def _print_saved(label: str, env_writes: dict, key_line: str, server: str | None = None) -> None: - 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") - - -_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- 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]: - """Parse CLI flags from argv. Returns dict of flag values.""" 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): if args[i] == "--dry-run": flags["dry_run"] = True - i += 1 elif args[i] in flag_map and i + 1 < len(args): flags[flag_map[args[i]]] = args[i + 1] - i += 2 - else: i += 1 + i += 1 return flags @@ -132,8 +107,7 @@ def _model_block(flags: dict, registry: dict, prefix: str) -> tuple[str, dict, d """Resolve (provider_id, provider_def, config) for an LLM/embedder section from flags.""" pid = flags.get(prefix, "openai") pdef = registry[pid] - model = flags.get(f"{prefix}_model") or pdef["default_model"] - cfg: dict[str, Any] = {"model": model} + cfg: dict[str, Any] = {"model": flags.get(f"{prefix}_model") or pdef["default_model"]} url = flags.get(f"{prefix}_url") or pdef.get("default_url") if url and pdef.get("base_url_key"): cfg[pdef["base_url_key"]] = url @@ -145,64 +119,35 @@ def build_oss_config(flags: dict[str, str]) -> tuple[dict, dict[str, str]]: llm_id, llm_def, llm_config = _model_block(flags, LLM_PROVIDERS, "oss_llm") if llm_id == "openai" and llm_config["model"] == "gpt-5-mini": llm_config["is_reasoning_model"] = True - embedder_id, embedder_def, embedder_config = _model_block(flags, EMBEDDER_PROVIDERS, "oss_embedder") 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, ()): - val = flags.get(f"oss_vector_{key}") - if val: + if val := flags.get(f"oss_vector_{key}"): 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}, - } - - env_writes: dict[str, str] = {} - if llm_def.get("needs_key") and flags.get("oss_llm_key"): - env_writes[llm_def["env_var"]] = flags["oss_llm_key"] - if embedder_def.get("needs_key"): - # An embedder sharing the LLM's provider reuses the LLM key when no embedder key was given. - key = flags.get("oss_embedder_key") or (flags.get("oss_llm_key") if embedder_id == llm_id else "") - if key: - env_writes[embedder_def["env_var"]] = key + 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 "" + emb_key = (flags.get("oss_embedder_key") or (flags.get("oss_llm_key") if embedder_id == llm_id else "")) if embedder_def.get("needs_key") else "" + env_writes = {d["env_var"]: k for d, k in ((llm_def, llm_key), (embedder_def, emb_key)) if k} return oss_config, env_writes def _write_env(env_path: Path, env_writes: dict[str, str]) -> None: - """Append or update env vars in .env file.""" env_path.parent.mkdir(parents=True, exist_ok=True) - # utf-8-sig like the canonical .env readers: locale decoding (cp1252/GBK) mangles - # non-ASCII values, and a BOM'd first line would miss the key match and get duplicated. + # utf-8-sig like the canonical .env readers: a BOM'd first line would miss the key match and get duplicated. existing_lines = env_path.read_text(encoding="utf-8-sig").splitlines() if env_path.exists() else [] - updated_keys: set[str] = set() - new_lines: list[str] = [] - for line in existing_lines: - key = line.split("=", 1)[0].strip() if "=" in line and not line.startswith("#") else None - if key and key in env_writes: - updated_keys.add(key) - line = f"{key}={env_writes[key]}" - new_lines.append(line) - new_lines += [f"{k}={v}" for k, v in env_writes.items() if k not in updated_keys] + keys = [line.split("=", 1)[0].strip() if "=" in line and not line.startswith("#") else None for line in existing_lines] + new_lines = [f"{k}={env_writes[k]}" if k in env_writes else line for k, line in zip(keys, existing_lines)] + new_lines += [f"{k}={v}" for k, v in env_writes.items() if k not in keys] env_path.write_text("\n".join(new_lines) + "\n", encoding="utf-8") -def _save_mem0_json(hermes_home: str, data: dict) -> None: - """Merge-write to mem0.json.""" - config_path = Path(hermes_home) / "mem0.json" - existing = _read_mem0_json(config_path) - existing.update(data) - config_path.write_text(json.dumps(existing, indent=2) + "\n", encoding="utf-8") - - def _activate_provider(config: dict) -> None: """Point config.yaml's memory.provider at mem0.""" from hermes_cli.config import save_config @@ -210,53 +155,40 @@ def _activate_provider(config: dict) -> None: save_config(config) -def _persist_provider_config(hermes_home: str, config: dict, provider_config: dict, env_writes: dict[str, str]) -> None: - """Shared platform/self-hosted tail: activate, write mem0.json (0600), then .env.""" +def _persist_provider_config(hermes_home: str, config: dict, provider_config: dict, env_writes: dict[str, str], label: str, key_line: str, server: str | None = None) -> None: + """Shared platform/self-hosted tail: activate, write mem0.json (0600), then .env, then a saved summary.""" _activate_provider(config) from plugins.memory.mem0 import Mem0MemoryProvider Mem0MemoryProvider().save_config(provider_config, hermes_home) if env_writes: _write_env(Path(hermes_home) / ".env", env_writes) + if server: + _check_selfhosted_server(server) + 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: """Platform mode setup — prompts for API key (secret -> .env), user/agent ids and rerank (-> mem0.json).""" provider_config = _read_mem0_json(Path(hermes_home) / "mem0.json") - print("\n Configuring mem0:\n") - env_writes = _api_key_writes(flags, "Mem0 Platform API key", url="https://app.mem0.ai") for key, desc, default in (("user_id", "User identifier", "hermes-user"), ("agent_id", "Agent identifier", "hermes")): - val = _prompt(desc, default=str(provider_config.get(key) or default)) - if val: + if val := _prompt(desc, default=str(provider_config.get(key) or default)): provider_config[key] = val choices = ["true", "false"] - current = provider_config.get("rerank", "false") - current_idx = choices.index(str(current).lower()) if current and str(current).lower() in choices else 0 - sel = _curses_select(" Enable reranking for recall", [(c, "") for c in choices], default=current_idx) - provider_config["rerank"] = choices[sel] - + 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 (_create_backend), so a stale - # self-hosted host must be cleared. Set "" rather than pop(): save_config - # merges into the existing mem0.json, so a popped key would survive. - provider_config["host"] = "" - # _load_config() also seeds ``host`` from MEM0_HOST (docs tell self-hosted - # users to put it in .env); the file clear can't help there, so warn. + # 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.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( - "\n ⚠ MEM0_HOST is set in your environment " - f"({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." - ) - - _persist_provider_config(hermes_home, config, provider_config, env_writes) - _print_saved("mem0", env_writes, "API keys saved to .env") + 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.") + _persist_provider_config(hermes_home, config, provider_config, env_writes, "mem0", "API keys saved to .env") def _check_selfhosted_server(host: str) -> None: @@ -274,56 +206,41 @@ def _check_selfhosted_server(host: str) -> None: def _setup_selfhosted(hermes_home: str, config: dict, flags: dict[str, str]) -> None: """Self-hosted mode — point at an existing Mem0 server: URL -> mem0.json, key -> .env (MEM0_API_KEY).""" provider_config = _read_mem0_json(Path(hermes_home) / "mem0.json") - print("\n Configuring mem0 (self-hosted server):\n") - host = flags.get("host") or _prompt("Mem0 server URL (e.g. http://localhost:8888)", default=provider_config.get("host") or None) if not host: print(" Error: a server URL is required for self-hosted mode.", file=sys.stderr) return host = host.rstrip("/") - 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 - provider_config.update(mode="platform", host=host, user_id=user_id, agent_id=agent_id) # routing: oss > host > platform - _persist_provider_config(hermes_home, config, provider_config, env_writes) - - _check_selfhosted_server(host) - _print_saved("mem0 (self-hosted)", env_writes, "API key saved to .env", server=host) + _persist_provider_config(hermes_home, config, provider_config, env_writes, "mem0 (self-hosted)", "API key saved to .env", server=host) 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: +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: """Shared OSS tail: write secrets + mem0.json, install deps, activate, check, summarize.""" if env_writes: _write_env(Path(hermes_home) / ".env", env_writes) - _save_mem0_json(hermes_home, {"mode": "oss", "user_id": user_id, "agent_id": agent_id, "oss": oss_config}) + config_path = Path(hermes_home) / "mem0.json" # merge-write, plain text (platform path uses save_config's 0600 atomic write) + config_path.write_text(json.dumps({**_read_mem0_json(config_path), "mode": "oss", "user_id": user_id, "agent_id": agent_id, "oss": oss_config}, indent=2) + "\n", encoding="utf-8") _install_provider_deps(oss_config["llm"]["provider"], oss_config["embedder"]["provider"], oss_config["vector_store"]["provider"]) if pgvector_config: _ensure_pgvector_extension(pgvector_config) @@ -338,80 +255,57 @@ def _setup_oss(hermes_home: str, config: dict, flags: dict[str, str]) -> None: _setup_oss_interactive(hermes_home, config) return oss_config, env_writes = build_oss_config(flags) - errors = validate_oss_config(oss_config) - if errors: - for e in errors: - print(f" Error: {e}", file=sys.stderr) + if errors := validate_oss_config(oss_config): + print("".join(f" Error: {e}\n" for e in errors), end="", file=sys.stderr) sys.exit(1) if flags.get("dry_run"): _print_oss_summary(oss_config, env_writes, dry_run=True) _run_connectivity_checks(oss_config) print(" [dry-run] No files written.\n") return - _finish_oss(hermes_home, config, oss_config, env_writes, flags.get("user_id") or os.getenv("USER", "hermes-user"), "hermes") -_PGVECTOR_CONTAINER = "hermes-pgvector" -_PGVECTOR_IMAGE = "pgvector/pgvector:pg17" -_PGVECTOR_PASSWORD = "hermes" - - def _docker(*args: str, timeout: int, **kwargs) -> subprocess.CompletedProcess: return subprocess.run(["docker", *args], capture_output=True, timeout=timeout, stdin=subprocess.DEVNULL, **kwargs) +def _pg_ready(host: str, port: int, wait: int) -> bool: + """Wait up to ``wait`` seconds for the port, then report whether PostgreSQL answers.""" + _wait_for_port(host, port, timeout=wait) + return _check_pgvector(host, port)[0] + + def _ensure_pgvector(host: str = "localhost", port: int = 5432) -> dict | None: - """Ensure pgvector is reachable; offer Docker setup if not. Returns the Docker - container's vector_config if one was started, None otherwise.""" + """Ensure pgvector is reachable, offering Docker if not; returns the started container's vector_config, else None.""" if _check_pgvector(host, port)[0]: print(f" ✓ PostgreSQL reachable at {host}:{port}") return None - print(f" PostgreSQL not reachable at {host}:{port}") 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: - result = _docker("inspect", _PGVECTOR_CONTAINER, "--format", "{{.State.Status}}", - timeout=10, text=True, encoding='utf-8', errors='replace') + 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...") _docker("start", _PGVECTOR_CONTAINER, timeout=15) - _wait_for_port(host, port, timeout=15) - if _check_pgvector(host, port)[0]: + 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.") - return None - - -def _start_pgvector_docker(host: str, port: int) -> dict | None: - """Pull and start pgvector Docker container.""" + 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) _docker("rm", "-f", _PGVECTOR_CONTAINER, timeout=10) # remove existing container if present print(f" Starting container '{_PGVECTOR_CONTAINER}' on port {port}...") - _docker( - "run", "-d", "--name", _PGVECTOR_CONTAINER, - "-e", f"POSTGRES_PASSWORD={_PGVECTOR_PASSWORD}", - "-p", f"{port}:5432", _PGVECTOR_IMAGE, - timeout=30, check=True, - ) - _wait_for_port(host, port, timeout=20) - if _check_pgvector(host, port)[0]: + _docker("run", "-d", "--name", _PGVECTOR_CONTAINER, "-e", f"POSTGRES_PASSWORD={_PGVECTOR_PASSWORD}", "-p", f"{port}:5432", _PGVECTOR_IMAGE, timeout=30, check=True) + if _pg_ready(host, port, 20): print(f" ✓ pgvector running on {host}:{port}") else: - print(" Warning: Container started but PostgreSQL not yet accepting connections.\n" - " It may need a few more seconds. Config will be saved; retry later.") + print(" Warning: Container started but PostgreSQL not yet accepting connections.\n It may need a few more seconds. Config will be saved; retry later.") return {"host": host, "port": port, "user": "postgres", "password": _PGVECTOR_PASSWORD, "dbname": "postgres"} except subprocess.CalledProcessError as e: print(f" Failed to start Docker container: {e}") @@ -421,33 +315,29 @@ def _start_pgvector_docker(host: str, port: int) -> dict | None: def _ensure_ollama(models: list[str]) -> bool: - """Ensure Ollama is running and required models are pulled. Returns False when - the user must handle it manually.""" - url = "http://localhost:11434" + """Ensure Ollama is running and ``models`` are pulled; False when the user must handle it manually.""" ollama_bin = shutil.which("ollama") - ok = _check_ollama(url)[0] - if not ok: + if not (ok := _check_ollama(_OLLAMA_URL)[0]): if not ollama_bin: - print(" Ollama not found. Install it:\n curl -fsSL https://ollama.com/install.sh | sh\n" - " Or on macOS: brew install ollama") + print(" Ollama not found. Install it:\n curl -fsSL https://ollama.com/install.sh | sh\n Or on macOS: brew install ollama") return False print(" Ollama installed but not running. Starting...") try: - subprocess.Popen( - [ollama_bin, "serve"], stdin=subprocess.DEVNULL, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL - ) + subprocess.Popen([ollama_bin, "serve"], stdin=subprocess.DEVNULL, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL) _wait_for_port("localhost", 11434, timeout=10) - ok = _check_ollama(url)[0] - if ok: + if ok := _check_ollama(_OLLAMA_URL)[0]: print(" ✓ Ollama started") except Exception as e: print(f" Could not start Ollama: {e}") if not ok: print(" Warning: Ollama not reachable. Models cannot be pulled.") return False - for model in models: - if _ollama_has_model(url, model): + 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)") @@ -459,27 +349,14 @@ def _ensure_ollama(models: list[str]) -> bool: return True -def _ollama_has_model(url: str, model: str) -> bool: - """Check if Ollama already has a model pulled.""" - try: - names = [m.get("name", "") for m in json.loads(_http_get(url, "/api/tags", 5).read()).get("models", [])] - base_model = model.split(":")[0] - return any(model in n or base_model in n for n in names) - except Exception: - return False - - def _ensure_pgvector_extension(pg_config: dict) -> None: - """Create the pgvector extension if it doesn't exist.""" try: import psycopg2 except ImportError: return - conn_params = {k: pg_config.get(k, d) for k, d in (("host", "localhost"), ("port", 5432), ("user", "postgres"), ("dbname", "postgres"))} - if pg_config.get("password"): - conn_params["password"] = pg_config["password"] + defaults = {"host": "localhost", "port": 5432, "user": "postgres", "dbname": "postgres"} try: - conn = psycopg2.connect(**conn_params) + conn = psycopg2.connect(**(defaults | {k: v for k, v in pg_config.items() if k in defaults or (k == "password" and v)})) conn.autocommit = True conn.cursor().execute("CREATE EXTENSION IF NOT EXISTS vector") conn.close() @@ -489,7 +366,6 @@ def _ensure_pgvector_extension(pg_config: dict) -> None: def _wait_for_port(host: str, port: int, timeout: int = 15) -> None: - """Wait until a TCP port is accepting connections.""" deadline = time.monotonic() + timeout while time.monotonic() < deadline: try: @@ -499,34 +375,20 @@ def _wait_for_port(host: str, port: int, timeout: int = 15) -> None: time.sleep(0.5) -def _provider_description(v: dict) -> str: - """Description for LLM/embedder picker: model + URL if applicable.""" - model, url = v.get("default_model", ""), v.get("default_url") - return f"{model} ({url})" if url else model +# Picker descriptions: LLM/embedder show model (+ URL); vector stores by provider id (default: the id itself). +_VECTOR_DESCRIPTIONS = {"qdrant": lambda cfg: cfg.get("path", "local storage"), "pgvector": lambda cfg: f"{cfg.get('host', 'localhost')}:{cfg.get('port', 5432)}"} -def _vector_description(pid: str, v: dict) -> str: - cfg = v.get("default_config", {}) - if pid == "qdrant": - return cfg.get("path", "local storage") - return f"{cfg.get('host', 'localhost')}:{cfg.get('port', 5432)}" if pid == "pgvector" else pid - - -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. - Returns (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()] +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"], 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") if pdef["needs_key"]: if llm is None or pid != llm[0]: - label = pdef["label"] if llm is None else f"{pdef['label']} embedder" - key = _prompt_api_key(label, pdef["env_var"], hermes_home) - if key: + if key := _prompt_api_key(pdef["label"] if llm is None else f"{pdef['label']} embedder", pdef["env_var"], hermes_home): env_writes[pdef["env_var"]] = key elif llm[1].get("env_var") in env_writes: env_writes[pdef["env_var"]] = env_writes[llm[1]["env_var"]] @@ -537,106 +399,71 @@ def _configure_model_provider( def _setup_oss_interactive(hermes_home: str, config: dict) -> None: - """Interactive OSS setup using curses pickers.""" 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_description(pid, v)) for pid, v in VECTOR_PROVIDERS.items()] + 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 + # 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) - - # Auto-setup: ensure pgvector is reachable (offer Docker if not) - pgvector_config = None - if vector_id == "pgvector": - pgvector_config = _ensure_pgvector() - if not pgvector_config: - # 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() - pgvector_config = {"host": pg["host"], "port": int(pg["port"]), "user": pg["user"], "dbname": pg["dbname"]} - if pg_password: - pgvector_config["password"] = pg_password - + 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 "", "oss_embedder": embedder_id, "oss_embedder_model": embedder_model, "oss_embedder_url": embedder_url or "", "oss_vector": vector_id, "user_id": user_id, } - if pgvector_config: - for key in ("host", "port", "user", "password", "dbname"): - if pgvector_config.get(key): - flags[f"oss_vector_{key}"] = str(pgvector_config[key]) - + flags.update({f"oss_vector_{key}": str(val) for key, val in (pgvector_config or {}).items() if val}) oss_config, _ = build_oss_config(flags) _finish_oss(hermes_home, config, oss_config, env_writes, user_id, agent_id, pgvector_config) def _install_provider_deps(llm_id: str, embedder_id: str, vector_id: str) -> None: - """Install all optional pip deps for selected providers.""" - deps = { - registry[pid]["pip_dep"] - for (_, registry), pid in zip(SECTION_REGISTRIES, (llm_id, embedder_id, vector_id)) - if registry.get(pid, {}).get("pip_dep") - } + deps = {registry[pid]["pip_dep"] for (_, registry), pid in zip(SECTION_REGISTRIES, (llm_id, embedder_id, vector_id)) if registry.get(pid, {}).get("pip_dep")} for dep in sorted(deps): + print(f" Installing {dep}...") try: - print(f" Installing {dep}...") - # Environment-aware install: sealed hosted venvs redirect to the - # durable data-volume target instead of /opt/hermes. + # Environment-aware install: sealed hosted venvs redirect to the durable data-volume target instead of /opt/hermes. from tools.lazy_deps import install_specs outcome = install_specs([dep], timeout=60) - if outcome.ok: - print(f" ✓ Installed {dep}") - elif outcome.blocked: - print(f" Warning: cannot install {dep}: {outcome.reason}") - else: - print(f" Warning: Could not install {dep}. Install manually: uv pip install {dep}") except Exception: - print(f" Warning: Could not install {dep}. Install manually: uv pip install {dep}") + outcome = None + 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() +def _probe(fn, ok: str, fail: str, exc=Exception) -> tuple[bool, str]: + """Run ``fn``; (True, ok) on success, (False, "fail: ") on ``exc``.""" + try: + fn() + return True, ok + except exc as e: + return False, f"{fail}: {e}" + + def _check_qdrant_path(path: str) -> tuple[bool, str]: """Check that qdrant local storage parent dir is writable.""" parent = Path(path).expanduser().parent - try: - parent.mkdir(parents=True, exist_ok=True) - return True, f"Directory writable: {parent}" - except OSError as e: - return False, f"Cannot write to {parent}: {e}" + return _probe(lambda: parent.mkdir(parents=True, exist_ok=True), f"Directory writable: {parent}", f"Cannot write to {parent}", OSError) def _check_ollama(url: str) -> tuple[bool, str]: - """Check Ollama is reachable via /api/tags.""" - try: - _http_get(url, "/api/tags", 3) - return True, "Ollama reachable" - except Exception as e: - return False, f"Ollama not reachable at {url}: {e}" + return _probe(lambda: _http_get(url, "/api/tags", 3), "Ollama reachable", f"Ollama not reachable at {url}") def _check_pgvector(host: str, port: int) -> tuple[bool, str]: - """Check PGVector via TCP socket.""" - try: - socket.create_connection((host, port), timeout=3).close() - return True, f"PGVector reachable at {host}:{port}" - except Exception as e: - return False, f"PGVector not reachable at {host}:{port}: {e}" + return _probe(lambda: socket.create_connection((host, port), timeout=3).close(), f"PGVector reachable at {host}:{port}", f"PGVector not reachable at {host}:{port}") def _warn_unless(check: tuple[bool, str]) -> None: @@ -646,7 +473,6 @@ def _warn_unless(check: tuple[bool, str]) -> None: def _run_connectivity_checks(oss_config: dict) -> None: - """Run connectivity checks and print warnings.""" vs = oss_config.get("vector_store", {}) cfg = vs.get("config", {}) if vs.get("provider") == "qdrant": @@ -654,50 +480,28 @@ def _run_connectivity_checks(oss_config: dict) -> None: if path: _warn_unless(_check_qdrant_path(path)) elif url: - try: - _http_get(url, "/healthz", 3) - except Exception as e: - print(f" Warning: Qdrant not reachable at {url}: {e}") + _warn_unless(_probe(lambda: _http_get(url, "/healthz", 3), "Qdrant reachable", f"Qdrant not reachable at {url}")) elif vs.get("provider") == "pgvector": _warn_unless(_check_pgvector(cfg.get("host", "localhost"), cfg.get("port", 5432))) - llm = oss_config.get("llm", {}) if llm.get("provider") == "ollama": - _warn_unless(_check_ollama(llm.get("config", {}).get("ollama_base_url", "http://localhost:11434"))) + _warn_unless(_check_ollama(llm.get("config", {}).get("ollama_base_url", _OLLAMA_URL))) -def _check_min_dep_version() -> None: - """Ensure mem0ai meets the minimum version from plugin.yaml.""" - try: - 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" - f" Run: uv pip install --python {sys.executable} 'mem0ai>=2.0.7'") - except Exception: - pass - - -_MODE_HANDLERS = { - "oss": _setup_oss, - "selfhosted": _setup_selfhosted, - "self-hosted": _setup_selfhosted, - "platform": _setup_platform, -} +_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 called by hermes memory setup framework. Routes on --mode - (platform / selfhosted / oss); with no flag shows a picker. OSS is - non-interactive only when the mode came from the flag.""" - _check_min_dep_version() + """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.""" + 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'") flags = parse_flags(sys.argv[1:]) handler = _MODE_HANDLERS.get(flags["mode"]) flags["_mode_from_flag"] = handler is not None diff --git a/plugins/memory/retaindb/__init__.py b/plugins/memory/retaindb/__init__.py index ed21555795..fe8f2972c9 100644 --- a/plugins/memory/retaindb/__init__.py +++ b/plugins/memory/retaindb/__init__.py @@ -2,9 +2,8 @@ Cross-session memory via the RetainDB cloud API: durable SQLite write-behind queue, semantic search + profile, context overlay, dialectic/agent self-model prefetch, shared file store tools. - -Config (env vars, or config.yaml ``memory.retaindb`` for the non-secret ones): RETAINDB_API_KEY (required), -RETAINDB_BASE_URL (default https://api.retaindb.com), RETAINDB_PROJECT (optional; defaults to "default"). +Config: RETAINDB_API_KEY (required, scoped secret), RETAINDB_BASE_URL (default https://api.retaindb.com), +RETAINDB_PROJECT (optional; defaults to "default"); the non-secret two also read config.yaml ``memory.retaindb``. """ from __future__ import annotations @@ -20,7 +19,7 @@ import time from contextlib import suppress from datetime import datetime, timezone from pathlib import Path -from typing import Any, Callable, Dict, List +from typing import Any, Callable from urllib.parse import quote from agent.memory_provider import MemoryProvider @@ -35,103 +34,83 @@ _ASYNC_SHUTDOWN = object() _TEXT_EXTS = (".txt", ".md", ".json", ".csv", ".yaml", ".yml", ".xml", ".html") -def _load_retaindb_config() -> Dict[str, Any]: +def _load_retaindb_config() -> dict[str, Any]: """``memory.retaindb`` block from config.yaml (empty on error): Dashboard-persisted base_url/project; api_key stays in scoped secrets.""" try: from hermes_cli.config import load_config_readonly - - provider_config = load_config_readonly().get("memory", {}).get("retaindb", {}) - return dict(provider_config) if isinstance(provider_config, dict) else {} + block = load_config_readonly().get("memory", {}).get("retaindb", {}) except Exception: - return {} - - -def _config_str(value: Any) -> str: - """Stripped string for a config value, else ``""``.""" - return value.strip() if isinstance(value, str) else "" + block = None + return dict(block) if isinstance(block, dict) else {} def _q(s: str) -> str: return quote(s, safe="") -# ── Tool schemas ───────────────────────────────────────────────────────────── +def _quiet(label: str, fn: Callable[[], Any]) -> Any: + """Run *fn*; on any exception log "RetainDB