feat: solution for 'Экзамен: RAG-агент с ChromaDB и веб-поиском'
This commit is contained in:
@@ -1,62 +1,110 @@
|
||||
# RAG Agent with ChromaDB and Web Search
|
||||
|
||||
This project implements a Retrieval-Augmented Generation (RAG) agent that uses **ChromaDB** as the vector database and the **OpenAI API** to generate responses based on retrieved documents. It also includes a simple web‑search component that fetches content from specified URLs for indexing.
|
||||
This project implements a Retrieval-Augmented Generation (RAG) agent that uses a local ChromaDB vector store for document retrieval and falls back to DuckDuckGo web search when the local store does not provide sufficient context.
|
||||
|
||||
## Features
|
||||
|
||||
- **Vector Store**: Uses ChromaDB to store embeddings of text chunks.
|
||||
- **OpenAI Integration**: Generates answers using GPT‑3.5‑Turbo.
|
||||
- **Web Search**: Fetches and parses HTML pages, splits them into manageable chunks.
|
||||
- **Command Line Interface**: Ask questions interactively.
|
||||
- **Local Retrieval** – Store and query embeddings in a persistent ChromaDB collection.
|
||||
- **Web Search Fallback** – If local retrieval fails to find relevant context, the agent performs a DuckDuckGo search and uses the snippets.
|
||||
- **OpenAI Integration** – Uses OpenAI embeddings (`text-embedding-ada-002`) and the `gpt-3.5-turbo` model for generation.
|
||||
- **CLI** – Simple command line interface for ingesting documents and asking questions.
|
||||
|
||||
## Prerequisites
|
||||
|
||||
- Node.js v18+ (supports native ES modules and `node-fetch` v2).
|
||||
- An OpenAI API key.
|
||||
- Python 3.9+
|
||||
- An OpenAI API key
|
||||
- (Optional) Internet access for web search
|
||||
|
||||
## Setup
|
||||
## Installation
|
||||
|
||||
1. **Clone the repository** (or copy the files into a directory).
|
||||
```bash
|
||||
# Clone the repository
|
||||
git clone https://git.brojs.ru/kuzakhmetovartur/ekzamen-rag-agent-s-chromadb-i-veb-poisk.git
|
||||
cd ekzamen-rag-agent-s-chromadb-i-veb-poisk
|
||||
|
||||
2. **Install dependencies**
|
||||
# Create a virtual environment (recommended)
|
||||
python -m venv .venv
|
||||
source .venv/bin/activate # On Windows use `.venv\Scripts\activate`
|
||||
|
||||
```bash
|
||||
npm install
|
||||
```
|
||||
# Install dependencies
|
||||
pip install -r requirements.txt
|
||||
```
|
||||
|
||||
3. **Configure environment**
|
||||
`requirements.txt` contains:
|
||||
|
||||
Create a `.env` file in the project root (or edit the existing one) and add your OpenAI API key:
|
||||
```
|
||||
openai
|
||||
chromadb
|
||||
duckduckgo-search
|
||||
beautifulsoup4
|
||||
requests
|
||||
```
|
||||
|
||||
```dotenv
|
||||
OPENAI_API_KEY=your_api_key_here
|
||||
```
|
||||
## Environment Variables
|
||||
|
||||
4. **Run the agent**
|
||||
| Variable | Description | Example |
|
||||
|----------|-------------|---------|
|
||||
| `OPENAI_API_KEY` | Your OpenAI API key | `sk-...` |
|
||||
| `CHROMA_DB_PATH` | Directory where ChromaDB stores data | `./chromadb` |
|
||||
| `CHROMA_COLLECTION_NAME` | Name of the collection | `rag_collection` |
|
||||
| `TOP_K` | Number of top documents to retrieve | `5` |
|
||||
| `SIMILARITY_THRESHOLD` | Minimum similarity to consider a document relevant | `0.5` |
|
||||
| `WEB_SEARCH_MAX_RESULTS` | Max number of web snippets to fetch | `3` |
|
||||
|
||||
```bash
|
||||
npm start
|
||||
```
|
||||
Set them in your shell or create a `.env` file and load with `dotenv` (optional).
|
||||
|
||||
The script will:
|
||||
- Fetch and index the example URLs.
|
||||
- Prompt you to enter questions.
|
||||
- Display answers generated by the RAG agent.
|
||||
## Usage
|
||||
|
||||
## Customization
|
||||
### Ingest Documents
|
||||
|
||||
- **Adding URLs**: Edit the `urls` array in `src/index.js` to index different web pages.
|
||||
- **Chunk Size**: Adjust the `size` parameter in `chunkText` inside `src/webSearch.js` if you need larger or smaller chunks.
|
||||
- **Model Parameters**: Modify temperature, max tokens, or model name in `src/agent.js`.
|
||||
Place your plain text files (`.txt`) in a folder, then run:
|
||||
|
||||
## Notes
|
||||
```bash
|
||||
python src/index.py ingest /path/to/text/files
|
||||
```
|
||||
|
||||
- The implementation strictly uses **ChromaDB** as the vector database; no other vector DBs are used.
|
||||
- All dependencies are declared in `package.json` and can be installed via `npm install`.
|
||||
- The OpenAI API key is loaded securely from the `.env` file using `dotenv`.
|
||||
The script will read all `.txt` files, split them into chunks, embed them, and store them in ChromaDB.
|
||||
|
||||
### Ask a Question
|
||||
|
||||
```bash
|
||||
python src/index.py ask "What is the capital of France?"
|
||||
```
|
||||
|
||||
The agent will:
|
||||
|
||||
1. Query the local vector store for relevant passages.
|
||||
2. If none are found above the similarity threshold, perform a DuckDuckGo search.
|
||||
3. Combine the retrieved context into a prompt.
|
||||
4. Call OpenAI’s `gpt-3.5-turbo` to generate an answer.
|
||||
|
||||
## Example
|
||||
|
||||
```bash
|
||||
$ python src/index.py ingest ./data
|
||||
INFO:root:Added 12 documents to collection 'rag_collection'.
|
||||
|
||||
$ python src/index.py ask "Explain the theory of relativity."
|
||||
Answer:
|
||||
The theory of relativity, developed by Albert Einstein, consists of two parts: special relativity and general relativity. ...
|
||||
```
|
||||
|
||||
## Testing
|
||||
|
||||
Unit tests are provided in the `tests/` directory. To run them:
|
||||
|
||||
```bash
|
||||
pytest tests/
|
||||
```
|
||||
|
||||
(If you don't have `pytest` installed, run `pip install pytest`.)
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
- **No documents ingested** – Ensure the folder path is correct and contains `.txt` files.
|
||||
- **OpenAI errors** – Verify that `OPENAI_API_KEY` is set and that you have sufficient quota.
|
||||
- **Web search fails** – Check your internet connection and that DuckDuckGo is reachable.
|
||||
|
||||
## License
|
||||
|
||||
MIT License
|
||||
---
|
||||
Enjoy building with RAG!
|
||||
+6
-5
@@ -1,5 +1,6 @@
|
||||
chromadb==0.4.22
|
||||
openai==1.12.0
|
||||
requests==2.31.0
|
||||
beautifulsoup4==4.12.3
|
||||
python-dotenv==1.0.1
|
||||
openai
|
||||
chromadb
|
||||
duckduckgo-search
|
||||
beautifulsoup4
|
||||
requests
|
||||
pytest
|
||||
+241
-92
@@ -1,137 +1,286 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
RAG Agent using ChromaDB for vector storage and Tavily for web search.
|
||||
RAG Agent with ChromaDB and Web Search
|
||||
|
||||
The agent can ingest web pages (or arbitrary text) into a Chroma collection
|
||||
and answer queries by retrieving relevant documents and passing them to an
|
||||
OpenAI LLM.
|
||||
This module implements a Retrieval-Augmented Generation (RAG) agent that
|
||||
uses a local ChromaDB vector store for document retrieval and falls back
|
||||
to DuckDuckGo web search when the local store does not provide sufficient
|
||||
context.
|
||||
|
||||
Prerequisites:
|
||||
- OpenAI API key set in the environment variable OPENAI_API_KEY
|
||||
- Tavily API key set in the environment variable TAVILY_API_KEY
|
||||
- Python 3.9+
|
||||
|
||||
Usage:
|
||||
python src/index.py ingest <url_or_text> # Ingest a URL or raw text
|
||||
python src/index.py query <question> # Query the agent
|
||||
|
||||
Example:
|
||||
python src/index.py ingest https://en.wikipedia.org/wiki/OpenAI
|
||||
python src/index.py query "What is OpenAI?"
|
||||
- OpenAI API key set in OPENAI_API_KEY
|
||||
- DuckDuckGo search library (pip install duckduckgo-search)
|
||||
- ChromaDB client (pip install chromadb)
|
||||
- OpenAI Python SDK (pip install openai)
|
||||
"""
|
||||
|
||||
import os
|
||||
import sys
|
||||
import argparse
|
||||
import json
|
||||
import textwrap
|
||||
import logging
|
||||
from pathlib import Path
|
||||
from typing import List
|
||||
from typing import List, Tuple, Optional
|
||||
|
||||
from langchain.embeddings import OpenAIEmbeddings
|
||||
from langchain.vectorstores import Chroma
|
||||
from langchain.llms import OpenAI
|
||||
from langchain.chains import RetrievalQA
|
||||
from tavily import TavilyClient
|
||||
import openai
|
||||
import chromadb
|
||||
from chromadb import Client
|
||||
from chromadb.config import Settings
|
||||
from duckduckgo_search import ddg
|
||||
|
||||
# Configure logging
|
||||
logging.basicConfig(
|
||||
level=logging.INFO,
|
||||
format="%(asctime)s [%(levelname)s] %(message)s",
|
||||
handlers=[logging.StreamHandler(sys.stdout)],
|
||||
)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Configuration
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
||||
# Directory where the ChromaDB data will be stored
|
||||
CHROMA_DATA_DIR = Path.home() / ".rag_agent" / "chromadb"
|
||||
# Environment variables
|
||||
OPENAI_API_KEY = os.getenv("OPENAI_API_KEY")
|
||||
if not OPENAI_API_KEY:
|
||||
logger.error("OPENAI_API_KEY environment variable is not set.")
|
||||
sys.exit(1)
|
||||
|
||||
# Name of the collection used for storing documents
|
||||
COLLECTION_NAME = "rag_collection"
|
||||
openai.api_key = OPENAI_API_KEY
|
||||
|
||||
# ChromaDB settings
|
||||
CHROMA_DB_PATH = os.getenv("CHROMA_DB_PATH", "./chromadb")
|
||||
CHROMA_COLLECTION_NAME = os.getenv("CHROMA_COLLECTION_NAME", "rag_collection")
|
||||
|
||||
# Retrieval settings
|
||||
TOP_K = int(os.getenv("TOP_K", "5"))
|
||||
SIMILARITY_THRESHOLD = float(os.getenv("SIMILARITY_THRESHOLD", "0.5"))
|
||||
MAX_CHUNK_SIZE = int(os.getenv("MAX_CHUNK_SIZE", "500")) # characters
|
||||
|
||||
# Web search settings
|
||||
WEB_SEARCH_MAX_RESULTS = int(os.getenv("WEB_SEARCH_MAX_RESULTS", "3"))
|
||||
WEB_SEARCH_TIMEOUT = int(os.getenv("WEB_SEARCH_TIMEOUT", "10")) # seconds
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Helper functions
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
||||
def get_chroma_collection() -> Chroma:
|
||||
def _split_text(text: str, max_chunk_size: int = MAX_CHUNK_SIZE) -> List[str]:
|
||||
"""
|
||||
Create or load a Chroma collection.
|
||||
Split a long text into smaller chunks of at most max_chunk_size characters.
|
||||
Splits on sentence boundaries when possible.
|
||||
"""
|
||||
embeddings = OpenAIEmbeddings()
|
||||
return Chroma(
|
||||
collection_name=COLLECTION_NAME,
|
||||
embedding_function=embeddings,
|
||||
persist_directory=str(CHROMA_DATA_DIR),
|
||||
sentences = text.replace("\n", " ").split(". ")
|
||||
chunks = []
|
||||
current = ""
|
||||
for sentence in sentences:
|
||||
if len(current) + len(sentence) + 1 <= max_chunk_size:
|
||||
current += sentence + ". "
|
||||
else:
|
||||
if current:
|
||||
chunks.append(current.strip())
|
||||
current = sentence + ". "
|
||||
if current:
|
||||
chunks.append(current.strip())
|
||||
return chunks
|
||||
|
||||
def _embed_text(text: str) -> List[float]:
|
||||
"""
|
||||
Embed a single text string using OpenAI embeddings.
|
||||
"""
|
||||
try:
|
||||
response = openai.Embedding.create(
|
||||
model="text-embedding-ada-002",
|
||||
input=text,
|
||||
)
|
||||
return response["data"][0]["embedding"]
|
||||
except Exception as e:
|
||||
logger.exception(f"Embedding failed for text: {text[:30]}...: {e}")
|
||||
return []
|
||||
|
||||
def ingest_text(text: str, collection: Chroma) -> None:
|
||||
def _fetch_web_content(url: str) -> Optional[str]:
|
||||
"""
|
||||
Add raw text to the Chroma collection.
|
||||
Fetch the textual content of a web page.
|
||||
"""
|
||||
collection.add_texts([text])
|
||||
|
||||
def ingest_url(url: str, collection: Chroma, tavily_client: TavilyClient) -> None:
|
||||
"""
|
||||
Fetch content from a URL using Tavily, embed it, and store it in Chroma.
|
||||
"""
|
||||
# Tavily's search returns a list of results; we use the first result's content.
|
||||
results = tavily_client.search(query=url, max_results=1)
|
||||
if not results:
|
||||
print(f"No results found for URL: {url}")
|
||||
return
|
||||
content = results[0].content
|
||||
if not content:
|
||||
print(f"No content extracted from URL: {url}")
|
||||
return
|
||||
collection.add_texts([content])
|
||||
print(f"Ingested content from {url}")
|
||||
|
||||
def query_agent(question: str, collection: Chroma) -> str:
|
||||
"""
|
||||
Retrieve relevant documents from Chroma and ask OpenAI to answer.
|
||||
"""
|
||||
llm = OpenAI(temperature=0)
|
||||
qa_chain = RetrievalQA.from_chain_type(
|
||||
llm=llm,
|
||||
chain_type="stuff",
|
||||
retriever=collection.as_retriever(search_kwargs={"k": 4}),
|
||||
)
|
||||
return qa_chain.run(question)
|
||||
try:
|
||||
import requests
|
||||
resp = requests.get(url, timeout=WEB_SEARCH_TIMEOUT)
|
||||
resp.raise_for_status()
|
||||
# Very naive extraction: strip HTML tags
|
||||
from bs4 import BeautifulSoup
|
||||
soup = BeautifulSoup(resp.text, "html.parser")
|
||||
text = soup.get_text(separator=" ", strip=True)
|
||||
return text
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to fetch {url}: {e}")
|
||||
return None
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Main entry point
|
||||
# ChromaDB wrapper
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(description="RAG Agent with ChromaDB & Tavily")
|
||||
class ChromaDBWrapper:
|
||||
def __init__(self, path: str = CHROMA_DB_PATH, collection_name: str = CHROMA_COLLECTION_NAME):
|
||||
self.client: Client = chromadb.Client(Settings(
|
||||
chroma_db_impl="duckdb+parquet",
|
||||
persist_directory=path,
|
||||
))
|
||||
self.collection = self.client.get_or_create_collection(name=collection_name)
|
||||
|
||||
def add_documents(self, documents: List[str], ids: List[str]) -> None:
|
||||
embeddings = [_embed_text(doc) for doc in documents]
|
||||
self.collection.add(
|
||||
documents=documents,
|
||||
embeddings=embeddings,
|
||||
ids=ids,
|
||||
)
|
||||
logger.info(f"Added {len(documents)} documents to collection '{self.collection.name}'.")
|
||||
|
||||
def query(self, query_text: str, k: int = TOP_K) -> List[Tuple[str, float]]:
|
||||
"""
|
||||
Return top-k (document, similarity) tuples for the query_text.
|
||||
"""
|
||||
query_embedding = _embed_text(query_text)
|
||||
if not query_embedding:
|
||||
return []
|
||||
|
||||
results = self.collection.query(
|
||||
query_embeddings=[query_embedding],
|
||||
n_results=k,
|
||||
include=["documents", "distances"],
|
||||
)
|
||||
docs = results["documents"][0]
|
||||
distances = results["distances"][0]
|
||||
# Convert distances to similarity (1 - distance)
|
||||
similarities = [1 - d for d in distances]
|
||||
return list(zip(docs, similarities))
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Web search fallback
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
||||
def web_search(query: str, max_results: int = WEB_SEARCH_MAX_RESULTS) -> List[str]:
|
||||
"""
|
||||
Perform a DuckDuckGo search and return a list of snippet texts.
|
||||
"""
|
||||
try:
|
||||
results = ddg(query, max_results=max_results)
|
||||
snippets = []
|
||||
for r in results:
|
||||
snippet = r.get("body") or r.get("snippet") or ""
|
||||
if snippet:
|
||||
snippets.append(snippet.strip())
|
||||
logger.info(f"Web search returned {len(snippets)} snippets for query '{query}'.")
|
||||
return snippets
|
||||
except Exception as e:
|
||||
logger.exception(f"Web search failed for query '{query}': {e}")
|
||||
return []
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# RAG Agent
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
||||
class RAGAgent:
|
||||
def __init__(self, db_wrapper: ChromaDBWrapper):
|
||||
self.db = db_wrapper
|
||||
|
||||
def ingest_folder(self, folder_path: str) -> None:
|
||||
"""
|
||||
Read all .txt files in folder_path, split into chunks, and store in ChromaDB.
|
||||
"""
|
||||
folder = Path(folder_path)
|
||||
if not folder.is_dir():
|
||||
logger.error(f"Folder {folder_path} does not exist.")
|
||||
return
|
||||
|
||||
documents = []
|
||||
ids = []
|
||||
for txt_file in folder.rglob("*.txt"):
|
||||
try:
|
||||
text = txt_file.read_text(encoding="utf-8")
|
||||
chunks = _split_text(text)
|
||||
documents.extend(chunks)
|
||||
ids.extend([f"{txt_file.stem}_{i}" for i in range(len(chunks))])
|
||||
logger.info(f"Processed {txt_file} into {len(chunks)} chunks.")
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to read {txt_file}: {e}")
|
||||
|
||||
if documents:
|
||||
self.db.add_documents(documents, ids)
|
||||
else:
|
||||
logger.warning("No documents found to ingest.")
|
||||
|
||||
def answer_query(self, query: str) -> str:
|
||||
"""
|
||||
Generate an answer to the user query using local retrieval and web fallback.
|
||||
"""
|
||||
# 1. Retrieve from local store
|
||||
retrieved = self.db.query(query, k=TOP_K)
|
||||
relevant_docs = [doc for doc, sim in retrieved if sim >= SIMILARITY_THRESHOLD]
|
||||
|
||||
# 2. If not enough context, perform web search
|
||||
if not relevant_docs:
|
||||
logger.info("No relevant local documents found; performing web search.")
|
||||
snippets = web_search(query)
|
||||
relevant_docs = snippets
|
||||
|
||||
# 3. Build prompt
|
||||
context = "\n\n".join(relevant_docs[:TOP_K])
|
||||
prompt = textwrap.dedent(
|
||||
f"""
|
||||
You are an AI assistant. Use the following context to answer the question.
|
||||
Context:
|
||||
{context}
|
||||
|
||||
Question: {query}
|
||||
Answer:
|
||||
"""
|
||||
)
|
||||
|
||||
# 4. Call OpenAI ChatCompletion
|
||||
try:
|
||||
response = openai.ChatCompletion.create(
|
||||
model="gpt-3.5-turbo",
|
||||
messages=[
|
||||
{"role": "system", "content": "You are a helpful assistant."},
|
||||
{"role": "user", "content": prompt},
|
||||
],
|
||||
temperature=0.2,
|
||||
max_tokens=512,
|
||||
)
|
||||
answer = response["choices"][0]["message"]["content"].strip()
|
||||
return answer
|
||||
except Exception as e:
|
||||
logger.exception(f"OpenAI completion failed: {e}")
|
||||
return "Sorry, I couldn't generate an answer at this time."
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# CLI
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
||||
def main():
|
||||
import argparse
|
||||
|
||||
parser = argparse.ArgumentParser(description="RAG Agent CLI")
|
||||
subparsers = parser.add_subparsers(dest="command", required=True)
|
||||
|
||||
ingest_parser = subparsers.add_parser("ingest", help="Ingest a URL or raw text")
|
||||
ingest_parser.add_argument("source", help="URL or raw text to ingest")
|
||||
ingest_parser = subparsers.add_parser("ingest", help="Ingest documents from a folder")
|
||||
ingest_parser.add_argument("folder", help="Path to folder containing .txt files")
|
||||
|
||||
query_parser = subparsers.add_parser("query", help="Ask a question")
|
||||
query_parser = subparsers.add_parser("ask", help="Ask a question")
|
||||
query_parser.add_argument("question", help="The question to ask the agent")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
# Ensure required environment variables are set
|
||||
if "OPENAI_API_KEY" not in os.environ:
|
||||
print("Error: OPENAI_API_KEY environment variable not set.")
|
||||
sys.exit(1)
|
||||
if "TAVILY_API_KEY" not in os.environ:
|
||||
print("Error: TAVILY_API_KEY environment variable not set.")
|
||||
sys.exit(1)
|
||||
|
||||
# Initialize Chroma collection
|
||||
collection = get_chroma_collection()
|
||||
|
||||
# Initialize Tavily client
|
||||
tavily_client = TavilyClient(api_key=os.environ["TAVILY_API_KEY"])
|
||||
db_wrapper = ChromaDBWrapper()
|
||||
agent = RAGAgent(db_wrapper)
|
||||
|
||||
if args.command == "ingest":
|
||||
source = args.source
|
||||
if source.startswith(("http://", "https://")):
|
||||
ingest_url(source, collection, tavily_client)
|
||||
else:
|
||||
ingest_text(source, collection)
|
||||
print("Ingested raw text.")
|
||||
elif args.command == "query":
|
||||
answer = query_agent(args.question, collection)
|
||||
print("\nAnswer:\n")
|
||||
print(answer)
|
||||
agent.ingest_folder(args.folder)
|
||||
elif args.command == "ask":
|
||||
answer = agent.answer_query(args.question)
|
||||
print("\nAnswer:\n" + answer)
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,128 @@
|
||||
import os
|
||||
import sys
|
||||
import json
|
||||
import tempfile
|
||||
import shutil
|
||||
from pathlib import Path
|
||||
from unittest import mock
|
||||
|
||||
# Ensure the environment variable is set before importing the module
|
||||
os.environ["OPENAI_API_KEY"] = "test-key"
|
||||
|
||||
# Import the module after setting the environment variable
|
||||
import src.index as index
|
||||
|
||||
# Helper to create a temporary text file
|
||||
def create_temp_txt(folder: Path, name: str, content: str):
|
||||
file_path = folder / name
|
||||
file_path.write_text(content, encoding="utf-8")
|
||||
return file_path
|
||||
|
||||
def test_split_text_basic():
|
||||
text = "Sentence one. Sentence two. Sentence three."
|
||||
chunks = index._split_text(text, max_chunk_size=50)
|
||||
assert len(chunks) == 3
|
||||
assert chunks[0] == "Sentence one."
|
||||
assert chunks[1] == "Sentence two."
|
||||
assert chunks[2] == "Sentence three."
|
||||
|
||||
def test_split_text_long_sentence():
|
||||
long_sentence = "A" * 200
|
||||
text = f"{long_sentence}. Another short sentence."
|
||||
chunks = index._split_text(text, max_chunk_size=100)
|
||||
# The long sentence should be split into two chunks
|
||||
assert len(chunks) == 2
|
||||
assert chunks[0].startswith("A" * 100)
|
||||
assert chunks[1].startswith("A" * 100)
|
||||
|
||||
def test_ingest_folder(monkeypatch):
|
||||
# Create a temporary directory with a single .txt file
|
||||
temp_dir = Path(tempfile.mkdtemp())
|
||||
try:
|
||||
content = "Hello world. This is a test."
|
||||
create_temp_txt(temp_dir, "test.txt", content)
|
||||
|
||||
# Mock the add_documents method to capture its arguments
|
||||
captured = {}
|
||||
def mock_add_documents(self, documents, ids):
|
||||
captured["documents"] = documents
|
||||
captured["ids"] = ids
|
||||
|
||||
monkeypatch.setattr(index.ChromaDBWrapper, "add_documents", mock_add_documents)
|
||||
|
||||
# Instantiate the agent with a dummy db wrapper
|
||||
dummy_db = index.ChromaDBWrapper()
|
||||
agent = index.RAGAgent(dummy_db)
|
||||
|
||||
# Run ingestion
|
||||
agent.ingest_folder(str(temp_dir))
|
||||
|
||||
# Verify that documents were split and added
|
||||
assert "documents" in captured
|
||||
assert "ids" in captured
|
||||
assert len(captured["documents"]) == 2 # two sentences
|
||||
assert captured["documents"][0] == "Hello world."
|
||||
assert captured["documents"][1] == "This is a test."
|
||||
assert len(captured["ids"]) == 2
|
||||
assert captured["ids"][0].startswith("test_")
|
||||
finally:
|
||||
shutil.rmtree(temp_dir)
|
||||
|
||||
def test_answer_query_local(monkeypatch):
|
||||
# Dummy database that returns a relevant document
|
||||
class DummyDB:
|
||||
def query(self, query_text, k=5):
|
||||
return [("Relevant context about Python.", 0.8)]
|
||||
|
||||
dummy_db = DummyDB()
|
||||
agent = index.RAGAgent(dummy_db)
|
||||
|
||||
# Mock the OpenAI ChatCompletion to return a predictable answer
|
||||
mock_response = {
|
||||
"choices": [
|
||||
{"message": {"content": "Python is a programming language."}}
|
||||
]
|
||||
}
|
||||
monkeypatch.setattr(index.openai.ChatCompletion, "create", lambda **kwargs: mock_response)
|
||||
|
||||
answer = agent.answer_query("What is Python?")
|
||||
assert answer == "Python is a programming language."
|
||||
|
||||
def test_answer_query_fallback(monkeypatch):
|
||||
# Dummy database that returns no relevant documents
|
||||
class DummyDB:
|
||||
def query(self, query_text, k=5):
|
||||
return []
|
||||
|
||||
dummy_db = DummyDB()
|
||||
agent = index.RAGAgent(dummy_db)
|
||||
|
||||
# Mock web_search to return snippets
|
||||
monkeypatch.setattr(index, "web_search", lambda query, max_results=3: ["Snippet about AI.", "Another snippet."])
|
||||
|
||||
# Mock the OpenAI ChatCompletion to return a predictable answer
|
||||
mock_response = {
|
||||
"choices": [
|
||||
{"message": {"content": "AI stands for Artificial Intelligence."}}
|
||||
]
|
||||
}
|
||||
monkeypatch.setattr(index.openai.ChatCompletion, "create", lambda **kwargs: mock_response)
|
||||
|
||||
answer = agent.answer_query("What does AI stand for?")
|
||||
assert answer == "AI stands for Artificial Intelligence."
|
||||
|
||||
def test_web_search_mock(monkeypatch):
|
||||
# Mock ddg to return predefined results
|
||||
mock_results = [
|
||||
{"body": "First snippet content."},
|
||||
{"body": "Second snippet content."},
|
||||
]
|
||||
monkeypatch.setattr(index.ddg, "__call__", lambda query, max_results=3: mock_results)
|
||||
|
||||
snippets = index.web_search("test query")
|
||||
assert snippets == ["First snippet content.", "Second snippet content."]
|
||||
|
||||
if __name__ == "__main__":
|
||||
# Run tests manually if executed as a script
|
||||
import pytest
|
||||
sys.exit(pytest.main([__file__]))
|
||||
Reference in New Issue
Block a user