merge(r3-27): group D
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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()))
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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):
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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,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
@@ -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
@@ -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)),
|
||||
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user