diff --git a/rag_tools.py b/rag_tools.py new file mode 100644 index 0000000..4622a10 --- /dev/null +++ b/rag_tools.py @@ -0,0 +1,41 @@ +from langchain.tools import tool +from qdrant_client import QdrantClient +from qdrant_client.http import models +from langchain_ollama import OllamaEmbeddings +from langchain_text_splitter import RecursiveCharacterTextSplitter + +# Initialize global store +client = QdrantClient(url="http://localhost:6333") +collection_name = "rag_collection" +# Ensure collection exists +if collection_name not in client.get_collections().collections: + client.recreate_collection( + collection_name=collection_name, + vectors_config=models.VectorParams(size=384, distance=models.Distance.COSINE), + ) + +embeddings = OllamaEmbeddings(model="nomic-embed-text") +splitter = RecursiveCharacterTextSplitter(chunk_size=1000, chunk_overlap=200) + +@tool("search_knowledge_base") +def search_knowledge_base(query: str, max_results: int = 5): + """Semantic search in the knowledge base.""" + query_vec = embeddings.embed_query(query) + results = client.search( + collection_name=collection_name, + query_vector=query_vec, + limit=max_results, + with_payload=True, + ) + return [r.payload for r in results] + +@tool("add_to_knowledge_base") +def add_to_knowledge_base(content: str, title: str): + """Add a document to the knowledge base.""" + chunks = splitter.split_text(content) + vectors = embeddings.embed_documents(chunks) + points = [] + for i, (chunk, vec) in enumerate(zip(chunks, vectors)): + points.append(models.PointStruct(id=i, vector=vec, payload={"title": title, "chunk": chunk})) + client.upsert(collection_name=collection_name, points=models.Batch(points=points)) + return f"Added {len(chunks)} chunks to the knowledge base."