diff --git a/.env.example b/.env.example new file mode 100644 index 0000000..ec03b22 --- /dev/null +++ b/.env.example @@ -0,0 +1,5 @@ +RAG_MODE=mock +EMBEDDING_PROVIDER=mock + +# This workshop repository intentionally requires no external API keys, +# secrets, or network access for standard local development. diff --git a/.github/CODEOWNERS b/.github/CODEOWNERS new file mode 100644 index 0000000..5311426 --- /dev/null +++ b/.github/CODEOWNERS @@ -0,0 +1,4 @@ +# TODO: Replace @ORG placeholders with the actual organization or team slug when available. +* @ORG/rag-eval-maintainers +/tests/ @ORG/rag-eval-maintainers +/.github/ @ORG/release-maintainers diff --git a/.github/ISSUE_TEMPLATE/bug.yml b/.github/ISSUE_TEMPLATE/bug.yml new file mode 100644 index 0000000..453e1c1 --- /dev/null +++ b/.github/ISSUE_TEMPLATE/bug.yml @@ -0,0 +1,19 @@ +name: Bug report +about: Report a bug in the workshop lab +title: "" +labels: [bug] +body: + - type: textarea + attributes: + label: Describe the bug + description: What happened and what was expected? + validations: + required: true + - type: textarea + attributes: + label: Reproduction steps + description: Include minimal commands or code to reproduce the issue. + - type: textarea + attributes: + label: Environment + description: Include Python version, OS, and relevant setup notes. diff --git a/.github/ISSUE_TEMPLATE/feature.yml b/.github/ISSUE_TEMPLATE/feature.yml new file mode 100644 index 0000000..39dfaec --- /dev/null +++ b/.github/ISSUE_TEMPLATE/feature.yml @@ -0,0 +1,19 @@ +name: Feature request +about: Suggest an improvement or new workshop feature +title: "" +labels: [enhancement] +body: + - type: textarea + attributes: + label: Problem + description: What problem does this feature solve? + validations: + required: true + - type: textarea + attributes: + label: Proposed solution + description: Describe the feature and any design constraints. + - type: textarea + attributes: + label: Scope + description: Note whether this is focused on retrieval, chunking, embeddings, MCP tooling, or docs. diff --git a/.github/pull_request_template.md b/.github/pull_request_template.md new file mode 100644 index 0000000..6df1e9e --- /dev/null +++ b/.github/pull_request_template.md @@ -0,0 +1,11 @@ +## Related issue + +## Summary + +## Tests + +## Mock/offline verification + +## Scope check + +## Secrets check diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml new file mode 100644 index 0000000..7484efc --- /dev/null +++ b/.github/workflows/ci.yml @@ -0,0 +1,32 @@ +name: CI + +on: + push: + branches: ["**"] + pull_request: + +permissions: + contents: read + +jobs: + test: + runs-on: ubuntu-latest + steps: + - name: Check out code + uses: actions/checkout@v4 + + - name: Set up Python 3.12 + uses: actions/setup-python@v5 + with: + python-version: "3.12" + + - name: Install dependencies + run: | + python -m pip install --upgrade pip + python -m pip install -r requirements.txt + + - name: Run Ruff + run: ruff check . + + - name: Run pytest + run: pytest -q diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..08935d1 --- /dev/null +++ b/.gitignore @@ -0,0 +1,22 @@ +**/__pycache__/ +*.py[cod] +*.so +*.egg-info/ +.venv/ +venv/ +env/ +.venv*/ +.pytest_cache/ +.ruff_cache/ +.coverage +htmlcov/ +.env +.env.local +.env.* +!.env.example +.DS_Store +build/ +dist/ +.idea/ +.vscode/ +*.log diff --git a/CODE_OF_CONDUCT.md b/CODE_OF_CONDUCT.md new file mode 100644 index 0000000..bdc058c --- /dev/null +++ b/CODE_OF_CONDUCT.md @@ -0,0 +1,22 @@ +# Contributor Covenant Code of Conduct + +## Our pledge + +We pledge to make participation in this project a harassment-free experience for everyone. + +## Our standards + +Examples of behavior that contributes to a positive environment include: + +- being respectful and welcoming +- assuming good intent +- being constructive in feedback +- focusing on the technical problem rather than the individual + +## Enforcement + +Project maintainers are expected to enforce this Code of Conduct fairly and consistently. + +## Contact + +If you need to report an issue, contact the maintainers through a private channel or the repository security policy. diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md new file mode 100644 index 0000000..984b010 --- /dev/null +++ b/CONTRIBUTING.md @@ -0,0 +1,70 @@ +# Contributing to Agentic RAG Lab + +Thank you for helping improve this workshop-focused repository. + +## Local setup + +- Use Python 3.12. +- Create a virtual environment: + +```bash +python3.12 -m venv .venv +source .venv/bin/activate +python -m pip install --upgrade pip +pip install -r requirements.txt +``` + +## Testing + +Run the local test suite before opening a PR: + +```bash +pytest -q +ruff check . +``` + +## Branching and commits + +- Use a feature branch such as `feat/add-chunker-optimizer` +- Keep commit messages clear and scoped +- Keep the repository offline-first and deterministic + +## Pull request expectations + +- Add or update tests for changed behavior +- Validate the offline/mock workflow +- Do not add paid APIs or cloud dependencies to core examples +- Avoid broad refactors not related to the change + +## Working on chunking + +When modifying chunking, check: + +- chunk size validation +- overlap validation +- deterministic ordering +- empty-document handling +- metadata preservation + +## Working on retrieval + +When modifying retrieval, check: + +- query validation +- top-k boundaries +- ordering by relevance +- empty-index handling +- deterministic output + +## Working on MCP integration + +Keep retrieval independent from the MCP transport layer. + +- validate JSON input +- return structured output +- do not couple the implementation to external services +- prefer small integration tests over broad mock-heavy tests + +## Offline-first policy + +This project should never depend on paid APIs or external services for core examples. If a new dependency is needed, prefer local, minimal packages and document the reason clearly. diff --git a/LICENSE b/LICENSE new file mode 100644 index 0000000..ebbbbd5 --- /dev/null +++ b/LICENSE @@ -0,0 +1,21 @@ +MIT License + +Copyright (c) 2026 Workshop contributors + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all +copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +SOFTWARE. diff --git a/README.md b/README.md index 6f4ac3b..07004d4 100644 --- a/README.md +++ b/README.md @@ -1 +1,154 @@ -# agentic-rag-lab +# Agentic RAG Lab + +A small, deterministic, offline-first playground for retrieval fundamentals. The project teaches the core pieces of a retrieval pipeline without relying on external APIs, paid services, or hosted runtimes. + +## Purpose + +This repository is intentionally focused on the foundations of retrieval-augmented generation: + +- deterministic document loading +- text chunking +- consistent mock embeddings +- FAISS-backed vector indexing +- top-k retrieval +- an MCP-compatible retrieval tool layer +- reproducible, local tests + +The goal is to make workshop participants comfortable with the mechanics behind retrieval before they move on to larger agent frameworks or hosted vector services. + +## Architecture + +Documents + ↓ +Loader + ↓ +Chunker + ↓ +Mock Embeddings + ↓ +FAISS + ↓ +Retriever + ↓ +MCP Retrieval Tool + +## Requirements + +- Python 3.12 +- Local package installation with pip +- No API keys or external services required + +## Setup + +Linux/macOS: + +```bash +python3.12 -m venv .venv +source .venv/bin/activate +python -m pip install --upgrade pip +pip install -r requirements.txt +python -m pip install -e . +``` + +Windows: + +```powershell +py -3.12 -m venv .venv +.venv\Scripts\activate +python -m pip install --upgrade pip +pip install -r requirements.txt +python -m pip install -e . +``` + +## Test + +```bash +pytest -q +``` + +## Run Example + +```bash +python examples/rag_demo.py +``` + +## No API Key Required + +This repository is offline-first and intentionally uses deterministic mock embeddings rather than paid external embedding providers. The examples and tests are designed to run entirely on local fixtures. No secret, credential, or hosted service is required for standard execution. + +## Architecture Explanation + +### Documents +The loader reads UTF-8 text files from a directory and converts them into a lightweight `Document` model. It sorts files deterministically, filters unsupported file types, and preserves metadata such as source and filename. + +### Loader +The loader is purposely small and explicit. It is responsible for turning raw files into consistent `Document` objects without performing any network access. + +### Chunker +The chunker splits a document into overlapping or contiguous text units. It validates chunk size and overlap to avoid invalid configurations and always produces non-empty chunks. + +### Mock Embeddings +The embeddings layer derives fixed-length vectors deterministically from string content using `hashlib.sha256`. This keeps retrieval reproducible across runs and avoids Python hash randomization. + +### FAISS +FAISS provides an in-memory vector index for local search. This project uses an `IndexFlatL2` index. Distances are lower-is-better, and the retriever converts them to a "score" using the negative distance so higher scores indicate stronger matches. + +### Retriever +The retriever combines a query embedding with a vector store and returns ordered retrieval results. Query strings must be non-empty; results are sorted by relevance and can be bounded by `top_k`. + +### MCP Retrieval Tool +The MCP integration is intentionally thin and keeps retrieval logic independent from the transport layer. The MCP handler validates input, invokes the retriever, and returns JSON-compatible structured results. + +## Why these design decisions? + +### Why mock embeddings? +Mock embeddings keep the workshop reproducible, deterministic, and fully offline. They teach the mechanics of retrieval without requiring any external model access. + +### Why FAISS? +FAISS is a widely used local vector indexing library and works well for a small educational lab. It keeps the focus on indexing and retrieval behavior without adding cloud dependencies. + +### Why deterministic tests? +The tests are designed to be local, stable, and machine-independent. That makes the repository suitable for workshop environments and CI without depending on randomness or external systems. + +### Why retrieval is independent from MCP? +Separation keeps the retrieval logic reusable and testable outside the MCP layer. The MCP tool simply adapts the same retrieve interface to structured JSON payloads. + +### Why the repository does not use LangChain or LlamaIndex? +Those frameworks are useful in larger systems, but they hide the mechanics that this workshop is meant to teach. This lab keeps the implementation explicit and understandable. + +### Why external APIs are optional rather than required? +The repository is designed to work entirely offline. This reduces setup friction for contributors and ensures that learner exercises are repeatable regardless of network access or account setup. + +## Contributor Guide + +See [CONTRIBUTING.md](CONTRIBUTING.md) for setup, testing, branch workflow, and contribution expectations. + +## Project Structure + +```text +agentic-rag-lab/ +├── src/ +│ └── rag/ +│ ├── __init__.py +│ ├── loaders.py +│ ├── chunkers.py +│ ├── embeddings.py +│ ├── vectorstore.py +│ ├── retriever.py +│ └── mcp_retrieval.py +├── tests/ +├── fixtures/ +├── schemas/ +├── examples/ +├── .github/ +├── .env.example +├── .gitignore +├── CONTRIBUTING.md +├── CODE_OF_CONDUCT.md +├── LICENSE +├── README.md +├── SECURITY.md +├── pyproject.toml +├── requirements.txt +└── .github/workflows/ci.yml +``` diff --git a/SECURITY.md b/SECURITY.md new file mode 100644 index 0000000..5b1a739 --- /dev/null +++ b/SECURITY.md @@ -0,0 +1,16 @@ +# Security Policy + +## Reporting security issues + +Please report security problems privately. Do not create public issues or pull requests for security vulnerabilities. + +## Expected practices + +- Never commit credentials, tokens, or API secrets +- Never put API keys in issues, PR descriptions, logs, or examples +- Keep workshop examples offline by default +- Prefer local fixtures and deterministic mock data for tests + +## Scope + +This repository is intended for workshop and educational use. It should stay offline-first and free of external service requirements unless a specific exercise explicitly calls for them. diff --git a/examples/rag_demo.py b/examples/rag_demo.py new file mode 100644 index 0000000..34d6965 --- /dev/null +++ b/examples/rag_demo.py @@ -0,0 +1,46 @@ +from __future__ import annotations + +import sys +from pathlib import Path + +REPO_ROOT = Path(__file__).resolve().parents[1] +SRC_DIR = REPO_ROOT / "src" +if str(SRC_DIR) not in sys.path: + sys.path.insert(0, str(SRC_DIR)) + +from rag.chunkers import chunk_document +from rag.embeddings import MockEmbeddingProvider +from rag.loaders import load_documents +from rag.retriever import Retriever +from rag.vectorstore import FaissVectorStore + + +def main() -> None: + fixtures_dir = Path(__file__).resolve().parents[1] / "fixtures" / "documents" + documents = load_documents(fixtures_dir) + chunks = [] + for document in documents: + chunks.extend(chunk_document(document, chunk_size=120, overlap=20)) + + embedding_provider = MockEmbeddingProvider() + embeddings = embedding_provider.embed([chunk.text for chunk in chunks]) + vector_store = FaissVectorStore() + vector_store.add(chunks, embeddings) + + retriever = Retriever( + embedding_provider=embedding_provider, + vector_store=vector_store, + chunks=chunks, + ) + + query = "MCP tools and retrieval basics" + results = retriever.retrieve(query, top_k=3) + print(f"Query: {query}") + print("Results:") + for index, result in enumerate(results, start=1): + print(f" {index}. {result.source} (score={result.score:.4f})") + print(f" {result.text[:120]}...") + + +if __name__ == "__main__": + main() diff --git a/fixtures/documents/mcp_basics.txt b/fixtures/documents/mcp_basics.txt new file mode 100644 index 0000000..474c9f3 --- /dev/null +++ b/fixtures/documents/mcp_basics.txt @@ -0,0 +1,3 @@ +MCP stands for Model Context Protocol. It helps clients connect to servers and invoke tools. The protocol defines a clean contract for machine-readable interactions between agents and local tools. + +A client sends a request. A server exposes tools, resources, and prompts. The protocol keeps the interaction explicit and structured so the agent can reason about what is available. diff --git a/fixtures/documents/mcp_tools.txt b/fixtures/documents/mcp_tools.txt new file mode 100644 index 0000000..50459bb --- /dev/null +++ b/fixtures/documents/mcp_tools.txt @@ -0,0 +1,3 @@ +Tool discovery is the process of learning which capabilities a server offers. The server describes each tool with metadata such as a name, description, and parameter schema. Clients can inspect those descriptions before invoking a tool. + +Tool invocation uses a structured payload. The client sends a query or task and the server returns a result that is easy for other systems to parse. This makes retrieval and direct function calls predictable. diff --git a/fixtures/documents/rag_basics.txt b/fixtures/documents/rag_basics.txt new file mode 100644 index 0000000..3ff08a8 --- /dev/null +++ b/fixtures/documents/rag_basics.txt @@ -0,0 +1,3 @@ +Retrieval-augmented generation combines a knowledge corpus with a retriever. A document loader reads source files. A chunker splits them into smaller segments. An embedding model converts each chunk into a vector representation. + +A vector store indexes those embeddings so a query can be compared against the corpus. Top-k retrieval selects the most relevant chunks. The result is then used as context for a downstream model or workflow. diff --git a/fixtures/mock_embeddings.json b/fixtures/mock_embeddings.json new file mode 100644 index 0000000..99a77a4 --- /dev/null +++ b/fixtures/mock_embeddings.json @@ -0,0 +1,5 @@ +{ + "mcp_basics.txt": [0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8], + "mcp_tools.txt": [0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9], + "rag_basics.txt": [0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9, 1.0] +} diff --git a/pyproject.toml b/pyproject.toml new file mode 100644 index 0000000..9545cb6 --- /dev/null +++ b/pyproject.toml @@ -0,0 +1,37 @@ +[build-system] +requires = ["setuptools>=68", "wheel"] +build-backend = "setuptools.build_meta" + +[project] +name = "agentic-rag-lab" +version = "0.1.0" +description = "Offline-first educational retrieval lab for deterministic RAG primitives" +readme = "README.md" +requires-python = ">=3.12,<3.13" +license = { text = "MIT" } +authors = [{ name = "Workshop contributors" }] +dependencies = [ + "numpy>=1.26", + "faiss-cpu>=1.8.0", +] + +[project.optional-dependencies] +dev = [ + "pytest>=8.0", + "jsonschema>=4.0", + "ruff>=0.6", +] + +[tool.pytest.ini_options] +pythonpath = ["src"] +testpaths = ["tests"] + +[tool.setuptools] +package-dir = {"" = "src"} + +[tool.setuptools.packages.find] +where = ["src"] + +[tool.ruff] +line-length = 120 +target-version = "py312" diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..96a2128 --- /dev/null +++ b/requirements.txt @@ -0,0 +1,5 @@ +numpy>=1.26 +faiss-cpu>=1.8.0 +jsonschema>=4.0 +pytest>=8.0 +ruff>=0.6 diff --git a/schemas/retrieval.schema.json b/schemas/retrieval.schema.json new file mode 100644 index 0000000..0da467a --- /dev/null +++ b/schemas/retrieval.schema.json @@ -0,0 +1,27 @@ +{ + "$schema": "https://json-schema.org/draft/2020-12/schema", + "title": "RetrievalResults", + "type": "object", + "required": ["results"], + "properties": { + "results": { + "type": "array", + "items": { + "type": "object", + "required": ["chunk_id", "text", "source", "score"], + "properties": { + "chunk_id": {"type": "string"}, + "text": {"type": "string"}, + "source": {"type": "string"}, + "score": {"type": "number"}, + "metadata": { + "type": "object", + "additionalProperties": true + } + }, + "additionalProperties": true + } + } + }, + "additionalProperties": false +} diff --git a/src/rag/__init__.py b/src/rag/__init__.py new file mode 100644 index 0000000..3ee240a --- /dev/null +++ b/src/rag/__init__.py @@ -0,0 +1,19 @@ +"""Core retrieval primitives for the Agentic RAG Lab.""" + +from rag.chunkers import Chunk, chunk_document +from rag.embeddings import EmbeddingProvider, MockEmbeddingProvider +from rag.loaders import Document, load_documents +from rag.retriever import RetrievalResult, Retriever +from rag.vectorstore import FaissVectorStore + +__all__ = [ + "Chunk", + "Document", + "EmbeddingProvider", + "FaissVectorStore", + "MockEmbeddingProvider", + "RetrievalResult", + "Retriever", + "chunk_document", + "load_documents", +] diff --git a/src/rag/chunkers.py b/src/rag/chunkers.py new file mode 100644 index 0000000..d5c9d61 --- /dev/null +++ b/src/rag/chunkers.py @@ -0,0 +1,64 @@ +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import Any + +from rag.loaders import Document + + +@dataclass(frozen=True) +class Chunk: + chunk_id: str + text: str + document_id: str + metadata: dict[str, Any] = field(default_factory=dict) + + +def chunk_document(document: Document, chunk_size: int, overlap: int) -> list[Chunk]: + """Split a document into deterministic chunks with optional overlap.""" + if not isinstance(chunk_size, int) or chunk_size <= 0: + raise TypeError("chunk_size must be a positive integer") + if not isinstance(overlap, int): + raise TypeError("overlap must be an integer") + if overlap < 0: + raise ValueError("overlap must be >= 0") + if overlap >= chunk_size: + raise ValueError("overlap must be smaller than chunk_size") + + if not document.text: + return [] + + if len(document.text) <= chunk_size: + return [ + Chunk( + chunk_id=f"{document.document_id}::0", + text=document.text, + document_id=document.document_id, + metadata=dict(document.metadata), + ) + ] + + step = chunk_size - overlap + chunks: list[Chunk] = [] + start = 0 + + while start < len(document.text): + end = min(start + chunk_size, len(document.text)) + chunk_text = document.text[start:end] + if not chunk_text: + break + + chunks.append( + Chunk( + chunk_id=f"{document.document_id}::{len(chunks)}", + text=chunk_text, + document_id=document.document_id, + metadata=dict(document.metadata), + ) + ) + + if end >= len(document.text): + break + start += step + + return chunks diff --git a/src/rag/embeddings.py b/src/rag/embeddings.py new file mode 100644 index 0000000..073ccee --- /dev/null +++ b/src/rag/embeddings.py @@ -0,0 +1,41 @@ +from __future__ import annotations + +import hashlib +from collections.abc import Sequence +from typing import Protocol + + +class EmbeddingProvider(Protocol): + def embed(self, texts: Sequence[str]) -> list[list[float]]: + """Return one deterministic embedding vector per input text.""" + + +class MockEmbeddingProvider: + """Deterministic mock embeddings for workshop retrieval exercises. + + This is intentionally not a semantically strong production embedding model. + The implementation uses a stable SHA-256 hash to keep results reproducible and + independent of Python's randomized hash seed. + """ + + def __init__(self, dimension: int = 8) -> None: + if not isinstance(dimension, int) or dimension <= 0: + raise ValueError("dimension must be a positive integer") + self.dimension = dimension + + def embed(self, texts: Sequence[str]) -> list[list[float]]: + if texts == []: + return [] + + vectors: list[list[float]] = [] + for text in texts: + if not isinstance(text, str): + raise TypeError("Each text value must be a string") + digest = hashlib.sha256(text.encode("utf-8")).digest() + vector: list[float] = [] + for index in range(self.dimension): + chunk = digest[(index * 4) : ((index + 1) * 4)] + value = int.from_bytes(chunk, byteorder="big", signed=False) / (2**32 - 1) + vector.append((value * 2.0) - 1.0) + vectors.append(vector) + return vectors diff --git a/src/rag/loaders.py b/src/rag/loaders.py new file mode 100644 index 0000000..021cc53 --- /dev/null +++ b/src/rag/loaders.py @@ -0,0 +1,51 @@ +from __future__ import annotations + +from dataclasses import dataclass, field +from pathlib import Path +from typing import Any + +SUPPORTED_TEXT_EXTENSIONS = {"txt", "md", "csv", "json", "yaml", "yml", "log", "rst"} + + +@dataclass(frozen=True) +class Document: + document_id: str + text: str + metadata: dict[str, Any] = field(default_factory=dict) + + @property + def source(self) -> str: + return str(self.metadata.get("source", "")) + + @property + def filename(self) -> str: + return str(self.metadata.get("filename", "")) + + +def load_documents(directory: str | Path) -> list[Document]: + """Load UTF-8 text documents from a directory in deterministic order.""" + path = Path(directory) + if not path.exists() or not path.is_dir(): + 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 + suffix = file_path.suffix.lower().lstrip(".") + if suffix not in SUPPORTED_TEXT_EXTENSIONS: + continue + + content = file_path.read_text(encoding="utf-8") + documents.append( + Document( + document_id=file_path.stem, + text=content, + metadata={ + "source": str(file_path), + "filename": file_path.name, + }, + ) + ) + + return documents diff --git a/src/rag/mcp_retrieval.py b/src/rag/mcp_retrieval.py new file mode 100644 index 0000000..72c5a78 --- /dev/null +++ b/src/rag/mcp_retrieval.py @@ -0,0 +1,88 @@ +from __future__ import annotations + +from typing import Any + +from rag.embeddings import MockEmbeddingProvider +from rag.retriever import Retriever +from rag.vectorstore import FaissVectorStore + + +def build_mcp_tool( + vector_store: FaissVectorStore, + chunks: list[Any], + embedding_provider: Any | None = None, +) -> dict[str, Any]: + """Create a lightweight MCP-style retrieval tool interface.""" + provider = embedding_provider or MockEmbeddingProvider() + + def handler(payload: dict[str, Any]) -> dict[str, Any]: + if not isinstance(payload, dict): + raise TypeError("payload must be a dictionary") + + query = str(payload.get("query", "")).strip() + if not query: + raise ValueError("query must be a non-empty string") + + top_k = int(payload.get("top_k", 3)) + retriever = Retriever(embedding_provider=provider, vector_store=vector_store, chunks=chunks) + results = retriever.retrieve(query, top_k=top_k) + + return { + "results": [ + { + "chunk_id": result.chunk_id, + "text": result.text, + "source": result.source, + "score": result.score, + "metadata": result.metadata, + } + for result in results + ] + } + + return { + "name": "retrieval_tool", + "description": "Retrieve the most relevant chunks for a user query.", + "input_schema": { + "type": "object", + "required": ["query"], + "properties": { + "query": {"type": "string", "description": "Search query for local retrieval"}, + "top_k": {"type": "integer", "minimum": 1, "default": 3}, + }, + "additionalProperties": True, + }, + "handler": handler, + } + + +def execute_retrieval_tool(payload: dict[str, Any], retriever: Retriever | None = None) -> dict[str, Any]: + """Convenience wrapper for direct tool invocation.""" + if not isinstance(payload, dict): + raise TypeError("payload must be a dictionary") + + if retriever is None: + retriever = Retriever( + embedding_provider=MockEmbeddingProvider(), + vector_store=FaissVectorStore(), + chunks=[], + ) + + query = str(payload.get("query", "")).strip() + if not query: + raise ValueError("query must be a non-empty string") + + top_k = int(payload.get("top_k", 3)) + results = retriever.retrieve(query, top_k=top_k) + return { + "results": [ + { + "chunk_id": result.chunk_id, + "text": result.text, + "source": result.source, + "score": result.score, + "metadata": result.metadata, + } + for result in results + ] + } diff --git a/src/rag/retriever.py b/src/rag/retriever.py new file mode 100644 index 0000000..2873dbd --- /dev/null +++ b/src/rag/retriever.py @@ -0,0 +1,60 @@ +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any + +from rag.embeddings import EmbeddingProvider, MockEmbeddingProvider +from rag.vectorstore import FaissVectorStore + + +@dataclass(frozen=True) +class RetrievalResult: + chunk_id: str + text: str + source: str + score: float | None + metadata: dict[str, Any] + + +class Retriever: + def __init__( + self, + embedding_provider: EmbeddingProvider, + vector_store: FaissVectorStore, + chunks: list[Any], + ) -> None: + self.embedding_provider = embedding_provider + self.vector_store = vector_store + self.chunks = list(chunks) + + def retrieve(self, query: str, top_k: int = 3) -> list[RetrievalResult]: + if not isinstance(query, str) or not query.strip(): + raise ValueError("query must be a non-empty string") + + normalized_query = query.strip() + if not self.chunks or self.vector_store.index is None or self.vector_store.index.ntotal == 0: + return [] + + query_embedding = self.embedding_provider.embed([normalized_query])[0] + results = self.vector_store.search(query_embedding, top_k=top_k) + + return [ + RetrievalResult( + chunk_id=result["chunk_id"], + text=result["text"], + source=result["source"], + score=result.get("score"), + metadata=result.get("metadata", {}), + ) + for result in results + ] + + +def make_default_retriever(chunks: list[Any] | None = None) -> Retriever: + """Convenience constructor used by examples and local tooling.""" + store = FaissVectorStore() + documents = chunks or [] + if documents: + embeddings = MockEmbeddingProvider().embed([chunk.text for chunk in documents]) + store.add(documents, embeddings) + return Retriever(embedding_provider=MockEmbeddingProvider(), vector_store=store, chunks=documents) diff --git a/src/rag/vectorstore.py b/src/rag/vectorstore.py new file mode 100644 index 0000000..2cd1e2a --- /dev/null +++ b/src/rag/vectorstore.py @@ -0,0 +1,136 @@ +from __future__ import annotations + +import json +from pathlib import Path +from typing import Any + +import faiss +import numpy as np + +from rag.chunkers import Chunk + + +class FaissVectorStore: + """Simple vector store built on FAISS for local retrieval experiments.""" + + def __init__(self) -> None: + self._index: faiss.Index | None = None + self._chunks: list[Chunk] = [] + self._dimension: int | None = None + + @property + def index(self) -> faiss.Index | None: + return self._index + + @property + def dimension(self) -> int | None: + return self._dimension + + def add(self, chunks: list[Chunk], embeddings: list[list[float]]) -> None: + if not chunks and not embeddings: + return + if len(chunks) != len(embeddings): + raise ValueError("chunks and embeddings must have the same length") + if len(chunks) == 0: + raise ValueError("chunks cannot be empty when embeddings are provided") + + first_embedding = embeddings[0] + if not isinstance(first_embedding, list): + raise TypeError("Each embedding must be a list of floats") + + dimension = len(first_embedding) + if dimension <= 0: + raise ValueError("embedding dimension must be positive") + + if self._index is None: + self._index = faiss.IndexFlatL2(dimension) + self._dimension = dimension + elif self._dimension != dimension: + raise ValueError(f"embedding dimensionality mismatch: expected {self._dimension}, got {dimension}") + + array = np.asarray(embeddings, dtype=np.float32) + if array.ndim != 2 or array.shape[1] != dimension: + raise ValueError("embeddings must be a 2D array with a consistent dimension") + + self._index.add(array) + self._chunks.extend(chunks) + + def search(self, query_embedding: list[float], top_k: int) -> list[dict[str, Any]]: + if not isinstance(top_k, int): + raise TypeError("top_k must be an integer") + if top_k <= 0: + raise ValueError("top_k must be greater than zero") + if self._index is None or self._index.ntotal == 0: + return [] + + query = np.asarray(query_embedding, dtype=np.float32).reshape(1, -1) + if self._dimension is None: + raise ValueError("vector store has no indexed dimension") + if query.shape[1] != self._dimension: + raise ValueError(f"query dimension mismatch: expected {self._dimension}, got {query.shape[1]}") + + limit = min(top_k, self._index.ntotal) + distances, indices = self._index.search(query, limit) + results: list[dict[str, Any]] = [] + + for distance_value, chunk_index in zip(distances[0], indices[0], strict=True): + if chunk_index < 0 or chunk_index >= len(self._chunks): + continue + chunk = self._chunks[int(chunk_index)] + distance = float(distance_value) + results.append( + { + "chunk_id": chunk.chunk_id, + "text": chunk.text, + "source": chunk.metadata.get("source", ""), + "score": -distance, + "distance": distance, + "metadata": dict(chunk.metadata), + } + ) + + return results + + def save(self, path: str | Path) -> None: + if self._index is None: + raise ValueError("cannot save an empty vector store") + + destination = Path(path) + destination.parent.mkdir(parents=True, exist_ok=True) + faiss.write_index(self._index, str(destination)) + + metadata_payload = { + "chunks": [ + { + "chunk_id": chunk.chunk_id, + "text": chunk.text, + "document_id": chunk.document_id, + "metadata": dict(chunk.metadata), + } + for chunk in self._chunks + ] + } + meta_path = destination.with_suffix(destination.suffix + ".meta.json") + meta_path.write_text(json.dumps(metadata_payload, ensure_ascii=False), encoding="utf-8") + + def load(self, path: str | Path) -> None: + source = Path(path) + if not source.exists(): + raise FileNotFoundError(f"Index file not found: {source}") + + self._index = faiss.read_index(str(source)) + self._dimension = self._index.d + self._chunks = [] + + meta_path = source.with_suffix(source.suffix + ".meta.json") + if meta_path.exists(): + payload = json.loads(meta_path.read_text(encoding="utf-8")) + for item in payload.get("chunks", []): + self._chunks.append( + Chunk( + chunk_id=item["chunk_id"], + text=item["text"], + document_id=item["document_id"], + metadata=item.get("metadata", {}), + ) + ) diff --git a/tests/conftest.py b/tests/conftest.py new file mode 100644 index 0000000..9a4408e --- /dev/null +++ b/tests/conftest.py @@ -0,0 +1,49 @@ +from __future__ import annotations + +from pathlib import Path + +import pytest + +from rag.chunkers import Chunk, chunk_document +from rag.loaders import Document, load_documents + + +@pytest.fixture +def temp_document_dir(tmp_path: Path) -> Path: + documents_dir = tmp_path / "documents" + documents_dir.mkdir() + (documents_dir / "alpha.txt").write_text("alpha beta gamma\n", encoding="utf-8") + (documents_dir / "beta.txt").write_text("delta epsilon zeta\n", encoding="utf-8") + return documents_dir + + +@pytest.fixture +def sample_documents(temp_document_dir: Path) -> list[Document]: + return load_documents(temp_document_dir) + + +@pytest.fixture +def sample_chunks(sample_documents: list[Document]) -> list[Chunk]: + chunks: list[Chunk] = [] + for document in sample_documents: + chunks.extend(chunk_document(document, chunk_size=10, overlap=2)) + return chunks + + +@pytest.fixture +def sample_embeddings() -> list[list[float]]: + return [ + [0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8], + [0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9], + [0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9, 1.0], + ] + + +@pytest.fixture +def temp_faiss_index(tmp_path: Path): + return tmp_path / "store.index" + + +@pytest.fixture +def schema_path() -> Path: + return Path(__file__).resolve().parents[1] / "schemas" / "retrieval.schema.json" diff --git a/tests/test_chunkers.py b/tests/test_chunkers.py new file mode 100644 index 0000000..d7fefe9 --- /dev/null +++ b/tests/test_chunkers.py @@ -0,0 +1,75 @@ +from __future__ import annotations + +import pytest + +from rag.chunkers import chunk_document +from rag.loaders import Document + + +@pytest.fixture +def sample_document() -> Document: + return Document( + document_id="doc-1", + text="abcdefghijklmno", + metadata={"source": "fixture", "filename": "sample.txt"}, + ) + + +def test_chunk_document_empty_text() -> None: + document = Document(document_id="doc-empty", text="", metadata={"source": "fixture", "filename": "empty.txt"}) + + assert chunk_document(document, chunk_size=10, overlap=2) == [] + + +def test_chunk_document_shorter_than_chunk_size(sample_document: Document) -> None: + chunks = chunk_document(sample_document, chunk_size=20, overlap=2) + + assert len(chunks) == 1 + assert chunks[0].text == sample_document.text + + +def test_chunk_document_exact_boundary(sample_document: Document) -> None: + chunks = chunk_document(sample_document, chunk_size=15, overlap=0) + + assert len(chunks) == 1 + assert [chunk.text for chunk in chunks] == [sample_document.text] + + +def test_chunk_document_with_overlap(sample_document: Document) -> None: + chunks = chunk_document(sample_document, chunk_size=10, overlap=2) + + assert len(chunks) == 2 + assert [chunk.text for chunk in chunks] == [ + "abcdefghij", + "ijklmno", + ] + + +def test_chunk_document_invalid_chunk_size(sample_document: Document) -> None: + with pytest.raises((TypeError, ValueError)): + chunk_document(sample_document, chunk_size=0, overlap=0) + + with pytest.raises((TypeError, ValueError)): + chunk_document(sample_document, chunk_size=-1, overlap=0) + + +def test_chunk_document_invalid_overlap(sample_document: Document) -> None: + with pytest.raises((TypeError, ValueError)): + chunk_document(sample_document, chunk_size=10, overlap=-1) + + with pytest.raises((TypeError, ValueError)): + chunk_document(sample_document, chunk_size=10, overlap=10) + + +def test_chunk_document_metadata_preserved(sample_document: Document) -> None: + chunks = chunk_document(sample_document, chunk_size=10, overlap=2) + + assert chunks[0].metadata["source"] == "fixture" + assert chunks[0].metadata["filename"] == "sample.txt" + + +def test_chunk_document_deterministic_repeated_execution(sample_document: Document) -> None: + first = chunk_document(sample_document, chunk_size=10, overlap=2) + second = chunk_document(sample_document, chunk_size=10, overlap=2) + + assert [chunk.text for chunk in first] == [chunk.text for chunk in second] diff --git a/tests/test_embeddings.py b/tests/test_embeddings.py new file mode 100644 index 0000000..79d6043 --- /dev/null +++ b/tests/test_embeddings.py @@ -0,0 +1,42 @@ +from __future__ import annotations + +import numpy as np + +from rag.embeddings import MockEmbeddingProvider + + +def test_mock_embedding_provider_empty_input() -> None: + provider = MockEmbeddingProvider() + + assert provider.embed([]) == [] + + +def test_mock_embedding_provider_single_text() -> None: + provider = MockEmbeddingProvider() + embedding = provider.embed(["hello world"])[0] + + assert len(embedding) == 8 + assert all(isinstance(value, float) for value in embedding) + + +def test_mock_embedding_provider_multiple_texts() -> None: + provider = MockEmbeddingProvider() + embeddings = provider.embed(["alpha", "beta", "gamma"]) + + assert len(embeddings) == 3 + assert all(len(vector) == 8 for vector in embeddings) + + +def test_mock_embedding_provider_is_deterministic() -> None: + provider = MockEmbeddingProvider() + + assert provider.embed(["same text"]) == provider.embed(["same text"]) + assert provider.embed(["alpha", "beta"]) == provider.embed(["alpha", "beta"]) + + +def test_mock_embedding_provider_matches_numpy_float32_shape() -> None: + provider = MockEmbeddingProvider() + embedding = provider.embed(["demo"])[0] + array = np.asarray(embedding, dtype=np.float32) + + assert array.shape == (8,) diff --git a/tests/test_loaders.py b/tests/test_loaders.py new file mode 100644 index 0000000..edcf742 --- /dev/null +++ b/tests/test_loaders.py @@ -0,0 +1,42 @@ +from __future__ import annotations + +from pathlib import Path + +import pytest + +from rag.loaders import load_documents + + +def test_load_documents_valid_directory(temp_document_dir: Path) -> None: + documents = load_documents(temp_document_dir) + + assert len(documents) == 2 + assert [doc.filename for doc in documents] == ["alpha.txt", "beta.txt"] + assert documents[0].text == "alpha beta gamma\n" + + +def test_load_documents_empty_directory(tmp_path: Path) -> None: + empty_dir = tmp_path / "empty" + empty_dir.mkdir() + + assert load_documents(empty_dir) == [] + + +def test_load_documents_missing_directory(tmp_path: Path) -> None: + missing_dir = tmp_path / "missing" + + with pytest.raises(FileNotFoundError): + load_documents(missing_dir) + + +def test_load_documents_utf8_and_unsupported_extensions(tmp_path: Path) -> None: + doc_dir = tmp_path / "utf8" + doc_dir.mkdir() + (doc_dir / "hello.txt").write_text("héllo\n", encoding="utf-8") + (doc_dir / "skip.bin").write_bytes(b"\x00\x01\x02") + (doc_dir / "notes.md").write_text("markdown\n", encoding="utf-8") + + documents = load_documents(doc_dir) + + assert [doc.filename for doc in documents] == ["hello.txt", "notes.md"] + assert documents[0].text == "héllo\n" diff --git a/tests/test_mcp_retrieval.py b/tests/test_mcp_retrieval.py new file mode 100644 index 0000000..018ba91 --- /dev/null +++ b/tests/test_mcp_retrieval.py @@ -0,0 +1,106 @@ +from __future__ import annotations + +import json +from pathlib import Path + +import jsonschema +import pytest + +from rag.chunkers import Chunk +from rag.embeddings import MockEmbeddingProvider +from rag.mcp_retrieval import build_mcp_tool, execute_retrieval_tool +from rag.vectorstore import FaissVectorStore + + +@pytest.fixture +def retrieval_fixture() -> tuple[list[Chunk], FaissVectorStore]: + chunks = [ + Chunk(chunk_id="c1", text="MCP tools are defined with schemas.", document_id="doc-1", metadata={"source": "mcp.txt", "filename": "mcp.txt"}), + Chunk(chunk_id="c2", text="Retrieval works with vector similarity.", document_id="doc-2", metadata={"source": "rag.txt", "filename": "rag.txt"}), + ] + store = FaissVectorStore() + store.add(chunks, MockEmbeddingProvider().embed([chunk.text for chunk in chunks])) + return chunks, store + + +def test_mcp_tool_valid_input(retrieval_fixture: tuple[list[Chunk], FaissVectorStore]) -> None: + chunks, store = retrieval_fixture + tool = build_mcp_tool(store, chunks) + + payload = {"query": "MCP tools", "top_k": 3} + response = tool["handler"](payload) + + assert response["results"] + assert isinstance(response["results"][0]["chunk_id"], str) + + +def test_mcp_tool_default_top_k(retrieval_fixture: tuple[list[Chunk], FaissVectorStore]) -> None: + _, store = retrieval_fixture + tool = build_mcp_tool(store, [ + Chunk(chunk_id="c1", text="one", document_id="d1", metadata={"source": "one.txt", "filename": "one.txt"}), + Chunk(chunk_id="c2", text="two", document_id="d2", metadata={"source": "two.txt", "filename": "two.txt"}), + ]) + response = tool["handler"]({"query": "two"}) + + assert len(response["results"]) <= 3 + + +def test_mcp_tool_custom_top_k(retrieval_fixture: tuple[list[Chunk], FaissVectorStore]) -> None: + _, store = retrieval_fixture + tool = build_mcp_tool(store, [ + Chunk(chunk_id="c1", text="one", document_id="d1", metadata={"source": "one.txt", "filename": "one.txt"}), + Chunk(chunk_id="c2", text="two", document_id="d2", metadata={"source": "two.txt", "filename": "two.txt"}), + ]) + response = tool["handler"]({"query": "two", "top_k": 1}) + + assert len(response["results"]) == 1 + + +def test_mcp_tool_empty_query(retrieval_fixture: tuple[list[Chunk], FaissVectorStore]) -> None: + _, store = retrieval_fixture + tool = build_mcp_tool(store, [ + Chunk(chunk_id="c1", text="one", document_id="d1", metadata={"source": "one.txt", "filename": "one.txt"}), + ]) + + with pytest.raises(ValueError): + tool["handler"]({"query": " "}) + + +def test_mcp_tool_invalid_top_k(retrieval_fixture: tuple[list[Chunk], FaissVectorStore]) -> None: + _, store = retrieval_fixture + tool = build_mcp_tool(store, [ + Chunk(chunk_id="c1", text="one", document_id="d1", metadata={"source": "one.txt", "filename": "one.txt"}), + ]) + + with pytest.raises(ValueError): + tool["handler"]({"query": "one", "top_k": 0}) + + +def test_mcp_tool_structured_output(retrieval_fixture: tuple[list[Chunk], FaissVectorStore]) -> None: + _, store = retrieval_fixture + tool = build_mcp_tool(store, [ + Chunk(chunk_id="c1", text="one", document_id="d1", metadata={"source": "one.txt", "filename": "one.txt"}), + Chunk(chunk_id="c2", text="two", document_id="d2", metadata={"source": "two.txt", "filename": "two.txt"}), + ]) + + response = tool["handler"]({"query": "one", "top_k": 2}) + + assert set(response.keys()) == {"results"} + assert set(response["results"][0].keys()) >= {"chunk_id", "text", "source", "score"} + + +def test_mcp_tool_schema_compatibility(retrieval_fixture: tuple[list[Chunk], FaissVectorStore], schema_path: Path) -> None: + _, store = retrieval_fixture + tool = build_mcp_tool(store, [ + Chunk(chunk_id="c1", text="one", document_id="d1", metadata={"source": "one.txt", "filename": "one.txt"}), + ]) + + response = tool["handler"]({"query": "one"}) + schema = json.loads(schema_path.read_text(encoding="utf-8")) + jsonschema.validate(instance=response, schema=schema) + + +def test_execute_retrieval_tool() -> None: + response = execute_retrieval_tool({"query": "test"}, retriever=None) + assert isinstance(response, dict) + assert "results" in response diff --git a/tests/test_retriever.py b/tests/test_retriever.py new file mode 100644 index 0000000..d11d48b --- /dev/null +++ b/tests/test_retriever.py @@ -0,0 +1,58 @@ +from __future__ import annotations + +import pytest + +from rag.chunkers import Chunk +from rag.embeddings import MockEmbeddingProvider +from rag.retriever import Retriever +from rag.vectorstore import FaissVectorStore + + +@pytest.fixture +def retriever() -> Retriever: + chunks = [ + Chunk(chunk_id="c1", text="MCP tools help clients call servers.", document_id="doc-1", metadata={"source": "mcp.txt", "filename": "mcp.txt"}), + Chunk(chunk_id="c2", text="Vector search uses embeddings to rank chunks.", document_id="doc-2", metadata={"source": "rag.txt", "filename": "rag.txt"}), + ] + store = FaissVectorStore() + store.add(chunks, MockEmbeddingProvider().embed([chunk.text for chunk in chunks])) + return Retriever(embedding_provider=MockEmbeddingProvider(), vector_store=store, chunks=chunks) + + +def test_retriever_basic_query(retriever: Retriever) -> None: + results = retriever.retrieve("MCP tools") + + assert len(results) >= 1 + assert results[0].chunk_id in {"c1", "c2"} + assert results[0].source in {"mcp.txt", "rag.txt"} + + +def test_retriever_top_k(retriever: Retriever) -> None: + results = retriever.retrieve("MCP tools", top_k=1) + + assert len(results) == 1 + + +def test_retriever_rejects_empty_query(retriever: Retriever) -> None: + with pytest.raises(ValueError): + retriever.retrieve(" ") + + +def test_retriever_deterministic(retriever: Retriever) -> None: + first = retriever.retrieve("MCP tools", top_k=2) + second = retriever.retrieve("MCP tools", top_k=2) + + assert [result.chunk_id for result in first] == [result.chunk_id for result in second] + + +def test_retriever_handles_no_indexed_documents() -> None: + retriever = Retriever(embedding_provider=MockEmbeddingProvider(), vector_store=FaissVectorStore(), chunks=[]) + + assert retriever.retrieve("what is retrieval") == [] + + +def test_retriever_metadata_in_results(retriever: Retriever) -> None: + result = retriever.retrieve("MCP tools", top_k=1)[0] + + assert result.metadata["filename"] in {"mcp.txt", "rag.txt"} + assert result.score is not None diff --git a/tests/test_schema.py b/tests/test_schema.py new file mode 100644 index 0000000..8cdf752 --- /dev/null +++ b/tests/test_schema.py @@ -0,0 +1,63 @@ +from __future__ import annotations + +import json + +import jsonschema +import pytest + + +@pytest.fixture +def schema_path() -> str: + return "schemas/retrieval.schema.json" + + +def test_retrieval_schema_accepts_valid_response() -> None: + with open("schemas/retrieval.schema.json", "r", encoding="utf-8") as schema_file: + schema = json.load(schema_file) + payload = { + "results": [ + { + "chunk_id": "c1", + "text": "hello", + "source": "demo.txt", + "score": 0.75, + } + ] + } + + jsonschema.validate(instance=payload, schema=schema) + + +def test_retrieval_schema_rejects_missing_results() -> None: + with open("schemas/retrieval.schema.json", "r", encoding="utf-8") as schema_file: + schema = json.load(schema_file) + + with pytest.raises(jsonschema.exceptions.ValidationError): + jsonschema.validate(instance={}, schema=schema) + + +def test_retrieval_schema_rejects_missing_chunk_id() -> None: + with open("schemas/retrieval.schema.json", "r", encoding="utf-8") as schema_file: + schema = json.load(schema_file) + payload = {"results": [{"text": "hello", "source": "demo.txt", "score": 0.75}]} + + with pytest.raises(jsonschema.exceptions.ValidationError): + jsonschema.validate(instance=payload, schema=schema) + + +def test_retrieval_schema_rejects_missing_source() -> None: + with open("schemas/retrieval.schema.json", "r", encoding="utf-8") as schema_file: + schema = json.load(schema_file) + payload = {"results": [{"chunk_id": "c1", "text": "hello", "score": 0.75}]} + + with pytest.raises(jsonschema.exceptions.ValidationError): + jsonschema.validate(instance=payload, schema=schema) + + +def test_retrieval_schema_rejects_invalid_score_type() -> None: + with open("schemas/retrieval.schema.json", "r", encoding="utf-8") as schema_file: + schema = json.load(schema_file) + payload = {"results": [{"chunk_id": "c1", "text": "hello", "source": "demo.txt", "score": "high"}]} + + with pytest.raises(jsonschema.exceptions.ValidationError): + jsonschema.validate(instance=payload, schema=schema) diff --git a/tests/test_vectorstore.py b/tests/test_vectorstore.py new file mode 100644 index 0000000..293f988 --- /dev/null +++ b/tests/test_vectorstore.py @@ -0,0 +1,73 @@ +from __future__ import annotations + +from pathlib import Path + +import pytest + +from rag.chunkers import Chunk +from rag.vectorstore import FaissVectorStore + + +@pytest.fixture +def chunks() -> list[Chunk]: + return [ + Chunk(chunk_id="chunk-1", text="one", document_id="doc-1", metadata={"source": "a.txt", "filename": "a.txt"}), + Chunk(chunk_id="chunk-2", text="two", document_id="doc-1", metadata={"source": "b.txt", "filename": "b.txt"}), + Chunk(chunk_id="chunk-3", text="three", document_id="doc-2", metadata={"source": "c.txt", "filename": "c.txt"}), + ] + + +def test_faiss_vector_store_add_and_query(chunks: list[Chunk], sample_embeddings: list[list[float]]) -> None: + store = FaissVectorStore() + store.add(chunks, sample_embeddings) + + results = store.search(sample_embeddings[0], top_k=2) + + assert len(results) == 2 + assert results[0]["chunk_id"] in {"chunk-1", "chunk-2", "chunk-3"} + assert "metadata" in results[0] + + +def test_faiss_vector_store_top_k_clamps_to_available(chunks: list[Chunk], sample_embeddings: list[list[float]]) -> None: + store = FaissVectorStore() + store.add(chunks, sample_embeddings) + + results = store.search(sample_embeddings[0], top_k=10) + assert len(results) == 3 + + +def test_faiss_vector_store_rejects_invalid_top_k(chunks: list[Chunk], sample_embeddings: list[list[float]]) -> None: + store = FaissVectorStore() + store.add(chunks, sample_embeddings) + + with pytest.raises(ValueError): + store.search(sample_embeddings[0], top_k=0) + + with pytest.raises(ValueError): + store.search(sample_embeddings[0], top_k=-1) + + +def test_faiss_vector_store_rejects_dimension_mismatch(chunks: list[Chunk]) -> None: + store = FaissVectorStore() + + with pytest.raises(ValueError): + store.add(chunks, [[0.1, 0.2]]) + + +def test_faiss_vector_store_empty_index_returns_empty_search() -> None: + store = FaissVectorStore() + + assert store.search([0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8], top_k=3) == [] + + +def test_faiss_vector_store_save_and_load_round_trip(chunks: list[Chunk], sample_embeddings: list[list[float]], tmp_path: Path) -> None: + store = FaissVectorStore() + store.add(chunks, sample_embeddings) + path = tmp_path / "retrieval.index" + + store.save(path) + restored = FaissVectorStore() + restored.load(path) + + assert restored.search(sample_embeddings[0], top_k=1)[0]["chunk_id"] == store.search(sample_embeddings[0], top_k=1)[0]["chunk_id"] + assert restored.index.ntotal == 3