diff --git a/src/rag/loaders.py b/src/rag/loaders.py index 9dd2468..6aa25eb 100644 --- a/src/rag/loaders.py +++ b/src/rag/loaders.py @@ -29,9 +29,8 @@ def load_documents(directory: str | Path) -> list[Document]: raise FileNotFoundError(f"Document directory not found: {path}") documents: list[Document] = [] - for file_path in sorted(path.iterdir(), key=lambda item: item.name): - if not file_path.is_file(): - continue + files = [file_path for file_path in path.rglob("*") if file_path.is_file()] + for file_path in sorted(files, key=lambda item: item.relative_to(path).as_posix()): suffix = file_path.suffix.lower().lstrip(".") if suffix not in SUPPORTED_TEXT_EXTENSIONS: continue diff --git a/tests/test_loaders.py b/tests/test_loaders.py index 94c80c3..5b7ffcb 100644 --- a/tests/test_loaders.py +++ b/tests/test_loaders.py @@ -62,3 +62,118 @@ def test_workshop_corpus_is_a_markdown_collection() -> None: assert len(documents) == 6 assert all(doc.filename.endswith(".md") for doc in documents) assert all(doc.metadata["source"].endswith(".md") for doc in documents) + + +def test_load_documents_root_level_document(tmp_path: Path) -> None: + root_dir = tmp_path / "root" + root_dir.mkdir() + (root_dir / "document.txt").write_text("root content\n", encoding="utf-8") + + documents = load_documents(root_dir) + + assert len(documents) == 1 + assert documents[0].filename == "document.txt" + assert documents[0].source == "document.txt" + assert documents[0].text == "root content\n" + + +def test_load_documents_one_level_subdirectory(tmp_path: Path) -> None: + root_dir = tmp_path / "root" + docs_dir = root_dir / "docs" + docs_dir.mkdir(parents=True) + (docs_dir / "document.txt").write_text("nested content\n", encoding="utf-8") + + documents = load_documents(root_dir) + + assert len(documents) == 1 + assert documents[0].filename == "document.txt" + assert documents[0].source == "docs/document.txt" + + +def test_load_documents_multiple_nested_levels(tmp_path: Path) -> None: + root_dir = tmp_path / "root" + nested_dir = root_dir / "docs" / "java" / "oop" + nested_dir.mkdir(parents=True) + file_path = nested_dir / "inheritance.txt" + file_path.write_text("class Inheritance {}\n", encoding="utf-8") + + documents = load_documents(root_dir) + + assert len(documents) == 1 + assert documents[0].filename == "inheritance.txt" + assert documents[0].source == "docs/java/oop/inheritance.txt" + + +def test_load_documents_unsupported_extensions(tmp_path: Path) -> None: + root_dir = tmp_path / "root" + root_dir.mkdir() + (root_dir / "valid.txt").write_text("valid txt\n", encoding="utf-8") + (root_dir / "valid.md").write_text("valid md\n", encoding="utf-8") + (root_dir / "image.png").write_bytes(b"\x89PNG\r\n\x1a\n") + (root_dir / "random.xyz").write_text("random\n", encoding="utf-8") + + documents = load_documents(root_dir) + + assert [doc.filename for doc in documents] == ["valid.md", "valid.txt"] + assert [doc.source for doc in documents] == ["valid.md", "valid.txt"] + + +def test_load_documents_mixed_root_and_nested_documents(tmp_path: Path) -> None: + root_dir = tmp_path / "root" + root_dir.mkdir() + (root_dir / "a.txt").write_text("a\n", encoding="utf-8") + (root_dir / "docs" / "java").mkdir(parents=True) + (root_dir / "docs" / "b.txt").write_text("b\n", encoding="utf-8") + (root_dir / "docs" / "java" / "c.txt").write_text("c\n", encoding="utf-8") + (root_dir / "notes").mkdir(parents=True) + (root_dir / "notes" / "d.md").write_text("d\n", encoding="utf-8") + + documents = load_documents(root_dir) + + sources = [doc.source for doc in documents] + assert sources == ["a.txt", "docs/b.txt", "docs/java/c.txt", "notes/d.md"] + + +def test_load_documents_deterministic_ordering(tmp_path: Path) -> None: + root_dir = tmp_path / "root" + root_dir.mkdir() + (root_dir / "z_root.txt").write_text("z\n", encoding="utf-8") + (root_dir / "b_dir").mkdir() + (root_dir / "b_dir" / "a_file.txt").write_text("ba\n", encoding="utf-8") + (root_dir / "a_dir" / "sub_dir").mkdir(parents=True) + (root_dir / "a_dir" / "sub_dir" / "x_file.txt").write_text("ax\n", encoding="utf-8") + + expected_sources = [ + "a_dir/sub_dir/x_file.txt", + "b_dir/a_file.txt", + "z_root.txt", + ] + + for _ in range(10): + documents = load_documents(root_dir) + assert [doc.source for doc in documents] == expected_sources + + +def test_load_documents_empty_nested_directories(tmp_path: Path) -> None: + root_dir = tmp_path / "root" + (root_dir / "empty1" / "empty2").mkdir(parents=True) + + documents = load_documents(root_dir) + + assert documents == [] + + +def test_load_documents_unsupported_files_in_nested_directories(tmp_path: Path) -> None: + root_dir = tmp_path / "root" + docs_dir = root_dir / "docs" + docs_dir.mkdir(parents=True) + (docs_dir / "image.png").write_bytes(b"png data") + (docs_dir / "data.xyz").write_text("xyz data\n", encoding="utf-8") + (docs_dir / "valid.txt").write_text("valid content\n", encoding="utf-8") + + documents = load_documents(root_dir) + + assert len(documents) == 1 + assert documents[0].filename == "valid.txt" + assert documents[0].source == "docs/valid.txt" +