From 5b342f89f14aaa48d48830695ffb0468ad1aa1d8 Mon Sep 17 00:00:00 2001 From: dingliang <2650876010@qq.com> Date: Wed, 8 Jul 2026 16:13:46 +0800 Subject: [PATCH 1/7] feat(api): align OpenMem v1 SDK with cloud API --- src/memos/api/client.py | 263 ++++++++++++++++++++++--- tests/api/test_client.py | 406 +++++++++++++++++++++++++++++++++++++++ 2 files changed, 647 insertions(+), 22 deletions(-) create mode 100644 tests/api/test_client.py diff --git a/src/memos/api/client.py b/src/memos/api/client.py index 818ce5e0d..1ace6ec8f 100644 --- a/src/memos/api/client.py +++ b/src/memos/api/client.py @@ -65,6 +65,32 @@ 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, @@ -102,7 +128,7 @@ def get_message( def add_message( self, messages: list[dict[str, Any]], - user_id: str, + user_id: str | list[str], conversation_id: str, info: dict[str, Any] | None = None, source: str | None = None, @@ -112,6 +138,7 @@ def add_message( 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 @@ -130,8 +157,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: @@ -151,7 +179,8 @@ def search_memory( self, query: str, user_id: str, - conversation_id: str, + 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,6 +189,11 @@ 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 @@ -170,12 +204,18 @@ def search_memory( "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 +235,28 @@ 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) 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, } @@ -313,7 +365,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 +373,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 +387,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 +458,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): @@ -486,17 +564,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 +616,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,20 +733,25 @@ 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, include_preference: bool = True, preference_limit_number: int = 6, memory_limit_number: int = 6, + stream: bool = False, + include_tool_memory: bool = False, + tool_memory_limit_number: int = 6, + relativity: float | None = None, ) -> MemOSChatResponse | None: """chat""" # Validate required parameters @@ -565,12 +779,17 @@ 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): diff --git a/tests/api/test_client.py b/tests/api/test_client.py new file mode 100644 index 000000000..61724b0d4 --- /dev/null +++ b/tests/api/test_client.py @@ -0,0 +1,406 @@ +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 + + +def _response_for(url: str) -> dict: + 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("/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("/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 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="user-1", + 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", + ) From f46c02ed7b9199e812edc03cd3f4fd7dae459dc9 Mon Sep 17 00:00:00 2001 From: dingliang <2650876010@qq.com> Date: Tue, 14 Jul 2026 17:18:25 +0800 Subject: [PATCH 2/7] feat(api): align OpenMem v1 SDK with cloud API --- src/memos/api/client.py | 89 +++++++--- src/memos/api/product_models.py | 46 +++++- tests/api/test_client.py | 282 +++++++++++++++++++++++++++++++- 3 files changed, 387 insertions(+), 30 deletions(-) diff --git a/src/memos/api/client.py b/src/memos/api/client.py index 1ace6ec8f..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 @@ -95,13 +97,13 @@ 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 = { @@ -128,12 +130,12 @@ def get_message( def add_message( self, messages: list[dict[str, Any]], - user_id: str | list[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, @@ -142,9 +144,9 @@ def add_message( ) -> 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 = { @@ -178,7 +180,7 @@ def add_message( def search_memory( self, query: str, - user_id: str, + user_id: str | None = None, conversation_id: str | None = None, agent_id: str | None = None, memory_limit_number: int = 6, @@ -197,7 +199,8 @@ def search_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 = { @@ -248,6 +251,8 @@ def get_memory( """get memories""" # Validate required parameters 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 = { @@ -275,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 = { @@ -524,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, @@ -534,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 = { @@ -743,8 +776,8 @@ def chat( 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, @@ -752,7 +785,7 @@ def chat( include_tool_memory: bool = False, tool_memory_limit_number: int = 6, relativity: float | None = None, - ) -> MemOSChatResponse | None: + ) -> MemOSChatResponse | Iterator[str] | None: """chat""" # Validate required parameters self._validate_required_params( @@ -795,9 +828,15 @@ def chat( 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/product_models.py b/src/memos/api/product_models.py index ff9d94859..86f6fb80c 100644 --- a/src/memos/api/product_models.py +++ b/src/memos/api/product_models.py @@ -1048,6 +1048,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" ) @@ -1059,6 +1068,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): @@ -1070,6 +1082,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): @@ -1098,6 +1126,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) ────────────────── @@ -1181,12 +1219,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/tests/api/test_client.py b/tests/api/test_client.py index 61724b0d4..2e911e4f4 100644 --- a/tests/api/test_client.py +++ b/tests/api/test_client.py @@ -43,7 +43,30 @@ 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, @@ -54,10 +77,18 @@ def _response_for(url: str) -> dict: 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"): @@ -102,6 +133,24 @@ def fake_post(url: str, **kwargs): 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") @@ -134,7 +183,7 @@ def test_search_memory_sends_updated_existing_request_fields( ) -> None: client.search_memory( query="hello", - user_id="user-1", + user_id=None, agent_id="agent-1", relativity=0.2, include_skill=True, @@ -404,3 +453,234 @@ def test_profile_subject_requires_exactly_one_user_or_agent(client: Any) -> None 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 From 9d2c9b75d2f462e8e05f7750e2266d7eea0a7a48 Mon Sep 17 00:00:00 2001 From: de1ty <7804799+de1tydev@users.noreply.github.com> Date: Tue, 14 Jul 2026 22:27:08 +0800 Subject: [PATCH 3/7] fix(memos-local): include json hint in user messages (#1756) Co-authored-by: shinetata <149466187+shinetata@users.noreply.github.com> --- apps/memos-local-plugin/core/llm/client.ts | 19 +++++++++++++++++-- .../tests/unit/llm/client.test.ts | 8 ++++++-- 2 files changed, 23 insertions(+), 4 deletions(-) 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/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 () => { From a000a1c7341fdfbe3d13ac091b83e5aec4cf9abd Mon Sep 17 00:00:00 2001 From: shinetata <149466187+shinetata@users.noreply.github.com> Date: Thu, 16 Jul 2026 14:35:44 +0800 Subject: [PATCH 4/7] docs: fix scheduler API examples to match server endpoints (#2113) Rewrite scheduler Quick Start examples to call the real /product/scheduler REST endpoints with requests instead of non-existent MemOSClient methods; align response fields with the actual handler output; rename the mis-named ' wait.md' to 'wait.md'. Closes #2083 Co-authored-by: sunqi Co-authored-by: Cursor --- .../open_source_api/scheduler/get_status.md | 48 ++++++++++------- .../scheduler/{ wait.md => wait.md} | 51 +++++++++++-------- .../open_source_api/scheduler/get_status.md | 48 ++++++++++------- .../open_source_api/scheduler/wait.md | 51 +++++++++++-------- 4 files changed, 124 insertions(+), 74 deletions(-) rename docs/cn/open_source/open_source_api/scheduler/{ wait.md => wait.md} (61%) 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 ``` From 2c60267e9672cf9e0e7af53e8e348be8c2c1d99a Mon Sep 17 00:00:00 2001 From: shinetata <149466187+shinetata@users.noreply.github.com> Date: Thu, 16 Jul 2026 14:36:07 +0800 Subject: [PATCH 5/7] fix(api): show structured add example in /docs for /product/add (#2112) fix(api): show structured add example in /docs (messages not string) Swagger UI rendered APIADDRequest.messages as "string" because it picks the leading `str` branch of the `str | MessageList | RawMessageList` union when generating the example. Add a model-level json_schema_extra example so /docs shows a copy-paste-ready payload with a structured messages list. Schema-only change: field types, validation and runtime behaviour are unchanged, so the core add pipeline is unaffected. Closes #1505 Co-authored-by: sunqi Co-authored-by: Cursor --- src/memos/api/product_models.py | 15 ++++++++++++ tests/api/test_product_models.py | 41 ++++++++++++++++++++++++++++++++ 2 files changed, 56 insertions(+) create mode 100644 tests/api/test_product_models.py diff --git a/src/memos/api/product_models.py b/src/memos/api/product_models.py index a2c83f622..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( 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 From 74d8598dfee14c8b17bfa3c0a3d4748aa5528b5c Mon Sep 17 00:00:00 2001 From: MemOS AutoDev Date: Thu, 16 Jul 2026 15:59:40 +0800 Subject: [PATCH 6/7] fix(chunkers): make SimpleTextSplitter fallback URL-safe (#2115) `SimpleTextSplitter._simple_split_text()` called `self.protect_urls` and `self.restore_urls`, methods defined only on `BaseChunker`. Since `SimpleTextSplitter` does not inherit `BaseChunker`, every call raised `AttributeError: 'SimpleTextSplitter' object has no attribute 'protect_urls'`. The multi-modal file-parsing pipeline swallowed the error and fell back to returning the whole text as a single chunk, producing ~5.8k noisy log rows on ACK where langchain_text_splitters is missing and the fallback branch is actually exercised. Extract the URL protect/restore helpers into a small `URLProtectionMixin` in `chunkers/base.py`; have both `BaseChunker` and `SimpleTextSplitter` inherit it. This preserves BaseChunker's public API (mixin methods are inherited transparently), keeps SimpleTextSplitter's constructor and return type unchanged, and shares a single URL regex between the two paths. Add regression tests in tests/chunkers/test_simple_chunker.py covering short/long input, empty input, no-URL text, and parametrised (chunk_size, overlap) combinations to ensure the fallback never raises again. --- src/memos/chunkers/base.py | 33 +++++++---- src/memos/chunkers/simple_chunker.py | 16 ++++- tests/chunkers/test_simple_chunker.py | 84 +++++++++++++++++++++++++++ 3 files changed, 120 insertions(+), 13 deletions(-) create mode 100644 tests/chunkers/test_simple_chunker.py diff --git a/src/memos/chunkers/base.py b/src/memos/chunkers/base.py index e858132e1..d7c5f42a8 100644 --- a/src/memos/chunkers/base.py +++ b/src/memos/chunkers/base.py @@ -14,16 +14,16 @@ def __init__(self, text: str, token_count: int, sentences: list[str]): self.sentences = sentences -class BaseChunker(ABC): - """Base class for all text chunkers.""" +class URLProtectionMixin: + """Shared URL protect/restore helpers used across chunkers. - @abstractmethod - def __init__(self, config: BaseChunkerConfig): - """Initialize the chunker with the given configuration.""" + Extracted so that lightweight fallbacks such as + :class:`memos.chunkers.simple_chunker.SimpleTextSplitter` can reuse the + same URL-aware splitting logic as :class:`BaseChunker` without inheriting + the full chunker contract (see issue #2115). + """ - @abstractmethod - def chunk(self, text: str) -> list[Chunk]: - """Chunk the given text into smaller chunks.""" + _URL_PATTERN = r'https?://[^\s<>"{}|\\^`\[\]]+' def protect_urls(self, text: str) -> tuple[str, dict[str, str]]: """ @@ -35,8 +35,7 @@ def protect_urls(self, text: str) -> tuple[str, dict[str, str]]: Returns: tuple: (Text with URLs replaced by placeholders, URL mapping dictionary) """ - url_pattern = r'https?://[^\s<>"{}|\\^`\[\]]+' - url_map = {} + url_map: dict[str, str] = {} def replace_url(match): url = match.group(0) @@ -44,7 +43,7 @@ def replace_url(match): url_map[placeholder] = url return placeholder - protected_text = re.sub(url_pattern, replace_url, text) + protected_text = re.sub(self._URL_PATTERN, replace_url, text) return protected_text, url_map def restore_urls(self, text: str, url_map: dict[str, str]) -> str: @@ -63,3 +62,15 @@ def restore_urls(self, text: str, url_map: dict[str, str]) -> str: restored_text = restored_text.replace(placeholder, url) return restored_text + + +class BaseChunker(URLProtectionMixin, ABC): + """Base class for all text chunkers.""" + + @abstractmethod + def __init__(self, config: BaseChunkerConfig): + """Initialize the chunker with the given configuration.""" + + @abstractmethod + def chunk(self, text: str) -> list[Chunk]: + """Chunk the given text into smaller chunks.""" diff --git a/src/memos/chunkers/simple_chunker.py b/src/memos/chunkers/simple_chunker.py index 58e12e2f1..c4bccd667 100644 --- a/src/memos/chunkers/simple_chunker.py +++ b/src/memos/chunkers/simple_chunker.py @@ -1,5 +1,17 @@ -class SimpleTextSplitter: - """Simple text splitter wrapper.""" +from memos.chunkers.base import URLProtectionMixin + + +class SimpleTextSplitter(URLProtectionMixin): + """Simple text splitter wrapper. + + Fallback used by :mod:`memos.mem_reader.read_multi_modal.utils` when the + optional ``langchain_text_splitters``-backed chunkers (``CharacterTextChunker`` + / ``MarkdownChunker``) cannot be constructed at import time. + + Inherits URL protect/restore helpers from :class:`URLProtectionMixin` + (see issue #2115: without the mixin, ``chunk()`` raised ``AttributeError`` + on every call that reached the fallback path). + """ def __init__(self, chunk_size: int, chunk_overlap: int): self.chunk_size = chunk_size diff --git a/tests/chunkers/test_simple_chunker.py b/tests/chunkers/test_simple_chunker.py new file mode 100644 index 000000000..3b8c94c39 --- /dev/null +++ b/tests/chunkers/test_simple_chunker.py @@ -0,0 +1,84 @@ +"""Regression tests for `SimpleTextSplitter` fallback (issue #2115). + +The fallback is exercised in production when `langchain_text_splitters` is +missing (ACK image drift). Prior to the fix, `SimpleTextSplitter.chunk()` +raised `AttributeError: 'SimpleTextSplitter' object has no attribute +'protect_urls'` because `_simple_split_text` referenced `self.protect_urls` +/ `self.restore_urls`, which are only defined on `BaseChunker`. +""" + +import pytest + +from memos.chunkers.simple_chunker import SimpleTextSplitter + + +def test_simple_text_splitter_short_text_with_url_returns_single_chunk(): + """Short text below chunk_size should return one chunk with the URL intact.""" + splitter = SimpleTextSplitter(chunk_size=512, chunk_overlap=128) + text = "This is a test document with a URL: https://example.com/path/to/resource" + + chunks = splitter.chunk(text) + + assert chunks == [text] + + +def test_simple_text_splitter_long_text_preserves_url(): + """A URL must never be split across chunks — it either appears whole or not at all in a chunk. + + The critical property (issue #2115): even after fallback splitting, we + must never see a chunk that contains only part of a URL. Overlap MAY + cause the same URL to appear in more than one chunk; that is by design + for retrieval quality and is not what the issue asks us to change. + """ + url = "https://example.com/very/long/path/segment?query=one&other=two#fragment" + prefix = "A" * 400 + suffix = "B" * 400 + text = f"{prefix} {url} {suffix}" + + splitter = SimpleTextSplitter(chunk_size=200, chunk_overlap=50) + chunks = splitter.chunk(text) + + assert len(chunks) > 1, "text should be split into multiple chunks" + # The URL must appear whole at least once. + assert any(url in c for c in chunks), ( + f"URL was fully lost after splitting; chunks (first 5)={chunks[:5]}" + ) + # No chunk should contain the placeholder marker leftover. + for c in chunks: + assert "__URL_" not in c, f"unresolved URL placeholder leaked into chunk: {c!r}" + # No chunk should contain a *partial* URL — i.e., if the chunk mentions + # "https://" it must contain the URL in full. + for c in chunks: + if "https://" in c: + assert url in c, f"chunk contains a partial URL: {c!r}" + + +def test_simple_text_splitter_empty_input_returns_empty_list(): + splitter = SimpleTextSplitter(chunk_size=100, chunk_overlap=20) + assert splitter.chunk("") == [] + assert splitter.chunk(" \n \t ") == [] + + +def test_simple_text_splitter_no_url_still_chunks(): + splitter = SimpleTextSplitter(chunk_size=50, chunk_overlap=10) + text = "Hello world. " * 20 # > 50 chars, no URL + chunks = splitter.chunk(text) + assert len(chunks) >= 2 + # Reassembling should recover all non-whitespace content. + joined = "".join(chunks) + for word in ["Hello", "world"]: + assert word in joined + + +@pytest.mark.parametrize( + ("chunk_size", "overlap"), + [(100, 20), (256, 64), (1024, 128)], +) +def test_simple_text_splitter_various_sizes_do_not_raise(chunk_size, overlap): + """The fallback used to raise AttributeError for *any* input containing a URL.""" + splitter = SimpleTextSplitter(chunk_size=chunk_size, chunk_overlap=overlap) + text = "prefix " + ("word " * 200) + "https://example.com/x " + ("tail " * 200) + # Must not raise. + chunks = splitter.chunk(text) + assert isinstance(chunks, list) + assert all(isinstance(c, str) for c in chunks) From e2b5e7423ac4bce7ee462f242dcf56043d109e09 Mon Sep 17 00:00:00 2001 From: MemOS AutoDev Date: Thu, 16 Jul 2026 16:11:20 +0800 Subject: [PATCH 7/7] refactor(chunkers): expose URL placeholder prefix as class constant Address OCR review on #2116: the placeholder-leak assertion in tests/chunkers/test_simple_chunker.py hardcoded the string `'__URL_'`, which duplicates an implementation detail of `URLProtectionMixin.protect_urls` (formatted as `f'__URL_{len(url_map)}__'`). If the prefix ever changed in `base.py`, the assertion would silently keep passing while no longer catching real placeholder leaks. Expose the prefix as a class-level constant `URLProtectionMixin._URL_PLACEHOLDER_PREFIX = "__URL_"`, use it inside `protect_urls`, and import it in the test so the two paths stay in sync automatically. No behavior change: placeholders keep the same textual form (`__URL___`), so both the current base chunker and the fallback `SimpleTextSplitter` produce identical output to before. --- src/memos/chunkers/base.py | 8 +++++++- tests/chunkers/test_simple_chunker.py | 5 ++++- 2 files changed, 11 insertions(+), 2 deletions(-) diff --git a/src/memos/chunkers/base.py b/src/memos/chunkers/base.py index d7c5f42a8..b78fad128 100644 --- a/src/memos/chunkers/base.py +++ b/src/memos/chunkers/base.py @@ -24,6 +24,12 @@ class URLProtectionMixin: """ _URL_PATTERN = r'https?://[^\s<>"{}|\\^`\[\]]+' + # Prefix used for the placeholders emitted by :meth:`protect_urls`. Exposed + # as a class-level constant so tests (and any other consumer that needs to + # detect leaked placeholders) can reference it without hardcoding the + # literal — keeping the assertion in sync with the implementation if the + # placeholder format ever changes. + _URL_PLACEHOLDER_PREFIX = "__URL_" def protect_urls(self, text: str) -> tuple[str, dict[str, str]]: """ @@ -39,7 +45,7 @@ def protect_urls(self, text: str) -> tuple[str, dict[str, str]]: def replace_url(match): url = match.group(0) - placeholder = f"__URL_{len(url_map)}__" + placeholder = f"{self._URL_PLACEHOLDER_PREFIX}{len(url_map)}__" url_map[placeholder] = url return placeholder diff --git a/tests/chunkers/test_simple_chunker.py b/tests/chunkers/test_simple_chunker.py index 3b8c94c39..405995a7e 100644 --- a/tests/chunkers/test_simple_chunker.py +++ b/tests/chunkers/test_simple_chunker.py @@ -9,6 +9,7 @@ import pytest +from memos.chunkers.base import URLProtectionMixin from memos.chunkers.simple_chunker import SimpleTextSplitter @@ -45,7 +46,9 @@ def test_simple_text_splitter_long_text_preserves_url(): ) # No chunk should contain the placeholder marker leftover. for c in chunks: - assert "__URL_" not in c, f"unresolved URL placeholder leaked into chunk: {c!r}" + assert URLProtectionMixin._URL_PLACEHOLDER_PREFIX not in c, ( + f"unresolved URL placeholder leaked into chunk: {c!r}" + ) # No chunk should contain a *partial* URL — i.e., if the chunk mentions # "https://" it must contain the URL in full. for c in chunks: