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
16 changes: 12 additions & 4 deletions jiuwen_memory/foundation/store/index/simple_memory_index.py
Original file line number Diff line number Diff line change
Expand Up @@ -666,9 +666,13 @@ async def _list_memories_via_vector_pushdown(

if mem_types:
type_order = {mt: i for i, mt in enumerate(mem_types)}
docs.sort(key=lambda d: (type_order.get(d.type, len(type_order)), -d.timestamp.timestamp()))
docs.sort(key=lambda d: (
type_order.get(d.type, len(type_order)),
-d.timestamp.timestamp(),
d.id,
))
else:
docs.sort(key=lambda d: d.timestamp, reverse=True)
docs.sort(key=lambda d: (-d.timestamp.timestamp(), d.id))
return docs[offset:offset + limit]

async def _list_memories_via_kv_scan(
Expand Down Expand Up @@ -708,9 +712,13 @@ async def _list_memories_via_kv_scan(
docs.append(doc)
if mem_types:
type_order = {mt: i for i, mt in enumerate(mem_types)}
docs.sort(key=lambda d: (type_order.get(d.type, len(type_order)), -d.timestamp.timestamp()))
docs.sort(key=lambda d: (
type_order.get(d.type, len(type_order)),
-d.timestamp.timestamp(),
d.id,
))
else:
docs.sort(key=lambda d: d.timestamp, reverse=True)
docs.sort(key=lambda d: (-d.timestamp.timestamp(), d.id))
if filters is not None:
docs = [d for d in docs if _apply_filter_group(d, filters)]
return docs[offset:offset + limit]
Expand Down
96 changes: 81 additions & 15 deletions jiuwen_memory/memory_core/long_term_memory.py
Original file line number Diff line number Diff line change
Expand Up @@ -2321,16 +2321,54 @@ async def get_user_mem_by_page(self,

start_idx = page_size * (page_idx - 1)
fetch_size = start_idx + page_size
search_data = await self.search_manager.list_user_mem(user_id=user_id, scope_id=scope_id,
nums=fetch_size, pages=1,
mem_type=None,
filters=filters)
mem_results = self._build_mem_info_list(search_data)
mem_results.extend(await self._list_middle_memories(user_id=user_id, scope_id=scope_id,
limit=fetch_size))
mem_results.sort(key=self._mem_info_timestamp_sort_key)

# Both indexed memories and middle-term messages are exposed by their
# stores newest-first. Sorting only an expanding newest-N prefix into
# chronological order makes the prefix shift on every page and causes
# already-returned records to appear again. Build one complete,
# deterministic snapshot first, then paginate that snapshot once.
mem_results = await self._list_all_indexed_memories(
user_id=user_id,
scope_id=scope_id,
filters=filters,
batch_size=max(fetch_size, 100),
)
mem_results.extend(await self._list_all_middle_memories(
user_id=user_id,
scope_id=scope_id,
initial_limit=max(fetch_size, 100),
))
mem_results.sort(key=self._mem_info_chronological_sort_key)
return mem_results[start_idx:start_idx + page_size]

async def _list_all_indexed_memories(
self,
user_id: str,
scope_id: str,
filters: "FilterGroup | None",
batch_size: int,
) -> list[MemInfo]:
"""Read a stable snapshot of all non-middle memories in bounded pages."""
page_size = max(batch_size, 1)
page_idx = 1
results_by_id: dict[str, MemInfo] = {}
while True:
search_data = await self.search_manager.list_user_mem(
user_id=user_id,
scope_id=scope_id,
nums=page_size,
pages=page_idx,
mem_type=None,
filters=filters,
)
batch = self._build_mem_info_list(search_data)
for mem_info in batch:
results_by_id.setdefault(mem_info.mem_id, mem_info)
if len(batch) < page_size:
break
page_idx += 1
return list(results_by_id.values())

@staticmethod
def _build_mem_info_list(search_data: list[dict] | None) -> list[MemInfo]:
if not search_data:
Expand All @@ -2349,12 +2387,14 @@ def _build_mem_info_list(search_data: list[dict] | None) -> list[MemInfo]:
return mem_results

@staticmethod
def _mem_info_timestamp_sort_key(mem_info: MemInfo) -> float:
def _mem_info_chronological_sort_key(mem_info: MemInfo) -> tuple[float, str]:
if mem_info.timestamp is None:
return 0.0
if mem_info.timestamp.tzinfo is None:
return mem_info.timestamp.replace(tzinfo=timezone.utc).timestamp()
return mem_info.timestamp.timestamp()
timestamp = 0.0
elif mem_info.timestamp.tzinfo is None:
timestamp = mem_info.timestamp.replace(tzinfo=timezone.utc).timestamp()
else:
timestamp = mem_info.timestamp.timestamp()
return timestamp, mem_info.mem_id

async def _list_middle_memories_by_page(
self,
Expand All @@ -2365,9 +2405,35 @@ async def _list_middle_memories_by_page(
) -> list[MemInfo]:
start_idx = page_size * (page_idx - 1)
fetch_size = start_idx + page_size
mem_results = await self._list_middle_memories(user_id=user_id, scope_id=scope_id, limit=fetch_size)
mem_results = await self._list_all_middle_memories(
user_id=user_id,
scope_id=scope_id,
initial_limit=max(fetch_size, 100),
)
mem_results.sort(key=self._mem_info_chronological_sort_key)
return mem_results[start_idx:start_idx + page_size]

async def _list_all_middle_memories(
self,
user_id: str,
scope_id: str,
initial_limit: int,
) -> list[MemInfo]:
"""Read all middle memories by growing the newest-prefix request."""
limit = max(initial_limit, 1)
while True:
mem_results = await self._list_middle_memories(
user_id=user_id,
scope_id=scope_id,
limit=limit,
)
if len(mem_results) < limit:
results_by_id: dict[str, MemInfo] = {}
for mem_info in mem_results:
results_by_id.setdefault(mem_info.mem_id, mem_info)
return list(results_by_id.values())
limit *= 2

async def _list_middle_memories(self, user_id: str, scope_id: str, limit: int) -> list[MemInfo]:
if limit <= 0 or not self.message_manager:
return []
Expand Down Expand Up @@ -3032,4 +3098,4 @@ async def _create_semantic_store_with_embedding(self, scope_id: str) -> Semantic
semantic_store.initialize_embedding_model(self._base_embed)
else:
pass
return semantic_store
return semantic_store
4 changes: 2 additions & 2 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -62,7 +62,7 @@ sqlite = ["aiosqlite>=0.22.1", "greenlet>=3.1.1"]
redis = ["redis>=7.1.0"]
mem0 = ["mem0ai>=1.0.0"]
agentarts = ["agentarts-sdk>=0.1.2,<0.2"]
server = ["uvicorn>=0.30.0", "fastapi>=0.110.0", "python-dotenv>=1.2.1", "mcp>=1.2.0,<2.0.0"]
server = ["uvicorn>=0.30.0", "fastapi>=0.110.0", "python-dotenv>=1.2.1", "mcp>=1.14.1,<2.0.0"]
all-storage = ["JiuwenMemory[sqlite,postgres,mysql,gaussdb,redis]"]
all-vector = ["JiuwenMemory[chromadb,gaussvector,elasticsearch]"]
file-index = ["sqlite-vec>=0.1.9", "watchdog>=4.0", "jieba>=0.42.1"]
Expand All @@ -76,7 +76,7 @@ memory-mcp = "jiuwen_memory.server.mcp_server:main"
[dependency-groups]
test = [
"pytest>=8.3.5", "pytest-asyncio>=1.0.0", "pytest-html>=4.1.1",
"pytest-mock>=3.14.0", "pytest-cov>=7.0.0", "coverage>=7.7.1",
"pytest-mock>=3.14.0", "pytest-cov>=7.0.0", "coverage>=7.7.1", "mcp>=1.14.1,<2.0.0",
]
lint = ["ruff>=0.11.2", "pylint>=3.0.0", "mypy>=1.12.0"]
dev = [{ include-group = "test" }, { include-group = "lint" }]
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -469,8 +469,132 @@ async def test_get_user_mem_by_page_unknown_merges_long_and_middle_memory():
assert memory.search_manager.calls == [{
"user_id": "u1",
"scope_id": "s1",
"nums": 10,
"nums": 100,
"pages": 1,
"mem_type": None,
"filters": None,
}]


@pytest.mark.asyncio
async def test_get_user_mem_by_page_unknown_has_no_cross_page_duplicates():
oldest_first = [
{
"id": f"long_{idx:03d}",
"mem": f"长期记忆 {idx}",
"mem_type": MemoryType.SEMANTIC_MEMORY.value,
"timestamp": datetime(2026, 7, 21, idx, 0, tzinfo=timezone.utc),
}
for idx in range(10)
]
memory = LongTermMemory()
# Production indexes list newest-first. The old implementation fetched
# an expanding newest-N prefix, reversed it, and returned duplicate pages.
memory.search_manager = FakeSearchManager(list(reversed(oldest_first)))
memory.message_manager = FakeMessageManager()

pages = [
await memory.get_user_mem_by_page(
user_id="u1",
scope_id="s1",
page_size=3,
page_idx=page_idx,
memory_type=MemoryType.UNKNOWN,
)
for page_idx in range(1, 6)
]

assert [[item.mem_id for item in page] for page in pages] == [
["long_000", "long_001", "long_002"],
["long_003", "long_004", "long_005"],
["long_006", "long_007", "long_008"],
["long_009"],
[],
]
all_ids = [item.mem_id for page in pages for item in page]
assert len(all_ids) == len(set(all_ids)) == 10


@pytest.mark.asyncio
async def test_get_user_mem_by_page_unknown_orders_same_timestamp_by_mem_id():
timestamp = datetime(2026, 7, 21, 12, 0, tzinfo=timezone.utc)
memory = LongTermMemory()
memory.search_manager = FakeSearchManager([
{
"id": "b-long",
"mem": "长期记忆 B",
"mem_type": MemoryType.SEMANTIC_MEMORY.value,
"timestamp": timestamp,
},
{
"id": "a-long",
"mem": "长期记忆 A",
"mem_type": MemoryType.SEMANTIC_MEMORY.value,
"timestamp": timestamp,
},
])
memory.message_manager = FakeMessageManager()
memory.message_manager.messages = [
(BaseMessage(role="user", content="中期记忆 D"), timestamp, "d-middle"),
(BaseMessage(role="user", content="中期记忆 C"), timestamp, "c-middle"),
]

pages = [
await memory.get_user_mem_by_page(
user_id="u1",
scope_id="s1",
page_size=2,
page_idx=page_idx,
memory_type=MemoryType.UNKNOWN,
)
for page_idx in (1, 2)
]

assert [[item.mem_id for item in page] for page in pages] == [
["a-long", "b-long"],
["c-middle", "d-middle"],
]
all_ids = [item.mem_id for page in pages for item in page]
assert len(all_ids) == len(set(all_ids)) == 4


@pytest.mark.asyncio
async def test_get_user_mem_by_page_middle_has_no_cross_page_duplicates():
class NewestPrefixMessageManager(FakeMessageManager):
async def get(self, user_id=None, scope_id=None, session_id=None, message_len=10):
return self.messages[-message_len:]

oldest_first = [
(
BaseMessage(role="user", content=f"中期记忆 {idx}"),
datetime(2026, 7, 21, idx, 0, tzinfo=timezone.utc),
f"middle_{idx:03d}",
)
for idx in range(10)
]
memory = LongTermMemory()
memory.search_manager = FakeSearchManager([])
memory.message_manager = NewestPrefixMessageManager()
# Production MessageManager returns the newest prefix in chronological
# order after reversing the database's DESC query.
memory.message_manager.messages = oldest_first

pages = [
await memory.get_user_mem_by_page(
user_id="u1",
scope_id="s1",
page_size=3,
page_idx=page_idx,
memory_type=MemoryType.MIDDLE_TERM_MEMORY,
)
for page_idx in range(1, 5)
]

assert [[item.mem_id for item in page] for page in pages] == [
["middle_000", "middle_001", "middle_002"],
["middle_003", "middle_004", "middle_005"],
["middle_006", "middle_007", "middle_008"],
["middle_009"],
]
all_ids = [item.mem_id for page in pages for item in page]
assert len(all_ids) == len(set(all_ids)) == 10
Loading