""" Simple Retrieval-Augmented Generation (RAG) agent. Usage: python rag_agent.py --data_dir /path/to/docs --query "Your question" The script will: 1. Load all .txt files from data_dir. 2. Create embeddings using OpenAI's text-embedding-ada-002. 3. Store them in a FAISS vector store. 4. Retrieve top‑k relevant documents for the query. 5. Generate an answer using OpenAI GPT‑3.5‑Turbo or GPT‑4 via LangChain. """ import os import argparse from pathlib import Path from langchain.document_loaders import TextLoader from langchain.embeddings import OpenAIEmbeddings from langchain.vectorstores import FAISS from langchain.llms import OpenAI from langchain.chains import RetrievalQA from langchain.prompts import PromptTemplate def load_documents(data_dir: str): """Load all .txt files from data_dir. Args: data_dir: Directory containing text files. Returns: List of langchain Document objects. """ docs = [] for file_path in Path(data_dir).rglob("*.txt"): loader = TextLoader(str(file_path)) docs.extend(loader.load()) return docs def main(): parser = argparse.ArgumentParser(description="Simple RAG agent using LangChain and OpenAI") parser.add_argument("--data_dir", required=True, help="Directory with .txt documents") parser.add_argument("--query", required=True, help="Query to answer") parser.add_argument("--model", default="gpt-3.5-turbo", help="OpenAI model for generation (gpt-3.5-turbo or gpt-4)") parser.add_argument("--k", type=int, default=5, help="Number of documents to retrieve") args = parser.parse_args() if "OPENAI_API_KEY" not in os.environ: raise ValueError("OPENAI_API_KEY environment variable not set") # 1. Load documents print("Loading documents...") docs = load_documents(args.data_dir) if not docs: raise ValueError("No documents found in the specified directory") # 2. Create embeddings and vector store print("Creating embeddings and vector store...") embeddings = OpenAIEmbeddings() vector_store = FAISS.from_documents(docs, embeddings) # 3. Build Retriever and QA chain retriever = vector_store.as_retriever(search_kwargs={"k": args.k}) prompt_template = PromptTemplate( input_variables=["context", "question"], template=""" You are an AI assistant. Use the following context to answer the question. Context:\n{context}\n\nQuestion: {question}\n\nAnswer:\n""" ) llm = OpenAI(model_name=args.model, temperature=0) qa_chain = RetrievalQA.from_chain_type( llm=llm, chain_type="stuff", retriever=retriever, return_source_documents=False, chain_type_kwargs={"prompt": prompt_template}, ) # 4. Run query print("Running query...") result = qa_chain(args.query) print("\nAnswer:\n", result["answer"]) if __name__ == "__main__": main()