Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 9 additions & 8 deletions evaluation/scripts/utils/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand All @@ -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"]


Expand Down Expand Up @@ -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(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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}")
Expand Down
2 changes: 2 additions & 0 deletions src/memos/multi_mem_cube/single_cube.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down Expand Up @@ -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
Expand Down
1 change: 1 addition & 0 deletions src/memos/search/search_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"),
)
Expand Down
77 changes: 77 additions & 0 deletions tests/evaluation/test_longmemeval_api_contract.py
Original file line number Diff line number Diff line change
@@ -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
10 changes: 9 additions & 1 deletion tests/memories/textual/test_tree_searcher.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
17 changes: 17 additions & 0 deletions tests/memories/textual/test_tree_task_goal_parser.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 """
{
Expand Down Expand Up @@ -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())

Expand Down
Loading