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
64 changes: 53 additions & 11 deletions src/rag.py
Original file line number Diff line number Diff line change
Expand Up @@ -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():
Expand All @@ -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()
Expand Down Expand Up @@ -116,22 +136,44 @@ 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)
continue
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

Expand Down
163 changes: 163 additions & 0 deletions tests/test_rag_chunking.py
Original file line number Diff line number Diff line change
@@ -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