Files
task-6a02e23da6fe2e4ac16acf65/main.py
T

128 lines
4.6 KiB
Python

# Main script implementing Qdrant-based RAG agent
import os
import sys
import json
from pathlib import Path
from typing import List, Dict, Any
from langchain_ollama import ChatOllama, OllamaEmbeddings
from langchain_qdrant import QdrantVectorStore
from langchain_text_splitters import RecursiveCharacterTextSplitter
from langchain.tools import tool, BaseTool
from langchain.agents import create_agent, AgentExecutor, AgentType
# Configuration
QDRANT_HOST = os.getenv("QDRANT_HOST", "localhost")
QDRANT_PORT = int(os.getenv("QDRANT_PORT", "6333"))
COLLECTION_NAME = "knowledge"
EMBEDDING_MODEL = "nomic-embed-text"
LLM_MODEL = "llama3"
# Vector store wrapper
class QdrantStore:
def __init__(self, host: str, port: int, collection: str):
self.store = QdrantVectorStore(
url=f"http://{host}:{port}",
collection_name=collection,
embedding=OllamaEmbeddings(model=EMBEDDING_MODEL),
)
def add_documents(self, documents: List[str], metadatas: List[Dict[str, Any]]):
self.store.add_texts(documents, metadatas=metadatas)
def similarity_search(self, query: str, k: int = 5) -> List[Dict[str, Any]]:
results = self.store.similarity_search(query, k=k)
return [
{
"content": doc.page_content,
"metadata": doc.metadata,
"score": doc.metadata.get("score", 0),
}
for doc in results
]
# Text splitter
text_splitter = RecursiveCharacterTextSplitter(chunk_size=1000, chunk_overlap=200)
# Store instance
store = QdrantStore(QDRANT_HOST, QDRANT_PORT, COLLECTION_NAME)
# Tools
@tool("search_knowledge_base", "Semantic search in the knowledge base")
def search_knowledge_base(query: str, max_results: int = 5) -> str:
results = store.similarity_search(query, k=max_results)
return json.dumps(results, ensure_ascii=False)
@tool("add_to_knowledge_base", "Add a document to the knowledge base")
def add_to_knowledge_base(content: str, title: str = "Untitled") -> str:
chunks = text_splitter.split_text(content)
metadatas = [{"title": title, "chunk_idx": i} for i in range(len(chunks))]
store.add_documents(chunks, metadatas)
return f"Added {len(chunks)} chunks titled '{title}'."
# Agent
llm = ChatOllama(model=LLM_MODEL)
agent = create_agent(
llm=llm,
tools=[search_knowledge_base, add_to_knowledge_base],
agent_type=AgentType.ZERO_SHOT_REACT_DESCRIPTION,
system_message="You are an assistant that can search and add information to a local knowledge base. Use the tools when appropriate.",
)
executor = AgentExecutor(agent=agent, tools=[search_knowledge_base, add_to_knowledge_base], verbose=True)
# CLI helpers
def load_documents_from_dir(dir_path: str):
for path in Path(dir_path).glob("**/*"):
if path.suffix.lower() in {".txt", ".md"}:
content = path.read_text(encoding="utf-8")
title = path.stem
add_to_knowledge_base(content, title)
print("Loading complete.")
def main():
if len(sys.argv) > 1 and sys.argv[1] == "load":
if len(sys.argv) < 3:
print("Usage: python main.py load <directory>")
sys.exit(1)
load_documents_from_dir(sys.argv[2])
sys.exit(0)
print("Interactive mode. Commands: /add <title> <file>, /search <query>, /quit")
while True:
try:
user_input = input("\n> ")
except (EOFError, KeyboardInterrupt):
break
if not user_input:
continue
if user_input.startswith("/quit"):
break
if user_input.startswith("/add"):
parts = user_input.split(maxsplit=2)
if len(parts) != 3:
print("Usage: /add <title> <file_path>")
continue
title, file_path = parts[1], parts[2]
try:
content = Path(file_path).read_text(encoding="utf-8")
except Exception as e:
print(f"Error reading file: {e}")
continue
print(add_to_knowledge_base(content, title))
continue
if user_input.startswith("/search"):
query = user_input[len("/search"):].strip()
if not query:
print("Provide a query.")
continue
results = search_knowledge_base(query)
print("Search results:")
for r in json.loads(results):
print(f"- {r['metadata'].get('title', 'Untitled')} (chunk {r['metadata'].get('chunk_idx')})\n {r['content'][:200]}...")
continue
response = executor.invoke({"input": user_input})
print(response["output"])
if __name__ == "__main__":
main()