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
39 changes: 37 additions & 2 deletions src/memos/chunkers/markdown_chunker.py
Original file line number Diff line number Diff line change
Expand Up @@ -75,12 +75,41 @@ def chunk(self, text: str, **kwargs) -> list[str] | list[Chunk]:
logger.debug(f"Generated {len(chunks)} chunks from input text")
return chunks

_FENCE_RE = re.compile(r"^\s{0,3}(`{3,}|~{3,})")

@classmethod
def _code_block_mask(cls, lines: list[str]) -> list[bool]:
"""Mark which lines sit inside a fenced code block.

Fenced code (``` or ~~~, CommonMark-style) is tracked so that ``#``
comments inside embedded code are never mistaken for markdown
headers. The fence lines themselves are masked as well.
"""
mask = []
fence_char = None
for line in lines:
fence_match = cls._FENCE_RE.match(line)
if fence_match:
char = fence_match.group(1)[0]
if fence_char is None:
fence_char = char
elif char == fence_char:
fence_char = None
mask.append(True)
else:
mask.append(fence_char is not None)
return mask

def _detect_malformed_headers(self, text: str) -> bool:
"""Detect if markdown has improper header hierarchy usage."""
# Extract all valid markdown header lines
header_levels = []
pattern = re.compile(r"^#{1,6}\s+.+")
for line in text.split("\n"):
lines = text.split("\n")
code_mask = self._code_block_mask(lines)
for line, in_code in zip(lines, code_mask, strict=True):
if in_code:
continue
stripped_line = line.strip()
if pattern.match(stripped_line):
hash_match = re.match(r"^(#+)", stripped_line)
Expand Down Expand Up @@ -123,10 +152,16 @@ def _fix_header_hierarchy(self, text: str) -> str:
"""
header_pattern = re.compile(r"^(#{1,6})\s+(.+)$")
lines = text.split("\n")
code_mask = self._code_block_mask(lines)
fixed_lines = []
first_valid_header = False

for line in lines:
for line, in_code in zip(lines, code_mask, strict=True):
if in_code:
# Fenced code must pass through untouched: `#` comments in
# embedded code are not headers.
fixed_lines.append(line)
continue
stripped_line = line.strip()
# Match valid header lines (invalid # lines kept as-is)
header_match = header_pattern.match(stripped_line)
Expand Down
76 changes: 76 additions & 0 deletions tests/chunkers/test_markdown_chunker_code_fence.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,76 @@
import unittest

from unittest.mock import patch

from memos.chunkers.markdown_chunker import MarkdownChunker


class TestMarkdownChunkerCodeFence(unittest.TestCase):
"""Fenced code blocks must never be treated as markdown headers.

``#``-comment lines inside embedded code (```python blocks, shell
scripts, ...) used to be counted as level-1 headers by
``_detect_malformed_headers``; enough of them triggered the
"malformed hierarchy" repair, and ``_fix_header_hierarchy`` rewrote the
comments into ``## ...`` lines inside the code block, corrupting it.
"""

def _chunker(self) -> MarkdownChunker:
with patch("langchain_text_splitters.MarkdownHeaderTextSplitter"):
return MarkdownChunker(config=None, auto_fix_headers=True)

def test_code_block_only_has_no_headers(self):
text = "```python\n# one\n# two\n# three\n# four\n# five\nx = 1\n```\n"
chunker = self._chunker()

self.assertFalse(chunker._detect_malformed_headers(text))

def test_fix_leaves_code_block_intact(self):
text = (
"# Title\n\n"
"Intro.\n\n"
"```python\n"
"# comment one\n"
"# comment two\n"
"# comment three\n"
"# comment four\n"
"# comment five\n"
"x = 1\n"
"```\n"
)
chunker = self._chunker()

# the fixer must leave fenced code byte-for-byte intact
self.assertEqual(chunker._fix_header_hierarchy(text), text)

def test_real_malformed_headers_still_fixed(self):
text = "# A\n# B\n# C\nbody\n"
chunker = self._chunker()

self.assertTrue(chunker._detect_malformed_headers(text))
fixed = chunker._fix_header_hierarchy(text)
self.assertIn("# A\n", fixed)
self.assertIn("## B\n", fixed)
self.assertIn("## C\n", fixed)

def test_tilde_fence_ignored_and_real_headers_fixed(self):
text = "~~~\n# not a header\n~~~\n\n# Real One\n# Real Two\n"
chunker = self._chunker()

# only the two real headers are counted, which is malformed
self.assertTrue(chunker._detect_malformed_headers(text))
fixed = chunker._fix_header_hierarchy(text)
self.assertIn("~~~\n# not a header\n~~~", fixed)
self.assertIn("# Real One\n", fixed)
self.assertIn("## Real Two\n", fixed)

def test_unclosed_fence_holds_to_end(self):
text = "```python\n# one\n# two\n# three\n# four\n# five\n"
chunker = self._chunker()

self.assertFalse(chunker._detect_malformed_headers(text))
self.assertEqual(chunker._fix_header_hierarchy(text), text)


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