refactor(memory/mem0): _try breaker wrapper for background paths, _pg_ready, vector-description table, compact comments
This commit is contained in:
@@ -1,11 +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_* env vars remain a 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
|
||||
@@ -76,27 +75,16 @@ def _schema(name: str, description: str, properties: dict[str, tuple[str, str]],
|
||||
return {"name": name, "description": description, "parameters": {"type": "object", "properties": props, "required": required}}
|
||||
|
||||
|
||||
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"],
|
||||
)]
|
||||
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"]),
|
||||
]
|
||||
|
||||
_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"
|
||||
@@ -160,9 +148,8 @@ class Mem0MemoryProvider(MemoryProvider):
|
||||
return template.format(vs=self._config.get("oss", {}).get("vector_store", {}).get("provider", default))
|
||||
|
||||
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.
|
||||
# Lazy-install the mem0 SDK before the backend imports it (honors security.allow_lazy_installs);
|
||||
# on failure the backend import raises the canonical error, captured below.
|
||||
try:
|
||||
from tools.lazy_deps import ensure as _lazy_ensure
|
||||
_lazy_ensure("memory.mem0", prompt=False)
|
||||
@@ -203,26 +190,34 @@ 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:
|
||||
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()
|
||||
self._record_success()
|
||||
return result
|
||||
except Exception as e:
|
||||
self._record_failure()
|
||||
log(msg, e)
|
||||
return None
|
||||
|
||||
def initialize(self, session_id: str, **kwargs) -> None:
|
||||
self._config = _load_config()
|
||||
self._mode = self._config.get("mode", "platform")
|
||||
self._api_key = self._config.get("api_key", "")
|
||||
self._host = self._config.get("host", "")
|
||||
# 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.
|
||||
# user_id precedence: operator-configured (env/mem0.json) > gateway-native id (kwargs) > _DEFAULT_USER_ID.
|
||||
# The literal placeholder counts as unset so wizard users still get gateway-native ids.
|
||||
configured = self._config.get("user_id")
|
||||
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.
|
||||
# Persisted rerank preference: default for mem0_search when the model omits ``rerank``. Platform-only.
|
||||
self._rerank_default = _truthy(self._config.get("rerank", False))
|
||||
self._channel = kwargs.get("platform") or "cli"
|
||||
self._backend = self._create_backend()
|
||||
@@ -231,9 +226,8 @@ class Mem0MemoryProvider(MemoryProvider):
|
||||
self._atexit_registered = True
|
||||
|
||||
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
|
||||
# (dashboard per-channel filtering) so narrower views remain possible at query time.
|
||||
# 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 _add(self, messages: list, infer: bool):
|
||||
@@ -241,8 +235,7 @@ class Mem0MemoryProvider(MemoryProvider):
|
||||
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 f"# Mem0 Memory\nActive. Mode: {mode_label}. User: {self._user_id}.\n{_PROMPT_BODY}{rerank_note}"
|
||||
@@ -268,15 +261,9 @@ class Mem0MemoryProvider(MemoryProvider):
|
||||
self._prefetch_query, self._prefetch_result, self._prefetch_done = query, "", False
|
||||
|
||||
def _run():
|
||||
body = ""
|
||||
try:
|
||||
lines = [r.get("memory", "") for r in (self._search(query, backend=backend) 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, self._prefetch_done = body, True
|
||||
@@ -305,14 +292,9 @@ class Mem0MemoryProvider(MemoryProvider):
|
||||
return
|
||||
|
||||
def _sync():
|
||||
if self._backend is None:
|
||||
return
|
||||
try:
|
||||
self._add([{"role": "user", "content": user_content}, {"role": "assistant", "content": assistant_content}], infer=True)
|
||||
self._record_success()
|
||||
except Exception as e:
|
||||
self._record_failure()
|
||||
logger.warning("Mem0 sync failed: %s", e)
|
||||
if self._backend is not None:
|
||||
messages = [{"role": "user", "content": user_content}, {"role": "assistant", "content": assistant_content}]
|
||||
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():
|
||||
@@ -362,9 +344,8 @@ class Mem0MemoryProvider(MemoryProvider):
|
||||
if tool_name not in self._TOOL_HANDLERS:
|
||||
return tool_error(f"Unknown tool: {tool_name}")
|
||||
required, label, body, on_client_error = self._TOOL_HANDLERS[tool_name]
|
||||
for k in required:
|
||||
if not args.get(k, ""):
|
||||
return tool_error(f"Missing required parameter: {k}")
|
||||
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)
|
||||
self._record_success()
|
||||
|
||||
@@ -23,10 +23,7 @@ def _unwrap_results(response: Any) -> list:
|
||||
|
||||
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]: ...
|
||||
@@ -74,19 +71,15 @@ 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.
|
||||
# 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:
|
||||
@@ -168,8 +161,7 @@ class OSSBackend(Mem0Backend):
|
||||
|
||||
config = {"vector_store": vector_store, "llm": _provider_block("llm"), "embedder": _provider_block("embedder"), "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)
|
||||
@@ -212,11 +204,7 @@ class OSSBackend(Mem0Backend):
|
||||
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,),
|
||||
)
|
||||
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)))
|
||||
@@ -243,13 +231,10 @@ class OSSBackend(Mem0Backend):
|
||||
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()
|
||||
# Memory, then its vector store, then the store's raw client; the first failure aborts the chain.
|
||||
for obj in filter(None, (self._memory, vs, getattr(vs, "client", None))):
|
||||
if hasattr(obj, "close"):
|
||||
obj.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
@@ -68,12 +68,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 {}
|
||||
|
||||
|
||||
@@ -115,12 +112,10 @@ def parse_flags(argv: list[str] | None = None) -> dict[str, str]:
|
||||
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
|
||||
|
||||
|
||||
@@ -171,8 +166,7 @@ def build_oss_config(flags: dict[str, str]) -> tuple[dict, dict[str, str]]:
|
||||
|
||||
def _write_env(env_path: Path, env_writes: dict[str, str]) -> None:
|
||||
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] = []
|
||||
@@ -227,12 +221,10 @@ def _setup_platform(hermes_home: str, config: dict, flags: dict[str, str]) -> No
|
||||
_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.
|
||||
# Routing checks ``host`` before platform, so clear a stale self-hosted host. "" rather than
|
||||
# pop(): save_config merges into the existing mem0.json, so a popped key would survive.
|
||||
provider_config["host"] = ""
|
||||
# _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.
|
||||
# _load_config() also seeds ``host`` from MEM0_HOST (.env); the file clear can't help there, so warn.
|
||||
if os.environ.get("MEM0_HOST", "").strip():
|
||||
print(f"\n ⚠ MEM0_HOST is set in your environment ({os.environ['MEM0_HOST']}). It overrides platform mode — remove it from ~/.hermes/.env (or unset it) or Hermes will keep routing to the self-hosted server.")
|
||||
_persist_provider_config(hermes_home, config, provider_config, env_writes)
|
||||
@@ -311,8 +303,7 @@ def _setup_oss(hermes_home: str, config: dict, flags: dict[str, str]) -> None:
|
||||
return
|
||||
oss_config, env_writes = build_oss_config(flags)
|
||||
if errors := validate_oss_config(oss_config):
|
||||
for e in errors:
|
||||
print(f" Error: {e}", file=sys.stderr)
|
||||
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)
|
||||
@@ -326,9 +317,14 @@ 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
|
||||
@@ -342,8 +338,7 @@ def _ensure_pgvector(host: str = "localhost", port: int = 5432) -> dict | None:
|
||||
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:
|
||||
@@ -361,8 +356,7 @@ def _start_pgvector_docker(host: str, port: int) -> dict | None:
|
||||
_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]:
|
||||
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.")
|
||||
@@ -375,11 +369,9 @@ 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."""
|
||||
"""Ensure Ollama is running and ``models`` are pulled; False when the user must handle it manually."""
|
||||
ollama_bin = shutil.which("ollama")
|
||||
ok = _check_ollama(_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")
|
||||
return False
|
||||
@@ -387,8 +379,7 @@ def _ensure_ollama(models: list[str]) -> bool:
|
||||
try:
|
||||
subprocess.Popen([ollama_bin, "serve"], stdin=subprocess.DEVNULL, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL)
|
||||
_wait_for_port("localhost", 11434, timeout=10)
|
||||
ok = _check_ollama(_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}")
|
||||
@@ -449,17 +440,20 @@ def _provider_description(v: dict) -> str:
|
||||
return f"{model} ({url})" if url else model
|
||||
|
||||
|
||||
# Vector-store picker description 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
|
||||
return _VECTOR_DESCRIPTIONS.get(pid, lambda cfg: pid)(v.get("default_config", {}))
|
||||
|
||||
|
||||
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."""
|
||||
"""Pick an LLM/embedder provider, collect its key, and (for Ollama) model + URL -> (id, definition, model, url).
|
||||
For the embedder (``llm`` given), a provider shared with the LLM reuses the LLM key instead of prompting again."""
|
||||
items = [(v["label"], _provider_description(v)) for v in registry.values()]
|
||||
pid = list(registry)[_curses_select(f"{kind} Provider", items, 0)]
|
||||
pdef = registry[pid]
|
||||
@@ -481,10 +475,7 @@ def _prompt_pgvector_config() -> dict:
|
||||
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
|
||||
return pgvector_config
|
||||
return {**pg, "port": int(pg["port"]), **({"password": pg_password} if pg_password else {})}
|
||||
|
||||
|
||||
def _setup_oss_interactive(hermes_home: str, config: dict) -> None:
|
||||
@@ -510,9 +501,7 @@ def _setup_oss_interactive(hermes_home: str, config: dict) -> None:
|
||||
"oss_embedder": embedder_id, "oss_embedder_model": embedder_model, "oss_embedder_url": embedder_url or "",
|
||||
"oss_vector": vector_id, "user_id": user_id,
|
||||
}
|
||||
for key, val in (pgvector_config or {}).items():
|
||||
if val:
|
||||
flags[f"oss_vector_{key}"] = str(val)
|
||||
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)
|
||||
|
||||
@@ -561,10 +550,6 @@ def _check_pgvector(host: str, port: int) -> tuple[bool, str]:
|
||||
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 _check_qdrant_url(url: str) -> tuple[bool, str]:
|
||||
return _probe(lambda: _http_get(url, "/healthz", 3), "Qdrant reachable", f"Qdrant not reachable at {url}")
|
||||
|
||||
|
||||
def _warn_unless(check: tuple[bool, str]) -> None:
|
||||
ok, msg = check
|
||||
if not ok:
|
||||
@@ -579,7 +564,7 @@ def _run_connectivity_checks(oss_config: dict) -> None:
|
||||
if path:
|
||||
_warn_unless(_check_qdrant_path(path))
|
||||
elif url:
|
||||
_warn_unless(_check_qdrant_url(url))
|
||||
_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", {})
|
||||
@@ -609,9 +594,8 @@ _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."""
|
||||
"""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."""
|
||||
_check_min_dep_version()
|
||||
flags = parse_flags(sys.argv[1:])
|
||||
handler = _MODE_HANDLERS.get(flags["mode"])
|
||||
|
||||
Reference in New Issue
Block a user