diff --git a/src/memos/mem_os/utils/reference_utils.py b/src/memos/mem_os/utils/reference_utils.py index 09b812207..a843cf024 100644 --- a/src/memos/mem_os/utils/reference_utils.py +++ b/src/memos/mem_os/utils/reference_utils.py @@ -1,3 +1,5 @@ +import re + from memos.memories.textual.item import ( TextualMemoryItem, ) @@ -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 @@ -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+:[^\]]+\]" diff --git a/tests/mem_os/test_reference_utils.py b/tests/mem_os/test_reference_utils.py new file mode 100644 index 000000000..c1c400b6a --- /dev/null +++ b/tests/mem_os/test_reference_utils.py @@ -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()