refactor(memory/mem0): inline single-use setup helpers, joined summary prints, blank-line squeeze
This commit is contained in:
@@ -40,13 +40,6 @@ def _is_client_error(exc: Exception) -> bool:
|
||||
return type(exc).__name__ in _CLIENT_ERROR_TYPES or any(s in err_str for s in ("404", "not found", "valid uuid"))
|
||||
|
||||
|
||||
def _truthy(value: Any, falsy_strings: tuple[str, ...] | None = None) -> bool:
|
||||
"""Coerce a flag; strings match ``falsy_strings`` (deny-list) when given, else a truthy allow-list."""
|
||||
if not isinstance(value, str):
|
||||
return bool(value)
|
||||
return value.lower() not in falsy_strings if falsy_strings else value.lower() in ("true", "1", "yes")
|
||||
|
||||
|
||||
def _read_mem0_json(config_path: Path) -> dict:
|
||||
"""Best-effort read of mem0.json; missing/corrupt file -> {}."""
|
||||
if config_path.exists():
|
||||
@@ -117,9 +110,7 @@ class Mem0MemoryProvider(MemoryProvider):
|
||||
"""Merge-write config to $HERMES_HOME/mem0.json."""
|
||||
from utils import atomic_json_write
|
||||
config_path = Path(hermes_home) / "mem0.json"
|
||||
existing = _read_mem0_json(config_path)
|
||||
existing.update(values)
|
||||
atomic_json_write(config_path, existing, mode=0o600)
|
||||
atomic_json_write(config_path, {**_read_mem0_json(config_path), **values}, mode=0o600)
|
||||
|
||||
def get_config_schema(self):
|
||||
api_key_required = _load_config().get("mode", "platform") != "oss"
|
||||
@@ -137,9 +128,7 @@ class Mem0MemoryProvider(MemoryProvider):
|
||||
|
||||
def _oss_hint(self, template: str, default: str = "vector store") -> str:
|
||||
"""OSS-only hint; ``{vs}`` is the configured vector-store provider. "" in other modes."""
|
||||
if self._mode != "oss":
|
||||
return ""
|
||||
return template.format(vs=self._config.get("oss", {}).get("vector_store", {}).get("provider", default))
|
||||
return template.format(vs=self._config.get("oss", {}).get("vector_store", {}).get("provider", default)) if self._mode == "oss" else ""
|
||||
|
||||
def _create_backend(self):
|
||||
# Lazy-install the mem0 SDK before the backend imports it (honors security.allow_lazy_installs);
|
||||
@@ -151,9 +140,7 @@ class Mem0MemoryProvider(MemoryProvider):
|
||||
from . import _backend
|
||||
if self._mode == "oss":
|
||||
return _backend.OSSBackend(self._config.get("oss", {}))
|
||||
if self._host:
|
||||
return _backend.SelfHostedBackend(self._api_key, self._host)
|
||||
return _backend.PlatformBackend(self._api_key)
|
||||
return _backend.SelfHostedBackend(self._api_key, self._host) if self._host else _backend.PlatformBackend(self._api_key)
|
||||
except Exception as e:
|
||||
logger.error("Mem0 backend failed to initialize (%s mode): %s", self._mode, e)
|
||||
self._init_error = str(e)
|
||||
@@ -191,12 +178,12 @@ class Mem0MemoryProvider(MemoryProvider):
|
||||
"""Background-path wrapper: run ``call`` under the breaker; on error log ``msg`` and return None."""
|
||||
try:
|
||||
result = call()
|
||||
self._record_success()
|
||||
return result
|
||||
except Exception as e:
|
||||
self._record_failure()
|
||||
log(msg, e)
|
||||
return None
|
||||
self._record_success()
|
||||
return result
|
||||
|
||||
def initialize(self, session_id: str, **kwargs) -> None:
|
||||
self._config = cfg = _load_config()
|
||||
@@ -206,7 +193,8 @@ class Mem0MemoryProvider(MemoryProvider):
|
||||
configured = cfg.get("user_id")
|
||||
self._user_id = (None if configured == _DEFAULT_USER_ID else configured) or kwargs.get("user_id") or _DEFAULT_USER_ID
|
||||
# Persisted rerank preference: default for mem0_search when the model omits ``rerank``. Platform-only.
|
||||
self._rerank_default = _truthy(cfg.get("rerank", False))
|
||||
_rr = cfg.get("rerank", False)
|
||||
self._rerank_default = _rr.lower() in ("true", "1", "yes") if isinstance(_rr, str) else bool(_rr)
|
||||
self._channel = kwargs.get("platform") or "cli"
|
||||
self._backend = self._create_backend()
|
||||
if self._backend and not self._atexit_registered:
|
||||
@@ -243,11 +231,6 @@ class Mem0MemoryProvider(MemoryProvider):
|
||||
backend = self._backend
|
||||
if not query or backend is None or self._is_breaker_open():
|
||||
return
|
||||
with self._prefetch_lock:
|
||||
# Same query already answered or still in flight: don't restart it.
|
||||
if self._prefetch_query == query and (self._prefetch_done or (self._prefetch_thread and self._prefetch_thread.is_alive())):
|
||||
return
|
||||
self._prefetch_query, self._prefetch_result, self._prefetch_done = query, "", False
|
||||
|
||||
def _run():
|
||||
results = self._try(lambda: self._search(query, backend=backend), logger.debug, "Mem0 prefetch failed: %s")
|
||||
@@ -257,9 +240,12 @@ class Mem0MemoryProvider(MemoryProvider):
|
||||
if self._prefetch_query == query:
|
||||
self._prefetch_result, self._prefetch_done = body, True
|
||||
|
||||
t = threading.Thread(target=_run, daemon=True, name="mem0-prefetch")
|
||||
with self._prefetch_lock:
|
||||
self._prefetch_thread = t
|
||||
# Same query already answered or still in flight: don't restart it.
|
||||
if self._prefetch_query == query and (self._prefetch_done or (self._prefetch_thread and self._prefetch_thread.is_alive())):
|
||||
return
|
||||
self._prefetch_query, self._prefetch_result, self._prefetch_done = query, "", False
|
||||
self._prefetch_thread = t = threading.Thread(target=_run, daemon=True, name="mem0-prefetch")
|
||||
t.start()
|
||||
|
||||
def prefetch(self, query: str, *, session_id: str = "") -> str:
|
||||
@@ -271,8 +257,7 @@ class Mem0MemoryProvider(MemoryProvider):
|
||||
thread = self._prefetch_thread if self._prefetch_query == query else None
|
||||
if thread:
|
||||
thread.join(timeout=_PREFETCH_WAIT_SECS)
|
||||
# Slow backend: skip injection; mem0_search tool remains the backstop.
|
||||
return self._consume_prefetch_result(query) or ""
|
||||
return self._consume_prefetch_result(query) or "" # slow backend: skip injection; mem0_search remains the backstop
|
||||
|
||||
def sync_turn(self, user_content: str, assistant_content: str, *, session_id: str = "") -> None:
|
||||
"""Send the turn to Mem0 for server-side fact extraction (non-blocking)."""
|
||||
@@ -302,7 +287,8 @@ class Mem0MemoryProvider(MemoryProvider):
|
||||
|
||||
def _tool_search(self, args: dict) -> str:
|
||||
top_k = max(1, min(int(args.get("top_k", 10)), 50))
|
||||
rerank = _truthy(args.get("rerank", self._rerank_default), falsy_strings=("false", "0", "no"))
|
||||
rerank_raw = args.get("rerank", self._rerank_default)
|
||||
rerank = rerank_raw.lower() not in ("false", "0", "no") if isinstance(rerank_raw, str) else bool(rerank_raw)
|
||||
results = self._search(args["query"], top_k, rerank)
|
||||
if not results:
|
||||
return json.dumps({"result": "No relevant memories found."})
|
||||
@@ -336,8 +322,6 @@ class Mem0MemoryProvider(MemoryProvider):
|
||||
return tool_error(f"Missing required parameter: {missing}")
|
||||
try:
|
||||
result = body(self, args)
|
||||
self._record_success()
|
||||
return result
|
||||
except Exception as e:
|
||||
client = _is_client_error(e)
|
||||
if client and on_client_error == "not_found":
|
||||
@@ -345,6 +329,8 @@ class Mem0MemoryProvider(MemoryProvider):
|
||||
if not client or on_client_error == "count":
|
||||
self._record_failure()
|
||||
return tool_error(self._format_error(label, e))
|
||||
self._record_success()
|
||||
return result
|
||||
|
||||
def _shutdown_backend(self):
|
||||
with suppress(Exception):
|
||||
|
||||
@@ -13,9 +13,7 @@ def _add_kwargs(user_id: str, agent_id: str, infer: bool, metadata: dict | None)
|
||||
|
||||
def _unwrap_results(response: Any) -> list:
|
||||
"""Normalize API response — extract results list from dict or pass through."""
|
||||
if isinstance(response, dict):
|
||||
return response.get("results", [])
|
||||
return response if isinstance(response, list) else []
|
||||
return response.get("results", []) if isinstance(response, dict) else response if isinstance(response, list) else []
|
||||
|
||||
|
||||
class Mem0Backend(ABC):
|
||||
@@ -105,7 +103,6 @@ def _register_direct_openai_provider() -> None:
|
||||
"""Register Hermes' OpenAI-only Mem0 LLM provider once per factory."""
|
||||
from mem0.configs.llms.openai import OpenAIConfig
|
||||
from mem0.utils.factory import LlmFactory
|
||||
|
||||
provider_map = getattr(LlmFactory, "provider_to_class", None)
|
||||
register_provider = getattr(LlmFactory, "register_provider", None)
|
||||
if not isinstance(provider_map, dict) or not callable(register_provider):
|
||||
@@ -143,7 +140,6 @@ class OSSBackend(Mem0Backend):
|
||||
vs_config["embedding_model_dims"] = dims
|
||||
self._recreate_collection_if_dims_changed(vector_store.get("provider", "qdrant"), vs_config, dims)
|
||||
vector_store["config"] = vs_config
|
||||
|
||||
config = {"vector_store": vector_store, "llm": _provider_block("llm", LLM_PROVIDERS), "embedder": _provider_block("embedder", EMBEDDER_PROVIDERS), "version": "v1.1"}
|
||||
if str(config["llm"].get("provider") or "").strip().lower() == "openai":
|
||||
# mem0 validates LlmConfig.provider before its factory lookup: build the supported OpenAI config, then swap the provider.
|
||||
|
||||
@@ -54,12 +54,9 @@ def _prompt_api_key(label: str, env_var: str, hermes_home: str) -> str:
|
||||
"""Prompt for API key, showing masked existing value if found."""
|
||||
existing = os.environ.get(env_var, "")
|
||||
env_path = Path(hermes_home) / ".env"
|
||||
if not existing and env_path.exists():
|
||||
# utf-8-sig: a Notepad BOM on line 1 would otherwise defeat the key match.
|
||||
for line in env_path.read_text(encoding="utf-8-sig", errors="replace").splitlines():
|
||||
if line.startswith(f"{env_var}="):
|
||||
existing = line.split("=", 1)[1].strip()
|
||||
break
|
||||
if not existing and env_path.exists(): # utf-8-sig: a Notepad BOM on line 1 would otherwise defeat the key match
|
||||
lines = env_path.read_text(encoding="utf-8-sig", errors="replace").splitlines()
|
||||
existing = next((line.split("=", 1)[1].strip() for line in lines if line.startswith(f"{env_var}=")), "")
|
||||
hint = f" (current: {_masked(existing)}, blank to keep)" if existing else ""
|
||||
return getpass.getpass(f" {label} API key{hint}: ").strip()
|
||||
|
||||
@@ -93,8 +90,7 @@ _FLAG_DEFAULTS = {"oss_llm": "openai", "oss_embedder": "openai", "oss_vector": "
|
||||
|
||||
def parse_flags(argv: list[str] | None = None) -> dict[str, str]:
|
||||
args = argv if argv is not None else sys.argv[1:]
|
||||
flags: dict[str, Any] = {k: _FLAG_DEFAULTS.get(k, "") for k in _FLAG_KEYS}
|
||||
flags["dry_run"] = False
|
||||
flags: dict[str, Any] = {**{k: _FLAG_DEFAULTS.get(k, "") for k in _FLAG_KEYS}, "dry_run": False}
|
||||
flag_map = {"--" + k.replace("_", "-"): k for k in _FLAG_KEYS}
|
||||
i = 0
|
||||
while i < len(args):
|
||||
@@ -127,7 +123,6 @@ def build_oss_config(flags: dict[str, str]) -> tuple[dict, dict[str, str]]:
|
||||
dims = KNOWN_DIMS.get(embedder_config["model"])
|
||||
if dims:
|
||||
embedder_config["embedding_dims"] = dims
|
||||
|
||||
vector_id = flags.get("oss_vector", "qdrant")
|
||||
vector_config = dict(VECTOR_PROVIDERS[vector_id]["default_config"])
|
||||
for key in _VECTOR_FLAG_KEYS.get(vector_id, ()):
|
||||
@@ -135,7 +130,6 @@ def build_oss_config(flags: dict[str, str]) -> tuple[dict, dict[str, str]]:
|
||||
vector_config[key] = int(val) if key == "port" else val
|
||||
if "url" in vector_config:
|
||||
vector_config.pop("path", None) # a remote Qdrant URL replaces local storage
|
||||
|
||||
oss_config = {"llm": {"provider": llm_id, "config": llm_config}, "embedder": {"provider": embedder_id, "config": embedder_config}, "vector_store": {"provider": vector_id, "config": vector_config}}
|
||||
# An embedder sharing the LLM's provider reuses the LLM key when no embedder key was given.
|
||||
llm_key = flags.get("oss_llm_key") if llm_def.get("needs_key") else ""
|
||||
@@ -170,14 +164,8 @@ def _persist_provider_config(hermes_home: str, config: dict, provider_config: di
|
||||
_write_env(Path(hermes_home) / ".env", env_writes)
|
||||
if server:
|
||||
_check_selfhosted_server(server)
|
||||
print(f"\n Memory provider: {label}")
|
||||
if server:
|
||||
print(f" Server: {server}")
|
||||
print(" Activation saved to config.yaml")
|
||||
print(" Provider config saved")
|
||||
if env_writes:
|
||||
print(f" {key_line}")
|
||||
print("\n Start a new session to activate.\n")
|
||||
print("\n".join(["", f" Memory provider: {label}", *([f" Server: {server}"] if server else []), " Activation saved to config.yaml", " Provider config saved",
|
||||
*([f" {key_line}"] if env_writes else []), "", " Start a new session to activate.", ""]))
|
||||
|
||||
|
||||
def _setup_platform(hermes_home: str, config: dict, flags: dict[str, str]) -> None:
|
||||
@@ -191,14 +179,12 @@ def _setup_platform(hermes_home: str, config: dict, flags: dict[str, str]) -> No
|
||||
choices = ["true", "false"]
|
||||
current = str(provider_config.get("rerank", "false") or "").lower()
|
||||
provider_config["rerank"] = choices[_curses_select(" Enable reranking for recall", [(c, "") for c in choices], default=choices.index(current) if current in choices else 0)]
|
||||
|
||||
if flags.get("dry_run"):
|
||||
_print_dry_run(str(provider_config), env_writes)
|
||||
return
|
||||
provider_config["mode"] = "platform"
|
||||
# Routing checks ``host`` before platform, so clear a stale self-hosted host. "" rather than
|
||||
# pop(): save_config merges into the existing mem0.json, so a popped key would survive.
|
||||
provider_config["host"] = ""
|
||||
provider_config.update(mode="platform", host="")
|
||||
# _load_config() also seeds ``host`` from MEM0_HOST (.env); the file clear can't help there, so warn.
|
||||
if os.environ.get("MEM0_HOST", "").strip():
|
||||
print(f"\n ⚠ MEM0_HOST is set in your environment ({os.environ['MEM0_HOST']}). It overrides platform mode — remove it from ~/.hermes/.env (or unset it) or Hermes will keep routing to the self-hosted server.")
|
||||
@@ -229,7 +215,6 @@ def _setup_selfhosted(hermes_home: str, config: dict, flags: dict[str, str]) ->
|
||||
env_writes = _api_key_writes(flags, "Server API key", fresh_label="Server API key (blank if AUTH_DISABLED)")
|
||||
user_id = flags.get("user_id") or _prompt("User identifier", default=provider_config.get("user_id") or "hermes-user")
|
||||
agent_id = _prompt("Agent identifier", default=provider_config.get("agent_id") or "hermes")
|
||||
|
||||
if flags.get("dry_run"):
|
||||
_print_dry_run(f"host={host}, user_id={user_id}, agent_id={agent_id}", env_writes, lambda: _check_selfhosted_server(host))
|
||||
return
|
||||
@@ -240,19 +225,14 @@ def _setup_selfhosted(hermes_home: str, config: dict, flags: dict[str, str]) ->
|
||||
def _print_oss_summary(oss_config: dict, env_writes: dict, dry_run: bool = False) -> None:
|
||||
llm, emb = oss_config["llm"], oss_config["embedder"]
|
||||
w = 0 if dry_run else 9 # final summary column-aligns the labels
|
||||
print("\n [dry-run] OSS config would be:" if dry_run else "\n ✓ Mem0 configured (OSS mode)")
|
||||
print(f" {'LLM:':<{w}} {llm['provider']} ({llm['config'].get('model', '')})")
|
||||
print(f" {'Embedder:':<{w}} {emb['provider']} ({emb['config'].get('model', '')})")
|
||||
print(f" {'Vector:':<{w}} {oss_config['vector_store']['provider']}")
|
||||
lines = ["", " [dry-run] OSS config would be:" if dry_run else " ✓ Mem0 configured (OSS mode)",
|
||||
f" {'LLM:':<{w}} {llm['provider']} ({llm['config'].get('model', '')})", f" {'Embedder:':<{w}} {emb['provider']} ({emb['config'].get('model', '')})",
|
||||
f" {'Vector:':<{w}} {oss_config['vector_store']['provider']}"]
|
||||
if dry_run:
|
||||
if env_writes:
|
||||
print(f" Env vars: {', '.join(env_writes.keys())}")
|
||||
return
|
||||
if env_writes:
|
||||
print(" API keys saved to .env")
|
||||
print(" Config saved to mem0.json")
|
||||
print(" Provider set in config.yaml")
|
||||
print("\n Start a new session to activate.\n")
|
||||
lines += [f" Env vars: {', '.join(env_writes.keys())}"] if env_writes else []
|
||||
else:
|
||||
lines += [*([" API keys saved to .env"] if env_writes else []), " Config saved to mem0.json", " Provider set in config.yaml", "", " Start a new session to activate.", ""]
|
||||
print("\n".join(lines))
|
||||
|
||||
|
||||
def _finish_oss(hermes_home: str, config: dict, oss_config: dict, env_writes: dict[str, str], user_id: str, agent_id: str, pgvector_config: dict | None = None) -> None:
|
||||
@@ -313,13 +293,9 @@ def _ensure_pgvector(host: str = "localhost", port: int = 5432) -> dict | None:
|
||||
if _pg_ready(host, port, 15):
|
||||
print(" ✓ PostgreSQL container restarted")
|
||||
return None
|
||||
if input(" Start pgvector via Docker? [Y/n]: ").strip().lower() in ("", "y", "yes"):
|
||||
return _start_pgvector_docker(host, port)
|
||||
print(" Skipping Docker setup. Make sure PostgreSQL with pgvector is running.")
|
||||
return None
|
||||
|
||||
|
||||
def _start_pgvector_docker(host: str, port: int) -> dict | None:
|
||||
if input(" Start pgvector via Docker? [Y/n]: ").strip().lower() not in ("", "y", "yes"):
|
||||
print(" Skipping Docker setup. Make sure PostgreSQL with pgvector is running.")
|
||||
return None
|
||||
try:
|
||||
print(f" Pulling {_PGVECTOR_IMAGE}...")
|
||||
_docker("pull", _PGVECTOR_IMAGE, timeout=120)
|
||||
@@ -357,7 +333,11 @@ def _ensure_ollama(models: list[str]) -> bool:
|
||||
print(" Warning: Ollama not reachable. Models cannot be pulled.")
|
||||
return False
|
||||
for model in models:
|
||||
if any(model in n or model.split(":")[0] in n for n in _ollama_models(_OLLAMA_URL)):
|
||||
try:
|
||||
names = [m.get("name", "") for m in json.loads(_http_get(_OLLAMA_URL, "/api/tags", 5).read()).get("models", [])]
|
||||
except Exception:
|
||||
names = []
|
||||
if any(model in n or model.split(":")[0] in n for n in names):
|
||||
print(f" ✓ Model '{model}' available")
|
||||
continue
|
||||
print(f" Pulling '{model}'... (this may take a few minutes)")
|
||||
@@ -369,13 +349,6 @@ def _ensure_ollama(models: list[str]) -> bool:
|
||||
return True
|
||||
|
||||
|
||||
def _ollama_models(url: str) -> list[str]:
|
||||
try:
|
||||
return [m.get("name", "") for m in json.loads(_http_get(url, "/api/tags", 5).read()).get("models", [])]
|
||||
except Exception:
|
||||
return []
|
||||
|
||||
|
||||
def _ensure_pgvector_extension(pg_config: dict) -> None:
|
||||
try:
|
||||
import psycopg2
|
||||
@@ -403,18 +376,13 @@ def _wait_for_port(host: str, port: int, timeout: int = 15) -> None:
|
||||
|
||||
|
||||
# Picker descriptions: LLM/embedder show model (+ URL); vector stores by provider id (default: the id itself).
|
||||
def _provider_description(v: dict) -> str:
|
||||
model, url = v.get("default_model", ""), v.get("default_url")
|
||||
return f"{model} ({url})" if url else model
|
||||
|
||||
|
||||
_VECTOR_DESCRIPTIONS = {"qdrant": lambda cfg: cfg.get("path", "local storage"), "pgvector": lambda cfg: f"{cfg.get('host', 'localhost')}:{cfg.get('port', 5432)}"}
|
||||
|
||||
|
||||
def _configure_model_provider(kind: str, registry: dict, hermes_home: str, env_writes: dict[str, str], llm: tuple[str, dict] | None = None) -> tuple[str, dict, str, str | None]:
|
||||
"""Pick an LLM/embedder provider, collect its key, and (for Ollama) model + URL -> (id, definition, model, url).
|
||||
For the embedder (``llm`` given), a provider shared with the LLM reuses the LLM key instead of prompting again."""
|
||||
items = [(v["label"], _provider_description(v)) for v in registry.values()]
|
||||
items = [(v["label"], f"{v.get('default_model', '')} ({v['default_url']})" if v.get("default_url") else v.get("default_model", "")) for v in registry.values()]
|
||||
pid = list(registry)[_curses_select(f"{kind} Provider", items, 0)]
|
||||
pdef = registry[pid]
|
||||
model, url = pdef["default_model"], pdef.get("default_url")
|
||||
@@ -430,31 +398,23 @@ def _configure_model_provider(kind: str, registry: dict, hermes_home: str, env_w
|
||||
return pid, pdef, model, url
|
||||
|
||||
|
||||
def _prompt_pgvector_config() -> dict:
|
||||
"""Native PostgreSQL — prompt for connection details (user first, matching the historical order)."""
|
||||
pg = {k: _input(f"PostgreSQL {label}", d) for k, label, d in (
|
||||
("user", "user", os.getenv("USER", "postgres")), ("host", "host", "localhost"), ("port", "port", "5432"), ("dbname", "database", "postgres"))}
|
||||
pg_password = getpass.getpass(" PostgreSQL password (blank if none): ").strip()
|
||||
return {**pg, "port": int(pg["port"]), **({"password": pg_password} if pg_password else {})}
|
||||
|
||||
|
||||
def _setup_oss_interactive(hermes_home: str, config: dict) -> None:
|
||||
env_writes: dict[str, str] = {}
|
||||
llm_id, llm_def, llm_model, llm_url = _configure_model_provider("LLM", LLM_PROVIDERS, hermes_home, env_writes)
|
||||
embedder_id, _, embedder_model, embedder_url = _configure_model_provider("Embedder", EMBEDDER_PROVIDERS, hermes_home, env_writes, llm=(llm_id, llm_def))
|
||||
vector_items = [(v["label"], _VECTOR_DESCRIPTIONS.get(pid, lambda cfg: pid)(v.get("default_config", {}))) for pid, v in VECTOR_PROVIDERS.items()]
|
||||
vector_id = list(VECTOR_PROVIDERS)[_curses_select("Vector Store", vector_items, 0)]
|
||||
|
||||
# Auto-setup: ensure Ollama is running and models are pulled; ensure pgvector is reachable (offer Docker if not).
|
||||
ollama_models = [m for pid, m in ((llm_id, llm_model), (embedder_id, embedder_model)) if pid == "ollama"]
|
||||
if ollama_models:
|
||||
_ensure_ollama(ollama_models)
|
||||
pgvector_config = None
|
||||
if vector_id == "pgvector":
|
||||
pgvector_config = _ensure_pgvector() or _prompt_pgvector_config()
|
||||
pgvector_config = _ensure_pgvector() if vector_id == "pgvector" else None
|
||||
if vector_id == "pgvector" and not pgvector_config: # native PostgreSQL: prompt for connection details (user first, historical order)
|
||||
pg = {k: _input(f"PostgreSQL {label}", d) for k, label, d in (("user", "user", os.getenv("USER", "postgres")), ("host", "host", "localhost"), ("port", "port", "5432"), ("dbname", "database", "postgres"))}
|
||||
pg_password = getpass.getpass(" PostgreSQL password (blank if none): ").strip()
|
||||
pgvector_config = {**pg, "port": int(pg["port"]), **({"password": pg_password} if pg_password else {})}
|
||||
user_id = _input("User ID", os.getenv("USER", "hermes-user"))
|
||||
agent_id = _input("Agent ID", "hermes")
|
||||
|
||||
flags = {
|
||||
"oss_llm": llm_id, "oss_llm_model": llm_model, "oss_llm_url": llm_url or "",
|
||||
"oss_llm_key": env_writes.get(llm_def["env_var"], "") if llm_def.get("env_var") else "",
|
||||
@@ -476,12 +436,8 @@ def _install_provider_deps(llm_id: str, embedder_id: str, vector_id: str) -> Non
|
||||
outcome = install_specs([dep], timeout=60)
|
||||
except Exception:
|
||||
outcome = None
|
||||
if outcome is not None and outcome.ok:
|
||||
print(f" ✓ Installed {dep}")
|
||||
elif outcome is not None and outcome.blocked:
|
||||
print(f" Warning: cannot install {dep}: {outcome.reason}")
|
||||
else:
|
||||
print(f" Warning: Could not install {dep}. Install manually: uv pip install {dep}")
|
||||
print(f" ✓ Installed {dep}" if outcome is not None and outcome.ok else f" Warning: cannot install {dep}: {outcome.reason}" if outcome is not None and outcome.blocked
|
||||
else f" Warning: Could not install {dep}. Install manually: uv pip install {dep}")
|
||||
if deps:
|
||||
import importlib
|
||||
importlib.invalidate_caches()
|
||||
|
||||
Reference in New Issue
Block a user