diff --git a/apps/memos-local-plugin/adapters/hermes/memos_provider/__init__.py b/apps/memos-local-plugin/adapters/hermes/memos_provider/__init__.py index 25546dfff..47392b8a7 100644 --- a/apps/memos-local-plugin/adapters/hermes/memos_provider/__init__.py +++ b/apps/memos-local-plugin/adapters/hermes/memos_provider/__init__.py @@ -65,13 +65,11 @@ if str(_PLUGIN_DIR) not in sys.path: sys.path.insert(0, str(_PLUGIN_DIR)) -from bridge_client import BridgeError, MemosBridgeClient, MemosHttpClient # noqa: E402 +from bridge_client import BridgeError, MemosBridgeClient # noqa: E402 from daemon_manager import ( # noqa: E402 ensure_bridge_running, ensure_viewer_daemon, kill_zombie_bridges, - probe_viewer_status, - startup_lock_active, ) @@ -279,7 +277,7 @@ class MemTensorProvider(MemoryProvider): """ def __init__(self) -> None: - self._bridge: MemosBridgeClient | MemosHttpClient | None = None + self._bridge: MemosBridgeClient | None = None self._reconnect_lock = threading.Lock() self._session_id: str = "" self._episode_id: str = "" @@ -329,23 +327,6 @@ def is_available(self) -> bool: # type: ignore[override] # ─── Lifecycle ──────────────────────────────────────────────────────── - def _connect_http_bridge(self, session_id: str, *, timeout: float = 60.0) -> bool: - """Try to connect via HTTP bridge. Sets self._bridge on success.""" - http_bridge: MemosHttpClient | None = None - try: - http_bridge = MemosHttpClient() - http_bridge.register_host_handler("host.llm.complete", self._handle_host_llm_complete) - self._bridge = http_bridge - self._open_session(session_id, timeout=timeout) - return True - except Exception as err: - logger.warning("MemOS: HTTP bridge failed, falling back to stdio — %s", err) - if http_bridge is not None: - with contextlib.suppress(Exception): - http_bridge.close() - self._bridge = None - return False - def initialize(self, session_id: str, **kwargs: Any) -> None: # type: ignore[override] """Called once at agent startup. @@ -414,32 +395,13 @@ def initialize(self, session_id: str, **kwargs: Any) -> None: # type: ignore[ov except Exception: pass - # If the daemon is already running on the viewer port, connect - # to it over HTTP instead of spawning a new stdio bridge. This - # eliminates zombie bridge accumulation. - viewer_status = probe_viewer_status() - if viewer_status == "running_memos": - if self._connect_http_bridge(session_id): - logger.info( - "MemOS: bridge ready (HTTP) session=%s platform=%s (episode deferred)", - self._session_id, - self._platform, - ) - else: - viewer_status = "free" # force stdio fallback below - elif viewer_status == "free": - # Re-probe after a short wait only when another process may be - # mid-startup (startup lock is held). On a cold first-launch the - # lock doesn't exist, so we skip the delay entirely. - if startup_lock_active(): - time.sleep(1.0) - viewer_status = probe_viewer_status() - if viewer_status == "running_memos" and self._connect_http_bridge(session_id): - logger.info( - "MemOS: bridge ready (HTTP, late probe) session=%s platform=%s (episode deferred)", - self._session_id, - self._platform, - ) + # NOTE: An HTTP bridge path used to live here that connected to a + # running viewer daemon over HTTP instead of spawning a stdio + # subprocess. It depended on ``MemosHttpClient`` which was never + # committed — the class was referenced by name only. Issue #2096 + # reverts the half-merged HTTP feature; the stdio path below is + # the sole connection mechanism until the HTTP client lands as a + # complete change. if self._bridge is None: try: @@ -1958,13 +1920,9 @@ def _reconnect_bridge(self, session_id: str = "", *, timeout: float = 30.0) -> N logger.info("MemOS: old bridge closed (pid=%s)", old_pid) ensure_bridge_running() - # Try HTTP first if daemon is running - viewer_status = probe_viewer_status() - if viewer_status == "running_memos" and self._connect_http_bridge( - session_id, timeout=timeout - ): - logger.info("MemOS: reconnected via HTTP") - return + # NOTE: HTTP bridge reconnect path was removed alongside issue + # #2096. See ``initialize`` for the rationale. Reconnect always + # spawns a fresh stdio bridge. try: ensure_viewer_daemon() diff --git a/apps/memos-local-plugin/core/llm/client.ts b/apps/memos-local-plugin/core/llm/client.ts index 41f31a439..6bedafa70 100644 --- a/apps/memos-local-plugin/core/llm/client.ts +++ b/apps/memos-local-plugin/core/llm/client.ts @@ -260,6 +260,21 @@ export function createLlmClientWithProvider( return [{ role: "system", content: systemInsert }, ...messages]; } + function ensureJsonWordInUserMessage(messages: LlmMessage[]): LlmMessage[] { + const lastUserIdx = messages.map((m) => m.role).lastIndexOf("user"); + if (lastUserIdx < 0) return [...messages, { role: "user", content: "Return valid json only." }]; + + const msg = messages[lastUserIdx]; + if (/\bjson\b/i.test(msg.content)) return messages; + + const out = messages.slice(); + out[lastUserIdx] = { + ...msg, + content: `${msg.content}\n\nReturn valid json only.`, + }; + return out; + } + function buildCallInput(opts: LlmCallOptions | undefined, jsonMode: boolean): ProviderCallInput { return { temperature: opts?.temperature ?? config.temperature, @@ -463,7 +478,7 @@ export function createLlmClientWithProvider( ): Promise { const messages = normalizeMessages(input); const msgsWithJsonHint = opts?.jsonMode - ? inject(messages, buildJsonSystemHint()) + ? ensureJsonWordInUserMessage(inject(messages, buildJsonSystemHint())) : messages; const call = buildCallInput(opts, opts?.jsonMode === true); const { completion } = await callWithFallback(msgsWithJsonHint, call, opts, opts?.op ?? "complete"); @@ -476,7 +491,7 @@ export function createLlmClientWithProvider( ): Promise> { const messages = normalizeMessages(input); const systemHint = buildJsonSystemHint(opts.schemaHint); - const msgs = inject(messages, systemHint); + const msgs = ensureJsonWordInUserMessage(inject(messages, systemHint)); const call = buildCallInput(opts, true); const op = opts.op ?? "complete.json"; const maxMalformedRetries = Math.max(0, opts.malformedRetries ?? 1); diff --git a/apps/memos-local-plugin/tests/python/test_hermes_provider_pipeline.py b/apps/memos-local-plugin/tests/python/test_hermes_provider_pipeline.py index 1c8cb6a6c..71e3b347b 100644 --- a/apps/memos-local-plugin/tests/python/test_hermes_provider_pipeline.py +++ b/apps/memos-local-plugin/tests/python/test_hermes_provider_pipeline.py @@ -68,6 +68,51 @@ def request(self, method: str, params: dict | None = None, **_kwargs: object) -> class HermesProviderPipelineTests(unittest.TestCase): + def test_module_imports_cleanly(self) -> None: + """Regression guard for #2096: asserts that ``MemosHttpClient`` is + NOT present in ``memos_provider``, since the class was referenced + before it was ever committed (see issue #2096). + + Note: the import itself is already validated at collection time — + the ``import memos_provider`` at the top of this file will raise + ``ImportError`` if a dangling reference is reintroduced, causing + the entire test file to fail to load. This test body only adds: + + * the explicit negative guard on ``MemosHttpClient`` below (unique + to this test), which covers both the ``memos_provider`` + re-export surface *and* ``bridge_client`` itself so a partial + re-add of the class only in ``bridge_client`` (with no + matching re-export) still fails the guard, and + * positive checks on ``MemTensorProvider`` (the class the Hermes + host actually instantiates) and on ``bridge_client``'s real + contract (``MemosBridgeClient`` / ``BridgeError``), rather than + on their incidental re-exports through ``memos_provider`` — the + latter only appear on the package namespace because + ``__init__.py`` uses a bare ``from bridge_client import ...``, + which is an implementation detail we don't want the test to + lock in. + """ + import importlib + + # MemTensorProvider is the class hermes-agent host instantiates. + self.assertTrue(hasattr(memos_provider, "MemTensorProvider")) + + # Assert the actual contract on bridge_client directly rather + # than on its re-exports through memos_provider. + bc = importlib.import_module("bridge_client") + self.assertTrue(hasattr(bc, "MemosBridgeClient")) + self.assertTrue(hasattr(bc, "BridgeError")) + + # ``MemosHttpClient`` was referenced by name in a half-merged HTTP + # bridge feature (see #2096). It must not reappear until the class + # itself is committed in ``bridge_client``. Guard both the + # ``memos_provider`` re-export (which is what the original + # ImportError travelled through) and ``bridge_client`` itself — + # otherwise a partial re-add of the class in ``bridge_client`` + # without a matching re-export would slip past this test. + self.assertFalse(hasattr(memos_provider, "MemosHttpClient")) + self.assertFalse(hasattr(bc, "MemosHttpClient")) + def test_lifecycle_persists_turn_and_closes_real_episode(self) -> None: bridge = FakeBridge() with ( diff --git a/apps/memos-local-plugin/tests/unit/llm/client.test.ts b/apps/memos-local-plugin/tests/unit/llm/client.test.ts index 7e904a2c6..dee0de228 100644 --- a/apps/memos-local-plugin/tests/unit/llm/client.test.ts +++ b/apps/memos-local-plugin/tests/unit/llm/client.test.ts @@ -96,12 +96,14 @@ describe("llm/client", () => { expect(fake.lastMessages).toEqual([{ role: "user", content: "hi there" }]); }); - it("injects a json system hint when jsonMode=true", async () => { + it("injects json hints into system and user messages when jsonMode=true", async () => { const fake = new FakeProvider("openai_compatible", () => ({ text: '{"ok":1}', durationMs: 1 })); const client = createLlmClientWithProvider(cfg(), fake); await client.complete("do it", { jsonMode: true }); expect(fake.lastMessages?.[0]?.role).toBe("system"); expect(fake.lastMessages?.[0]?.content).toMatch(/single valid JSON value/i); + expect(fake.lastMessages?.at(-1)?.role).toBe("user"); + expect(fake.lastMessages?.at(-1)?.content).toMatch(/valid json only/i); expect(fake.lastInput?.jsonMode).toBe(true); }); @@ -270,7 +272,9 @@ describe("llm/client", () => { expect(fake.lastMessages?.[0]?.role).toBe("system"); expect(fake.lastMessages?.[0]?.content).toMatch(/You are strict\./); expect(fake.lastMessages?.[0]?.content).toMatch(/single valid JSON value/); - expect(fake.lastMessages?.[1]).toEqual({ role: "user", content: "go" }); + expect(fake.lastMessages?.[1]?.role).toBe("user"); + expect(fake.lastMessages?.[1]?.content).toMatch(/^go/); + expect(fake.lastMessages?.[1]?.content).toMatch(/valid json only/i); }); it("rejects empty messages array", async () => { diff --git a/docs/cn/open_source/open_source_api/scheduler/get_status.md b/docs/cn/open_source/open_source_api/scheduler/get_status.md index 87e1a4a5a..0a60d6fac 100644 --- a/docs/cn/open_source/open_source_api/scheduler/get_status.md +++ b/docs/cn/open_source/open_source_api/scheduler/get_status.md @@ -64,34 +64,48 @@ desc: 监控 MemOS 异步任务的生命周期,提供包括任务进度、队 ## 4. 快速上手示例 -使用 SDK 轮询任务状态直至完成: +这些接口由开源版 Server(`server_api`,路由前缀 `/product`)直接提供,使用标准 HTTP 请求即可访问。以下示例轮询任务状态直至完成: ```python -from memos.api.client import MemOSClient import time -client = MemOSClient(api_key="...", base_url="...") +import requests + +# 自部署 MemOS Server 的地址(如启用了鉴权,请自行补充 Authorization 请求头) +base_url = "http://localhost:8000" # 1. 系统级概览:查看整个 MemOS 系统的运行健康度 -global_res = client.get_all_scheduler_status() -if global_res: - print(f"系统运行概况: {global_res.data['scheduler_summary']}") +resp = requests.get(f"{base_url}/product/scheduler/allstatus", timeout=10) +resp.raise_for_status() +global_res = resp.json() +print(f"系统运行概况: {global_res['data']['scheduler_summary']}") # 2. 队列指标监控:检查特定用户的任务积压情况 -queue_res = client.get_task_queue_status(user_id="dev_user_01") -if queue_res: - print(f"待处理任务数: {queue_res.data['remaining_tasks_count']}") - print(f"已下发未完成任务数: {queue_res.data['pending_tasks_count']}") +resp = requests.get( + f"{base_url}/product/scheduler/task_queue_status", + params={"user_id": "dev_user_01"}, + timeout=10, +) +resp.raise_for_status() +queue_res = resp.json() +print(f"排队中任务数: {queue_res['data']['remaining_tasks_count']}") +print(f"已下发未确认任务数: {queue_res['data']['pending_tasks_count']}") # 3. 任务进度追踪:轮询特定任务直至结束 task_id = "task_888999" +active_states = {"waiting", "pending", "in_progress"} while True: - res = client.get_task_status(user_id="dev_user_01", task_id=task_id) - if res and res.code == 200: - current_status = res.data[0]['status'] # data 为状态列表 - print(f"任务 {task_id} 当前状态: {current_status}") - - if current_status in ['completed', 'failed', 'cancelled']: - break + resp = requests.get( + f"{base_url}/product/scheduler/status", + params={"user_id": "dev_user_01", "task_id": task_id}, + timeout=10, + ) + resp.raise_for_status() + items = resp.json().get("data", []) # data 为状态列表:[{"task_id": ..., "status": ...}] + statuses = {item["status"] for item in items} + print(f"任务 {task_id} 当前状态: {statuses or '空'}") + + if not statuses or statuses.isdisjoint(active_states): + break time.sleep(2) ``` diff --git a/docs/cn/open_source/open_source_api/scheduler/ wait.md b/docs/cn/open_source/open_source_api/scheduler/wait.md similarity index 61% rename from docs/cn/open_source/open_source_api/scheduler/ wait.md rename to docs/cn/open_source/open_source_api/scheduler/wait.md index 9849ffe68..52b0d79f6 100644 --- a/docs/cn/open_source/open_source_api/scheduler/ wait.md +++ b/docs/cn/open_source/open_source_api/scheduler/wait.md @@ -42,36 +42,47 @@ desc: 提供阻塞等待与流式进度观测能力,确保在执行后续操 ## 4. 快速上手示例 -使用开源版 SDK 进行阻塞式等待: +这些接口由开源版 Server(`server_api`,路由前缀 `/product`)直接提供,使用标准 HTTP 请求即可访问。注意:`user_name`、`timeout_seconds`、`poll_interval` 均为查询参数(Query),而非请求体(Body)。以下示例进行阻塞式等待: ```python -from memos.api.client import MemOSClient +import json -client = MemOSClient(api_key="...", base_url="...") +import requests + +# 自部署 MemOS Server 的地址(如启用了鉴权,请自行补充 Authorization 请求头) +base_url = "http://localhost:8000" user_name = "dev_user_01" # --- 场景 A:同步阻塞等待 (常用于 Python 自动化脚本) --- print(f"正在等待用户 {user_name} 的任务队列清空...") -res = client.wait_until_idle( - user_name=user_name, - timeout_seconds=300, - poll_interval=2 +resp = requests.post( + f"{base_url}/product/scheduler/wait", + params={"user_name": user_name, "timeout_seconds": 300, "poll_interval": 2}, + timeout=310, # HTTP 超时应大于 timeout_seconds ) -if res and res.code == 200: +resp.raise_for_status() +result = resp.json() # {"message": "idle" | "timeout", "data": {...}} +if result["message"] == "idle": print("✅ 任务已全部完成。") +else: + print(f"⚠️ 等待超时,仍有 {result['data']['running_tasks']} 个任务在执行。") # --- 场景 B:流式进度观测 (常用于前端进度条渲染) --- print("开始监听任务实时进度流...") -# 注意:SSE 接口在 SDK 中通常返回一个生成器 (Generator) -progress_stream = client.stream_scheduler_progress( - user_name=user_name, - timeout_seconds=300 -) - -for event in progress_stream: - # 实时打印剩余任务数 - print(f"当前排队任务数: {event['remaining_tasks_count']}") - if event['status'] == 'idle': - print("🎉 调度器已空闲") - break +with requests.get( + f"{base_url}/product/scheduler/wait/stream", + params={"user_name": user_name, "timeout_seconds": 300}, + stream=True, + timeout=310, +) as resp: + resp.raise_for_status() + for line in resp.iter_lines(decode_unicode=True): + if not line or not line.startswith("data:"): + continue + event = json.loads(line.removeprefix("data:").strip()) + # 实时打印仍在执行的任务数 + print(f"当前活跃任务数: {event['active_tasks']},状态: {event['status']}") + if event["status"] in ("idle", "timeout"): + print("🎉 调度器已空闲" if event["status"] == "idle" else "⚠️ 监听超时") + break ``` diff --git a/docs/en/open_source/open_source_api/scheduler/get_status.md b/docs/en/open_source/open_source_api/scheduler/get_status.md index f2014d9e5..7d1565858 100644 --- a/docs/en/open_source/open_source_api/scheduler/get_status.md +++ b/docs/en/open_source/open_source_api/scheduler/get_status.md @@ -66,34 +66,48 @@ When you send a status request, **SchedulerHandler** performs the following oper ## 4. Quick Start -Poll task status with the SDK until completion: +These endpoints are served directly by the open-source Server (`server_api`, router prefix `/product`) and can be called with plain HTTP requests. The example below polls task status until completion: ```python -from memos.api.client import MemOSClient import time -client = MemOSClient(api_key="...", base_url="...") +import requests + +# Address of your self-hosted MemOS Server (add an Authorization header if auth is enabled) +base_url = "http://localhost:8000" # 1. System overview: inspect overall MemOS health. -global_res = client.get_all_scheduler_status() -if global_res: - print(f"System summary: {global_res.data['scheduler_summary']}") +resp = requests.get(f"{base_url}/product/scheduler/allstatus", timeout=10) +resp.raise_for_status() +global_res = resp.json() +print(f"System summary: {global_res['data']['scheduler_summary']}") # 2. Queue metrics: inspect backlog for a specific user. -queue_res = client.get_task_queue_status(user_id="dev_user_01") -if queue_res: - print(f"Remaining tasks: {queue_res.data['remaining_tasks_count']}") - print(f"Pending tasks: {queue_res.data['pending_tasks_count']}") +resp = requests.get( + f"{base_url}/product/scheduler/task_queue_status", + params={"user_id": "dev_user_01"}, + timeout=10, +) +resp.raise_for_status() +queue_res = resp.json() +print(f"Remaining tasks: {queue_res['data']['remaining_tasks_count']}") +print(f"Pending tasks: {queue_res['data']['pending_tasks_count']}") # 3. Task progress: poll a specific task until it finishes. task_id = "task_888999" +active_states = {"waiting", "pending", "in_progress"} while True: - res = client.get_task_status(user_id="dev_user_01", task_id=task_id) - if res and res.code == 200: - current_status = res.data[0]['status'] # data is a status list - print(f"Task {task_id} status: {current_status}") - - if current_status in ['completed', 'failed', 'cancelled']: - break + resp = requests.get( + f"{base_url}/product/scheduler/status", + params={"user_id": "dev_user_01", "task_id": task_id}, + timeout=10, + ) + resp.raise_for_status() + items = resp.json().get("data", []) # data is a status list: [{"task_id": ..., "status": ...}] + statuses = {item["status"] for item in items} + print(f"Task {task_id} status: {statuses or 'empty'}") + + if not statuses or statuses.isdisjoint(active_states): + break time.sleep(2) ``` diff --git a/docs/en/open_source/open_source_api/scheduler/wait.md b/docs/en/open_source/open_source_api/scheduler/wait.md index 9de0ff4be..6f356341a 100644 --- a/docs/en/open_source/open_source_api/scheduler/wait.md +++ b/docs/en/open_source/open_source_api/scheduler/wait.md @@ -41,36 +41,47 @@ Both endpoints share the following query parameters: ## 4. Quick Start -Use the open-source SDK for a blocking wait: +These endpoints are served directly by the open-source Server (`server_api`, router prefix `/product`) and can be called with plain HTTP requests. Note that `user_name`, `timeout_seconds`, and `poll_interval` are query parameters, not a request body. The example below performs a blocking wait: ```python -from memos.api.client import MemOSClient +import json -client = MemOSClient(api_key="...", base_url="...") +import requests + +# Address of your self-hosted MemOS Server (add an Authorization header if auth is enabled) +base_url = "http://localhost:8000" user_name = "dev_user_01" # Scenario A: blocking wait, commonly used in Python automation scripts. print(f"Waiting for user {user_name}'s task queue to drain...") -res = client.wait_until_idle( - user_name=user_name, - timeout_seconds=300, - poll_interval=2 +resp = requests.post( + f"{base_url}/product/scheduler/wait", + params={"user_name": user_name, "timeout_seconds": 300, "poll_interval": 2}, + timeout=310, # HTTP timeout should be larger than timeout_seconds ) -if res and res.code == 200: +resp.raise_for_status() +result = resp.json() # {"message": "idle" | "timeout", "data": {...}} +if result["message"] == "idle": print("All tasks have completed.") +else: + print(f"Timed out with {result['data']['running_tasks']} task(s) still running.") # Scenario B: streaming progress, commonly used by frontend progress bars. print("Listening to the live task progress stream...") -# The SSE endpoint usually returns a generator from the SDK. -progress_stream = client.stream_scheduler_progress( - user_name=user_name, - timeout_seconds=300 -) - -for event in progress_stream: - # Print the remaining queued tasks in real time. - print(f"Remaining queued tasks: {event['remaining_tasks_count']}") - if event['status'] == 'idle': - print("Scheduler is idle") - break +with requests.get( + f"{base_url}/product/scheduler/wait/stream", + params={"user_name": user_name, "timeout_seconds": 300}, + stream=True, + timeout=310, +) as resp: + resp.raise_for_status() + for line in resp.iter_lines(decode_unicode=True): + if not line or not line.startswith("data:"): + continue + event = json.loads(line.removeprefix("data:").strip()) + # Print the number of active tasks in real time. + print(f"Active tasks: {event['active_tasks']}, status: {event['status']}") + if event["status"] in ("idle", "timeout"): + print("Scheduler is idle" if event["status"] == "idle" else "Stream timed out") + break ``` diff --git a/pyproject.toml b/pyproject.toml index c7297c7d1..00d240ba8 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,7 +4,7 @@ ############################################################################## name = "MemoryOS" -version = "2.0.23" +version = "2.0.24" description = "Intelligence Begins with Memory" license = {text = "Apache-2.0"} readme = "README.md" diff --git a/src/memos/api/client.py b/src/memos/api/client.py index 818ce5e0d..b8055b5b6 100644 --- a/src/memos/api/client.py +++ b/src/memos/api/client.py @@ -2,7 +2,9 @@ import mimetypes import os +from collections.abc import Iterator from typing import Any +from urllib.parse import quote import requests @@ -65,17 +67,43 @@ def _validate_required_params(self, **params): if not param_value: raise ValueError(f"{param_name} is required") + def _validate_profile_subject(self, user_id: str | None, agent_id: str | None) -> None: + if bool(user_id) == bool(agent_id): + raise ValueError("exactly one of user_id or agent_id is required") + + def _post_json_dict( + self, endpoint: str, payload: dict[str, Any], operation: str + ) -> dict[str, Any] | None: + url = f"{self.base_url}/{endpoint}" + for retry in range(MAX_RETRY_COUNT): + try: + response = requests.post( + url, data=json.dumps(payload), headers=self.headers, timeout=30 + ) + response.raise_for_status() + return response.json() + except Exception as e: + logger.error( + "Failed to %s (retry %s/%s): %s", + operation, + retry + 1, + MAX_RETRY_COUNT, + e, + ) + if retry == MAX_RETRY_COUNT - 1: + raise + def get_message( self, user_id: str, conversation_id: str | None = None, - conversation_limit_number: int = 6, - message_limit_number: int = 6, + conversation_limit_number: int | None = None, + message_limit_number: int | None = None, source: str | None = None, ) -> MemOSGetMessagesResponse | None: """Get message""" # Validate required parameters - self._validate_required_params(user_id=user_id) + self._validate_required_params(user_id=user_id, conversation_id=conversation_id) url = f"{self.base_url}/get/message" payload = { @@ -102,22 +130,23 @@ def get_message( def add_message( self, messages: list[dict[str, Any]], - user_id: str, - conversation_id: str, + user_id: str | list[str] | None = None, + conversation_id: str | None = None, info: dict[str, Any] | None = None, source: str | None = None, app_id: str | None = None, - agent_id: str | None = None, + agent_id: str | list[str] | None = None, async_mode: bool = True, tags: list[str] | None = None, allow_public: bool = False, allow_knowledgebase_ids: list[str] | None = None, + allow_memory_view: list[str] | None = None, ) -> MemOSAddResponse | None: """Add message""" # Validate required parameters - self._validate_required_params( - messages=messages, user_id=user_id, conversation_id=conversation_id - ) + self._validate_required_params(messages=messages) + if not user_id and not agent_id: + raise ValueError("user_id or agent_id is required") url = f"{self.base_url}/add/message" payload = { @@ -130,8 +159,9 @@ def add_message( "agent_id": agent_id, "allow_public": allow_public, "allow_knowledgebase_ids": allow_knowledgebase_ids, + "allow_memory_view": allow_memory_view, "tags": tags, - "asyncMode": async_mode, + "async_mode": async_mode, } for retry in range(MAX_RETRY_COUNT): try: @@ -150,8 +180,9 @@ def add_message( def search_memory( self, query: str, - user_id: str, - conversation_id: str, + user_id: str | None = None, + conversation_id: str | None = None, + agent_id: str | None = None, memory_limit_number: int = 6, include_preference: bool = True, knowledgebase_ids: list[str] | None = None, @@ -160,22 +191,34 @@ def search_memory( include_tool_memory: bool = False, preference_limit_number: int = 6, tool_memory_limit_number: int = 6, + relativity: float | None = None, + include_skill: bool = False, + skill_limit_number: int = 6, + include_memory_view: list[str] | None = None, + context_format: str = "memory", ) -> MemOSSearchResponse | None: """Search memories""" # Validate required parameters - self._validate_required_params(query=query, user_id=user_id) + self._validate_required_params(query=query) + self._validate_profile_subject(user_id, agent_id) url = f"{self.base_url}/search/memory" payload = { "query": query, "user_id": user_id, "conversation_id": conversation_id, + "agent_id": agent_id, "memory_limit_number": memory_limit_number, "include_preference": include_preference, "knowledgebase_ids": knowledgebase_ids, "filter": filter, "preference_limit_number": preference_limit_number, "tool_memory_limit_number": tool_memory_limit_number, + "relativity": relativity, + "include_skill": include_skill, + "skill_limit_number": skill_limit_number, + "include_memory_view": include_memory_view, + "context_format": context_format, "source": source, "include_tool_memory": include_tool_memory, } @@ -195,16 +238,30 @@ def search_memory( raise def get_memory( - self, user_id: str, include_preference: bool = True, page: int = 1, size: int = 10 + self, + user_id: str | None = None, + include_preference: bool = True, + page: int = 1, + size: int = 10, + agent_id: str | None = None, + include_tool_memory: bool = True, + include_memory_view: list[str] | None = None, + filter: dict[str, Any] | None = None, ) -> MemOSGetMemoryResponse | None: """get memories""" # Validate required parameters - self._validate_required_params(include_preference=include_preference, user_id=user_id) + self._validate_profile_subject(user_id, agent_id) + if size > 50: + raise ValueError("size must be less than or equal to 50") url = f"{self.base_url}/get/memory" payload = { "include_preference": include_preference, "user_id": user_id, + "agent_id": agent_id, + "include_tool_memory": include_tool_memory, + "include_memory_view": include_memory_view, + "filter": filter, "page": page, "size": size, } @@ -223,17 +280,47 @@ def get_memory( if retry == MAX_RETRY_COUNT - 1: raise + @staticmethod + def _iter_sse_data(response: requests.Response) -> Iterator[str]: + """Yield decoded data payloads from a Server-Sent Events response.""" + try: + for line in response.iter_lines(decode_unicode=True): + if isinstance(line, bytes): + line = line.decode("utf-8") + if not line or not line.startswith("data:"): + continue + yield line.removeprefix("data:").lstrip() + finally: + response.close() + + def get_memory_by_id(self, memid: str) -> dict[str, Any] | None: + """Get one memory detail by its memory ID.""" + self._validate_required_params(memid=memid) + + url = f"{self.base_url}/get/memory/{quote(memid, safe='')}" + for retry in range(MAX_RETRY_COUNT): + try: + response = requests.get(url, headers=self.headers, timeout=30) + response.raise_for_status() + return response.json() + except Exception as e: + logger.error( + "Failed to get memory by ID (retry %s/%s): %s", + retry + 1, + MAX_RETRY_COUNT, + e, + ) + if retry == MAX_RETRY_COUNT - 1: + raise + def create_knowledgebase( - self, knowledgebase_name: str, knowledgebase_description: str + self, knowledgebase_name: str, knowledgebase_description: str | None = None ) -> MemOSCreateKnowledgebaseResponse | None: """ Create knowledgebase """ # Validate required parameters - self._validate_required_params( - knowledgebase_name=knowledgebase_name, - knowledgebase_description=knowledgebase_description, - ) + self._validate_required_params(knowledgebase_name=knowledgebase_name) url = f"{self.base_url}/create/knowledgebase" payload = { @@ -313,7 +400,7 @@ def add_knowledgebase_file_json( raise def add_knowledgebase_file_form( - self, knowledgebase_id: str, files: list[str] + self, knowledgebase_id: str, files: list[str], type: str | None = None ) -> MemOSAddKnowledgebaseFileResponse | None: """ add knowledgebase-file from form @@ -321,12 +408,12 @@ def add_knowledgebase_file_form( # Validate required parameters self._validate_required_params(knowledgebase_id=knowledgebase_id, files=files) - def build_file_form_param(file_path): + def build_file_form_param(file_path: str): """ form-Automatically generate the structure required for the `files` parameter in requests based on the local file path """ if not os.path.isfile(file_path): - logger.warning(f"File {file_path} does not exist") + logger.warning("File %s does not exist", file_path) return None filename = os.path.basename(file_path) @@ -335,31 +422,47 @@ def build_file_form_param(file_path): mime_type = "application/octet-stream" return ("file", (filename, open(file_path, "rb"), mime_type)) + def build_file_form_params() -> list: + file_params = [ + file_param + for file_path in files + if (file_param := build_file_form_param(file_path)) is not None + ] + if not file_params: + raise ValueError("files must contain at least one valid file path") + return file_params + url = f"{self.base_url}/add/knowledgebase-file" payload = { "knowledgebase_id": knowledgebase_id, } + if type is not None: + payload["type"] = type headers = { "Authorization": f"Token {self.api_key}", } for retry in range(MAX_RETRY_COUNT): + file_params = [] try: + file_params = build_file_form_params() response = requests.post( url, params=payload, headers=headers, timeout=30, - files=[build_file_form_param(file_path) for file_path in files], + files=file_params, ) response.raise_for_status() response_data = response.json() - print(response_data) return MemOSAddKnowledgebaseFileResponse(**response_data) except Exception as e: logger.error(f"Failed to add knowledgebase-file form (retry {retry + 1}/3): {e}") if retry == MAX_RETRY_COUNT - 1: raise + finally: + for file_param in file_params: + file_param[1][1].close() def delete_knowledgebase_file( self, file_ids: list[str] @@ -390,17 +493,27 @@ def delete_knowledgebase_file( raise def get_knowledgebase_file( - self, file_ids: list[str] + self, + file_ids: list[str] | None = None, + knowledgebase_id: str | None = None, + type: str | None = None, + page: int | None = None, + page_size: int | None = None, ) -> MemOSGetKnowledgebaseFileResponse | None: """ get knowledgebase-file """ # Validate required parameters - self._validate_required_params(file_ids=file_ids) + if bool(file_ids) == bool(knowledgebase_id): + raise ValueError("exactly one of file_ids or knowledgebase_id is required") url = f"{self.base_url}/get/knowledgebase-file" payload = { "file_ids": file_ids, + "knowledgebase_id": knowledgebase_id, + "type": type, + "page": page, + "page_size": page_size, } for retry in range(MAX_RETRY_COUNT): @@ -446,8 +559,8 @@ def get_task_status(self, task_id: str) -> MemOSGetTaskStatusResponse | None: def add_feedback( self, user_id: str, - conversation_id: str, - feedback_content: str, + conversation_id: str | None = None, + feedback_content: str | None = None, agent_id: str | None = None, app_id: str | None = None, feedback_time: str | None = None, @@ -456,9 +569,7 @@ def add_feedback( ) -> MemOSAddFeedBackResponse | None: """Add feedback""" # Validate required parameters - self._validate_required_params( - feedback_content=feedback_content, user_id=user_id, conversation_id=conversation_id - ) + self._validate_required_params(feedback_content=feedback_content, user_id=user_id) url = f"{self.base_url}/add/feedback" payload = { @@ -486,17 +597,43 @@ def add_feedback( raise def delete_memory( - self, user_ids: list[str], memory_ids: list[str] + self, + user_ids: list[str] | None = None, + memory_ids: list[str] | None = None, + *, + user_id: str | None = None, + agent_id: str | None = None, + filter: dict[str, Any] | None = None, + memory_type: str | None = None, ) -> MemOSDeleteMemoryResponse | None: """delete_memory memories""" - # Validate required parameters - self._validate_required_params(user_ids=user_ids, memory_ids=memory_ids) + if user_id is None and user_ids: + if len(user_ids) != 1 and not memory_ids: + raise ValueError("current API supports a single user_id, not multiple user_ids") + if not memory_ids: + user_id = user_ids[0] + + delete_modes = [ + bool(memory_ids), + bool(user_id), + bool(agent_id), + filter is not None, + ] + if sum(delete_modes) != 1: + raise ValueError("exactly one delete condition is required") url = f"{self.base_url}/delete/memory" - payload = { - "user_ids": user_ids, - "memory_ids": memory_ids, - } + payload: dict[str, Any] = {} + if memory_ids: + payload["memory_ids"] = memory_ids + if user_id: + payload["user_id"] = user_id + if agent_id: + payload["agent_id"] = agent_id + if filter is not None: + payload["filter"] = filter + if memory_type is not None: + payload["memory_type"] = memory_type for retry in range(MAX_RETRY_COUNT): try: @@ -512,6 +649,111 @@ def delete_memory( if retry == MAX_RETRY_COUNT - 1: raise + def update_memory( + self, + memory_id: str, + content: str | None = None, + title: str | None = None, + status: str | None = None, + ) -> dict[str, Any] | None: + """Update an existing memory.""" + self._validate_required_params(memory_id=memory_id) + if not content and not title and not status: + raise ValueError("content, title or status is required") + + payload = { + "memory_id": memory_id, + "content": content, + "title": title, + "status": status, + } + return self._post_json_dict("update/memory", payload, "update memory") + + def extract_memory( + self, + messages: list[dict[str, Any]], + extraction_types: list[str] | None = None, + model: str | None = None, + ) -> dict[str, Any] | None: + """Extract memory candidates from conversation messages.""" + self._validate_required_params(messages=messages) + + payload = { + "messages": messages, + "extraction_types": extraction_types, + "model": model, + } + return self._post_json_dict("extract/memory", payload, "extract memory") + + def rerank( + self, + query: str, + documents: list[str], + model: str | None = None, + top_n: int | None = None, + ) -> dict[str, Any] | None: + """Rerank documents for a query.""" + self._validate_required_params(query=query, documents=documents) + if top_n is not None and top_n <= 0: + raise ValueError("top_n must be greater than 0") + + payload = { + "query": query, + "documents": documents, + "model": model, + "top_n": top_n, + } + return self._post_json_dict("rerank", payload, "rerank documents") + + def bind_profile_template(self, bind_list: list[dict[str, Any]]) -> dict[str, Any] | None: + """Bind profile templates to user or agent subjects.""" + self._validate_required_params(bind_list=bind_list) + + payload = { + "bind_list": bind_list, + } + return self._post_json_dict("bind/profile_template", payload, "bind profile template") + + def edit_profile( + self, + profile_template_id: str, + user_id: str | None = None, + agent_id: str | None = None, + metadata: dict[str, Any] | None = None, + remove_fields: list[str] | None = None, + ) -> dict[str, Any] | None: + """Edit a profile instance.""" + self._validate_required_params(profile_template_id=profile_template_id) + self._validate_profile_subject(user_id, agent_id) + if metadata is None and not remove_fields: + raise ValueError("metadata or remove_fields is required") + + payload = { + "user_id": user_id, + "agent_id": agent_id, + "profile_template_id": profile_template_id, + "metadata": metadata, + "remove_fields": remove_fields, + } + return self._post_json_dict("edit/profile", payload, "edit profile") + + def delete_profile( + self, + profile_template_id: str, + user_id: str | None = None, + agent_id: str | None = None, + ) -> dict[str, Any] | None: + """Delete a profile instance.""" + self._validate_required_params(profile_template_id=profile_template_id) + self._validate_profile_subject(user_id, agent_id) + + payload = { + "user_id": user_id, + "agent_id": agent_id, + "profile_template_id": profile_template_id, + } + return self._post_json_dict("delete/profile", payload, "delete profile") + def chat( self, user_id: str, @@ -524,21 +766,26 @@ def chat( system_prompt: str | None = None, model_name: str | None = None, knowledgebase_ids: list[str] | None = None, - filter: dict[str:Any] | None = None, - add_message_on_answer: bool = False, + filter: dict[str, Any] | None = None, + add_message_on_answer: bool = True, app_id: str | None = None, agent_id: str | None = None, async_mode: bool = True, tags: list[str] | None = None, - info: dict[str:Any] | None = None, + info: dict[str, Any] | None = None, allow_public: bool = False, + allow_knowledgebase_ids: list[str] | None = None, max_tokens: int = 8192, - temperature: float | None = None, - top_p: float | None = None, + temperature: float | None = 0.7, + top_p: float | None = 0.95, include_preference: bool = True, preference_limit_number: int = 6, memory_limit_number: int = 6, - ) -> MemOSChatResponse | None: + stream: bool = False, + include_tool_memory: bool = False, + tool_memory_limit_number: int = 6, + relativity: float | None = None, + ) -> MemOSChatResponse | Iterator[str] | None: """chat""" # Validate required parameters self._validate_required_params( @@ -565,20 +812,31 @@ def chat( "tags": tags, "info": info, "allow_public": allow_public, + "allow_knowledgebase_ids": allow_knowledgebase_ids, "max_tokens": max_tokens, "temperature": temperature, "top_p": top_p, "include_preference": include_preference, "preference_limit_number": preference_limit_number, "memory_limit_number": memory_limit_number, + "stream": stream, + "include_tool_memory": include_tool_memory, + "tool_memory_limit_number": tool_memory_limit_number, + "relativity": relativity, } for retry in range(MAX_RETRY_COUNT): try: response = requests.post( - url, data=json.dumps(payload), headers=self.headers, timeout=30 + url, + data=json.dumps(payload), + headers=self.headers, + timeout=30, + stream=stream, ) response.raise_for_status() + if stream: + return self._iter_sse_data(response) response_data = response.json() return MemOSChatResponse(**response_data) diff --git a/src/memos/api/lifecycle.py b/src/memos/api/lifecycle.py new file mode 100644 index 000000000..e05a2f777 --- /dev/null +++ b/src/memos/api/lifecycle.py @@ -0,0 +1,26 @@ +from collections.abc import Mapping +from typing import Any + +from memos.log import get_logger + + +logger = get_logger(__name__) + + +def shutdown_components(components: Mapping[str, Any] | None) -> None: + """Release long-lived API components before the logging system shuts down.""" + if not components: + return + + mem_scheduler = components.get("mem_scheduler") + if mem_scheduler is None: + return + + for method_name in ("stop", "rabbitmq_close"): + method = getattr(mem_scheduler, method_name, None) + if not callable(method): + continue + try: + method() + except Exception: + logger.exception("Failed to run mem_scheduler.%s during API shutdown", method_name) diff --git a/src/memos/api/product_models.py b/src/memos/api/product_models.py index 2db6d0f75..1769b0ee6 100644 --- a/src/memos/api/product_models.py +++ b/src/memos/api/product_models.py @@ -608,6 +608,21 @@ def _convert_deprecated_fields(self) -> "APISearchRequest": class APIADDRequest(BaseRequest): """Request model for creating memories.""" + # Model-level example so the interactive docs (/docs) show a copy-paste-ready + # payload. Without it, Swagger UI renders the leading `str` branch of the + # `messages` union as `"string"` (see issue #1505). This only affects the + # generated OpenAPI schema, not validation or runtime behaviour. + model_config = { + "json_schema_extra": { + "example": { + "user_id": "8736b16e-1d20-4163-980b-a5063c3facdc", + "writable_cube_ids": ["b32d0977-435d-4828-a86f-4f47f8b55bca"], + "messages": [{"role": "user", "content": "I am learning ggplot2 in R."}], + "async_mode": "async", + } + } + } + # ==== Basic identifiers ==== user_id: str = Field(None, description="User ID") session_id: str | None = Field( @@ -1055,6 +1070,15 @@ class SearchMemoryData(BaseModel): alias="tool_memory_detail_list", description="List of tool_memor details (usually None)", ) + skill_detail_list: list[MemoryDetail] | None = Field( + None, alias="skill_detail_list", description="List of skill memory details" + ) + profile_detail_list: list[MemoryDetail] | None = Field( + None, alias="profile_detail_list", description="List of profile memory details" + ) + event_detail_list: list[MemoryDetail] | None = Field( + None, alias="event_detail_list", description="List of event memory details" + ) preference_note: str = Field( None, alias="preference_note", description="String of preference_note" ) @@ -1066,6 +1090,9 @@ class GetKnowledgebaseFileData(BaseModel): file_detail_list: list[FileDetail] = Field( default_factory=list, alias="file_detail_list", description="List of files details" ) + total: int | None = Field(None, description="Total number of matching files") + page: int | None = Field(None, description="Current page number") + page_size: int | None = Field(None, alias="page_size", description="Page size") class GetMemoryData(BaseModel): @@ -1077,6 +1104,22 @@ class GetMemoryData(BaseModel): preference_detail_list: list[MessageDetail] | None = Field( None, alias="preference_detail_list", description="List of preference detail" ) + tool_memory_detail_list: list[MemoryDetail] | None = Field( + None, alias="tool_memory_detail_list", description="List of tool memory details" + ) + profile_detail_list: list[MemoryDetail] | None = Field( + None, alias="profile_detail_list", description="List of profile memory details" + ) + event_detail_list: list[MemoryDetail] | None = Field( + None, alias="event_detail_list", description="List of event memory details" + ) + skill_detail_list: list[MemoryDetail] | None = Field( + None, alias="skill_detail_list", description="List of skill memory details" + ) + total: int | None = Field(None, description="Total number of memories") + size: int | None = Field(None, description="Page size") + current: int | None = Field(None, description="Current page number") + pages: int | None = Field(None, description="Total number of pages") class AddMessageData(BaseModel): @@ -1105,6 +1148,16 @@ class GetTaskStatusMessageData(BaseModel): status: str = Field(..., description="Operation task status") +class GetTaskStatusData(BaseModel): + """Current OpenMem task status response data.""" + + task_id: str = Field(..., description="Task identifier") + status: str = Field(..., description="Operation task status") + memory_views: dict[str, Any] | None = Field( + None, alias="memory_views", description="Memory view changes produced by the task" + ) + + # ─── MemOS Response Models (Similar to OpenAI ChatCompletion) ────────────────── @@ -1188,12 +1241,12 @@ class MemOSGetTaskStatusResponse(BaseModel): code: int = Field(..., description="Response status code") message: str = Field(..., description="Response message") - data: list[GetTaskStatusMessageData] = Field(..., description="Task status data") + data: GetTaskStatusData = Field(..., description="Task status data") @property - def messages(self) -> list[GetTaskStatusMessageData]: - """Convenient access to task status messages.""" - return self.data + def messages(self) -> list[GetTaskStatusData]: + """Backward-compatible list access to task status data.""" + return [self.data] class MemOSCreateKnowledgebaseResponse(BaseModel): diff --git a/src/memos/api/server_api.py b/src/memos/api/server_api.py index a9afe554c..1c6d93b9f 100644 --- a/src/memos/api/server_api.py +++ b/src/memos/api/server_api.py @@ -1,14 +1,18 @@ import logging import os +from collections.abc import AsyncIterator +from contextlib import asynccontextmanager + from dotenv import load_dotenv from fastapi import FastAPI, HTTPException from fastapi.exceptions import RequestValidationError from starlette.staticfiles import StaticFiles from memos.api.exceptions import APIExceptionHandler +from memos.api.lifecycle import shutdown_components from memos.api.middleware.request_context import RequestContextMiddleware -from memos.api.routers.server_router import router as server_router +from memos.api.routers import server_router as server_router_module from memos.plugins.manager import plugin_manager @@ -25,17 +29,25 @@ os.getenv("MEMSCHEDULER_REDIS_STREAM_KEY_PREFIX"), ) + +@asynccontextmanager +async def lifespan(_app: FastAPI) -> AsyncIterator[None]: + yield + shutdown_components(server_router_module.components) + + app = FastAPI( title="MemOS Server REST APIs", description="A REST API for managing multiple users with MemOS Server.", version="1.0.1", + lifespan=lifespan, ) app.mount("/download", StaticFiles(directory=os.getenv("FILE_LOCAL_PATH")), name="static_mapping") app.add_middleware(RequestContextMiddleware, source="server_api") # Include routers -app.include_router(server_router) +app.include_router(server_router_module.router) @app.get("/health") diff --git a/src/memos/api/server_api_ext.py b/src/memos/api/server_api_ext.py index 8c457e362..7b5aafc46 100644 --- a/src/memos/api/server_api_ext.py +++ b/src/memos/api/server_api_ext.py @@ -18,6 +18,9 @@ import logging import os +from collections.abc import AsyncIterator +from contextlib import asynccontextmanager + from fastapi import FastAPI from fastapi.exceptions import RequestValidationError from fastapi.middleware.cors import CORSMiddleware @@ -26,11 +29,12 @@ from starlette.responses import Response # Import Krolik extensions +from memos.api.lifecycle import shutdown_components from memos.api.middleware.rate_limit import RateLimitMiddleware -from memos.api.routers.admin_router import router as admin_router # Import base routers from MemOS -from memos.api.routers.server_router import router as server_router +from memos.api.routers import server_router as server_router_module +from memos.api.routers.admin_router import router as admin_router # Try to import exception handlers (may vary between MemOS versions) @@ -58,11 +62,18 @@ async def dispatch(self, request: Request, call_next) -> Response: return response +@asynccontextmanager +async def lifespan(_app: FastAPI) -> AsyncIterator[None]: + yield + shutdown_components(server_router_module.components) + + # Create FastAPI app app = FastAPI( title="MemOS Server REST APIs (Krolik Extended)", description="MemOS API with authentication, rate limiting, and admin endpoints.", version="2.0.3-krolik", + lifespan=lifespan, ) # CORS configuration @@ -94,7 +105,7 @@ async def dispatch(self, request: Request, call_next) -> Response: logger.info("Rate limiting enabled") # Include routers -app.include_router(server_router) +app.include_router(server_router_module.router) app.include_router(admin_router) # Exception handlers diff --git a/src/memos/log.py b/src/memos/log.py index c0bb5bf31..c18bd2118 100644 --- a/src/memos/log.py +++ b/src/memos/log.py @@ -27,6 +27,8 @@ load_dotenv() selected_log_level = logging.DEBUG if settings.DEBUG else logging.WARNING +_LOGGING_CONFIG_LOCK = threading.RLock() +_LOGGING_CONFIGURED_PID: int | None = None def _setup_logfile() -> Path: @@ -224,12 +226,31 @@ def close(self): } +def _get_current_pid() -> int: + return os.getpid() + + +def configure_logging(force: bool = False) -> None: + """Configure process-local logging once. + + Re-running dictConfig replaces and closes existing handlers. Guarding it avoids races with + background threads that may be emitting log records while other modules import loggers. + """ + global _LOGGING_CONFIGURED_PID + + with _LOGGING_CONFIG_LOCK: + current_pid = _get_current_pid() + if force or current_pid != _LOGGING_CONFIGURED_PID: + dictConfig(LOGGING_CONFIG) + _LOGGING_CONFIGURED_PID = current_pid + + def get_logger(name: str | None = None) -> logging.Logger: """returns the project logger, scoped to a child name if provided Args: name: will define a child logger """ - dictConfig(LOGGING_CONFIG) + configure_logging() parent_logger = logging.getLogger("") if name: diff --git a/tests/api/test_client.py b/tests/api/test_client.py new file mode 100644 index 000000000..2e911e4f4 --- /dev/null +++ b/tests/api/test_client.py @@ -0,0 +1,686 @@ +import json +import sys +import types + +from pathlib import Path +from typing import Any + +import pytest + + +SRC_DIR = Path(__file__).resolve().parents[2] / "src" / "memos" + + +def _install_memos_package_stub() -> None: + if "memos" not in sys.modules: + memos_pkg = types.ModuleType("memos") + memos_pkg.__path__ = [str(SRC_DIR)] + sys.modules["memos"] = memos_pkg + + if "memos.api" not in sys.modules: + api_pkg = types.ModuleType("memos.api") + api_pkg.__path__ = [str(SRC_DIR / "api")] + sys.modules["memos.api"] = api_pkg + sys.modules["memos"].api = api_pkg + + +def _load_client_module() -> Any: + _install_memos_package_stub() + + import memos.api.client as client_module + + return client_module + + +class DummyResponse: + def __init__(self, payload: dict): + self.payload = payload + + def raise_for_status(self) -> None: + return None + + def json(self) -> dict: + return self.payload + + +class DummyStreamResponse: + def __init__(self, lines: list[str]): + self.lines = lines + self.closed = False + self.json_called = False + + def raise_for_status(self) -> None: + return None + + def json(self) -> dict: + self.json_called = True + raise AssertionError("streaming responses must not be parsed as JSON") + + def iter_lines(self, decode_unicode: bool = False): + assert decode_unicode is True + yield from self.lines + + def close(self) -> None: + self.closed = True + + +def _response_for(url: str) -> dict: + if url.endswith("/get/message"): + return {"code": 200, "message": "ok", "data": {"message_detail_list": []}} + if url.endswith("/add/message"): + return { + "code": 200, + "message": "ok", + "data": {"success": True, "task_id": "task-1", "status": "completed"}, + } + if url.endswith("/search/memory"): + return {"code": 200, "message": "ok", "data": {"memory_detail_list": []}} + if url.endswith("/get/memory"): + return {"code": 200, "message": "ok", "data": {"memory_detail_list": []}} + if url.endswith("/create/knowledgebase"): + return {"code": 200, "message": "ok", "data": {"id": "kb-1"}} + if url.endswith("/get/knowledgebase-file"): + return {"code": 200, "message": "ok", "data": {"file_detail_list": []}} + if url.endswith("/delete/memory"): + return {"code": 200, "message": "ok", "data": {"success": True}} + if url.endswith("/add/feedback"): + return { + "code": 200, + "message": "ok", + "data": {"success": True, "task_id": "task-1", "status": "running"}, + } + if url.endswith("/chat"): + return {"code": 200, "message": "ok", "data": {"response": "answer"}} + if url.endswith("/add/knowledgebase-file"): + return {"code": 200, "message": "ok", "data": []} + if url.endswith("/update/memory"): + return {"code": 200, "message": "ok", "data": {"success": True}} + if url.endswith("/extract/memory"): + return { + "code": 200, + "message": "ok", + "data": { + "success": True, + "memory_detail_list": [], + "preference_detail_list": [], + }, + } + if url.endswith("/rerank"): + return {"code": 200, "message": "ok", "data": {"id": "rerank-1", "results": []}} + if url.endswith("/bind/profile_template"): + return {"code": 200, "message": "ok", "data": {"success": True}} + if url.endswith("/edit/profile"): + return {"code": 200, "message": "ok", "data": {"success": True}} + if url.endswith("/delete/profile"): + return {"code": 200, "message": "ok", "data": {"success": True}} + raise AssertionError(f"Unexpected URL: {url}") + + +@pytest.fixture +def client_module() -> Any: + return _load_client_module() + + +@pytest.fixture +def posted_requests(monkeypatch, client_module): + calls: list[dict] = [] + + def fake_post(url: str, **kwargs): + calls.append({"url": url, **kwargs}) + return DummyResponse(_response_for(url)) + + monkeypatch.setattr(client_module.requests, "post", fake_post) + return calls + + +@pytest.fixture +def fetched_requests(monkeypatch, client_module): + calls: list[dict] = [] + + def fake_get(url: str, **kwargs): + calls.append({"url": url, **kwargs}) + return DummyResponse( + { + "code": 200, + "message": "ok", + "data": {"id": "memory-1", "memory_type": "LongTermMemory"}, + } + ) + + monkeypatch.setattr(client_module.requests, "get", fake_get) + return calls + + +@pytest.fixture +def client(client_module) -> Any: + return client_module.MemOSClient(api_key="test-key", base_url="https://example.test/openmem/v1") + + +def _json_payload(call: dict) -> dict: + return json.loads(call["data"]) + + +def test_add_message_uses_snake_case_async_mode_and_memory_view( + client: Any, posted_requests: list[dict] +) -> None: + client.add_message( + messages=[{"role": "user", "content": "hello"}], + user_id="user-1", + conversation_id="conversation-1", + async_mode=False, + allow_memory_view=["kb-1"], + ) + + payload = _json_payload(posted_requests[0]) + + assert payload["async_mode"] is False + assert "asyncMode" not in payload + assert payload["allow_memory_view"] == ["kb-1"] + + +def test_search_memory_sends_updated_existing_request_fields( + client: Any, posted_requests: list[dict] +) -> None: + client.search_memory( + query="hello", + user_id=None, + agent_id="agent-1", + relativity=0.2, + include_skill=True, + skill_limit_number=4, + include_memory_view=["kb-1"], + context_format="json", + ) + + payload = _json_payload(posted_requests[0]) + + assert payload["conversation_id"] is None + assert payload["agent_id"] == "agent-1" + assert payload["relativity"] == 0.2 + assert payload["include_skill"] is True + assert payload["skill_limit_number"] == 4 + assert payload["include_memory_view"] == ["kb-1"] + assert payload["context_format"] == "json" + + +def test_get_memory_can_scope_by_agent_and_include_updated_filters( + client: Any, posted_requests: list[dict] +) -> None: + memory_filter = {"and": [{"memory_type": "LongTermMemory"}]} + + client.get_memory( + user_id=None, + agent_id="agent-1", + include_tool_memory=False, + include_memory_view=["kb-1"], + filter=memory_filter, + page=2, + size=20, + ) + + payload = _json_payload(posted_requests[0]) + + assert payload["user_id"] is None + assert payload["agent_id"] == "agent-1" + assert payload["include_tool_memory"] is False + assert payload["include_memory_view"] == ["kb-1"] + assert payload["filter"] == memory_filter + assert payload["page"] == 2 + assert payload["size"] == 20 + + +def test_get_memory_rejects_multiple_subjects(client: Any) -> None: + with pytest.raises(ValueError, match="exactly one of user_id or agent_id"): + client.get_memory(user_id="user-1", agent_id="agent-1") + + +def test_get_knowledgebase_file_supports_listing_by_knowledgebase( + client: Any, posted_requests: list[dict] +) -> None: + client.get_knowledgebase_file( + knowledgebase_id="kb-1", + type="doc", + page=2, + page_size=50, + ) + + payload = _json_payload(posted_requests[0]) + + assert payload == { + "file_ids": None, + "knowledgebase_id": "kb-1", + "type": "doc", + "page": 2, + "page_size": 50, + } + + +def test_delete_memory_keeps_legacy_memory_id_call_but_sends_current_contract( + client: Any, posted_requests: list[dict] +) -> None: + client.delete_memory(user_ids=["legacy-user"], memory_ids=["memory-1"]) + + payload = _json_payload(posted_requests[0]) + + assert payload == {"memory_ids": ["memory-1"]} + + +def test_delete_memory_supports_quick_delete_by_user_id( + client: Any, posted_requests: list[dict] +) -> None: + client.delete_memory(user_id="user-1") + + payload = _json_payload(posted_requests[0]) + + assert payload == {"user_id": "user-1"} + + +def test_chat_sends_updated_existing_request_fields( + client: Any, posted_requests: list[dict] +) -> None: + client.chat( + user_id="user-1", + conversation_id="conversation-1", + query="hello", + stream=True, + allow_knowledgebase_ids=["kb-1"], + include_tool_memory=True, + tool_memory_limit_number=3, + relativity=0.1, + ) + + payload = _json_payload(posted_requests[0]) + + assert payload["stream"] is True + assert payload["allow_knowledgebase_ids"] == ["kb-1"] + assert payload["include_tool_memory"] is True + assert payload["tool_memory_limit_number"] == 3 + assert payload["relativity"] == 0.1 + assert payload["add_message_on_answer"] is True + + +def test_add_knowledgebase_file_form_sends_type_and_closes_files( + client: Any, posted_requests: list[dict], tmp_path +) -> None: + file_path = tmp_path / "note.txt" + file_path.write_text("hello", encoding="utf-8") + + client.add_knowledgebase_file_form( + knowledgebase_id="kb-1", + files=[str(file_path)], + type="doc", + ) + + call = posted_requests[0] + uploaded_file = call["files"][0][1][1] + + assert call["params"] == {"knowledgebase_id": "kb-1", "type": "doc"} + assert uploaded_file.closed + + +def test_update_memory_sends_selected_fields(client: Any, posted_requests: list[dict]) -> None: + response = client.update_memory( + memory_id="memory-1", + content="new content", + title="new title", + status="activated", + ) + + payload = _json_payload(posted_requests[0]) + + assert posted_requests[0]["url"].endswith("/update/memory") + assert payload == { + "memory_id": "memory-1", + "content": "new content", + "title": "new title", + "status": "activated", + } + assert response["data"]["success"] is True + + +def test_update_memory_requires_a_change(client: Any) -> None: + with pytest.raises(ValueError, match="content, title or status is required"): + client.update_memory(memory_id="memory-1") + + +def test_extract_memory_sends_messages_and_options( + client: Any, posted_requests: list[dict] +) -> None: + messages = [{"role": "user", "content": "I like tea", "chat_time": "2026-07-06"}] + + client.extract_memory( + messages=messages, + extraction_types=["memory", "preference"], + model="extract-model", + ) + + payload = _json_payload(posted_requests[0]) + + assert posted_requests[0]["url"].endswith("/extract/memory") + assert payload == { + "messages": messages, + "extraction_types": ["memory", "preference"], + "model": "extract-model", + } + + +def test_rerank_sends_query_documents_and_options(client: Any, posted_requests: list[dict]) -> None: + client.rerank( + query="memory query", + documents=["doc a", "doc b"], + model="rerank-model", + top_n=1, + ) + + payload = _json_payload(posted_requests[0]) + + assert posted_requests[0]["url"].endswith("/rerank") + assert payload == { + "query": "memory query", + "documents": ["doc a", "doc b"], + "model": "rerank-model", + "top_n": 1, + } + + +def test_rerank_rejects_non_positive_top_n(client: Any) -> None: + with pytest.raises(ValueError, match="top_n must be greater than 0"): + client.rerank(query="memory query", documents=["doc a"], top_n=0) + + +def test_bind_profile_template_sends_bind_list(client: Any, posted_requests: list[dict]) -> None: + bind_list = [{"profile_template_id": "profile-template-1", "user_id": "user-1"}] + + client.bind_profile_template(bind_list=bind_list) + + payload = _json_payload(posted_requests[0]) + + assert posted_requests[0]["url"].endswith("/bind/profile_template") + assert payload == {"bind_list": bind_list} + + +def test_edit_profile_sends_metadata_and_remove_fields( + client: Any, posted_requests: list[dict] +) -> None: + metadata = {"basic": {"city": "Hangzhou"}} + + client.edit_profile( + profile_template_id="profile-template-1", + user_id="user-1", + metadata=metadata, + remove_fields=["basic.job"], + ) + + payload = _json_payload(posted_requests[0]) + + assert posted_requests[0]["url"].endswith("/edit/profile") + assert payload == { + "user_id": "user-1", + "agent_id": None, + "profile_template_id": "profile-template-1", + "metadata": metadata, + "remove_fields": ["basic.job"], + } + + +def test_edit_profile_requires_metadata_or_remove_fields(client: Any) -> None: + with pytest.raises(ValueError, match="metadata or remove_fields is required"): + client.edit_profile(profile_template_id="profile-template-1", user_id="user-1") + + +def test_delete_profile_sends_profile_template_and_subject( + client: Any, posted_requests: list[dict] +) -> None: + client.delete_profile(profile_template_id="profile-template-1", agent_id="agent-1") + + payload = _json_payload(posted_requests[0]) + + assert posted_requests[0]["url"].endswith("/delete/profile") + assert payload == { + "user_id": None, + "agent_id": "agent-1", + "profile_template_id": "profile-template-1", + } + + +def test_profile_subject_requires_exactly_one_user_or_agent(client: Any) -> None: + with pytest.raises(ValueError, match="exactly one of user_id or agent_id is required"): + client.delete_profile(profile_template_id="profile-template-1") + + with pytest.raises(ValueError, match="exactly one of user_id or agent_id is required"): + client.delete_profile( + profile_template_id="profile-template-1", + user_id="user-1", + agent_id="agent-1", + ) + + +def test_task_status_response_parses_current_object_shape(client_module: Any) -> None: + response = client_module.MemOSGetTaskStatusResponse( + code=200, + message="ok", + data={ + "task_id": "task-1", + "status": "running", + "memory_views": {"added": 1}, + }, + ) + + assert response.data.task_id == "task-1" + assert response.data.status == "running" + assert response.data.memory_views == {"added": 1} + + +def test_search_response_keeps_all_current_memory_view_lists(client_module: Any) -> None: + response = client_module.MemOSSearchResponse( + code=200, + message="ok", + data={ + "memory_detail_list": [], + "skill_detail_list": [{"id": "skill-1"}], + "profile_detail_list": [{"id": "profile-1"}], + "event_detail_list": [{"id": "event-1"}], + }, + ) + + assert response.data.skill_detail_list[0].id == "skill-1" + assert response.data.profile_detail_list[0].id == "profile-1" + assert response.data.event_detail_list[0].id == "event-1" + + +def test_get_memory_response_keeps_views_and_pagination(client_module: Any) -> None: + response = client_module.MemOSGetMemoryResponse( + code=200, + message="ok", + data={ + "memory_detail_list": [], + "tool_memory_detail_list": [{"id": "tool-1"}], + "profile_detail_list": [{"id": "profile-1"}], + "event_detail_list": [{"id": "event-1"}], + "skill_detail_list": [{"id": "skill-1"}], + "total": 21, + "size": 10, + "current": 2, + "pages": 3, + }, + ) + + assert response.data.tool_memory_detail_list[0].id == "tool-1" + assert response.data.profile_detail_list[0].id == "profile-1" + assert response.data.event_detail_list[0].id == "event-1" + assert response.data.skill_detail_list[0].id == "skill-1" + assert response.data.total == 21 + assert response.data.size == 10 + assert response.data.current == 2 + assert response.data.pages == 3 + + +def test_get_knowledgebase_file_response_keeps_pagination(client_module: Any) -> None: + response = client_module.MemOSGetKnowledgebaseFileResponse( + code=200, + message="ok", + data={ + "file_detail_list": [], + "total": 8, + "page": 2, + "page_size": 5, + }, + ) + + assert response.data.total == 8 + assert response.data.page == 2 + assert response.data.page_size == 5 + + +def test_get_message_requires_conversation_id(client: Any, posted_requests: list[dict]) -> None: + with pytest.raises(ValueError, match="conversation_id is required"): + client.get_message(user_id="user-1") + + assert posted_requests == [] + + +def test_get_message_uses_playground_default_limits( + client: Any, posted_requests: list[dict] +) -> None: + client.get_message(user_id="user-1", conversation_id="conversation-1") + + payload = _json_payload(posted_requests[0]) + + assert payload["conversation_limit_number"] is None + assert payload["message_limit_number"] is None + + +def test_add_message_allows_agent_only_and_generated_conversation( + client: Any, posted_requests: list[dict] +) -> None: + client.add_message( + messages=[{"role": "user", "content": "hello"}], + user_id=None, + agent_id="agent-1", + conversation_id=None, + ) + + payload = _json_payload(posted_requests[0]) + + assert payload["user_id"] is None + assert payload["agent_id"] == "agent-1" + assert payload["conversation_id"] is None + + +def test_search_memory_allows_agent_only(client: Any, posted_requests: list[dict]) -> None: + client.search_memory(query="hello", user_id=None, agent_id="agent-1") + + payload = _json_payload(posted_requests[0]) + + assert payload["user_id"] is None + assert payload["agent_id"] == "agent-1" + + +def test_search_memory_rejects_multiple_subjects(client: Any, posted_requests: list[dict]) -> None: + with pytest.raises(ValueError, match="exactly one of user_id or agent_id"): + client.search_memory(query="hello", user_id="user-1", agent_id="agent-1") + + assert posted_requests == [] + + +def test_create_knowledgebase_allows_empty_description( + client: Any, posted_requests: list[dict] +) -> None: + client.create_knowledgebase(knowledgebase_name="Knowledge Base") + + payload = _json_payload(posted_requests[0]) + + assert payload == { + "knowledgebase_name": "Knowledge Base", + "knowledgebase_description": None, + } + + +def test_add_feedback_allows_generated_conversation( + client: Any, posted_requests: list[dict] +) -> None: + client.add_feedback(user_id="user-1", feedback_content="helpful") + + payload = _json_payload(posted_requests[0]) + + assert payload["conversation_id"] is None + assert payload["feedback_content"] == "helpful" + + +def test_chat_uses_playground_sampling_defaults(client: Any, posted_requests: list[dict]) -> None: + client.chat(user_id="user-1", conversation_id="conversation-1", query="hello") + + payload = _json_payload(posted_requests[0]) + + assert payload["temperature"] == 0.7 + assert payload["top_p"] == 0.95 + + +def test_get_memory_rejects_size_above_playground_limit( + client: Any, posted_requests: list[dict] +) -> None: + with pytest.raises(ValueError, match="size must be less than or equal to 50"): + client.get_memory(user_id="user-1", size=51) + + assert posted_requests == [] + + +def test_get_memory_by_id_uses_detail_get_endpoint( + client: Any, fetched_requests: list[dict] +) -> None: + response = client.get_memory_by_id("memory-1") + + assert fetched_requests == [ + { + "url": "https://example.test/openmem/v1/get/memory/memory-1", + "headers": client.headers, + "timeout": 30, + } + ] + assert response == { + "code": 200, + "message": "ok", + "data": {"id": "memory-1", "memory_type": "LongTermMemory"}, + } + + +def test_get_memory_by_id_requires_memid(client: Any, fetched_requests: list[dict]) -> None: + with pytest.raises(ValueError, match="memid is required"): + client.get_memory_by_id("") + + assert fetched_requests == [] + + +def test_chat_stream_yields_sse_data_and_closes_response(monkeypatch, client_module: Any) -> None: + calls: list[dict] = [] + stream_response = DummyStreamResponse( + [ + "event: message", + 'data: {"response":"first"}', + "", + "data: [DONE]", + ] + ) + + def fake_post(url: str, **kwargs): + calls.append({"url": url, **kwargs}) + return stream_response + + monkeypatch.setattr(client_module.requests, "post", fake_post) + client = client_module.MemOSClient( + api_key="test-key", base_url="https://example.test/openmem/v1" + ) + + chunks = list( + client.chat( + user_id="user-1", + conversation_id="conversation-1", + query="hello", + stream=True, + ) + ) + + assert calls[0]["stream"] is True + assert chunks == ['{"response":"first"}', "[DONE]"] + assert stream_response.json_called is False + assert stream_response.closed is True diff --git a/tests/api/test_lifecycle.py b/tests/api/test_lifecycle.py new file mode 100644 index 000000000..b2f7f7064 --- /dev/null +++ b/tests/api/test_lifecycle.py @@ -0,0 +1,38 @@ +from memos.api.lifecycle import shutdown_components + + +class SchedulerStub: + def __init__(self, fail_stop: bool = False): + self.fail_stop = fail_stop + self.calls = [] + + def stop(self): + self.calls.append("stop") + if self.fail_stop: + raise RuntimeError("stop failed") + + def rabbitmq_close(self): + self.calls.append("rabbitmq_close") + + +def test_shutdown_components_stops_scheduler_before_rabbitmq_close(): + scheduler = SchedulerStub() + + shutdown_components({"mem_scheduler": scheduler}) + + assert scheduler.calls == ["stop", "rabbitmq_close"] + + +def test_shutdown_components_still_closes_rabbitmq_when_stop_fails(): + scheduler = SchedulerStub(fail_stop=True) + + shutdown_components({"mem_scheduler": scheduler}) + + assert scheduler.calls == ["stop", "rabbitmq_close"] + + +def test_shutdown_components_skips_missing_methods(): + class MinimalScheduler: + pass + + shutdown_components({"mem_scheduler": MinimalScheduler()}) diff --git a/tests/api/test_product_models.py b/tests/api/test_product_models.py new file mode 100644 index 000000000..177f8713d --- /dev/null +++ b/tests/api/test_product_models.py @@ -0,0 +1,41 @@ +"""Unit tests for API request-model OpenAPI schemas. + +These tests lock the OpenAPI schema behaviour of request models so the +interactive docs (``/docs``) stay consistent with the documented contract. + +Regression guard for issue #1505: the ``/product/add`` example must render +``messages`` as a structured message list instead of a bare ``"string"``. +Because ``messages`` is typed as ``str | MessageList | RawMessageList``, Swagger +UI would otherwise pick the leading ``str`` branch of the ``anyOf`` and show +``"messages": "string"``, which misleads users into sending plain text. +""" + +from memos.api.product_models import APIADDRequest + + +def test_add_request_exposes_model_level_example(): + """APIADDRequest must ship a model-level example for the interactive docs.""" + schema = APIADDRequest.model_json_schema() + + assert "example" in schema, "APIADDRequest should define a model-level example" + + +def test_add_request_example_messages_is_structured_list(): + """The example's ``messages`` must be a non-empty list of role/content items.""" + example = APIADDRequest.model_json_schema()["example"] + + messages = example.get("messages") + assert isinstance(messages, list), "messages example must be a list, not a bare string" + assert messages, "messages example should not be empty" + + first = messages[0] + assert first.get("role"), "each example message needs a role" + assert first.get("content"), "each example message needs content" + + +def test_add_request_example_covers_core_fields(): + """The example should be a copy-paste-ready payload for the core add flow.""" + example = APIADDRequest.model_json_schema()["example"] + + assert "user_id" in example + assert "writable_cube_ids" in example diff --git a/tests/test_log.py b/tests/test_log.py index fbd8791ee..5387adc34 100644 --- a/tests/test_log.py +++ b/tests/test_log.py @@ -28,3 +28,34 @@ def test_get_logger_returns_logger(): assert any(isinstance(h, logging.StreamHandler) for h in logger.parent.handlers) or any( isinstance(h, logging.FileHandler) for h in logger.parent.handlers ) + + +def test_get_logger_configures_logging_once_per_process(monkeypatch): + calls = [] + + monkeypatch.setattr(log, "_LOGGING_CONFIGURED_PID", None) + monkeypatch.setattr(log, "_get_current_pid", lambda: 123) + monkeypatch.setattr(log, "dictConfig", lambda config: calls.append(config)) + + log.get_logger("first") + log.get_logger("second") + + assert len(calls) == 1 + + +def test_get_logger_reconfigures_after_process_fork(monkeypatch): + calls = [] + pid = 123 + + def getpid(): + return pid + + monkeypatch.setattr(log, "_LOGGING_CONFIGURED_PID", None) + monkeypatch.setattr(log, "_get_current_pid", getpid) + monkeypatch.setattr(log, "dictConfig", lambda config: calls.append(config)) + + log.get_logger("parent") + pid = 456 + log.get_logger("child") + + assert len(calls) == 2