diff --git a/src/rag.py b/src/rag.py index 471540c3..1e523a3d 100644 --- a/src/rag.py +++ b/src/rag.py @@ -46,7 +46,7 @@ def _get_collection(): # --- Helpers ------------------------------------------------------------- -HEADING_RE = re.compile(r"^(#{1,4})\s+(.+)$", re.MULTILINE) +HEADING_RE = re.compile(r"^(#{1,4})[ \t]+(.+)$", re.MULTILINE) def _resolve_knowledge_dir(): @@ -65,15 +65,35 @@ def _decode_metta(s): # --- Chunking ------------------------------------------------------------ -def _chunk_markdown(text, filename): - """Heading-aware markdown chunking with breadcrumb tracking.""" - matches = list(HEADING_RE.finditer(text)) - if not matches: - return [{"text": text.strip(), "breadcrumb": filename}] - +def _cut_to_limit(text, breadcrumb): + """Split text that is still over MAX_CHUNK_CHARS.""" + pieces = [] + text = text.strip() + while len(text) > MAX_CHUNK_CHARS: + window = text[:MAX_CHUNK_CHARS] + # Cut mid-token only when there is no whitespace. + cut = max(window.rfind("\n"), window.rfind(" ")) + if cut <= 0: + cut = MAX_CHUNK_CHARS + head = text[:cut].strip() + if head: + pieces.append({"text": head, "breadcrumb": breadcrumb}) + text = text[cut:].strip() + if text: + pieces.append({"text": text, "breadcrumb": breadcrumb}) + return pieces + + +def _sections_from_headings(text, filename, matches): + """Split on headings and merge sections under MIN_CHUNK_CHARS.""" sections = [] stack = {} # level -> heading text + # Text before the first heading belongs to no section, so add it as one. + preamble = text[:matches[0].start()].strip() + if preamble: + sections.append({"text": preamble, "breadcrumb": filename, "heading": ""}) + for i, m in enumerate(matches): level = len(m.group(1)) heading = m.group(2).strip() @@ -116,8 +136,22 @@ def _chunk_markdown(text, filename): else: merged.append({"text": carry, "breadcrumb": carry_bc}) - # Split large sections on paragraph boundaries + return merged + + +def _chunk_markdown(text, filename): + """Heading-aware markdown chunking with breadcrumb tracking.""" + matches = list(HEADING_RE.finditer(text)) + if matches: + merged = _sections_from_headings(text, filename, matches) + else: + # No headings: one section, still goes through the size pass below. + logger.warning(f"{filename}: no headings, chunking on paragraphs") + merged = [{"text": text.strip(), "breadcrumb": filename}] + + # Split large sections on paragraph boundaries, then on characters if needed. final = [] + hard_cuts = 0 for s in merged: if len(s["text"]) <= MAX_CHUNK_CHARS: final.append(s) @@ -125,13 +159,21 @@ def _chunk_markdown(text, filename): paragraphs = s["text"].split("\n\n") chunk_text = "" for p in paragraphs: - if chunk_text and len(chunk_text) + len(p) > MAX_CHUNK_CHARS: - final.append({"text": chunk_text.strip(), "breadcrumb": s["breadcrumb"]}) + # Count the "\n\n" join, else the chunk overshoots and leaves a fragment. + if chunk_text and len(chunk_text) + 2 + len(p) > MAX_CHUNK_CHARS: + pieces = _cut_to_limit(chunk_text, s["breadcrumb"]) + hard_cuts += len(pieces) - 1 + final.extend(pieces) chunk_text = p else: chunk_text = (chunk_text + "\n\n" + p).strip() if chunk_text.strip(): - final.append({"text": chunk_text.strip(), "breadcrumb": s["breadcrumb"]}) + pieces = _cut_to_limit(chunk_text, s["breadcrumb"]) + hard_cuts += len(pieces) - 1 + final.extend(pieces) + + if hard_cuts: + logger.warning(f"{filename}: {hard_cuts} chunk(s) cut to MAX_CHUNK_CHARS") return final diff --git a/tests/test_rag_chunking.py b/tests/test_rag_chunking.py new file mode 100644 index 00000000..1bc31a06 --- /dev/null +++ b/tests/test_rag_chunking.py @@ -0,0 +1,163 @@ +import importlib.util +import logging +import sys +import types +from pathlib import Path + +import pytest + +REPO_ROOT = Path(__file__).resolve().parents[1] + + +@pytest.fixture +def rag(monkeypatch): + """Load src/rag.py with embedding and storage deps stubbed.""" + logger_mod = types.ModuleType("src.logger") + logger_mod.get_logger = lambda name: logging.getLogger(name) + monkeypatch.setitem(sys.modules, "src.logger", logger_mod) + + chromadb_mod = types.ModuleType("chromadb") + chromadb_mod.PersistentClient = lambda **kwargs: None + monkeypatch.setitem(sys.modules, "chromadb", chromadb_mod) + + openai_mod = types.ModuleType("openai") + openai_mod.OpenAI = object + monkeypatch.setitem(sys.modules, "openai", openai_mod) + + llm_mod = types.ModuleType("lib_llm_ext") + llm_mod.initLocalEmbedding = lambda: None + llm_mod.useLocalEmbedding = lambda text: [] + monkeypatch.setitem(sys.modules, "lib_llm_ext", llm_mod) + + config_mod = types.ModuleType("config") + config_mod.config_get_by_key = lambda key, default=None: default + monkeypatch.setitem(sys.modules, "config", config_mod) + + embedding_mod = types.ModuleType("embedding_models") + embedding_mod.embedding_model = lambda provider, model: model + monkeypatch.setitem(sys.modules, "embedding_models", embedding_mod) + + spec = importlib.util.spec_from_file_location( + "rag_under_test", REPO_ROOT / "src" / "rag.py" + ) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module + + +def largest(chunks): + return max(len(c["text"]) for c in chunks) + + +def test_headingless_file_is_not_returned_whole(rag): + text = ("word " * 3000 + "\n\n") * 8 + chunks = rag._chunk_markdown(text, "plain.md") + + assert len(chunks) > 1 + assert largest(chunks) <= rag.MAX_CHUNK_CHARS + + +def test_single_heading_does_not_change_the_bound(rag): + text = ("word " * 3000 + "\n\n") * 8 + + without = rag._chunk_markdown(text, "plain.md") + with_heading = rag._chunk_markdown("# Chapter\n\n" + text, "with.md") + + assert largest(without) <= rag.MAX_CHUNK_CHARS + assert largest(with_heading) <= rag.MAX_CHUNK_CHARS + + +def test_paragraph_longer_than_the_cap_is_cut(rag): + # One paragraph, no blank lines. + chunks = rag._chunk_markdown("# H\n\n" + "word " * 30000, "nl.md") + + assert largest(chunks) <= rag.MAX_CHUNK_CHARS + + +def test_text_with_no_whitespace_is_cut_mid_token(rag): + chunks = rag._chunk_markdown("x" * (rag.MAX_CHUNK_CHARS * 3), "blob.md") + + assert largest(chunks) <= rag.MAX_CHUNK_CHARS + assert "".join(c["text"] for c in chunks) == "x" * (rag.MAX_CHUNK_CHARS * 3) + + +def test_cuts_prefer_whitespace_boundaries(rag): + chunks = rag._chunk_markdown("word " * 30000, "words.md") + + assert largest(chunks) <= rag.MAX_CHUNK_CHARS + for chunk in chunks: + assert "wor d" not in chunk["text"] + assert chunk["text"].startswith("word") + assert chunk["text"].endswith("word") + + +def test_headingless_chunks_keep_the_filename_breadcrumb(rag): + chunks = rag._chunk_markdown(("word " * 3000 + "\n\n") * 8, "plain.md") + + assert {c["breadcrumb"] for c in chunks} == {"plain.md"} + + +def test_short_headingless_file_stays_one_chunk(rag): + chunks = rag._chunk_markdown("just a short note", "note.md") + + assert chunks == [{"text": "just a short note", "breadcrumb": "note.md"}] + + +def test_heading_structure_and_breadcrumbs_are_preserved(rag): + text = ( + "# Top\n\n" + "a" * 200 + "\n\n" + "## Nested\n\n" + "b" * 200 + "\n\n" + "# Second\n\n" + "c" * 200 + "\n" + ) + chunks = rag._chunk_markdown(text, "doc.md") + + assert [c["breadcrumb"] for c in chunks] == [ + "doc.md > Top", + "doc.md > Top > Nested", + "doc.md > Second", + ] + + +def test_paragraph_join_does_not_leave_a_fragment(rag): + # Two paragraphs landing within the "\n\n" join of the cap. + a = " ".join(["word"] * 600) + b = " ".join(["word"] * 600) + "s" + chunks = rag._chunk_markdown("# H\n\n" + a + "\n\n" + b, "t.md") + + assert [len(c["text"]) for c in chunks] == [2999, 3000] + + +def test_text_before_the_first_heading_is_kept(rag): + doc = "Text before the first heading.\n\n# Heading\n\nText after the heading.\n" + chunks = rag._chunk_markdown(doc, "t.md") + + assert "Text before the first heading." in " ".join(c["text"] for c in chunks) + + +def test_long_preamble_becomes_its_own_section(rag): + preamble = "a" * 500 + chunks = rag._chunk_markdown(preamble + "\n\n# Heading\n\n" + "b" * 500, "t.md") + + assert chunks[0] == {"text": preamble, "breadcrumb": "t.md"} + assert chunks[1]["breadcrumb"] == "t.md > Heading" + + +def test_bare_hash_line_is_not_a_heading(rag): + chunks = rag._chunk_markdown("Body text.\n\n#\n35\n", "t.md") + + assert {c["breadcrumb"] for c in chunks} == {"t.md"} + assert "35" in " ".join(c["text"] for c in chunks) + + +def test_warns_when_a_file_has_no_headings(rag, caplog): + with caplog.at_level(logging.WARNING): + rag._chunk_markdown(("word " * 3000 + "\n\n") * 8, "plain.md") + + assert "no headings, chunking on paragraphs" in caplog.text + + +def test_warns_when_a_character_level_cut_is_needed(rag, caplog): + with caplog.at_level(logging.WARNING): + rag._chunk_markdown("# H\n\n" + "word " * 30000, "nl.md") + + assert "cut to MAX_CHUNK_CHARS" in caplog.text