From d3c7237c30f153129a0dc8e9e5187883646fe3d2 Mon Sep 17 00:00:00 2001 From: wangxinyu Date: Fri, 25 Sep 2026 22:18:19 +0800 Subject: [PATCH] fix: preserve reference time in LongMemEval search --- evaluation/scripts/utils/client.py | 17 ++-- .../tree_text_memory/retrieve/searcher.py | 4 +- .../retrieve/task_goal_parser.py | 6 ++ src/memos/multi_mem_cube/single_cube.py | 2 + src/memos/search/search_service.py | 1 + .../test_longmemeval_api_contract.py | 77 +++++++++++++++++++ tests/memories/textual/test_tree_searcher.py | 10 ++- .../textual/test_tree_task_goal_parser.py | 17 ++++ 8 files changed, 124 insertions(+), 10 deletions(-) create mode 100644 tests/evaluation/test_longmemeval_api_contract.py diff --git a/evaluation/scripts/utils/client.py b/evaluation/scripts/utils/client.py index 157c3f8ea..4576173cc 100644 --- a/evaluation/scripts/utils/client.py +++ b/evaluation/scripts/utils/client.py @@ -165,13 +165,13 @@ def add(self, messages, user_id, conv_id, batch_size: int = 9999): ) response = requests.request("POST", url, data=payload, headers=self.headers) assert response.status_code == 200, response.text - assert json.loads(response.text)["message"] == "Memory added successfully", ( - response.text - ) + assert ( + json.loads(response.text)["message"] == "Memory added successfully" + ), response.text added_memories += json.loads(response.text)["data"] return added_memories - def search(self, query, user_id, top_k): + def search(self, query, user_id, top_k, reference_time=None): """Search memories.""" url = f"{self.memos_url}/product/search" payload = json.dumps( @@ -184,14 +184,15 @@ def search(self, query, user_id, top_k): "mode": os.getenv("SEARCH_MODE", "fast"), "include_preference": True, "pref_top_k": 6, + "reference_time": reference_time, }, ensure_ascii=False, ) response = requests.request("POST", url, data=payload, headers=self.headers) assert response.status_code == 200, response.text - assert json.loads(response.text)["message"] == "Search completed successfully", ( - response.text - ) + assert ( + json.loads(response.text)["message"] == "Search completed successfully" + ), response.text return json.loads(response.text)["data"] @@ -225,7 +226,7 @@ def add(self, messages, user_id, conv_id=None, batch_size: int = 9999): else: raise e - def search(self, query, user_id, top_k): + def search(self, query, user_id, top_k, reference_time=None): """Search memories.""" url = f"{self.memos_url}/search/memory" payload = json.dumps( diff --git a/src/memos/memories/textual/tree_text_memory/retrieve/searcher.py b/src/memos/memories/textual/tree_text_memory/retrieve/searcher.py index 5c9cce78e..9a10e5cd6 100644 --- a/src/memos/memories/textual/tree_text_memory/retrieve/searcher.py +++ b/src/memos/memories/textual/tree_text_memory/retrieve/searcher.py @@ -335,13 +335,15 @@ def _parse_task( """ # parse goal using LLM + parser_kwargs = dict(kwargs) + parser_kwargs.setdefault("reference_time", info.get("reference_time")) parsed_goal = self.task_goal_parser.parse( task_description=query, context="\n".join(context), conversation=info.get("chat_history", []), mode=mode, use_fast_graph=self.use_fast_graph, - **kwargs, + **parser_kwargs, ) query = parsed_goal.rephrased_query or query diff --git a/src/memos/memories/textual/tree_text_memory/retrieve/task_goal_parser.py b/src/memos/memories/textual/tree_text_memory/retrieve/task_goal_parser.py index 3b160a56e..3a1b85d22 100644 --- a/src/memos/memories/textual/tree_text_memory/retrieve/task_goal_parser.py +++ b/src/memos/memories/textual/tree_text_memory/retrieve/task_goal_parser.py @@ -96,6 +96,12 @@ def _parse_fine( prompt = Template(TASK_PARSE_PROMPT).substitute( task=query.strip(), context=context, conversation=conversation_prompt ) + reference_time = kwargs.get("reference_time") + if reference_time: + prompt += ( + "\nReference time for resolving relative dates in the user query: " + f"{reference_time}\n" + ) logger.info(f"Parsing Goal... LLM input is {prompt}") response = self.llm.generate(messages=[{"role": "user", "content": prompt}]) logger.info(f"Parsing Goal... LLM Response is {response}") diff --git a/src/memos/multi_mem_cube/single_cube.py b/src/memos/multi_mem_cube/single_cube.py index 27399259a..ec4544b22 100644 --- a/src/memos/multi_mem_cube/single_cube.py +++ b/src/memos/multi_mem_cube/single_cube.py @@ -232,6 +232,7 @@ def _deep_search( "user_id": search_req.user_id, "session_id": target_session_id, "chat_history": search_req.chat_history, + "reference_time": search_req.reference_time, } enhanced_memories = self.searcher.deep_search( @@ -295,6 +296,7 @@ def _fine_search( "user_id": search_req.user_id, "session_id": target_session_id, "chat_history": search_req.chat_history, + "reference_time": search_req.reference_time, } # Fine retrieve diff --git a/src/memos/search/search_service.py b/src/memos/search/search_service.py index db9f9a6aa..9b1611fe7 100644 --- a/src/memos/search/search_service.py +++ b/src/memos/search/search_service.py @@ -31,6 +31,7 @@ def build_search_context( "user_id": search_req.user_id, "session_id": target_session_id, "chat_history": search_req.chat_history, + "reference_time": search_req.reference_time, }, plugin=bool(search_req.source is not None and search_req.source == "plugin"), ) diff --git a/tests/evaluation/test_longmemeval_api_contract.py b/tests/evaluation/test_longmemeval_api_contract.py new file mode 100644 index 000000000..aaa97a023 --- /dev/null +++ b/tests/evaluation/test_longmemeval_api_contract.py @@ -0,0 +1,77 @@ +import importlib.util +import json + +from pathlib import Path +from unittest.mock import Mock, patch + +from memos.api.product_models import APISearchRequest +from memos.search.search_service import build_search_context + + +_CLIENT_PATH = Path(__file__).parents[2] / "evaluation" / "scripts" / "utils" / "client.py" +_CLIENT_SPEC = importlib.util.spec_from_file_location("longmemeval_client", _CLIENT_PATH) +_CLIENT_MODULE = importlib.util.module_from_spec(_CLIENT_SPEC) +assert _CLIENT_SPEC.loader is not None +_CLIENT_SPEC.loader.exec_module(_CLIENT_MODULE) + +MemosApiClient = _CLIENT_MODULE.MemosApiClient +MemosApiOnlineClient = _CLIENT_MODULE.MemosApiOnlineClient + + +def _response(payload: dict) -> Mock: + response = Mock(status_code=200) + response.text = json.dumps(payload) + return response + + +def test_memos_api_search_forwards_reference_time(monkeypatch): + monkeypatch.setenv("MEMOS_URL", "http://memos.test") + client = MemosApiClient() + reference_time = "2023-04-01T00:00:00Z" + + with patch.object( + _CLIENT_MODULE.requests, + "request", + return_value=_response({"message": "Search completed successfully", "data": {}}), + ) as request: + client.search("What happened yesterday?", "user-1", 5, reference_time=reference_time) + + payload = json.loads(request.call_args.kwargs["data"]) + assert payload["reference_time"] == reference_time + + +def test_online_search_accepts_but_does_not_send_reference_time(monkeypatch): + monkeypatch.setenv("MEMOS_ONLINE_URL", "http://memos-online.test") + client = MemosApiOnlineClient() + reference_time = "2023-04-01T00:00:00Z" + response_payload = { + "message": "ok", + "data": { + "memory_detail_list": [], + "preference_detail_list": [], + "preference_note": "", + }, + } + + with patch.object( + _CLIENT_MODULE.requests, + "request", + return_value=_response(response_payload), + ) as request: + client.search("What happened yesterday?", "user-1", 5, reference_time=reference_time) + + payload = json.loads(request.call_args.kwargs["data"]) + assert "reference_time" not in payload + + +def test_search_context_preserves_reference_time(): + reference_time = "2023-04-01T00:00:00Z" + search_request = APISearchRequest( + query="What happened yesterday?", + user_id="user-1", + reference_time=reference_time, + ) + + context = build_search_context(search_request) + + assert context.info["reference_time"] == reference_time diff --git a/tests/memories/textual/test_tree_searcher.py b/tests/memories/textual/test_tree_searcher.py index b79958ca1..cae1239e5 100644 --- a/tests/memories/textual/test_tree_searcher.py +++ b/tests/memories/textual/test_tree_searcher.py @@ -72,10 +72,18 @@ def retrieve_side_effect(*args, **kwargs): ] result = mock_searcher.search( - query=query, top_k=2, info={"test": True}, mode="fast", memory_type="All" + query=query, + top_k=2, + info={"test": True, "reference_time": "2023-04-01T00:00:00Z"}, + mode="fast", + memory_type="All", ) assert mock_searcher.task_goal_parser.parse.called + assert ( + mock_searcher.task_goal_parser.parse.call_args.kwargs["reference_time"] + == "2023-04-01T00:00:00Z" + ) mock_searcher.embedder.embed.assert_called_once() assert len(result) <= 2 diff --git a/tests/memories/textual/test_tree_task_goal_parser.py b/tests/memories/textual/test_tree_task_goal_parser.py index 899e2454b..f9b5e51b6 100644 --- a/tests/memories/textual/test_tree_task_goal_parser.py +++ b/tests/memories/textual/test_tree_task_goal_parser.py @@ -5,7 +5,11 @@ class MockLLM: + def __init__(self): + self.messages = [] + def generate(self, messages): + self.messages.append(messages) # Just return a fake JSON string return """ { @@ -35,6 +39,19 @@ def test_parse_fine_calls_llm_and_parses(): assert result.goal_type == "fact" +def test_parse_fine_includes_reference_time_for_temporal_queries(): + mock_llm = MockLLM() + parser = TaskGoalParser(llm=mock_llm) + + parser.parse( + "What happened yesterday?", + mode="fine", + reference_time="2023-04-01T00:00:00Z", + ) + + assert "2023-04-01T00:00:00Z" in mock_llm.messages[0][0]["content"] + + def test_parse_response_invalid_json(): parser = TaskGoalParser(llm=MockLLM())