merge(r3-27): group D

This commit is contained in:
Teknium
2026-09-02 22:52:17 -07:00
12 changed files with 1077 additions and 2244 deletions
+57 -102
View File
@@ -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:
+91 -196
View File
@@ -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()))
+40 -80
View File
@@ -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
+104 -225
View File
@@ -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
+88 -181
View File
@@ -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
+136 -269
View File
@@ -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.<method>(*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):
+35 -99
View File
@@ -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()
+9 -16
View File
@@ -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:
+6 -25
View File
@@ -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")
+127 -323
View File
@@ -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-<key> flags accepted per vector store (also the pgvector key order).
_VECTOR_FLAG_KEYS = {"qdrant": ("path", "url"), "pgvector": ("host", "port", "user", "password", "dbname")}
_FLAG_KEYS = ("mode", "api_key", "host", *(f"oss_{s}{k}" for s in ("llm", "embedder") for k in ("", "_key", "_model", "_url")),
"oss_vector", *(f"oss_vector_{k}" for ks in _VECTOR_FLAG_KEYS.values() for k in ks), "user_id")
_FLAG_DEFAULTS = {"oss_llm": "openai", "oss_embedder": "openai", "oss_vector": "qdrant"}
def parse_flags(argv: list[str] | None = None) -> dict[str, str]:
"""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: <error>") 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
+147 -259
View File
@@ -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 <label> failed" at debug and return None."""
try:
return fn()
except Exception as exc:
logger.debug("RetainDB %s failed: %s", label, exc)
return None
def _schema(name: str, description: str, properties: dict | None = None, required: tuple = ()) -> dict:
return {
"name": name,
"description": description,
"parameters": {"type": "object", "properties": properties or {}, "required": list(required)},
}
return {"name": name, "description": description,
"parameters": {"type": "object", "properties": properties or {}, "required": list(required)}}
def _prop(type_: str, description: str, **extra) -> dict:
def _p(description: str, type_: str = "string", **extra) -> dict:
return {"type": type_, **extra, "description": description}
def _s(description: str, **extra) -> dict:
return _prop("string", description, **extra)
PROFILE_SCHEMA = _schema(
"retaindb_profile", "Get the user's stable profile — preferences, facts, and patterns recalled from long-term memory.")
SEARCH_SCHEMA = _schema(
"retaindb_search", "Semantic search across stored memories. Returns ranked results with relevance scores.",
{"query": _s("What to search for."), "top_k": _prop("integer", "Max results (default: 8, max: 20).")}, ("query",))
CONTEXT_SCHEMA = _schema(
"retaindb_context", "Synthesized context block — what matters most for the current task, pulled from long-term memory.",
{"query": _s("Current task or question.")}, ("query",))
REMEMBER_SCHEMA = _schema(
"retaindb_remember", "Persist an explicit fact, preference, or decision to long-term memory.",
{"content": _s("The fact to remember."),
"memory_type": _s("Category (default: factual).", enum=["factual", "preference", "goal", "instruction", "event", "opinion"]),
"importance": _prop("number", "Importance 0-1 (default: 0.7).")}, ("content",))
FORGET_SCHEMA = _schema("retaindb_forget", "Delete a specific memory by ID.", {"memory_id": _s("Memory ID to delete.")}, ("memory_id",))
FILE_UPLOAD_SCHEMA = _schema(
"retaindb_upload_file", "Upload a file to the shared RetainDB file store. Returns an rdb:// URI any agent can reference.",
{"local_path": _s("Local file path to upload."), "remote_path": _s("Destination path, e.g. /reports/q1.pdf"),
"scope": _s("Access scope (default: PROJECT).", enum=["USER", "PROJECT", "ORG"]),
"ingest": _prop("boolean", "Also extract memories from file after upload (default: false).")}, ("local_path",))
FILE_LIST_SCHEMA = _schema(
"retaindb_list_files", "List files in the shared file store.",
{"prefix": _s("Path prefix to filter by, e.g. /reports/"), "limit": _prop("integer", "Max results (default: 50).")})
FILE_READ_SCHEMA = _schema(
"retaindb_read_file", "Read the text content of a stored file by its file ID.",
{"file_id": _s("File ID returned from upload or list.")}, ("file_id",))
FILE_INGEST_SCHEMA = _schema(
"retaindb_ingest_file", "Chunk, embed, and extract memories from a stored file. Makes its contents searchable.",
{"file_id": _s("File ID to ingest.")}, ("file_id",))
FILE_DELETE_SCHEMA = _schema("retaindb_delete_file", "Delete a stored file.", {"file_id": _s("File ID to delete.")}, ("file_id",))
_SCHEMAS = (
PROFILE_SCHEMA, SEARCH_SCHEMA, CONTEXT_SCHEMA, REMEMBER_SCHEMA, FORGET_SCHEMA,
FILE_UPLOAD_SCHEMA, FILE_LIST_SCHEMA, FILE_READ_SCHEMA, FILE_INGEST_SCHEMA, FILE_DELETE_SCHEMA,
_schema("retaindb_profile", "Get the user's stable profile — preferences, facts, and patterns recalled from long-term memory."),
_schema("retaindb_search", "Semantic search across stored memories. Returns ranked results with relevance scores.",
{"query": _p("What to search for."), "top_k": _p("Max results (default: 8, max: 20).", "integer")}, ("query",)),
_schema("retaindb_context", "Synthesized context block — what matters most for the current task, pulled from long-term memory.",
{"query": _p("Current task or question.")}, ("query",)),
_schema("retaindb_remember", "Persist an explicit fact, preference, or decision to long-term memory.",
{"content": _p("The fact to remember."),
"memory_type": _p("Category (default: factual).", enum=["factual", "preference", "goal", "instruction", "event", "opinion"]),
"importance": _p("Importance 0-1 (default: 0.7).", "number")}, ("content",)),
_schema("retaindb_forget", "Delete a specific memory by ID.", {"memory_id": _p("Memory ID to delete.")}, ("memory_id",)),
_schema("retaindb_upload_file", "Upload a file to the shared RetainDB file store. Returns an rdb:// URI any agent can reference.",
{"local_path": _p("Local file path to upload."), "remote_path": _p("Destination path, e.g. /reports/q1.pdf"),
"scope": _p("Access scope (default: PROJECT).", enum=["USER", "PROJECT", "ORG"]),
"ingest": _p("Also extract memories from file after upload (default: false).", "boolean")}, ("local_path",)),
_schema("retaindb_list_files", "List files in the shared file store.",
{"prefix": _p("Path prefix to filter by, e.g. /reports/"), "limit": _p("Max results (default: 50).", "integer")}),
_schema("retaindb_read_file", "Read the text content of a stored file by its file ID.",
{"file_id": _p("File ID returned from upload or list.")}, ("file_id",)),
_schema("retaindb_ingest_file", "Chunk, embed, and extract memories from a stored file. Makes its contents searchable.",
{"file_id": _p("File ID to ingest.")}, ("file_id",)),
_schema("retaindb_delete_file", "Delete a stored file.", {"file_id": _p("File ID to delete.")}, ("file_id",)),
)
# ── HTTP client ──────────────────────────────────────────────────────────────
class _Client:
"""Thin HTTP client over the RetainDB REST API (lazy ``requests`` import)."""
def __init__(self, api_key: str, base_url: str, project: str):
self.api_key = api_key
self.base_url = re.sub(r"/+$", "", base_url)
self.project = project
self.api_key, self.base_url, self.project = api_key, re.sub(r"/+$", "", base_url), project
def _headers(self, path: str, json_body: bool = True) -> dict:
token = self.api_key.replace("Bearer ", "").strip()
return {
"Authorization": f"Bearer {token}", "x-sdk-runtime": "hermes-plugin",
**({"Content-Type": "application/json"} if json_body else {}),
# memory/context routes also accept the key as X-API-Key
**({"X-API-Key": token} if path.startswith(("/v1/memory", "/v1/context")) else {}),
}
return {"Authorization": f"Bearer {token}", "x-sdk-runtime": "hermes-plugin",
**({"Content-Type": "application/json"} if json_body else {}),
**({"X-API-Key": token} if path.startswith(("/v1/memory", "/v1/context")) else {})} # memory/context also accept X-API-Key
def _http(self, method: str, path: str, *, json_body: bool = True, timeout: float = 30, **kwargs):
import requests
return requests.request(method, f"{self.base_url}{path}", headers=self._headers(path, json_body), timeout=timeout, **kwargs)
def request(self, method: str, path: str, *, params=None, json_body=None, timeout: float = 8.0) -> Any:
import requests
"""JSON request; raises RuntimeError carrying the server message on a non-2xx response."""
method = method.upper()
resp = requests.request(
method, f"{self.base_url}{path}", params=params, json=json_body if method not in {"GET", "DELETE"} else None,
headers=self._headers(path), timeout=timeout,
)
resp = self._http(method, path, params=params, json=json_body if method not in {"GET", "DELETE"} else None, timeout=timeout)
try:
payload = resp.json()
except Exception:
@@ -141,6 +120,11 @@ class _Client:
raise RuntimeError(f"RetainDB {method} {path} failed ({resp.status_code}): {msg or payload}")
return payload
def _raw(self, method: str, path: str, **kwargs) -> Any:
"""Non-JSON request (multipart upload / binary download); raises on HTTP error."""
resp = self._http(method, path, json_body=False, **kwargs)
return resp.raise_for_status() or resp
@staticmethod
def _with_fallback(primary: Callable[[], dict], fallback: Callable[[], dict]) -> dict:
"""Try the current API route; on any error retry via the legacy route."""
@@ -152,144 +136,100 @@ class _Client:
def _scoped(self, user_id: str, session_id: str, **extra) -> dict:
return {"project": self.project, "user_id": user_id, "session_id": session_id, **extra}
# Memory
# Memory routes (one endpoint per method; bodies are the wire payloads)
def query_context(self, user_id: str, session_id: str, query: str, max_tokens: int = 1200) -> dict:
body = self._scoped(user_id, session_id, query=query, include_memories=True, max_tokens=max_tokens)
return self.request("POST", "/v1/context/query", json_body=body)
return self.request("POST", "/v1/context/query", json_body=self._scoped(user_id, session_id, query=query, include_memories=True, max_tokens=max_tokens))
def search(self, user_id: str, session_id: str, query: str, top_k: int = 8) -> dict:
body = self._scoped(user_id, session_id, query=query, top_k=top_k, include_pending=True)
return self.request("POST", "/v1/memory/search", json_body=body)
return self.request("POST", "/v1/memory/search", json_body=self._scoped(user_id, session_id, query=query, top_k=top_k, include_pending=True))
def get_profile(self, user_id: str) -> dict:
return self._with_fallback(
lambda: self.request("GET", f"/v1/memory/profile/{_q(user_id)}", params={"project": self.project, "include_pending": "true"}),
lambda: self.request("GET", "/v1/memories", params={"project": self.project, "user_id": user_id, "limit": "200"}),
)
lambda: self.request("GET", "/v1/memories", params={"project": self.project, "user_id": user_id, "limit": "200"}))
def add_memory(self, user_id: str, session_id: str, content: str, memory_type: str = "factual", importance: float = 0.7) -> dict:
body = self._scoped(user_id, session_id, content=content, memory_type=memory_type, importance=importance)
return self._with_fallback(
lambda: self.request("POST", "/v1/memory", json_body={**body, "write_mode": "sync"}, timeout=5.0),
lambda: self.request("POST", "/v1/memories", json_body=body, timeout=5.0),
)
lambda: self.request("POST", "/v1/memories", json_body=body, timeout=5.0))
def delete_memory(self, memory_id: str) -> dict:
return self._with_fallback(
lambda: self.request("DELETE", f"/v1/memory/{_q(memory_id)}", timeout=5.0),
lambda: self.request("DELETE", f"/v1/memories/{_q(memory_id)}", timeout=5.0),
)
lambda: self.request("DELETE", f"/v1/memories/{_q(memory_id)}", timeout=5.0))
def ingest_session(self, user_id: str, session_id: str, messages: list, timeout: float = 15.0) -> dict:
body = self._scoped(user_id, session_id, messages=messages, write_mode="sync")
return self.request("POST", "/v1/memory/ingest/session", json_body=body, timeout=timeout)
return self.request("POST", "/v1/memory/ingest/session", json_body=self._scoped(user_id, session_id, messages=messages, write_mode="sync"), timeout=timeout)
def ask_user(self, user_id: str, query: str, reasoning_level: str = "low") -> dict:
body = {"project": self.project, "query": query, "reasoning_level": reasoning_level}
return self.request("POST", f"/v1/memory/profile/{_q(user_id)}/ask", json_body=body, timeout=8.0)
return self.request("POST", f"/v1/memory/profile/{_q(user_id)}/ask", json_body={"project": self.project, "query": query, "reasoning_level": reasoning_level}, timeout=8.0)
def get_agent_model(self, agent_id: str) -> dict:
return self.request("GET", f"/v1/memory/agent/{_q(agent_id)}/model", params={"project": self.project}, timeout=4.0)
def seed_agent_identity(self, agent_id: str, content: str, source: str = "soul_md") -> dict:
body = {"project": self.project, "content": content, "source": source}
return self.request("POST", f"/v1/memory/agent/{_q(agent_id)}/seed", json_body=body, timeout=20.0)
# Files
def _raw(self, method: str, path: str, **kwargs) -> Any:
"""Non-JSON request (multipart upload / binary download); raises on HTTP error."""
import requests
resp = requests.request(method, f"{self.base_url}{path}", headers=self._headers(path, json_body=False), timeout=30, **kwargs)
resp.raise_for_status()
return resp
return self.request("POST", f"/v1/memory/agent/{_q(agent_id)}/seed", json_body={"project": self.project, "content": content, "source": source}, timeout=20.0)
# File routes
def upload_file(self, data: bytes, filename: str, remote_path: str, mime_type: str, scope: str, project_id: str | None) -> dict:
import io
fields = {"path": remote_path, "scope": scope.upper(), **({"project_id": project_id} if project_id else {})}
return self._raw("POST", "/v1/files", files={"file": (filename, io.BytesIO(data), mime_type)}, data=fields).json()
def list_files(self, prefix: str | None = None, limit: int = 50) -> dict:
return self.request("GET", "/v1/files", params={"limit": limit, **({"prefix": prefix} if prefix else {})})
def get_file(self, file_id: str) -> dict:
return self.request("GET", f"/v1/files/{_q(file_id)}")
def read_file_content(self, file_id: str) -> bytes:
return self._raw("GET", f"/v1/files/{_q(file_id)}/content", allow_redirects=True).content
def ingest_file(self, file_id: str, user_id: str | None = None, agent_id: str | None = None) -> dict:
body = {k: v for k, v in (("user_id", user_id), ("agent_id", agent_id)) if v}
return self.request("POST", f"/v1/files/{_q(file_id)}/ingest", json_body=body, timeout=60.0)
def delete_file(self, file_id: str) -> dict:
return self.request("DELETE", f"/v1/files/{_q(file_id)}", timeout=5.0)
# ── Durable write-behind queue ───────────────────────────────────────────────
class _WriteQueue:
"""SQLite-backed async write queue. Survives crashes — pending rows replay on startup."""
def __init__(self, client: _Client, db_path: Path):
self._client = client
self._db_path = db_path
self._q: queue.Queue = queue.Queue()
self._client, self._db_path, self._q = client, db_path, queue.Queue()
self._thread = threading.Thread(target=self._loop, name="retaindb-writer", daemon=True)
self._db_path.parent.mkdir(parents=True, exist_ok=True)
self._local = threading.local() # one cached connection per thread
db_path.parent.mkdir(parents=True, exist_ok=True)
self._local = threading.local() # one cached connection per thread, all tracked in _connections
self._connections: set[sqlite3.Connection] = set()
self._connections_lock = threading.Lock()
self._shutdown_lock = threading.Lock()
self._shutdown = False
conn = self._execute(
"CREATE TABLE IF NOT EXISTS pending (id INTEGER PRIMARY KEY AUTOINCREMENT, user_id TEXT, "
"session_id TEXT, messages_json TEXT, created_at TEXT, last_error TEXT)"
).connection
self._connections_lock, self._shutdown_lock, self._shutdown = threading.Lock(), threading.Lock(), False
conn = self._execute("CREATE TABLE IF NOT EXISTS pending (id INTEGER PRIMARY KEY AUTOINCREMENT, user_id TEXT, "
"session_id TEXT, messages_json TEXT, created_at TEXT, last_error TEXT)").connection
self._thread.start()
# Replay any rows left from a previous crash
for row_id, user_id, session_id, msgs_json in conn.execute(
"SELECT id, user_id, session_id, messages_json FROM pending ORDER BY id ASC LIMIT 200"
).fetchall():
replay = conn.execute("SELECT id, user_id, session_id, messages_json FROM pending ORDER BY id ASC LIMIT 200").fetchall()
for row_id, user_id, session_id, msgs_json in replay: # rows left from a previous crash
self._q.put((row_id, user_id, session_id, json.loads(msgs_json)))
def _get_conn(self) -> sqlite3.Connection:
"""Return a cached connection for the current thread."""
conn = getattr(self._local, "conn", None)
if conn is None:
conn = sqlite3.connect(str(self._db_path), timeout=30, check_same_thread=False)
conn = self._local.conn = sqlite3.connect(str(self._db_path), timeout=30, check_same_thread=False)
conn.row_factory = sqlite3.Row
self._local.conn = conn
with self._connections_lock:
self._connections.add(conn)
return conn
def _execute(self, sql: str, params: tuple = ()) -> sqlite3.Cursor:
"""Execute + commit on this thread's connection."""
cur = self._get_conn().execute(sql, params)
cur.connection.commit()
return cur
def _close(self, *conns: sqlite3.Connection) -> None:
for conn in conns:
with self._connections_lock:
self._connections.discard(conn)
with suppress(Exception):
conn.close()
def _close_thread_conn(self) -> None:
conn = getattr(self._local, "conn", None)
if conn is None:
return
self._local.conn = None
with self._connections_lock:
self._connections.discard(conn)
with suppress(Exception):
conn.close()
conn, self._local.conn = getattr(self._local, "conn", None), None
self._close(*([conn] if conn is not None else []))
def enqueue(self, user_id: str, session_id: str, messages: list) -> None:
now = datetime.now(timezone.utc).isoformat()
with self._shutdown_lock:
if self._shutdown:
return
cur = self._execute(
"INSERT INTO pending (user_id, session_id, messages_json, created_at) VALUES (?,?,?,?)",
(user_id, session_id, json.dumps(messages, ensure_ascii=False), now),
)
cur = self._execute("INSERT INTO pending (user_id, session_id, messages_json, created_at) VALUES (?,?,?,?)",
(user_id, session_id, json.dumps(messages, ensure_ascii=False), now))
self._q.put((cur.lastrowid, user_id, session_id, messages))
def _flush_row(self, row_id: int, user_id: str, session_id: str, messages: list) -> None:
@@ -319,18 +259,12 @@ class _WriteQueue:
self._q.put(_ASYNC_SHUTDOWN)
self._close_thread_conn() # caller thread owns the connection opened in __init__
self._thread.join(timeout=10)
if not self._thread.is_alive():
# Executor workers that already exited may have left tracked handles;
# check_same_thread=False lets shutdown close them deterministically.
if not self._thread.is_alive(): # exited executor workers may have left tracked handles (check_same_thread=False)
with self._connections_lock:
connections, self._connections = list(self._connections), set()
for conn in connections:
with suppress(Exception):
conn.close()
stragglers = list(self._connections)
self._close(*stragglers)
# ── Overlay formatter ────────────────────────────────────────────────────────
def _compact(s: str) -> str:
return re.sub(r"\s+", " ", str(s or "")).strip()[:320]
@@ -352,18 +286,13 @@ def _build_overlay(profile: dict, query_result: dict, local_entries: list[str] |
out.append(c)
return out
profile_items = _dedupe((profile or {}).get("memories"))
query_items = _dedupe((query_result or {}).get("results"))
profile_items, query_items = _dedupe((profile or {}).get("memories")), _dedupe((query_result or {}).get("results"))
if not profile_items and not query_items:
return ""
return "\n".join(
["[RetainDB Context]", "Profile:"] + ([f"- {i}" for i in profile_items] or ["- None"])
+ ["Relevant memories:"] + ([f"- {i}" for i in query_items] or ["- None"])
)
return "\n".join(["[RetainDB Context]", "Profile:"] + ([f"- {i}" for i in profile_items] or ["- None"])
+ ["Relevant memories:"] + ([f"- {i}" for i in query_items] or ["- None"]))
# ── Provider ─────────────────────────────────────────────────────────────────
# Agent self-model keys -> prefetch line formatter, in display order.
_AGENT_MODEL_FIELDS = (
("persona", lambda v: f"Persona: {v}"),
@@ -379,11 +308,9 @@ class RetainDBMemoryProvider(MemoryProvider):
self._client: _Client | None = None
self._queue: _WriteQueue | None = None
self._user_id, self._session_id, self._agent_id = "default", "", "hermes"
self._lock = threading.Lock()
# Prefetch caches + thread tracking (prevents accumulation on rapid calls)
self._context_result = self._dialectic_result = ""
self._agent_model: dict = {}
self._prefetch_threads: list[threading.Thread] = []
self._lock = threading.Lock() # guards the prefetch caches below
self._context_result, self._dialectic_result, self._agent_model = "", "", {}
self._prefetch_threads: list[threading.Thread] = [] # tracked so rapid turns don't pile up threads
@property
def name(self) -> str:
@@ -392,7 +319,7 @@ class RetainDBMemoryProvider(MemoryProvider):
def is_available(self) -> bool:
return bool(get_secret("RETAINDB_API_KEY"))
def get_config_schema(self) -> List[Dict[str, Any]]:
def get_config_schema(self) -> list[dict[str, Any]]:
return [
{"key": "api_key", "description": "RetainDB API key", "secret": True, "required": True, "env_var": "RETAINDB_API_KEY", "url": "https://retaindb.com"},
{"key": "base_url", "description": "API endpoint", "default": _DEFAULT_BASE_URL},
@@ -401,118 +328,88 @@ class RetainDBMemoryProvider(MemoryProvider):
def initialize(self, session_id: str, **kwargs) -> None:
# Non-secret fields resolve env -> config.yaml (written by the Dashboard) -> default.
provider_config = _load_retaindb_config()
base_url = re.sub(r"/+$", "", os.environ.get("RETAINDB_BASE_URL") or _config_str(provider_config.get("base_url")) or _DEFAULT_BASE_URL)
cfg = {k: v.strip() for k, v in _load_retaindb_config().items() if isinstance(v, str)}
base_url = re.sub(r"/+$", "", os.environ.get("RETAINDB_BASE_URL") or cfg.get("base_url") or _DEFAULT_BASE_URL)
# Project: RETAINDB_PROJECT > config.yaml > hermes-<profile> > "default" (API auto-creates "default").
project = os.environ.get("RETAINDB_PROJECT") or _config_str(provider_config.get("project"))
project = os.environ.get("RETAINDB_PROJECT") or cfg.get("project")
if not project:
profile_name = os.path.basename(str(kwargs.get("hermes_home", "")))
project = f"hermes-{profile_name}" if profile_name not in {"", ".hermes"} else "default"
self._client = _Client(get_secret("RETAINDB_API_KEY", "") or "", base_url, project)
self._session_id = session_id
self._user_id = kwargs.get("user_id", "default") or "default"
self._session_id, self._user_id = session_id, kwargs.get("user_id", "default") or "default"
self._agent_id = kwargs.get("agent_id", "hermes") or "hermes"
from hermes_constants import get_hermes_home
hermes_home_path = get_hermes_home()
self._queue = _WriteQueue(self._client, hermes_home_path / "retaindb_queue.db")
# Seed agent identity from SOUL.md in background
soul_path = hermes_home_path / "SOUL.md"
soul_content = soul_path.read_text(encoding="utf-8", errors="replace").strip() if soul_path.exists() else ""
if soul_content:
threading.Thread(target=self._seed_soul, args=(soul_content,), name="retaindb-soul-seed", daemon=True).start()
def _seed_soul(self, content: str) -> None:
try:
self._client.seed_agent_identity(self._agent_id, content, source="soul_md")
except Exception as exc:
logger.debug("RetainDB soul seed failed: %s", exc)
home = get_hermes_home()
self._queue = _WriteQueue(self._client, home / "retaindb_queue.db")
soul = (home / "SOUL.md").read_text(encoding="utf-8", errors="replace").strip() if (home / "SOUL.md").exists() else ""
if soul: # seed agent identity from SOUL.md in background
seed = lambda: self._client.seed_agent_identity(self._agent_id, soul, source="soul_md") # noqa: E731
threading.Thread(target=_quiet, args=("soul seed", seed), name="retaindb-soul-seed", daemon=True).start()
def system_prompt_block(self) -> str:
project = self._client.project if self._client else "retaindb"
return (
f"# RetainDB Memory\nActive. Project: {project}.\n"
"Use retaindb_search to find memories, retaindb_remember to store facts, "
"retaindb_profile for a user overview, retaindb_context for current-task context."
)
# Background prefetch (fires at turn-end, consumed next turn-start)
return (f"# RetainDB Memory\nActive. Project: {project}.\nUse retaindb_search to find memories, retaindb_remember to store facts, "
"retaindb_profile for a user overview, retaindb_context for current-task context.")
def queue_prefetch(self, query: str, *, session_id: str = "") -> None:
"""Fire context + dialectic + agent model prefetches in background."""
"""Fire context + dialectic + agent model prefetches in background (turn-end); prefetch() consumes them next turn."""
if not self._client:
return
# Wait for the previous batch so threads don't accumulate on rapid turns.
for t in self._prefetch_threads:
for t in self._prefetch_threads: # wait for the previous batch so threads don't accumulate on rapid turns
t.join(timeout=2.0)
if any(t.is_alive() for t in self._prefetch_threads):
logger.debug("RetainDB prefetch still running; skipping new batch")
return
jobs = (
("retaindb-ctx", "context", lambda: ("_context_result", self._context_overlay(query)["context"])),
("retaindb-dialectic", "dialectic", lambda: self._fetch_dialectic(query)),
("retaindb-agent-model", "agent model", self._fetch_agent_model),
jobs = ( # (thread name, log label, cache attr, fetch) — fetch returns None to leave the cache untouched
("retaindb-ctx", "context", "_context_result", lambda: self._context_overlay(query)["context"]),
("retaindb-dialectic", "dialectic", "_dialectic_result", lambda: str(
self._client.ask_user(self._user_id, query, reasoning_level=self._reasoning_level(query)).get("answer") or "") or None),
("retaindb-agent-model", "agent model", "_agent_model", lambda: self._agent_model_or_none(self._client.get_agent_model(self._agent_id))),
)
threads = [threading.Thread(target=self._store, args=(label, fetch), name=name, daemon=True) for name, label, fetch in jobs]
self._prefetch_threads = threads
for t in threads:
self._prefetch_threads = [threading.Thread(target=self._store, args=(label, attr, fetch), name=name, daemon=True)
for name, label, attr, fetch in jobs]
for t in self._prefetch_threads:
t.start()
def _context_overlay(self, query: str) -> dict:
query_result = self._client.query_context(self._user_id, self._session_id, query)
profile = self._client.get_profile(self._user_id)
return {"context": _build_overlay(profile, query_result), "raw": query_result}
return {"context": _build_overlay(self._client.get_profile(self._user_id), query_result), "raw": query_result}
def _fetch_dialectic(self, query: str) -> tuple[str, str | None]:
result = self._client.ask_user(self._user_id, query, reasoning_level=self._reasoning_level(query))
return "_dialectic_result", str(result.get("answer") or "") or None
@staticmethod
def _agent_model_or_none(model: dict) -> dict | None:
return model if model.get("memory_count", 0) > 0 else None
def _fetch_agent_model(self) -> tuple[str, dict | None]:
model = self._client.get_agent_model(self._agent_id)
return "_agent_model", model if model.get("memory_count", 0) > 0 else None
def _store(self, label: str, fetch: Callable[[], tuple[str, Any]]) -> None:
"""Run one prefetch job; store (attr, value) under the lock unless value is None; log failures at debug."""
try:
attr, value = fetch()
if value is not None:
with self._lock:
setattr(self, attr, value)
except Exception as exc:
logger.debug("RetainDB %s prefetch failed: %s", label, exc)
def _store(self, label: str, attr: str, fetch: Callable[[], Any]) -> None:
"""Run one prefetch job; cache its value under the lock unless None (failures log at debug)."""
value = _quiet(f"{label} prefetch", fetch)
if value is not None:
with self._lock:
setattr(self, attr, value)
@staticmethod
def _reasoning_level(query: str) -> str:
n = len(query)
return "low" if n < 120 else "medium" if n < 400 else "high"
return "low" if len(query) < 120 else "medium" if len(query) < 400 else "high"
def prefetch(self, query: str, *, session_id: str = "") -> str:
"""Consume prefetched results and return them as a context block."""
with self._lock:
context, dialectic, agent_model = self._context_result, self._dialectic_result, self._agent_model
self._context_result = self._dialectic_result = ""
self._agent_model = {}
parts = [context] if context else []
if dialectic:
parts.append(f"[RetainDB User Synthesis]\n{dialectic}")
if agent_model.get("memory_count", 0) > 0:
model_lines = [fmt(agent_model[k]) for k, fmt in _AGENT_MODEL_FIELDS if agent_model.get(k)]
if model_lines:
parts.append("[RetainDB Agent Self-Model]\n" + "\n".join(model_lines))
return "\n\n".join(parts)
self._context_result, self._dialectic_result, self._agent_model = "", "", {}
model_lines = [fmt(agent_model[k]) for k, fmt in _AGENT_MODEL_FIELDS if agent_model.get(k)] if agent_model.get("memory_count", 0) > 0 else []
parts = [context, dialectic and f"[RetainDB User Synthesis]\n{dialectic}",
model_lines and "[RetainDB Agent Self-Model]\n" + "\n".join(model_lines)]
return "\n\n".join(p for p in parts if p)
def sync_turn(self, user_content: str, assistant_content: str, *, session_id: str = "") -> None:
"""Queue turn for async ingest. Returns immediately."""
if not self._queue or not user_content:
return
now = datetime.now(timezone.utc).isoformat()
self._queue.enqueue(self._user_id, session_id or self._session_id, [
{"role": "user", "content": user_content, "timestamp": now},
{"role": "assistant", "content": assistant_content, "timestamp": now},
])
self._queue.enqueue(self._user_id, session_id or self._session_id,
[{"role": "user", "content": user_content, "timestamp": now},
{"role": "assistant", "content": assistant_content, "timestamp": now}])
def get_tool_schemas(self) -> List[Dict[str, Any]]:
def get_tool_schemas(self) -> list[dict[str, Any]]:
return list(_SCHEMAS)
def handle_tool_call(self, tool_name: str, args: dict, **kwargs) -> str:
@@ -524,14 +421,11 @@ class RetainDBMemoryProvider(MemoryProvider):
return tool_error(str(exc))
def _dispatch(self, tool_name: str, args: dict) -> Any:
entry = _TOOLS.get(tool_name)
if entry is None:
required, handler = _TOOLS.get(tool_name, (None, None))
if handler is None:
return {"error": f"Unknown tool: {tool_name}"}
required, handler = entry
value = args.get(required, "") if required else None
if required and not value:
return {"error": f"{required} is required"}
return handler(self, args, value)
return {"error": f"{required} is required"} if required and not value else handler(self, args, value)
def _tool_upload_file(self, args: dict, local_path: str) -> Any:
path_obj = Path(local_path)
@@ -542,19 +436,17 @@ class RetainDBMemoryProvider(MemoryProvider):
except ValueError as exc:
return {"error": str(exc)}
import mimetypes
mime = mimetypes.guess_type(path_obj.name)[0] or "application/octet-stream"
result = self._client.upload_file(path_obj.read_bytes(), path_obj.name, args.get("remote_path") or f"/{path_obj.name}",
mime, args.get("scope", "PROJECT"), None)
mimetypes.guess_type(path_obj.name)[0] or "application/octet-stream", args.get("scope", "PROJECT"), None)
if args.get("ingest") and result.get("file", {}).get("id"):
result["ingest"] = self._ingest(result["file"]["id"])
return result
def _tool_read_file(self, args: dict, file_id: str) -> Any:
file_info = self._client.get_file(file_id).get("file") or {}
mime = (file_info.get("mime_type") or "").lower()
raw = self._client.read_file_content(file_id)
out = {"file_id": file_id, "rdb_uri": file_info.get("rdb_uri"), "name": file_info.get("name")}
if not (mime.startswith("text/") or file_info.get("name", "").endswith(_TEXT_EXTS)):
if not ((file_info.get("mime_type") or "").lower().startswith("text/") or file_info.get("name", "").endswith(_TEXT_EXTS)):
return {**out, "content": None, "note": "Binary file — use retaindb_ingest_file to extract text into memory."}
text = raw.decode("utf-8", errors="replace")
return {**out, "content": text[:32000], "truncated": len(text) > 32000}
@@ -566,23 +458,19 @@ class RetainDBMemoryProvider(MemoryProvider):
"""Mirror built-in memory writes to RetainDB."""
if action != "add" or not content or not self._client:
return
try:
memory_type = "preference" if target == "user" else "factual"
self._client.add_memory(self._user_id, self._session_id, content, memory_type=memory_type)
except Exception as exc:
logger.debug("RetainDB memory mirror failed: %s", exc)
_quiet("memory mirror", lambda: self._client.add_memory(
self._user_id, self._session_id, content, memory_type="preference" if target == "user" else "factual"))
def shutdown(self) -> None:
for t in self._prefetch_threads:
t.join(timeout=3.0)
self._prefetch_threads = []
queue_obj, self._queue, self._client = self._queue, None, None
queue_obj, self._prefetch_threads, self._queue, self._client = self._queue, [], None, None
if queue_obj:
queue_obj.shutdown()
# tool name -> (required arg or None, handler(provider, args, required_value)); missing arg -> "<arg> is required"
_TOOLS: Dict[str, tuple[str | None, Callable[..., Any]]] = {
_TOOLS: dict[str, tuple[str | None, Callable[..., Any]]] = {
"retaindb_profile": (None, lambda p, a, _: p._client.get_profile(p._user_id)),
"retaindb_search": ("query", lambda p, a, q: p._client.search(p._user_id, p._session_id, q, top_k=min(int(a.get("top_k", 8)), 20))),
"retaindb_context": ("query", lambda p, a, q: p._context_overlay(q)),
+237 -469
View File
@@ -1,8 +1,8 @@
"""Supermemory memory plugin (MemoryProvider): profile recall, semantic search, explicit memory tools,
cleaned turn capture, and session-end conversation ingest."""
"""Supermemory memory plugin (MemoryProvider): profile recall, semantic search, memory tools, turn capture, session ingest."""
from __future__ import annotations
import importlib
import json
import logging
import os
@@ -20,22 +20,13 @@ from tools.registry import tool_error
logger = logging.getLogger(__name__)
_DEFAULT_CONTAINER_TAG = "hermes"
_DEFAULT_MAX_RECALL_RESULTS = 10
_DEFAULT_PROFILE_FREQUENCY = 50
_DEFAULT_CAPTURE_MODE = "all"
_DEFAULT_SEARCH_MODE = "hybrid"
_VALID_SEARCH_MODES = ("hybrid", "memories", "documents")
_DEFAULT_API_TIMEOUT = 5.0
_MAX_ENTITY_CONTEXT_LENGTH = 1500
_DEFAULT_BASE_URL = "https://api.supermemory.ai"
_API_KEY_URL = "http://app.supermemory.ai/integrations?connect=hermes"
# Strips injected <supermemory-context> / <supermemory-containers> blocks before capture.
_INJECTED_BLOCK_RE = re.compile(
r"<supermemory-(context|containers)>[\s\S]*?</supermemory-\1>\s*", re.DOTALL
)
_INJECTED_BLOCK_RE = re.compile(r"<supermemory-(context|containers)>[\s\S]*?</supermemory-\1>\s*", re.DOTALL)
_DEFAULT_ENTITY_CONTEXT = (
"User-assistant conversation. Format: [role: user]...[user:end] and "
"[role: assistant]...[assistant:end].\n\n"
"User-assistant conversation. Format: [role: user]...[user:end] and [role: assistant]...[assistant:end].\n\n"
"Only extract things useful in future conversations. Most messages are not worth remembering.\n\n"
"Remember lasting personal facts, preferences, routines, tools, ongoing projects, working context, "
"and explicit requests to remember something.\n\n"
@@ -43,18 +34,20 @@ _DEFAULT_ENTITY_CONTEXT = (
"When in doubt, store less."
)
# snake_case tool name -> kebab-case alias exposed alongside it.
_KEBAB_ALIASES = {
"supermemory_store": "supermemory-save",
"supermemory_search": "supermemory-search",
"supermemory_forget": "supermemory-forget",
"supermemory_profile": "supermemory-profile",
}
_KEBAB_ALIASES = {"supermemory_store": "supermemory-save", "supermemory_search": "supermemory-search",
"supermemory_forget": "supermemory-forget", "supermemory_profile": "supermemory-profile"}
_ALIAS_TO_TOOL = {kebab: snake for snake, kebab in _KEBAB_ALIASES.items()}
_BOOL_WORDS = {**dict.fromkeys(("true", "1", "yes", "y", "on"), True), **dict.fromkeys(("false", "0", "no", "n", "off"), False)}
def _default_config() -> dict:
"""Fresh copy of every config default (lists are copied so callers can mutate safely)."""
return {k: (list(d) if isinstance(d, list) else d) for k, (d, _) in _CONFIG_SPEC.items()}
def _quietly(fn: Callable[[], Any], fail_msg: str = "", *args: Any, level: int = logging.DEBUG, default: Any = None) -> Any:
"""Run ``fn()``; on any exception log ``fail_msg`` (if given) with traceback and return ``default``."""
try:
return fn()
except Exception:
if fail_msg:
logger.log(level, fail_msg, *args, exc_info=True)
return default
def _sanitize_tag(raw: str) -> str:
@@ -62,35 +55,23 @@ def _sanitize_tag(raw: str) -> str:
def _resolve_base_url(config_value: Any = "") -> str:
"""Resolve the API base URL: config > SUPERMEMORY_BASE_URL env var > default (self-hosted support)."""
"""config > SUPERMEMORY_BASE_URL env var > default (self-hosted support)."""
raw = str(config_value or "").strip() or os.environ.get("SUPERMEMORY_BASE_URL", "").strip()
return (raw or _DEFAULT_BASE_URL).rstrip("/") or _DEFAULT_BASE_URL
def _clamp_entity_context(text: str) -> str:
return text.strip()[:_MAX_ENTITY_CONTEXT_LENGTH] if text else _DEFAULT_ENTITY_CONTEXT
_BOOL_WORDS = {**dict.fromkeys(("true", "1", "yes", "y", "on"), True), **dict.fromkeys(("false", "0", "no", "n", "off"), False)}
return text.strip()[:1500] if text else _DEFAULT_ENTITY_CONTEXT
def _as_bool(value: Any, default: bool) -> bool:
"""bool passthrough; common true/false words parsed; anything else (incl. ints) -> default."""
if isinstance(value, bool):
return value
return _BOOL_WORDS.get(value.strip().lower(), default) if isinstance(value, str) else default
def _valid_search_mode(mode: str) -> str:
return mode if mode in _VALID_SEARCH_MODES else _DEFAULT_SEARCH_MODE
return value if isinstance(value, bool) else _BOOL_WORDS.get(value.strip().lower(), default) if isinstance(value, str) else default
def _clamp_number(value: Any, default, lo, hi, cast):
"""Cast ``value`` and clamp it to [lo, hi]; fall back to ``default`` on any conversion error."""
try:
return max(lo, min(hi, cast(value)))
except Exception:
return default
return _quietly(lambda: max(lo, min(hi, cast(value))), default=default)
# config key -> (default, normalizer applied to the raw/merged value). Order = supermemory.json layout.
@@ -100,14 +81,13 @@ _CONFIG_SPEC: Dict[str, tuple] = {
"container_tag": (_DEFAULT_CONTAINER_TAG, lambda v: str(v).strip() or _DEFAULT_CONTAINER_TAG),
"auto_recall": (True, lambda v: _as_bool(v, True)),
"auto_capture": (True, lambda v: _as_bool(v, True)),
"max_recall_results": (_DEFAULT_MAX_RECALL_RESULTS, lambda v: _clamp_number(v, _DEFAULT_MAX_RECALL_RESULTS, 1, 20, int)),
"profile_frequency": (_DEFAULT_PROFILE_FREQUENCY, lambda v: _clamp_number(v, _DEFAULT_PROFILE_FREQUENCY, 1, 500, int)),
"capture_mode": (_DEFAULT_CAPTURE_MODE, lambda v: "everything" if v == "everything" else "all"),
"search_mode": (_DEFAULT_SEARCH_MODE, lambda v: _valid_search_mode(str(v).strip().lower())),
"max_recall_results": (10, lambda v: _clamp_number(v, 10, 1, 20, int)),
"profile_frequency": (50, lambda v: _clamp_number(v, 50, 1, 500, int)),
"capture_mode": ("all", lambda v: "everything" if v == "everything" else "all"),
"search_mode": ("hybrid", lambda v: v if (v := str(v).strip().lower()) in _VALID_SEARCH_MODES else "hybrid"),
"entity_context": (_DEFAULT_ENTITY_CONTEXT, lambda v: _clamp_entity_context(str(v))),
"api_timeout": (_DEFAULT_API_TIMEOUT, lambda v: _clamp_number(v, _DEFAULT_API_TIMEOUT, 0.5, 15.0, float)),
"api_timeout": (5.0, lambda v: _clamp_number(v, 5.0, 0.5, 15.0, float)),
"base_url": ("", lambda v: str(v or "").strip()),
# Multi-container support
"enable_custom_container_tags": (False, lambda v: _as_bool(v, False)),
"custom_containers": ([], lambda v: [_sanitize_tag(str(t)) for t in v if t] if isinstance(v, list) else []),
"custom_container_instructions": ("", lambda v: str(v).strip()),
@@ -115,202 +95,129 @@ _CONFIG_SPEC: Dict[str, tuple] = {
def _read_json_dict(path: Path) -> dict:
"""Return the JSON object stored at ``path`` or {} if missing/invalid."""
if path.exists():
try:
raw = json.loads(path.read_text(encoding="utf-8"))
if isinstance(raw, dict):
return raw
except Exception:
logger.debug("Failed to parse %s", path, exc_info=True)
return {}
raw = _quietly(lambda: json.loads(path.read_text(encoding="utf-8")), "Failed to parse %s", path) if path.exists() else None
return raw if isinstance(raw, dict) else {}
def _load_supermemory_config(hermes_home: str) -> dict:
config = _default_config()
config.update({k: v for k, v in _read_json_dict(Path(hermes_home) / "supermemory.json").items() if v is not None})
def _load_supermemory_config(hermes_home: Optional[str] = None) -> dict:
"""Defaults overlaid with $hermes_home/supermemory.json (None = defaults only), every key normalized."""
config = {k: (list(d) if isinstance(d, list) else d) for k, (d, _) in _CONFIG_SPEC.items()}
if hermes_home is not None:
config.update({k: v for k, v in _read_json_dict(Path(hermes_home) / "supermemory.json").items() if v is not None})
for key, (_, normalize) in _CONFIG_SPEC.items():
config[key] = normalize(config[key])
return config
def _save_supermemory_config(values: dict, hermes_home: str) -> None:
config_path = Path(hermes_home) / "supermemory.json"
existing = _read_json_dict(config_path)
existing.update(values)
from utils import atomic_json_write
atomic_json_write(config_path, existing, mode=0o600, sort_keys=True)
# Ordered: first matching pattern wins.
_CATEGORY_PATTERNS = (
("preference", r"prefer|like|love|hate|want"),
("decision", r"decided|will use|going with"),
("fact", r"\bis\b|\bare\b|\bhas\b|\bhave\b"),
)
config_path = Path(hermes_home) / "supermemory.json"
atomic_json_write(config_path, {**_read_json_dict(config_path), **values}, mode=0o600, sort_keys=True)
def _detect_category(text: str) -> str:
lowered = text.lower()
return next((cat for cat, pat in _CATEGORY_PATTERNS if re.search(pat, lowered)), "other")
lowered = text.lower() # first matching pattern wins
return next((cat for cat, pat in (("preference", r"prefer|like|love|hate|want"), ("decision", r"decided|will use|going with"),
("fact", r"\bis\b|\bare\b|\bhas\b|\bhave\b")) if re.search(pat, lowered)), "other")
def _format_relative_time(iso_timestamp: str) -> str:
try:
dt = datetime.fromisoformat(iso_timestamp.replace("Z", "+00:00"))
now = datetime.now(timezone.utc)
"""'just now' / '5m ago' / '3h ago' / '2d ago' / '%d %b[ %Y]'; '' when unparseable."""
def _fmt():
dt, now = datetime.fromisoformat(iso_timestamp.replace("Z", "+00:00")), datetime.now(timezone.utc)
seconds = (now - dt).total_seconds()
if seconds < 1800:
return "just now"
for limit, unit, label in ((3600, 60, "m"), (86400, 3600, "h"), (604800, 86400, "d")):
for limit, unit, label in ((1800, 0, "just now"), (3600, 60, "m ago"), (86400, 3600, "h ago"), (604800, 86400, "d ago")):
if seconds < limit:
return f"{int(seconds / unit)}{label} ago"
return f"{int(seconds / unit)}{label}" if unit else label
return dt.strftime("%d %b" if dt.year == now.year else "%d %b %Y")
except Exception:
return ""
def _deduplicate_recall(static_facts: list, dynamic_facts: list, search_results: list) -> tuple[list, list, list]:
"""Drop empties and repeats across the three lists; earlier lists win (profile facts beat search hits)."""
seen: set = set()
def _unique(items, key=lambda x: x):
out = []
for item in items or []:
k = key(item)
if k and k not in seen:
seen.add(k)
out.append(item)
return out
return _unique(static_facts), _unique(dynamic_facts), _unique(search_results, key=lambda i: i.get("memory", ""))
def _bullets(title: str, items: list) -> str:
return f"## {title}\n" + "\n".join(f"- {item}" for item in items)
return _quietly(_fmt, default="")
def _similarity_pct(value: Any) -> Optional[int]:
"""0..1 similarity -> whole percent; None when absent or unparseable."""
if value is None:
return None
try:
return round(float(value) * 100)
except Exception:
return None
return _quietly(lambda: None if value is None else round(float(value) * 100))
def _profile_sections(static_facts: list, dynamic_facts: list) -> list[str]:
return ([_bullets("User Profile (Persistent)", static_facts)] if static_facts else []) + \
([_bullets("Recent Context", dynamic_facts)] if dynamic_facts else [])
return [f"## {title}\n" + "\n".join(f"- {item}" for item in items)
for title, items in (("User Profile (Persistent)", static_facts), ("Recent Context", dynamic_facts)) if items]
def _format_prefetch_context(static_facts: list, dynamic_facts: list, search_results: list, max_results: int) -> str:
statics, dynamics, search = (lst[:max_results] for lst in _deduplicate_recall(static_facts, dynamic_facts, search_results))
sections = _profile_sections(statics, dynamics)
"""Dedupe across the three lists (earlier lists win: profile facts beat search hits), cap each, render."""
seen: set = set()
def _unique(items, key=lambda x: x): # set.add() returns None, so `not seen.add(k)` records k and keeps the item
return [i for i in items or [] if (k := key(i)) and k not in seen and not seen.add(k)][:max_results]
sections = _profile_sections(_unique(static_facts), _unique(dynamic_facts))
lines = []
for item in search: # dedupe already dropped items without a memory string
for item in _unique(search_results, key=lambda i: i.get("memory", "")):
rel = _format_relative_time(item.get("updated_at") or item.get("updatedAt") or "")
pct = _similarity_pct(item.get("similarity"))
prefix_bits = ([f"[{rel}]"] if rel else []) + ([f"[{pct}%]"] if pct is not None else [])
lines.append(f"- {' '.join(prefix_bits)} {item['memory']}".strip())
if lines:
sections.append("## Relevant Memories\n" + "\n".join(lines))
if not sections:
return ""
intro = ("The following is background context from long-term memory. Use it silently when relevant. "
"Do not force memories into the conversation.")
return f"<supermemory-context>\n{intro}\n\n" + "\n\n".join(sections) + "\n</supermemory-context>"
lines.append(f"- {' '.join(([f'[{rel}]'] if rel else []) + ([f'[{pct}%]'] if pct is not None else []))} {item['memory']}".strip())
sections += ["## Relevant Memories\n" + "\n".join(lines)] if lines else []
intro = "The following is background context from long-term memory. Use it silently when relevant. Do not force memories into the conversation."
return f"<supermemory-context>\n{intro}\n\n" + "\n\n".join(sections) + "\n</supermemory-context>" if sections else ""
def _clean_text_for_capture(text: str) -> str:
return _INJECTED_BLOCK_RE.sub("", text or "").strip()
def _updated_at(item: Any) -> Any:
return getattr(item, "updated_at", None) or getattr(item, "updatedAt", None)
def _memory_fields(item: Any, *keys: str) -> dict:
"""Pick SDK result attrs into a plain dict; ``updated_at`` also accepts camelCase ``updatedAt``."""
defaults = {"id": "", "memory": "", "similarity": None, "metadata": None}
return {k: getattr(item, "updated_at", None) or getattr(item, "updatedAt", None) if k == "updated_at" else getattr(item, k, defaults[k])
for k in keys}
class _SupermemoryClient:
def __init__(self, api_key: str, timeout: float, container_tag: str,
search_mode: str = "hybrid", base_url: str = ""):
# Lazy-install the SDK on demand (honors security.allow_lazy_installs and
# sealed Docker venvs). On failure fall through so the raw import below
# produces the canonical ImportError message.
try:
from tools.lazy_deps import ensure as _lazy_ensure
_lazy_ensure("memory.supermemory", prompt=False)
except Exception:
pass
# Lazy-install the SDK on demand (honors security.allow_lazy_installs and sealed Docker
# venvs). On failure fall through so the raw import produces the canonical ImportError.
_quietly(lambda: importlib.import_module("tools.lazy_deps").ensure("memory.supermemory", prompt=False))
from supermemory import Supermemory
self._api_key = api_key
self._container_tag = container_tag
self._search_mode = _valid_search_mode(search_mode)
self._timeout = timeout
self._api_key, self._container_tag, self._timeout = api_key, container_tag, timeout
self._search_mode = search_mode if search_mode in _VALID_SEARCH_MODES else "hybrid"
self._base_url = _resolve_base_url(base_url)
self._client = Supermemory(api_key=api_key, base_url=self._base_url, timeout=timeout, max_retries=0,
default_headers={"x-sm-source": "hermes"})
def _merge_metadata(self, metadata: Optional[dict]) -> dict:
# sm_source routes Hermes writes into the "Hermes" Space in the Supermemory
# app so the user can filter / bulk-manage them per source agent (a
# functional routing key for the user, not vendor telemetry).
# sm_source routes Hermes writes into the "Hermes" Space in the Supermemory app so the user
# can filter / bulk-manage them per source agent (a routing key for the user, not telemetry).
merged = {"sm_source": "hermes", **(metadata or {})}
legacy_source = merged.pop("source", None)
if legacy_source and "type" not in merged:
if (legacy_source := merged.pop("source", None)) and "type" not in merged:
merged["type"] = str(legacy_source)
return merged
def add_memory(self, content: str, metadata: Optional[dict] = None, *, entity_context: str = "",
container_tag: Optional[str] = None, custom_id: Optional[str] = None) -> dict:
kwargs: dict[str, Any] = {"content": content.strip(), "container_tags": [container_tag or self._container_tag]}
if metadata:
kwargs["metadata"] = self._merge_metadata(metadata)
if entity_context:
kwargs["entity_context"] = _clamp_entity_context(entity_context)
if custom_id:
kwargs["custom_id"] = custom_id
result = self._client.documents.add(**kwargs)
return {"id": getattr(result, "id", "")}
kwargs: dict[str, Any] = {"content": content.strip(), "container_tags": [container_tag or self._container_tag],
**({"metadata": self._merge_metadata(metadata)} if metadata else {}),
**({"entity_context": _clamp_entity_context(entity_context)} if entity_context else {}),
**({"custom_id": custom_id} if custom_id else {})}
return {"id": getattr(self._client.documents.add(**kwargs), "id", "")}
def search_memories(self, query: str, *, limit: int = 5, container_tag: Optional[str] = None,
search_mode: Optional[str] = None) -> list[dict]:
mode = search_mode or self._search_mode
kwargs: dict[str, Any] = {"q": query, "container_tag": container_tag or self._container_tag, "limit": limit}
if mode in _VALID_SEARCH_MODES:
kwargs["search_mode"] = mode
kwargs: dict[str, Any] = {"q": query, "container_tag": container_tag or self._container_tag, "limit": limit,
**({"search_mode": mode} if mode in _VALID_SEARCH_MODES else {})}
response = self._client.search.memories(**kwargs)
return [
{
"id": getattr(item, "id", ""),
"memory": getattr(item, "memory", "") or "",
"similarity": getattr(item, "similarity", None),
"updated_at": _updated_at(item),
"metadata": getattr(item, "metadata", None),
}
for item in (getattr(response, "results", None) or [])
]
return [{**_memory_fields(item, "id", "memory", "similarity", "updated_at", "metadata"), "memory": getattr(item, "memory", "") or ""}
for item in (getattr(response, "results", None) or [])]
def get_profile(self, query: Optional[str] = None, *, container_tag: Optional[str] = None) -> dict:
kwargs: dict[str, Any] = {"container_tag": container_tag or self._container_tag}
if query:
kwargs["q"] = query
response = self._client.profile(**kwargs)
response = self._client.profile(container_tag=container_tag or self._container_tag, **({"q": query} if query else {}))
profile_data = getattr(response, "profile", None)
search_data = getattr(response, "search_results", None) or getattr(response, "searchResults", None)
raw_results = getattr(search_data, "results", None) or search_data or []
return {
"static": (getattr(profile_data, "static", []) or []) if profile_data else [],
"dynamic": (getattr(profile_data, "dynamic", []) or []) if profile_data else [],
"search_results": [
item if isinstance(item, dict) else {
"memory": getattr(item, "memory", ""),
"updated_at": _updated_at(item),
"similarity": getattr(item, "similarity", None),
}
for item in raw_results
] if isinstance(raw_results, list) else [],
**{k: (getattr(profile_data, k, []) or []) if profile_data else [] for k in ("static", "dynamic")},
"search_results": [item if isinstance(item, dict) else _memory_fields(item, "memory", "updated_at", "similarity")
for item in raw_results] if isinstance(raw_results, list) else [],
}
def forget_memory(self, memory_id: str, *, container_tag: Optional[str] = None) -> None:
@@ -318,30 +225,27 @@ class _SupermemoryClient:
def forget_by_query(self, query: str, *, container_tag: Optional[str] = None) -> dict:
results = self.search_memories(query, limit=5, container_tag=container_tag)
if not results:
return {"success": False, "message": "No matching memory found to forget."}
memory_id = results[0].get("id", "")
memory_id = results[0].get("id", "") if results else ""
if not memory_id:
return {"success": False, "message": "Best matching memory has no id."}
return {"success": False, "message": "Best matching memory has no id." if results else "No matching memory found to forget."}
self.forget_memory(memory_id, container_tag=container_tag)
return {"success": True, "message": f'Forgot: "{(results[0].get("memory") or "")[:100]}"', "id": memory_id}
def ingest_conversation(self, session_id: str, messages: list[dict], metadata: dict | None = None) -> None:
payload: dict = {"conversationId": session_id, "messages": messages, "containerTags": [self._container_tag]}
if metadata:
payload["metadata"] = self._merge_metadata(metadata)
req = urllib.request.Request(
f"{self._base_url}/v4/conversations",
data=json.dumps(payload).encode("utf-8"),
headers={"Authorization": f"Bearer {self._api_key}", "Content-Type": "application/json",
"x-sm-source": "hermes"},
method="POST",
)
payload: dict = {"conversationId": session_id, "messages": messages, "containerTags": [self._container_tag],
**({"metadata": self._merge_metadata(metadata)} if metadata else {})}
req = urllib.request.Request(f"{self._base_url}/v4/conversations", data=json.dumps(payload).encode("utf-8"), method="POST",
headers={"Authorization": f"Bearer {self._api_key}", "Content-Type": "application/json",
"x-sm-source": "hermes"})
with urllib.request.urlopen(req, timeout=self._timeout + 3):
return
def _build_client(api_key: str, config: dict, container_tag: str) -> _SupermemoryClient:
return _SupermemoryClient(api_key=api_key, timeout=config["api_timeout"], container_tag=container_tag,
search_mode=config["search_mode"], base_url=_resolve_base_url(config["base_url"]))
def _resolve_container_tag(config_tag: str, identity: str) -> str:
"""SUPERMEMORY_CONTAINER_TAG env > config > default; {identity} expands to the agent identity, then sanitize."""
raw_tag = os.environ.get("SUPERMEMORY_CONTAINER_TAG", "").strip() or config_tag
@@ -350,105 +254,71 @@ def _resolve_container_tag(config_tag: str, identity: str) -> str:
def _probe_supermemory_connection(api_key: str, hermes_home: str, *, identity: str = "default") -> dict:
config = _load_supermemory_config(hermes_home)
status = {
"ok": False, "error": "", "profile_facts": 0,
"container_tag": _resolve_container_tag(config["container_tag"], identity),
"auto_recall": bool(config["auto_recall"]), "auto_capture": bool(config["auto_capture"]),
}
status = {"ok": False, "error": "", "profile_facts": 0, "container_tag": _resolve_container_tag(config["container_tag"], identity),
"auto_recall": bool(config["auto_recall"]), "auto_capture": bool(config["auto_capture"])}
if not (api_key or "").strip():
status["error"] = "SUPERMEMORY_API_KEY not set"
return status
return {**status, "error": "SUPERMEMORY_API_KEY not set"}
try:
__import__("supermemory")
except ImportError:
status["error"] = "supermemory package not installed"
return status
return {**status, "error": "supermemory package not installed"}
try:
client = _SupermemoryClient(api_key=api_key.strip(), timeout=config["api_timeout"],
container_tag=status["container_tag"], search_mode=config["search_mode"],
base_url=_resolve_base_url(config["base_url"]))
profile = client.get_profile()
status["profile_facts"] = sum(
1 for f in (profile.get("static") or []) + (profile.get("dynamic") or []) if f and str(f).strip()
)
status["ok"] = True
profile = _build_client(api_key.strip(), config, status["container_tag"]).get_profile()
except Exception as exc:
status["error"] = str(exc).strip()[:160] or "connection failed"
return status
return {**status, "error": str(exc).strip()[:160] or "connection failed"}
facts = sum(1 for f in (profile.get("static") or []) + (profile.get("dynamic") or []) if f and str(f).strip())
return {**status, "ok": True, "profile_facts": facts}
def _format_connection_summary(status: dict) -> str:
container = status.get("container_tag") or _DEFAULT_CONTAINER_TAG
flags = (f"auto_recall {'on' if status.get('auto_recall') else 'off'} · "
f"auto_capture {'on' if status.get('auto_capture') else 'off'}")
flags = f"auto_recall {'on' if status.get('auto_recall') else 'off'} · auto_capture {'on' if status.get('auto_capture') else 'off'}"
if status.get("ok"):
facts = int(status.get("profile_facts") or 0)
return f"✓ Connected · container: {container} · {facts} profile {'fact' if facts == 1 else 'facts'} · {flags}"
return f"✗ {status.get('error') or 'connection failed'} · container: {container} · {flags}"
def _schema(name: str, description: str, properties: dict, required: Optional[list] = None) -> dict:
parameters: dict[str, Any] = {"type": "object", "properties": properties}
if required:
parameters["required"] = required
return {"name": name, "description": description, "parameters": parameters}
def _str_prop(description: str) -> dict:
return {"type": "string", "description": description}
STORE_SCHEMA, SEARCH_SCHEMA, FORGET_SCHEMA, PROFILE_SCHEMA = _BASE_SCHEMAS = [
_schema("supermemory_store", "Store an explicit memory for future recall.", {
"content": _str_prop("The memory content to store."),
"metadata": {"type": "object", "description": "Optional metadata attached to the memory."},
}, required=["content"]),
_schema("supermemory_search", "Search long-term memory by semantic similarity.", {
"query": _str_prop("What to search for."),
"limit": {"type": "integer", "description": "Maximum results to return, 1 to 20."},
}, required=["query"]),
_schema("supermemory_forget", "Forget a memory by exact id or by best-match query.", {
"id": _str_prop("Exact memory id to delete."),
"query": _str_prop("Query used to find the memory to forget."),
}),
_schema("supermemory_profile", "Retrieve persistent profile facts and recent memory context.", {
"query": _str_prop("Optional query to focus the profile response."),
}),
# (name, description, ((prop, type, description), ...), required) -> tool schema; kebab aliases are added in get_tool_schemas().
_BASE_SCHEMAS = [
{"name": name, "description": description,
"parameters": {"type": "object", "properties": {p: {"type": t, "description": d} for p, t, d in props}, **({"required": req} if req else {})}}
for name, description, props, req in (
("supermemory_store", "Store an explicit memory for future recall.",
(("content", "string", "The memory content to store."), ("metadata", "object", "Optional metadata attached to the memory.")), ["content"]),
("supermemory_search", "Search long-term memory by semantic similarity.",
(("query", "string", "What to search for."), ("limit", "integer", "Maximum results to return, 1 to 20.")), ["query"]),
("supermemory_forget", "Forget a memory by exact id or by best-match query.",
(("id", "string", "Exact memory id to delete."), ("query", "string", "Query used to find the memory to forget.")), None),
("supermemory_profile", "Retrieve persistent profile facts and recent memory context.",
(("query", "string", "Optional query to focus the profile response."),), None),
)
]
def _turns_to_messages(turns: List[Dict[str, str]]) -> list[dict]:
return [{"role": role, "content": turn[role]} for turn in turns for role in ("user", "assistant") if turn.get(role)]
class _TagError(Exception):
"""Tool call named a container_tag outside the whitelist."""
def _tagged(resp: dict, tag: Optional[str]) -> dict:
return {**resp, "container_tag": tag} if tag else resp
class SupermemoryMemoryProvider(MemoryProvider):
def __init__(self):
self._api_key = self._session_id = self._hermes_home = self._prefetch_result = ""
self._api_key = self._session_id = self._hermes_home = ""
self._client: Optional[_SupermemoryClient] = None
self._container_tag = _DEFAULT_CONTAINER_TAG
self._turn_count = 0
self._prefetch_lock = threading.Lock()
self._prefetch_thread: Optional[threading.Thread] = None
self._sync_thread: Optional[threading.Thread] = None
self._write_thread: Optional[threading.Thread] = None
self._write_enabled = True
self._active = False
self._container_tag, self._turn_count, self._write_enabled, self._active = _DEFAULT_CONTAINER_TAG, 0, True, False
self._prefetch_thread = self._sync_thread = self._write_thread = None # only _write_thread is ever started
self._session_turns: List[Dict[str, str]] = []
self._apply_config(_default_config())
self._base_url = _DEFAULT_BASE_URL # env var is only consulted in initialize()
self._allowed_containers = []
self._apply_config(_load_supermemory_config())
self._base_url, self._allowed_containers = _DEFAULT_BASE_URL, [] # env var is only consulted in initialize()
def _apply_config(self, config: dict) -> None:
self._config = config
for key in ("auto_recall", "auto_capture", "max_recall_results", "profile_frequency", "capture_mode",
"search_mode", "entity_context", "api_timeout"):
"search_mode", "entity_context", "api_timeout", "custom_containers", "custom_container_instructions"):
setattr(self, f"_{key}", config[key])
# Base URL: config > SUPERMEMORY_BASE_URL env var > api.supermemory.ai (self-hosted support).
self._base_url = _resolve_base_url(config["base_url"])
# Multi-container support
self._enable_custom_containers = config["enable_custom_container_tags"]
self._custom_containers: List[str] = config["custom_containers"]
self._custom_container_instructions = config["custom_container_instructions"]
self._base_url, self._enable_custom_containers = _resolve_base_url(config["base_url"]), config["enable_custom_container_tags"]
self._allowed_containers: List[str] = [self._container_tag] + list(self._custom_containers)
@property
@@ -456,314 +326,212 @@ class SupermemoryMemoryProvider(MemoryProvider):
return "supermemory"
def is_available(self) -> bool:
# Key presence only, no SDK import check: the SDK is lazy-installed when the
# client is first constructed in initialize(), so gating on importability here
# would be a chicken-and-egg trap on sealed venvs. Mirrors honcho/mem0.
# Key presence only, no SDK import check: the SDK is lazy-installed in initialize(), so gating on
# importability here is a chicken-and-egg trap on sealed venvs. Mirrors honcho/mem0.
return bool(get_secret("SUPERMEMORY_API_KEY", ""))
def get_config_schema(self):
# Only prompt for the API key during `hermes memory setup`; other options
# live in $HERMES_HOME/supermemory.json or SUPERMEMORY_CONTAINER_TAG.
return [
{"key": "api_key", "description": "Supermemory API key", "secret": True, "required": True, "env_var": "SUPERMEMORY_API_KEY", "url": _API_KEY_URL},
]
# Only the API key is prompted during `hermes memory setup`; other options live in supermemory.json / env.
return [{"key": "api_key", "description": "Supermemory API key", "secret": True, "required": True, "env_var": "SUPERMEMORY_API_KEY", "url": _API_KEY_URL}]
def save_config(self, values, hermes_home):
sanitized = dict(values or {})
if "container_tag" in sanitized:
sanitized["container_tag"] = _sanitize_tag(str(sanitized["container_tag"]))
if "entity_context" in sanitized:
sanitized["entity_context"] = _clamp_entity_context(str(sanitized["entity_context"]))
for key, fix in (("container_tag", _sanitize_tag), ("entity_context", _clamp_entity_context)):
if key in sanitized:
sanitized[key] = fix(str(sanitized[key]))
_save_supermemory_config(sanitized, hermes_home)
def get_status_config(self, provider_config: dict) -> dict:
from hermes_constants import get_hermes_home
del provider_config
api_key = get_secret("SUPERMEMORY_API_KEY", "") or ""
status = _probe_supermemory_connection(api_key, str(get_hermes_home()))
return {"summary": _format_connection_summary(status)}
return {"summary": _format_connection_summary(_probe_supermemory_connection(get_secret("SUPERMEMORY_API_KEY", "") or "", str(get_hermes_home())))}
def post_setup(self, hermes_home: str, config: dict) -> None:
from hermes_cli.config import save_config
from hermes_cli.memory_setup import _prompt, _write_env_vars
print(f"\n Configuring supermemory:\n\n Get your API key at {_API_KEY_URL}\n")
existing = os.environ.get("SUPERMEMORY_API_KEY", "")
masked = f"...{existing[-4:]}" if len(existing) > 4 else "set"
val = _prompt(f"Supermemory API key (current: {masked}, blank to keep)" if existing else "Supermemory API key",
secret=True)
env_writes = {"SUPERMEMORY_API_KEY": val} if val else {}
if not isinstance(config.get("memory"), dict):
config["memory"] = {}
config["memory"]["provider"] = self.name
val = _prompt(f"Supermemory API key (current: {masked}, blank to keep)" if existing else "Supermemory API key", secret=True)
memory = config["memory"] = config["memory"] if isinstance(config.get("memory"), dict) else {}
memory["provider"] = self.name
save_config(config)
if env_writes:
_write_env_vars(env_writes, hermes_home=hermes_home)
if val:
_write_env_vars({"SUPERMEMORY_API_KEY": val}, hermes_home=hermes_home)
api_key = val or existing
# Make the freshly-entered key visible to the probe below. Single-profile
# only: under a multiplexed gateway, writing to the process-global environ
# would leak the key to sibling profiles and their subprocesses.
# Make the freshly-entered key visible to the probe below. Single-profile only: under a multiplexed
# gateway, writing to the process-global environ would leak the key to sibling profiles and their subprocesses.
if api_key and not is_multiplex_active() and os.environ.get("SUPERMEMORY_API_KEY") != api_key:
os.environ["SUPERMEMORY_API_KEY"] = api_key
status = _probe_supermemory_connection(api_key, hermes_home)
print(f"\n {_format_connection_summary(status)}\n\n Memory provider: supermemory\n Activation saved to config.yaml")
if env_writes:
if val:
print(" API keys saved to .env")
print("\n Start a new session to activate.\n")
def initialize(self, session_id: str, **kwargs) -> None:
from hermes_constants import get_hermes_home
self._hermes_home = kwargs.get("hermes_home") or str(get_hermes_home())
self._session_id = session_id
self._turn_count = 0
self._session_id, self._turn_count, self._session_turns = session_id, 0, []
config = _load_supermemory_config(self._hermes_home)
self._api_key = get_secret("SUPERMEMORY_API_KEY", "") or ""
self._container_tag = _resolve_container_tag(config["container_tag"], kwargs.get("agent_identity", "default"))
self._apply_config(config)
self._session_turns = []
self._write_enabled = kwargs.get("agent_context", "") not in {"cron", "flush", "subagent"}
self._active = bool(self._api_key)
self._client = None
if self._active:
try:
self._client = _SupermemoryClient(api_key=self._api_key, timeout=self._api_timeout,
container_tag=self._container_tag, search_mode=self._search_mode,
base_url=self._base_url)
except Exception:
logger.warning("Supermemory initialization failed", exc_info=True)
self._active = False
self._client = None
self._client = _quietly(lambda: _build_client(self._api_key, config, self._container_tag),
"Supermemory initialization failed", level=logging.WARNING) if self._api_key else None
self._active = self._client is not None
def on_turn_start(self, turn_number: int, message: str, **kwargs) -> None:
self._turn_count = max(turn_number, 0)
def system_prompt_block(self) -> str:
if not self._active:
return ""
lines = [
"# Supermemory",
f"Active. Container: {self._container_tag}.",
"Use supermemory-search, supermemory-save, supermemory-forget, and supermemory-profile (aliases: supermemory_search, supermemory_store, supermemory_forget, supermemory_profile).",
]
lines = ["# Supermemory", f"Active. Container: {self._container_tag}.",
"Use supermemory-search, supermemory-save, supermemory-forget, and supermemory-profile (aliases: supermemory_search, supermemory_store, supermemory_forget, supermemory_profile)."]
if self._enable_custom_containers and self._custom_containers:
lines.append(f"\nMulti-container mode enabled. Available containers: {', '.join(self._allowed_containers)}.")
lines.append("Pass an optional container_tag to supermemory_search, supermemory_store, supermemory_forget, and supermemory_profile to target a specific container.")
if self._custom_container_instructions:
lines.append(f"\n{self._custom_container_instructions}")
return "\n".join(lines)
lines += [f"\nMulti-container mode enabled. Available containers: {', '.join(self._allowed_containers)}.",
"Pass an optional container_tag to supermemory_search, supermemory_store, supermemory_forget, and supermemory_profile to target a specific container."]
lines += [f"\n{self._custom_container_instructions}"] if self._custom_container_instructions else []
return "\n".join(lines) if self._active else ""
def _can_write(self) -> bool:
return bool(self._active and self._write_enabled and self._client)
def prefetch(self, query: str, *, session_id: str = "") -> str:
if not self._active or not self._auto_recall or not self._client or not query.strip():
return ""
try:
def _recall():
profile = self._client.get_profile(query=query[:200])
include_profile = self._turn_count <= 1 or (self._turn_count % self._profile_frequency == 0)
return _format_prefetch_context(
static_facts=profile["static"] if include_profile else [],
dynamic_facts=profile["dynamic"] if include_profile else [],
search_results=profile["search_results"],
max_results=self._max_recall_results,
)
except Exception:
logger.debug("Supermemory prefetch failed", exc_info=True)
return ""
return _format_prefetch_context(profile["static"] if include_profile else [], profile["dynamic"] if include_profile else [],
profile["search_results"], self._max_recall_results)
return _quietly(_recall, "Supermemory prefetch failed", default="")
def sync_turn(self, user_content: str, assistant_content: str, *, session_id: str = "") -> None:
if not self._active or not self._auto_capture or not self._write_enabled or not self._client:
if not self._can_write() or not self._auto_capture:
return
clean_user = _clean_text_for_capture(user_content)
clean_assistant = _clean_text_for_capture(assistant_content)
if clean_user or clean_assistant:
# Buffer every turn for the single full-session document written at end/switch/shutdown
self._session_turns.append({"user": clean_user, "assistant": clean_assistant})
turn = {"user": _clean_text_for_capture(user_content), "assistant": _clean_text_for_capture(assistant_content)}
if any(turn.values()): # buffered for the single full-session document written at end/switch/shutdown
self._session_turns.append(turn)
def _ingest_session(self, session_id: str, messages: list[dict], metadata: dict,
fail_msg: str, level: int = logging.DEBUG) -> None:
try:
self._client.ingest_conversation(
session_id, messages, metadata={"type": "full_session", "session_id": session_id, **metadata}
)
except Exception:
logger.log(level, fail_msg, exc_info=True)
def _ingest(self, session_id: str, messages: list[dict], metadata: dict, fail_msg: str, level: int = logging.DEBUG) -> None:
metadata = {"type": "full_session", "session_id": session_id, **metadata}
_quietly(lambda: self._client.ingest_conversation(session_id, messages, metadata=metadata), fail_msg, level=level)
def _ingest_buffered_turns(self, session_id: str, *, partial: bool, fail_msg: str) -> None:
# message_count reports 2 per buffered turn regardless of empty sides.
self._ingest_session(session_id, _turns_to_messages(self._session_turns),
{"message_count": len(self._session_turns) * 2, "partial": partial}, fail_msg)
def _flush_turns(self, session_id: str, *, partial: bool, fail_msg: str) -> None:
turns = self._session_turns # message_count reports 2 per buffered turn regardless of empty sides
messages = [{"role": role, "content": t[role]} for t in turns for role in ("user", "assistant") if t.get(role)]
self._ingest(session_id, messages, {"message_count": len(turns) * 2, "partial": partial}, fail_msg)
def on_session_end(self, messages: List[Dict[str, Any]]) -> None:
if not self._active or not self._write_enabled or not self._client or not self._session_id:
if not self._can_write() or not self._session_id:
return
cleaned = [
{"role": m.get("role"), "content": _clean_text_for_capture(str(m.get("content", "")))}
for m in messages or [] if m.get("role") in {"user", "assistant"}
]
cleaned = [m for m in cleaned if m["content"]]
cleaned = [{"role": m.get("role"), "content": content} for m in messages or []
if m.get("role") in {"user", "assistant"} and (content := _clean_text_for_capture(str(m.get("content", ""))))]
if not cleaned or (len(cleaned) == 1 and len(cleaned[0]["content"]) < 20):
return
self._ingest_session(self._session_id, cleaned, {"message_count": len(cleaned)},
"Supermemory session ingest failed", level=logging.WARNING)
# Clear buffer so shutdown() doesn't duplicate on normal exit
self._session_turns = []
self._ingest(self._session_id, cleaned, {"message_count": len(cleaned)}, "Supermemory session ingest failed", level=logging.WARNING)
self._session_turns = [] # so shutdown() doesn't duplicate on normal exit
def on_session_switch(self, new_session_id: str, *, parent_session_id: str = "", reset: bool = False,
**kwargs) -> None:
def on_session_switch(self, new_session_id: str, *, parent_session_id: str = "", reset: bool = False, **kwargs) -> None:
"""Flush any buffered turns from the old session as one document, then reset for the new session."""
if not self._active or not self._write_enabled or not self._client:
self._session_id = str(new_session_id or "").strip() or self._session_id
self._session_turns = []
return
old_session_id = self._session_id
if self._session_turns and old_session_id:
self._ingest_buffered_turns(old_session_id, partial=not reset,
fail_msg="Supermemory session-switch ingest failed")
if self._can_write():
if self._session_turns and old_session_id:
self._flush_turns(old_session_id, partial=not reset, fail_msg="Supermemory session-switch ingest failed")
self._turn_count = 0
self._session_id = str(new_session_id or "").strip() or old_session_id
self._session_turns = []
self._turn_count = 0
def on_memory_write(self, action: str, target: str, content: str) -> None:
if not self._active or not self._write_enabled or not self._client:
if not self._can_write() or action != "add" or not (content or "").strip():
return
if action != "add" or not (content or "").strip():
return
def _run():
try:
self._client.add_memory(content.strip(), metadata={"target": target, "type": "explicit_memory"},
entity_context=self._entity_context)
except Exception:
logger.debug("Supermemory on_memory_write failed", exc_info=True)
if self._write_thread and self._write_thread.is_alive():
self._write_thread.join(timeout=2.0)
self._write_thread = threading.Thread(target=_run, daemon=False, name="supermemory-memory-write")
self._write_thread = threading.Thread(
target=_quietly, daemon=False, name="supermemory-memory-write",
args=(lambda: self._client.add_memory(content.strip(), metadata={"target": target, "type": "explicit_memory"},
entity_context=self._entity_context), "Supermemory on_memory_write failed"))
self._write_thread.start()
def shutdown(self) -> None:
# Emergency fallback (crashes only). Buffer is cleared on normal on_session_end().
if self._active and self._write_enabled and self._client and self._session_turns and self._session_id:
if self._can_write() and self._session_turns and self._session_id:
logger.warning("Supermemory: Saving session via shutdown (session=%s, turns=%d)", self._session_id, len(self._session_turns))
self._ingest_buffered_turns(self._session_id, partial=True, fail_msg="Supermemory shutdown ingest failed")
for attr_name in ("_prefetch_thread", "_sync_thread", "_write_thread"):
thread = getattr(self, attr_name, None)
if thread and thread.is_alive():
thread.join(timeout=5.0)
setattr(self, attr_name, None)
def _resolve_tool_container_tag(self, args: dict) -> Optional[str]:
"""Return the validated container_tag from args, None for primary; raise ValueError if not whitelisted."""
tag = str(args.get("container_tag") or "").strip() if self._enable_custom_containers else ""
if not tag:
return None
sanitized = _sanitize_tag(tag)
if sanitized not in self._allowed_containers:
raise ValueError(f"Container tag '{sanitized}' is not allowed. Allowed: {', '.join(self._allowed_containers)}")
return sanitized
self._flush_turns(self._session_id, partial=True, fail_msg="Supermemory shutdown ingest failed")
if self._write_thread and self._write_thread.is_alive():
self._write_thread.join(timeout=5.0)
self._prefetch_thread = self._sync_thread = self._write_thread = None
def get_tool_schemas(self) -> List[Dict[str, Any]]:
schemas = [json.loads(json.dumps(base)) for base in _BASE_SCHEMAS] # deep copies
if self._enable_custom_containers:
for schema in schemas: # multi-container mode: every tool takes an optional container_tag
schema["parameters"]["properties"]["container_tag"] = {
"type": "string",
"description": f"Optional container tag. Allowed: {', '.join(self._allowed_containers)}. Defaults to primary ({self._container_tag}).",
}
for schema in schemas if self._enable_custom_containers else (): # multi-container mode: every tool takes container_tag
schema["parameters"]["properties"]["container_tag"] = {
"type": "string", "description": f"Optional container tag. Allowed: {', '.join(self._allowed_containers)}. Defaults to primary ({self._container_tag})."}
# Kebab-case aliases are appended after all snake_case schemas (deep-copied, name swapped).
return schemas + [{**json.loads(json.dumps(s)), "name": _KEBAB_ALIASES[s["name"]]} for s in schemas]
def _run_tool(self, args: dict, fail_prefix: str, fn: Callable[[Optional[str]], Any], *,
tag_in_response: bool = True) -> str:
"""Resolve container_tag, run ``fn(tag)``, and JSON-encode its result; errors become tool_error()."""
try:
tag = self._resolve_tool_container_tag(args)
except ValueError as exc:
return tool_error(str(exc))
try:
resp = fn(tag)
if tag and tag_in_response:
resp["container_tag"] = tag
return json.dumps(resp)
except Exception as exc:
return tool_error(f"{fail_prefix}: {exc}")
def _tool_container_tag(self, args: dict) -> Optional[str]:
"""Validated container_tag from args; None = primary. Raises _TagError when not whitelisted."""
raw = str(args.get("container_tag") or "").strip() if self._enable_custom_containers else ""
tag = _sanitize_tag(raw) if raw else None
if tag and tag not in self._allowed_containers:
raise _TagError(f"Container tag '{tag}' is not allowed. Allowed: {', '.join(self._allowed_containers)}")
return tag
def _tool_store(self, args: dict) -> str:
def _tool_store(self, args: dict) -> dict | str:
content = str(args.get("content") or "").strip()
if not content:
return tool_error("content is required")
metadata = args.get("metadata") or {}
if not isinstance(metadata, dict):
metadata = {}
metadata = args.get("metadata") if isinstance(args.get("metadata"), dict) else {}
metadata.setdefault("type", _detect_category(content))
metadata.pop("source", None)
tag = self._tool_container_tag(args)
result = self._client.add_memory(content, metadata=metadata, entity_context=self._entity_context, container_tag=tag)
return _tagged({"saved": True, "id": result.get("id", ""), "preview": content[:80] + ("..." if len(content) > 80 else "")}, tag)
def _store(tag):
result = self._client.add_memory(content, metadata=metadata, entity_context=self._entity_context, container_tag=tag)
preview = content[:80] + ("..." if len(content) > 80 else "")
return {"saved": True, "id": result.get("id", ""), "preview": preview}
return self._run_tool(args, "Failed to store memory", _store)
def _tool_search(self, args: dict) -> str:
def _tool_search(self, args: dict) -> dict | str:
query = str(args.get("query") or "").strip()
if not query:
return tool_error("query is required")
limit = _clamp_number(args.get("limit", 5) or 5, 5, 1, 20, int)
tag = self._tool_container_tag(args)
results = [{"id": i.get("id", ""), "content": i.get("memory", ""), **({"similarity": pct} if (pct := _similarity_pct(i.get("similarity"))) is not None else {})}
for i in self._client.search_memories(query, limit=limit, container_tag=tag)]
return _tagged({"results": results, "count": len(results)}, tag)
def _search(tag):
formatted = []
for item in self._client.search_memories(query, limit=limit, container_tag=tag):
pct = _similarity_pct(item.get("similarity"))
formatted.append({"id": item.get("id", ""), "content": item.get("memory", ""),
**({"similarity": pct} if pct is not None else {})})
return {"results": formatted, "count": len(formatted)}
return self._run_tool(args, "Search failed", _search)
def _tool_forget(self, args: dict) -> str:
memory_id = str(args.get("id") or "").strip()
query = str(args.get("query") or "").strip()
def _tool_forget(self, args: dict) -> dict | str:
memory_id, query = str(args.get("id") or "").strip(), str(args.get("query") or "").strip()
if not memory_id and not query:
return tool_error("Provide either id or query")
def _forget(tag):
if memory_id:
self._client.forget_memory(memory_id, container_tag=tag)
return {"forgotten": True, "id": memory_id}
tag = self._tool_container_tag(args) # not echoed in the response
if not memory_id:
return self._client.forget_by_query(query, container_tag=tag)
self._client.forget_memory(memory_id, container_tag=tag)
return {"forgotten": True, "id": memory_id}
return self._run_tool(args, "Forget failed", _forget, tag_in_response=False)
def _tool_profile(self, args: dict) -> str:
query = str(args.get("query") or "").strip() or None
def _profile(tag):
profile = self._client.get_profile(query=query, container_tag=tag)
return {"profile": "\n\n".join(_profile_sections(profile["static"], profile["dynamic"])),
"static_count": len(profile["static"]), "dynamic_count": len(profile["dynamic"])}
return self._run_tool(args, "Profile failed", _profile)
def _tool_profile(self, args: dict) -> dict:
tag = self._tool_container_tag(args)
profile = self._client.get_profile(query=str(args.get("query") or "").strip() or None, container_tag=tag)
return _tagged({"profile": "\n\n".join(_profile_sections(profile["static"], profile["dynamic"])),
"static_count": len(profile["static"]), "dynamic_count": len(profile["dynamic"])}, tag)
def handle_tool_call(self, tool_name: str, args: Dict[str, Any], **kwargs) -> str:
"""Handlers return a tool_error() string for bad args or a dict to JSON-encode; client failures get ``fail_prefix``."""
if not self._active or not self._client:
return tool_error("Supermemory is not configured")
tool_name = _ALIAS_TO_TOOL.get(tool_name, tool_name)
handler = self._TOOL_HANDLERS.get(tool_name)
return handler(self, args) if handler else tool_error(f"Unknown tool: {tool_name}")
if tool_name not in self._TOOL_HANDLERS:
return tool_error(f"Unknown tool: {tool_name}")
handler, fail_prefix = self._TOOL_HANDLERS[tool_name]
try:
resp = handler(self, args)
except Exception as exc:
return tool_error(str(exc) if isinstance(exc, _TagError) else f"{fail_prefix}: {exc}")
return resp if isinstance(resp, str) else json.dumps(resp)
# snake_case tool name -> handler; kebab aliases are folded in via _ALIAS_TO_TOOL first.
_TOOL_HANDLERS = {"supermemory_store": _tool_store, "supermemory_search": _tool_search,
"supermemory_forget": _tool_forget, "supermemory_profile": _tool_profile}
# snake_case tool name -> (handler, error prefix); kebab aliases are folded in via _ALIAS_TO_TOOL first.
_TOOL_HANDLERS = {"supermemory_store": (_tool_store, "Failed to store memory"), "supermemory_search": (_tool_search, "Search failed"),
"supermemory_forget": (_tool_forget, "Forget failed"), "supermemory_profile": (_tool_profile, "Profile failed")}
def register(ctx):