diff --git a/README.md b/README.md index 4e06874..24e4500 100644 --- a/README.md +++ b/README.md @@ -8,7 +8,8 @@ Local-first Retrieval-Augmented Generation (RAG) toolkit built entirely on top o - **Deterministic and offline** – default embedding model is local; no network calls unless explicitly configured. - **Hybrid retrieval** – combines BM25 via FTS5 with cosine similarity over stored vectors. - **Python and CLI** – flexible API plus a friendly Typer-based CLI for scripting. -- **Extensible** – pluggable parsers, chunkers, embedding backends, and adapters for LangChain / LlamaIndex. +- **Extensible** – pluggable parsers, chunkers, embedding backends, rerankers, and adapters for LangChain / LlamaIndex. +- **Multi-user ready** – optional FastAPI server turns the SQLite database into a shared retrieval service. ## Installation @@ -45,6 +46,10 @@ raglite query "example text" --db knowledge.db -k 8 --hybrid 0.6 ```bash raglite query "How do I deploy?" --db knowledge.db -k 5 --hybrid 0.7 ``` +- Apply a cross-encoder reranker at query time: + ```bash + raglite query "release checklist" --db knowledge.db --reranker cross-encoder --reranker-option model_name=cross-encoder/ms-marco-MiniLM-L-6-v2 + ``` - Filter by tag or doc id: ```bash raglite query "release checklist" --db knowledge.db --filter tag=internal --filter doc_id=release-notes @@ -61,6 +66,32 @@ raglite query "example text" --db knowledge.db -k 8 --hybrid 0.6 ```bash raglite vacuum --db knowledge.db ``` +- Serve a REST API for multi-user access (requires `pip install raglite-sqlite[server]`): + ```bash + raglite serve --db knowledge.db --host 0.0.0.0 --port 8080 + ``` + +## REST API + +RagLite can expose a FastAPI-powered REST server for shared deployments. Install the optional extras first: + +```bash +pip install raglite-sqlite[server] +``` + +Then launch the server with `raglite serve`. Key endpoints: + +- `GET /health` – basic health check. +- `GET /stats` – retrieve database statistics. +- `POST /index` – index new files on disk. +- `POST /query` – perform hybrid search with optional reranking. +- `POST /delete` – remove a document by `doc_id`. + +All endpoints operate directly on the shared SQLite database, making it easy to expose retrieval to multiple users without extra infrastructure. + +## Document formats & OCR + +Beyond plain text and Markdown, RagLite understands PDF, HTML, DOCX, CSV, JSON, PPTX, and image files. Image ingestion uses OCR via `pytesseract`/`Pillow`—install them with `pip install raglite-sqlite[ocr]`. Custom parser options can be passed via `RagLite.index(..., parser_opts={...})` or the REST API. ## Python API diff --git a/pyproject.toml b/pyproject.toml index 7b55292..139e019 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -29,6 +29,14 @@ dependencies = [ [project.optional-dependencies] openai = ["openai>=1.0.0"] +ocr = [ + "pytesseract>=0.3.10", + "Pillow>=10.0.0", +] +server = [ + "fastapi>=0.110.0", + "uvicorn[standard]>=0.29.0", +] dev = [ "pytest>=8.1.0", "pytest-cov>=4.1.0", diff --git a/raglite_sqlite/__init__.py b/raglite_sqlite/__init__.py index 42882bb..db07740 100644 --- a/raglite_sqlite/__init__.py +++ b/raglite_sqlite/__init__.py @@ -2,4 +2,9 @@ from .api import RagLite -__all__ = ["RagLite"] +try: # pragma: no cover - optional dependency + from .server import create_app +except Exception: # pragma: no cover - FastAPI not installed + create_app = None # type: ignore[assignment] + +__all__ = ["RagLite", "create_app"] diff --git a/raglite_sqlite/api.py b/raglite_sqlite/api.py index f9ca192..02d6064 100644 --- a/raglite_sqlite/api.py +++ b/raglite_sqlite/api.py @@ -10,11 +10,15 @@ from .parsers.csv import CSVParser from .parsers.docx import DocxParser from .parsers.html import HTMLParser +from .parsers.image import ImageParser +from .parsers.json import JSONParser from .parsers.md import MarkdownParser from .parsers.pdf import PDFParser +from .parsers.pptx import PptxParser from .parsers.txt import TextParser from .search import assemble_results, bm25_search, hybrid_fuse, vector_search from .typing import SearchResult +from .rerank import Reranker, get_reranker from .utils import detect_mime, iter_files, loads_json, now_ts, normalize_text, sha256_file, sha256_text PARSER_REGISTRY = { @@ -24,6 +28,11 @@ "application/pdf": PDFParser(), "application/vnd.openxmlformats-officedocument.wordprocessingml.document": DocxParser(), "text/csv": CSVParser(), + "application/json": JSONParser(), + "application/vnd.openxmlformats-officedocument.presentationml.presentation": PptxParser(), + "image/png": ImageParser(), + "image/jpeg": ImageParser(), + "image/tiff": ImageParser(), } @@ -160,6 +169,8 @@ def search( embedding_backend: EmbeddingBackend | None = None, max_per_doc: int = 3, with_snippets: bool = True, + reranker: Reranker | str | None = None, + reranker_options: dict[str, object] | None = None, ) -> List[SearchResult]: norm_query = normalize_text(query) lexical = bm25_search(self.db, norm_query, k, filters) @@ -168,7 +179,20 @@ def search( if backend is not None: semantic = vector_search(self.db, backend, norm_query, model_name, k) fused = hybrid_fuse(lexical, semantic, hybrid_weight, k) - return assemble_results(self.db, fused, k, with_snippets=with_snippets, max_per_doc=max_per_doc) + results = assemble_results( + self.db, fused, k, with_snippets=with_snippets, max_per_doc=max_per_doc + ) + reranker_instance: Reranker | None + if isinstance(reranker, str): + try: + reranker_instance = get_reranker(reranker, **(reranker_options or {})) + except KeyError as exc: + raise ValueError(f"Unknown reranker '{reranker}'") from exc + else: + reranker_instance = reranker + if reranker_instance is not None: + results = list(reranker_instance.rerank(query, results)) + return results def delete(self, doc_id: str) -> None: self.db.delete_document(doc_id) diff --git a/raglite_sqlite/cli.py b/raglite_sqlite/cli.py index 69e9c55..7d10086 100644 --- a/raglite_sqlite/cli.py +++ b/raglite_sqlite/cli.py @@ -68,6 +68,10 @@ def query( hybrid: float = typer.Option(0.6, min=0.0, max=1.0, help="Hybrid weight"), max_per_doc: int = typer.Option(3, help="Max results per document"), filters: Optional[list[str]] = typer.Option(None, "--filter", help="Filter key=value"), + reranker: Optional[str] = typer.Option(None, help="Name of reranker to apply"), + reranker_option: Optional[list[str]] = typer.Option( + None, "--reranker-option", help="Reranker option key=value" + ), ) -> None: rag = get_rag(db) filter_dict: dict[str, str] | None = None @@ -78,7 +82,26 @@ def query( raise typer.BadParameter("Filters must be in key=value format") key, value = item.split("=", 1) filter_dict[key] = value - results = rag.search(text, k=k, hybrid_weight=hybrid, max_per_doc=max_per_doc, filters=filter_dict) + reranker_opts: dict[str, object] | None = None + if reranker_option: + reranker_opts = {} + for item in reranker_option: + if "=" not in item: + raise typer.BadParameter("Reranker options must be in key=value format") + key, value = item.split("=", 1) + reranker_opts[key] = value + try: + results = rag.search( + text, + k=k, + hybrid_weight=hybrid, + max_per_doc=max_per_doc, + filters=filter_dict, + reranker=reranker, + reranker_options=reranker_opts, + ) + except ValueError as exc: + raise typer.BadParameter(str(exc)) from exc table = Table(show_header=True, header_style="bold magenta") table.add_column("Score", justify="right") table.add_column("Doc ID") @@ -119,3 +142,29 @@ def vacuum(db: Path = typer.Option(..., help="Database path")) -> None: rag = get_rag(db) rag.vacuum() console.print("VACUUM completed") + + +@app.command() +def serve( + db: Path = typer.Option(..., help="Database path"), + host: str = typer.Option("127.0.0.1", help="Host to bind"), + port: int = typer.Option(8000, help="Port to bind"), + reload: bool = typer.Option(False, help="Enable auto-reload"), +) -> None: + """Run the optional REST server.""" + + try: + from .server import create_app + except ImportError as exc: # pragma: no cover - optional dependency missing + raise typer.BadParameter( + "The REST server requires FastAPI. Install raglite-sqlite[server]." + ) from exc + try: + import uvicorn + except ImportError as exc: # pragma: no cover - optional dependency missing + raise typer.BadParameter( + "Running the REST server requires uvicorn. Install raglite-sqlite[server]." + ) from exc + + app_instance = create_app(str(db)) + uvicorn.run(app_instance, host=host, port=port, reload=reload) diff --git a/raglite_sqlite/parsers/__init__.py b/raglite_sqlite/parsers/__init__.py index 6415c1a..195426d 100644 --- a/raglite_sqlite/parsers/__init__.py +++ b/raglite_sqlite/parsers/__init__.py @@ -6,6 +6,9 @@ from .pdf import PDFParser from .docx import DocxParser from .csv import CSVParser +from .json import JSONParser +from .pptx import PptxParser +from .image import ImageParser __all__ = [ "TextParser", @@ -14,4 +17,7 @@ "PDFParser", "DocxParser", "CSVParser", + "JSONParser", + "PptxParser", + "ImageParser", ] diff --git a/raglite_sqlite/parsers/image.py b/raglite_sqlite/parsers/image.py new file mode 100644 index 0000000..8c47d1f --- /dev/null +++ b/raglite_sqlite/parsers/image.py @@ -0,0 +1,36 @@ +from __future__ import annotations + +from pathlib import Path +from typing import Iterable + +from .base import BaseParser +from ..typing import ParsedBlock +from ..utils import normalize_text + + +class ImageParser(BaseParser): + """Parse image files using pytesseract-based OCR.""" + + def __init__(self, default_lang: str = "eng") -> None: + self.default_lang = default_lang + + def parse(self, path: str, **options: object) -> Iterable[ParsedBlock]: + lang = str(options.get("lang", self.default_lang)) + config = options.get("tesseract_config") + try: + from PIL import Image + except ImportError as exc: # pragma: no cover - dependency missing + raise RuntimeError("Pillow is required for OCR parsing") from exc + try: + import pytesseract + except ImportError as exc: # pragma: no cover - dependency missing + raise RuntimeError("pytesseract is required for OCR parsing") from exc + image_path = Path(path) + with Image.open(image_path) as image: + try: + text = pytesseract.image_to_string(image, lang=lang, config=config) + except pytesseract.pytesseract.TesseractNotFoundError as exc: # pragma: no cover - environment specific + raise RuntimeError( + "Tesseract OCR binary not found. Install it or adjust TESSDATA_PREFIX." + ) from exc + yield ParsedBlock(text=normalize_text(text), section=None) diff --git a/raglite_sqlite/parsers/json.py b/raglite_sqlite/parsers/json.py new file mode 100644 index 0000000..be159df --- /dev/null +++ b/raglite_sqlite/parsers/json.py @@ -0,0 +1,20 @@ +from __future__ import annotations + +import json +from pathlib import Path +from typing import Iterable + +from .base import BaseParser +from ..typing import ParsedBlock +from ..utils import normalize_text + + +class JSONParser(BaseParser): + """Parse JSON documents by rendering them as a stable, human-readable string.""" + + def parse(self, path: str, **options: object) -> Iterable[ParsedBlock]: + indent = int(options.get("indent", 2) or 0) + sort_keys = bool(options.get("sort_keys", True)) + data = json.loads(Path(path).read_text(encoding="utf-8")) + text = json.dumps(data, indent=indent or None, sort_keys=sort_keys, ensure_ascii=False) + yield ParsedBlock(text=normalize_text(text), section=None) diff --git a/raglite_sqlite/parsers/pptx.py b/raglite_sqlite/parsers/pptx.py new file mode 100644 index 0000000..20dac7c --- /dev/null +++ b/raglite_sqlite/parsers/pptx.py @@ -0,0 +1,45 @@ +from __future__ import annotations + +from pathlib import Path +from typing import Iterable +from xml.etree import ElementTree +from zipfile import ZipFile + +from .base import BaseParser +from ..typing import ParsedBlock +from ..utils import normalize_text + + +class PptxParser(BaseParser): + """Lightweight PPTX parser that extracts text from slide XML payloads.""" + + SLIDE_PREFIX = "ppt/slides/" + SLIDE_SUFFIX = ".xml" + + def parse(self, path: str, **options: object) -> Iterable[ParsedBlock]: + pptx_path = Path(path) + texts: list[str] = [] + with ZipFile(pptx_path) as archive: + slide_names = sorted( + name + for name in archive.namelist() + if name.startswith(self.SLIDE_PREFIX) and name.endswith(self.SLIDE_SUFFIX) + ) + for slide_name in slide_names: + with archive.open(slide_name) as handle: + xml_bytes = handle.read() + try: + root = ElementTree.fromstring(xml_bytes) + except ElementTree.ParseError: + continue + namespaces = { + "a": "http://schemas.openxmlformats.org/drawingml/2006/main", + } + slide_text: list[str] = [] + for node in root.findall('.//a:t', namespaces): + if node.text: + slide_text.append(node.text) + if slide_text: + texts.append(" ".join(slide_text)) + combined = normalize_text("\n\n".join(texts)) + yield ParsedBlock(text=combined, section=None) diff --git a/raglite_sqlite/rerank.py b/raglite_sqlite/rerank.py index f2aa880..85cae6a 100644 --- a/raglite_sqlite/rerank.py +++ b/raglite_sqlite/rerank.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import Iterable, Protocol +from typing import Callable, Iterable, Protocol from .typing import SearchResult @@ -13,3 +13,53 @@ def rerank(self, query: str, results: Iterable[SearchResult]) -> Iterable[Search class NoopReranker: def rerank(self, query: str, results: Iterable[SearchResult]) -> Iterable[SearchResult]: return results + + +RerankerFactory = Callable[..., Reranker] + +_RERANKER_REGISTRY: dict[str, RerankerFactory] = { + "none": lambda **_: NoopReranker(), +} + + +def register_reranker(name: str, factory: RerankerFactory) -> None: + """Register a reranker factory.""" + + _RERANKER_REGISTRY[name] = factory + + +def get_reranker(name: str, **options: object) -> Reranker: + """Instantiate a reranker by name.""" + + factory = _RERANKER_REGISTRY.get(name) + if factory is None: + raise KeyError(name) + return factory(**options) + + +class CrossEncoderReranker: + """Rerank using a sentence-transformers CrossEncoder.""" + + def __init__(self, model_name: str = "cross-encoder/ms-marco-MiniLM-L-6-v2", batch_size: int = 16) -> None: + try: + from sentence_transformers import CrossEncoder + except ImportError as exc: # pragma: no cover - dependency missing + raise RuntimeError( + "sentence-transformers is required for the cross-encoder reranker" + ) from exc + + self.model = CrossEncoder(model_name) + self.batch_size = batch_size + + def rerank(self, query: str, results: Iterable[SearchResult]) -> Iterable[SearchResult]: + result_list = list(results) + if not result_list: + return result_list + pairs = [(query, item["text"]) for item in result_list] + scores = self.model.predict(pairs, batch_size=self.batch_size) + for item, score in zip(result_list, scores): + item["rerank_score"] = float(score) + return sorted(result_list, key=lambda item: item.get("rerank_score", 0.0), reverse=True) + + +register_reranker("cross-encoder", lambda **opts: CrossEncoderReranker(**opts)) diff --git a/raglite_sqlite/server.py b/raglite_sqlite/server.py new file mode 100644 index 0000000..661aa01 --- /dev/null +++ b/raglite_sqlite/server.py @@ -0,0 +1,130 @@ +from __future__ import annotations + +from threading import RLock +from typing import Any, Dict, Optional + +from fastapi import FastAPI, HTTPException +from pydantic import BaseModel, Field + +from .api import RagLite +from .embeddings.base import EmbeddingBackend +from .rerank import Reranker, get_reranker + + +def _default_backend() -> EmbeddingBackend: + from .embeddings.sentence_transformers_backend import SentenceTransformersBackend + + return SentenceTransformersBackend() + + +class QueryPayload(BaseModel): + query: str + k: int = 8 + hybrid_weight: float = Field(default=0.6, ge=0.0, le=1.0) + max_per_doc: int = 3 + filters: Optional[Dict[str, str]] = None + model_name: Optional[str] = None + with_snippets: bool = True + use_semantic: bool = True + reranker: Optional[str] = None + reranker_options: Optional[Dict[str, Any]] = None + + +class IndexPayload(BaseModel): + paths: list[str] + tags: Optional[str] = None + parser_opts: Optional[Dict[str, Any]] = None + chunker: str = "recursive" + chunk_size_tokens: int = 512 + chunk_overlap_tokens: int = 64 + model_name: Optional[str] = None + skip_unchanged: bool = True + recurse: bool = True + glob: Optional[str] = None + + +class DeletePayload(BaseModel): + doc_id: str + + +def create_app( + db_path: str, + *, + embedding_backend: EmbeddingBackend | None = None, + rerankers: dict[str, Reranker] | None = None, +) -> FastAPI: + """Create a FastAPI application exposing the RagLite API.""" + + rag = RagLite(db_path) + backend = embedding_backend or _default_backend() + reranker_cache = rerankers or {} + lock = RLock() + + app = FastAPI(title="RagLite SQLite", version="1.0") + + @app.get("/health") + def health() -> dict[str, str]: + return {"status": "ok"} + + @app.get("/stats") + def stats() -> dict[str, Any]: + with lock: + return rag.stats() + + @app.post("/query") + def query(payload: QueryPayload) -> dict[str, Any]: + with lock: + resolved_backend = backend if payload.use_semantic else None + reranker_instance: Reranker | None = None + if payload.reranker: + reranker_instance = reranker_cache.get(payload.reranker) + if reranker_instance is None: + try: + reranker_instance = get_reranker( + payload.reranker, **(payload.reranker_options or {}) + ) + except KeyError as exc: + raise HTTPException(status_code=400, detail=str(exc)) from exc + reranker_cache[payload.reranker] = reranker_instance + results = rag.search( + payload.query, + k=payload.k, + hybrid_weight=payload.hybrid_weight, + filters=payload.filters, + model_name=payload.model_name, + embedding_backend=resolved_backend, + max_per_doc=payload.max_per_doc, + with_snippets=payload.with_snippets, + reranker=reranker_instance, + ) + return {"results": results} + + @app.post("/index") + def index(payload: IndexPayload) -> dict[str, Any]: + with lock: + stats = rag.index( + payload.paths, + tags=payload.tags, + parser_opts=payload.parser_opts, + chunker=payload.chunker, + chunk_size_tokens=payload.chunk_size_tokens, + chunk_overlap_tokens=payload.chunk_overlap_tokens, + embedding_backend=backend, + model_name=payload.model_name, + skip_unchanged=payload.skip_unchanged, + recurse=payload.recurse, + glob=payload.glob, + ) + return stats + + @app.post("/delete") + def delete(payload: DeletePayload) -> dict[str, str]: + with lock: + rag.delete(payload.doc_id) + return {"status": "deleted", "doc_id": payload.doc_id} + + @app.on_event("shutdown") + def shutdown() -> None: + rag.close() + + return app diff --git a/raglite_sqlite/typing.py b/raglite_sqlite/typing.py index f46924e..84c7f94 100644 --- a/raglite_sqlite/typing.py +++ b/raglite_sqlite/typing.py @@ -29,6 +29,7 @@ class SearchResult(TypedDict, total=False): tags: Optional[str] bm25_score: float vector_score: float + rerank_score: float class Parser(Protocol): diff --git a/raglite_sqlite/utils.py b/raglite_sqlite/utils.py index 097666d..0dd02fd 100644 --- a/raglite_sqlite/utils.py +++ b/raglite_sqlite/utils.py @@ -63,6 +63,15 @@ def detect_mime(path: Path) -> str: ".pdf": "application/pdf", ".docx": "application/vnd.openxmlformats-officedocument.wordprocessingml.document", ".csv": "text/csv", + ".json": "application/json", + ".yml": "application/x-yaml", + ".yaml": "application/x-yaml", + ".pptx": "application/vnd.openxmlformats-officedocument.presentationml.presentation", + ".png": "image/png", + ".jpg": "image/jpeg", + ".jpeg": "image/jpeg", + ".tif": "image/tiff", + ".tiff": "image/tiff", }.get(ext, "application/octet-stream") diff --git a/tests/test_parsers.py b/tests/test_parsers.py new file mode 100644 index 0000000..dcb91fc --- /dev/null +++ b/tests/test_parsers.py @@ -0,0 +1,65 @@ +from __future__ import annotations + +from pathlib import Path +from zipfile import ZipFile + +import pytest + +from raglite_sqlite.parsers.image import ImageParser +from raglite_sqlite.parsers.json import JSONParser +from raglite_sqlite.parsers.pptx import PptxParser + + +def test_json_parser(tmp_path: Path) -> None: + data_path = tmp_path / "sample.json" + data_path.write_text("{" "\"title\": \"Hello\", \"items\": [1, 2]}", encoding="utf-8") + parser = JSONParser() + blocks = list(parser.parse(str(data_path))) + assert blocks + assert "Hello" in blocks[0]["text"] + assert "items" in blocks[0]["text"] + + +def test_pptx_parser(tmp_path: Path) -> None: + pptx_path = tmp_path / "sample.pptx" + slide_xml = """ + + + + + + Hello PPTX + + + + + + """ + with ZipFile(pptx_path, "w") as archive: + archive.writestr("ppt/slides/slide1.xml", slide_xml) + parser = PptxParser() + blocks = list(parser.parse(str(pptx_path))) + assert blocks + assert "Hello PPTX" in blocks[0]["text"] + + +def test_image_parser(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> None: + pytest.importorskip("PIL") + pytest.importorskip("pytesseract") + from PIL import Image + + image_path = tmp_path / "hello.png" + image = Image.new("RGB", (10, 10), color="white") + image.save(image_path) + + import pytesseract + + def fake_ocr(image, lang="eng", config=None): # type: ignore[no-untyped-def] + return "Hello OCR" + + monkeypatch.setattr(pytesseract, "image_to_string", fake_ocr) + parser = ImageParser() + blocks = list(parser.parse(str(image_path))) + assert blocks + assert blocks[0]["text"] == "Hello OCR" diff --git a/tests/test_rerank.py b/tests/test_rerank.py new file mode 100644 index 0000000..97b5eb3 --- /dev/null +++ b/tests/test_rerank.py @@ -0,0 +1,40 @@ +from __future__ import annotations + +from pathlib import Path + +from raglite_sqlite.rerank import register_reranker + + +class ReverseReranker: + def rerank(self, query: str, results): # type: ignore[override] + return list(reversed(list(results))) + + +def test_custom_reranker(temp_db: tuple[Path, object, object]) -> None: + db_path, rag, backend = temp_db + data_dir = Path(__file__).parent / "data" + rag.index([str(data_dir)], embedding_backend=backend) + + class PrefixReranker: + def rerank(self, query: str, results): # type: ignore[override] + ordered = sorted(results, key=lambda item: item["doc_id"], reverse=True) + for idx, item in enumerate(ordered): + item["rerank_score"] = float(len(ordered) - idx) + return ordered + + results = rag.search("sample", embedding_backend=backend, reranker=PrefixReranker()) + assert results + assert results[0]["rerank_score"] >= results[-1]["rerank_score"] + + +def test_registry_reranker(temp_db: tuple[Path, object, object]) -> None: + db_path, rag, backend = temp_db + data_dir = Path(__file__).parent / "data" + rag.index([str(data_dir)], embedding_backend=backend) + + register_reranker("reverse", lambda **_: ReverseReranker()) + baseline = rag.search("sample", embedding_backend=backend) + results = rag.search("sample", embedding_backend=backend, reranker="reverse") + assert results + if len(baseline) > 1: + assert results[0]["chunk_id"] == baseline[-1]["chunk_id"] diff --git a/tests/test_server.py b/tests/test_server.py new file mode 100644 index 0000000..fc52d9f --- /dev/null +++ b/tests/test_server.py @@ -0,0 +1,24 @@ +from __future__ import annotations + +from pathlib import Path + +import pytest + +pytest.importorskip("fastapi", reason="FastAPI is required for server tests") +from fastapi.testclient import TestClient + +from raglite_sqlite.server import create_app + + +def test_server_query_endpoint(temp_db: tuple[Path, object, object]) -> None: + db_path, rag, backend = temp_db + data_dir = Path(__file__).parent / "data" + rag.index([str(data_dir)], embedding_backend=backend) + + app = create_app(str(db_path), embedding_backend=backend) + client = TestClient(app) + response = client.post("/query", json={"query": "sample", "use_semantic": False}) + assert response.status_code == 200 + data = response.json() + assert "results" in data + assert data["results"], "Expected at least one result from the REST API"