diff --git a/cli/node/src/backend/base.ts b/cli/node/src/backend/base.ts index c4da82f54..c0c9ae4f9 100644 --- a/cli/node/src/backend/base.ts +++ b/cli/node/src/backend/base.ts @@ -89,6 +89,8 @@ export interface Backend { deleteEntities(opts: EntityIds): Promise>; + ping(): Promise>; + status(opts?: { userId?: string; agentId?: string }): Promise< Record >; diff --git a/cli/node/src/index.ts b/cli/node/src/index.ts index 8c0acb17f..41de42cfd 100644 --- a/cli/node/src/index.ts +++ b/cli/node/src/index.ts @@ -8,10 +8,10 @@ import fs from "node:fs"; import path from "node:path"; import { fileURLToPath } from "node:url"; import { Command } from "commander"; -import { type Backend, getBackend } from "./backend/index.js"; -import { colors, printError } from "./branding.js"; +import { AuthError, type Backend, getBackend } from "./backend/index.js"; +import { colors, printError, printWarning } from "./branding.js"; import type { Mem0Config } from "./config.js"; -import { loadConfig } from "./config.js"; +import { loadConfig, saveConfig } from "./config.js"; import { richFormatHelp } from "./help.js"; import { setAgentMode } from "./state.js"; import { captureEvent } from "./telemetry.js"; @@ -19,12 +19,16 @@ import { CLI_VERSION } from "./version.js"; const program = new Command(); +// ── Validated user identity (set by getBackendAndConfig) ───────────────── + +let _validatedUserEmail: string | undefined; + // ── Helpers ────────────────────────────────────────────────────────────── -function getBackendAndConfig( +async function getBackendAndConfig( apiKey?: string, baseUrl?: string, -): { backend: Backend; config: Mem0Config } { +): Promise<{ backend: Backend; config: Mem0Config }> { const config = loadConfig(); if (apiKey) config.platform.apiKey = apiKey; @@ -38,11 +42,51 @@ function getBackendAndConfig( process.exit(1); } - return { backend: getBackend(config), config }; + const backend = getBackend(config); + + // Validate the API key upfront with a fast timeout + try { + const pingData = (await Promise.race([ + backend.ping(), + new Promise((_, reject) => + setTimeout(() => reject(new Error("timeout")), 5000), + ), + ])) as Record; + + const email = pingData?.user_email as string | undefined; + if (email) { + _validatedUserEmail = email; + if (config.platform.userEmail !== email) { + config.platform.userEmail = email; + try { + saveConfig(config); + } catch { + /* ignore */ + } + } + } + } catch (e) { + if (e instanceof AuthError) { + printError( + "Invalid or expired API key.", + "Run 'mem0 init' or set MEM0_API_KEY environment variable.", + ); + process.exit(1); + } + // Network error / timeout — warn but proceed + printWarning( + "Could not validate API key (network issue). Proceeding anyway.", + ); + } + + return { backend, config }; } -function getBackendOnly(apiKey?: string, baseUrl?: string): Backend { - return getBackendAndConfig(apiKey, baseUrl).backend; +async function getBackendOnly( + apiKey?: string, + baseUrl?: string, +): Promise { + return (await getBackendAndConfig(apiKey, baseUrl)).backend; } function checkAgentMode(): boolean { @@ -135,10 +179,14 @@ program.hook("preAction", (_thisCommand, actionCommand) => { ? `${parentName}.${commandName}` : commandName; const isAgent = !!(program.opts().json || program.opts().agent); - captureEvent(`cli.${fullCommand}`, { - command: fullCommand, - is_agent: isAgent, - }); + captureEvent( + `cli.${fullCommand}`, + { + command: fullCommand, + is_agent: isAgent, + }, + _validatedUserEmail, + ); } catch { /* silently swallow */ } @@ -200,7 +248,10 @@ program .action(async (text, opts) => { const { cmdAdd } = await import("./commands/memory.js"); const isAgent = checkAgentMode(); - const { backend, config } = getBackendAndConfig(opts.apiKey, opts.baseUrl); + const { backend, config } = await getBackendAndConfig( + opts.apiKey, + opts.baseUrl, + ); const ids = resolveIds(config, opts); const enableGraph = resolveGraph(config, opts); const output = isAgent ? "agent" : opts.output; @@ -254,7 +305,10 @@ program } const { cmdSearch } = await import("./commands/memory.js"); const isAgent = checkAgentMode(); - const { backend, config } = getBackendAndConfig(opts.apiKey, opts.baseUrl); + const { backend, config } = await getBackendAndConfig( + opts.apiKey, + opts.baseUrl, + ); const ids = resolveIds(config, opts); const enableGraph = resolveGraph(config, opts); const output = isAgent ? "agent" : opts.output; @@ -286,7 +340,7 @@ program .action(async (memoryId, opts) => { const { cmdGet } = await import("./commands/memory.js"); const isAgent = checkAgentMode(); - const backend = getBackendOnly(opts.apiKey, opts.baseUrl); + const backend = await getBackendOnly(opts.apiKey, opts.baseUrl); const output = isAgent ? "agent" : opts.output; await cmdGet(backend, memoryId, { output }); }); @@ -322,7 +376,10 @@ program .action(async (opts) => { const { cmdList } = await import("./commands/memory.js"); const isAgent = checkAgentMode(); - const { backend, config } = getBackendAndConfig(opts.apiKey, opts.baseUrl); + const { backend, config } = await getBackendAndConfig( + opts.apiKey, + opts.baseUrl, + ); const ids = resolveIds(config, opts); const enableGraph = resolveGraph(config, opts); const output = isAgent ? "agent" : opts.output; @@ -358,7 +415,7 @@ program } const { cmdUpdate } = await import("./commands/memory.js"); const isAgent = checkAgentMode(); - const backend = getBackendOnly(opts.apiKey, opts.baseUrl); + const backend = await getBackendOnly(opts.apiKey, opts.baseUrl); const output = isAgent ? "agent" : opts.output; await cmdUpdate(backend, memoryId, resolvedText, { metadata: opts.metadata, @@ -428,7 +485,7 @@ program // ── Dispatch: single memory ── if (memoryId) { const { cmdDelete } = await import("./commands/memory.js"); - const backend = getBackendOnly(opts.apiKey, opts.baseUrl); + const backend = await getBackendOnly(opts.apiKey, opts.baseUrl); await cmdDelete(backend, memoryId, { output, dryRun: opts.dryRun, @@ -440,7 +497,7 @@ program // ── Dispatch: --all ── if (opts.all) { const { cmdDeleteAll } = await import("./commands/memory.js"); - const { backend, config } = getBackendAndConfig( + const { backend, config } = await getBackendAndConfig( opts.apiKey, opts.baseUrl, ); @@ -465,7 +522,7 @@ program // ── Dispatch: --entity ── if (opts.entity) { const { cmdEntitiesDelete } = await import("./commands/entities.js"); - const backend = getBackendOnly(opts.apiKey, opts.baseUrl); + const backend = await getBackendOnly(opts.apiKey, opts.baseUrl); await cmdEntitiesDelete(backend, { ...opts, output }); return; } @@ -540,7 +597,7 @@ entityCmd .action(async (entityType, opts) => { const { cmdEntitiesList } = await import("./commands/entities.js"); const isAgent = checkAgentMode(); - const backend = getBackendOnly(opts.apiKey, opts.baseUrl); + const backend = await getBackendOnly(opts.apiKey, opts.baseUrl); const output = isAgent ? "agent" : opts.output; await cmdEntitiesList(backend, entityType, { output }); }); @@ -564,7 +621,7 @@ entityCmd .action(async (opts) => { const { cmdEntitiesDelete } = await import("./commands/entities.js"); const isAgent = checkAgentMode(); - const backend = getBackendOnly(opts.apiKey, opts.baseUrl); + const backend = await getBackendOnly(opts.apiKey, opts.baseUrl); const output = isAgent ? "agent" : opts.output; await cmdEntitiesDelete(backend, { ...opts, output }); }); @@ -590,7 +647,7 @@ eventCmd .action(async (opts) => { const { cmdEventList } = await import("./commands/events.js"); const isAgent = checkAgentMode(); - const backend = getBackendOnly(opts.apiKey, opts.baseUrl); + const backend = await getBackendOnly(opts.apiKey, opts.baseUrl); const output = isAgent ? "agent" : opts.output; await cmdEventList(backend, { output }); }); @@ -608,7 +665,7 @@ eventCmd .action(async (eventId, opts) => { const { cmdEventStatus } = await import("./commands/events.js"); const isAgent = checkAgentMode(); - const backend = getBackendOnly(opts.apiKey, opts.baseUrl); + const backend = await getBackendOnly(opts.apiKey, opts.baseUrl); const output = isAgent ? "agent" : opts.output; await cmdEventStatus(backend, eventId, { output }); }); @@ -625,7 +682,10 @@ program .action(async (opts) => { const { cmdStatus } = await import("./commands/utils.js"); const isAgent = checkAgentMode(); - const { backend, config } = getBackendAndConfig(opts.apiKey, opts.baseUrl); + const { backend, config } = await getBackendAndConfig( + opts.apiKey, + opts.baseUrl, + ); const output = isAgent ? "agent" : opts.output; await cmdStatus(backend, { userId: config.defaults.userId || undefined, @@ -649,7 +709,10 @@ program .action(async (filePath, opts) => { const { cmdImport } = await import("./commands/utils.js"); const isAgent = checkAgentMode(); - const { backend, config } = getBackendAndConfig(opts.apiKey, opts.baseUrl); + const { backend, config } = await getBackendAndConfig( + opts.apiKey, + opts.baseUrl, + ); const ids = resolveIds(config, opts); const output = isAgent ? "agent" : opts.output; await cmdImport(backend, filePath, { diff --git a/cli/node/src/telemetry.ts b/cli/node/src/telemetry.ts index 9a22fd2dd..c5680bf1d 100644 --- a/cli/node/src/telemetry.ts +++ b/cli/node/src/telemetry.ts @@ -53,16 +53,21 @@ function getDistinctId(): string { /** * Fire a PostHog event (non-blocking, returns void, never throws). * Spawns telemetry-sender.cjs as a detached subprocess. + * + * When `preResolvedEmail` is provided (e.g. from an upfront ping + * validation), it is used directly as the PostHog distinct ID and the + * subprocess skips its own `/v1/ping/` call. */ export function captureEvent( eventName: string, properties: Record = {}, + preResolvedEmail?: string, ): void { if (!isTelemetryEnabled()) return; try { const config = loadConfig(); - const distinctId = getDistinctId(); + const distinctId = preResolvedEmail || getDistinctId(); const payload = { api_key: POSTHOG_API_KEY, diff --git a/cli/python/src/mem0_cli/app.py b/cli/python/src/mem0_cli/app.py index ff2881267..bdfb75647 100644 --- a/cli/python/src/mem0_cli/app.py +++ b/cli/python/src/mem0_cli/app.py @@ -2,6 +2,7 @@ from __future__ import annotations +import contextlib import json as _json import os import stat as _stat_mod @@ -12,7 +13,7 @@ import typer from rich.console import Console from mem0_cli import __version__ -from mem0_cli.branding import BRAND_COLOR, print_error +from mem0_cli.branding import BRAND_COLOR, print_error, print_warning console = Console() err_console = Console(stderr=True) @@ -55,6 +56,10 @@ event_app = typer.Typer( # entity_app and event_app registered after Memory commands to control panel ordering +# ── Validated user identity (set by _get_backend_and_config) ────────────── + +_validated_user_email: str | None = None + # ── Telemetry helper ───────────────────────────────────────────────────── @@ -66,7 +71,7 @@ def _fire_telemetry(command_name: str, extra: dict | None = None) -> None: props = {"command": command_name} if extra: props.update(extra) - capture_event(f"cli.{command_name}", props) + capture_event(f"cli.{command_name}", props, pre_resolved_email=_validated_user_email) except Exception: pass @@ -96,9 +101,16 @@ def _get_backend_and_config( api_key: str | None = None, base_url: str | None = None, ): - """Build and return the Platform backend plus the loaded config.""" + """Build and return the Platform backend plus the loaded config. + + Validates the API key upfront via ``/v1/ping/`` and caches the + resolved user email for telemetry. + """ + global _validated_user_email + from mem0_cli.backend import get_backend - from mem0_cli.config import load_config + from mem0_cli.backend.platform import AuthError + from mem0_cli.config import load_config, save_config config = load_config() @@ -115,7 +127,29 @@ def _get_backend_and_config( ) raise typer.Exit(1) - return get_backend(config), config + backend = get_backend(config) + + # Validate the API key upfront with a fast timeout + try: + ping_data = backend.ping(timeout=5.0) + email = ping_data.get("user_email") if isinstance(ping_data, dict) else None + if email: + _validated_user_email = email + if config.platform.user_email != email: + config.platform.user_email = email + with contextlib.suppress(Exception): + save_config(config) + except AuthError: + print_error( + err_console, + "Invalid or expired API key.", + hint="Run 'mem0 init' or set MEM0_API_KEY environment variable.", + ) + raise typer.Exit(1) from None + except Exception: + print_warning(err_console, "Could not validate API key (network issue). Proceeding anyway.") + + return backend, config def _get_backend( diff --git a/cli/python/src/mem0_cli/backend/platform.py b/cli/python/src/mem0_cli/backend/platform.py index 0114d1f26..9224412e8 100644 --- a/cli/python/src/mem0_cli/backend/platform.py +++ b/cli/python/src/mem0_cli/backend/platform.py @@ -288,8 +288,18 @@ class PlatformBackend(Backend): result = self._request("DELETE", f"/v2/entities/{entity_type}/{entity_id}/") return result - def ping(self) -> dict: - """Call the ping endpoint and return the raw response.""" + def ping(self, timeout: float | None = None) -> dict: + """Call the ping endpoint and return the raw response. + + When *timeout* is given it overrides the client-level timeout so that + validation pings can fail fast without blocking the user. + """ + if timeout is not None: + resp = self._client.get("/v1/ping/", timeout=timeout) + if resp.status_code == 401: + raise AuthError("Authentication failed. Your API key may be invalid or expired.") + resp.raise_for_status() + return resp.json() return self._request("GET", "/v1/ping/") def status( diff --git a/cli/python/src/mem0_cli/telemetry.py b/cli/python/src/mem0_cli/telemetry.py index d7cebdde6..4ca15d986 100644 --- a/cli/python/src/mem0_cli/telemetry.py +++ b/cli/python/src/mem0_cli/telemetry.py @@ -45,8 +45,17 @@ def _get_distinct_id() -> str: return "anonymous-cli" -def capture_event(event_name: str, properties: dict[str, Any] | None = None) -> None: - """Fire a PostHog event via a detached subprocess (non-blocking).""" +def capture_event( + event_name: str, + properties: dict[str, Any] | None = None, + pre_resolved_email: str | None = None, +) -> None: + """Fire a PostHog event via a detached subprocess (non-blocking). + + When *pre_resolved_email* is provided (e.g. from an upfront ping + validation), it is used directly as the PostHog distinct ID and the + subprocess skips its own ``/v1/ping/`` call. + """ if not _is_telemetry_enabled(): return @@ -56,7 +65,7 @@ def capture_event(event_name: str, properties: dict[str, Any] | None = None) -> from mem0_cli.state import is_agent_mode config = load_config() - distinct_id = _get_distinct_id() + distinct_id = pre_resolved_email or _get_distinct_id() payload = { "api_key": POSTHOG_API_KEY,