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
19 changes: 15 additions & 4 deletions src/memos/mem_os/utils/reference_utils.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,5 @@
import re

from memos.memories.textual.item import (
TextualMemoryItem,
)
Expand Down Expand Up @@ -41,8 +43,19 @@ def split_continuous_references(text: str) -> str:
# Check if there's a comma between brackets
if "," not in content_between_brackets:
return text
text = text.replace(content_between_brackets, content_between_brackets.replace(", ", "]["))
text = text.replace(content_between_brackets, content_between_brackets.replace(",", "]["))
# Only reference tags (a numeric id before a colon in every element) are
# split; ordinary bracketed text such as "[x, y]" must pass through
# untouched.
if not re.fullmatch(r"\s*\d+:[^,\s]+(?:,\s*\d+:[^,\s]+)*\s*", content_between_brackets):
return text
# Split on every comma regardless of the whitespace that follows it: LLM
# output mixes "a, b" and "a,b" freely, and the previous two-pass
# str.replace handled only one style per call, leaving earlier references
# merged when both styles appeared in the same tag.
text = text.replace(
content_between_brackets,
re.sub(r",\s*", "][", content_between_brackets),
)

return text

Expand All @@ -57,8 +70,6 @@ def process_streaming_references_complete(text_buffer: str) -> tuple[str, str]:
Returns:
tuple[str, str]: (processed_text, remaining_buffer)
"""
import re

# Pattern to match complete reference tags: [refid:memoriesID]
complete_pattern = r"\[\d+:[^\]]+\]"

Expand Down
58 changes: 58 additions & 0 deletions tests/mem_os/test_reference_utils.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,58 @@
import unittest

from memos.mem_os.utils.reference_utils import split_continuous_references


class TestSplitContinuousReferences(unittest.TestCase):
def test_spaced_comma(self):
self.assertEqual(
split_continuous_references("[1:92ff35fb, 4:bfe6f044]"),
"[1:92ff35fb][4:bfe6f044]",
)

def test_tight_comma(self):
self.assertEqual(
split_continuous_references("[1:92ff35fb,4:bfe6f044]"),
"[1:92ff35fb][4:bfe6f044]",
)

def test_mixed_separators(self):
"""Mixed ', ' and ',' separators used to leave earlier refs merged.

The old two-pass str.replace applied one separator style per call;
the first successful pass removed the substring the second pass
searched for, so with "[a,b, c]" only the last boundary was split.
"""
self.assertEqual(
split_continuous_references("[1:92ff35fb,4:bfe6f044, 7:abcd1234]"),
"[1:92ff35fb][4:bfe6f044][7:abcd1234]",
)

def test_ordinary_bracketed_text_untouched(self):
# "[x, y]" has no reference-tag shape (numeric id + colon); it must
# not be rewritten into "[x][y]" on the streaming chat path
self.assertEqual(
split_continuous_references("interval [x, y] ends"),
"interval [x, y] ends",
)
self.assertEqual(split_continuous_references("cite [1, 2] here"), "cite [1, 2] here")

def test_surrounding_text_preserved(self):
self.assertEqual(
split_continuous_references("see refs [1:aaaa, 2:bbbb] for details"),
"see refs [1:aaaa][2:bbbb] for details",
)

def test_non_reference_text_untouched(self):
for text in (
"",
"plain text",
"[no commas here]",
"two [brackets] twice [here]",
"backwards ]here[",
):
self.assertEqual(split_continuous_references(text), text)


if __name__ == "__main__":
unittest.main()