Files
task-6a02e23da6fe2e4ac16acf65/main.py
T

187 lines
7.3 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
# main.py RAGagent with Qdrant, OpenRouter, and LangChain
# -----------------------------------------------------------------
# This script implements a simple RAG agent that can search and add
# documents to a Qdrant vector store. The agent is built with
# LangChain's `create_agent` and uses OpenRouter for both the LLM and
# embeddings. The code follows the "Исправить" section of the
# assignment and includes detailed comments explaining design choices.
# -----------------------------------------------------------------
import os
import asyncio
import argparse
from pathlib import Path
from langchain_openai import ChatOpenAI, OpenAIEmbeddings
from langchain_core.messages import HumanMessage
from langchain.tools import tool
from langchain_community.document_loaders import TextLoader
from langchain_community.document_loaders import DirectoryLoader
from langchain_community.vectorstores import Qdrant
from langchain_text_splitters import RecursiveCharacterTextSplitter
from langchain.agents import create_agent, AgentExecutor, AgentType
# -----------------------------------------------------------------
# Configuration all secrets are read from environment variables.
# -----------------------------------------------------------------
OPENAI_API_KEY = os.getenv("OPENAI_API_KEY")
if not OPENAI_API_KEY:
raise RuntimeError("OPENAI_API_KEY environment variable is required")
# LLM OpenRouter gpt-oss-20b:free (free tier)
llm = ChatOpenAI(
model="openai/gpt-oss-20b:free",
base_url="https://openrouter.ai/api/v1",
api_key=OPENAI_API_KEY,
temperature=0.0,
)
# Embeddings OpenAI text-embedding-3-small via OpenRouter
embeddings = OpenAIEmbeddings(
model="text-embedding-3-small",
base_url="https://openrouter.ai/api/v1",
api_key=OPENAI_API_KEY,
)
# Qdrant client assumes a local Qdrant instance running on default port
qdrant_url = os.getenv("QDRANT_URL", "http://localhost:6333")
vector_store = Qdrant(
client=None, # will be created lazily by Qdrant wrapper
collection_name="knowledge",
embeddings=embeddings,
url=qdrant_url,
)
# -----------------------------------------------------------------
# Tool definitions these are the only tools the agent can use.
# -----------------------------------------------------------------
@tool
def search_knowledge_base(query: str, max_results: int = 3) -> str:
"""Search the knowledge base for relevant information.
Parameters
----------
query: str
The search query.
max_results: int, optional
Number of top results to return (default 3).
"""
docs = vector_store.similarity_search(query, k=max_results)
if not docs:
return "No results found."
return "\n\n---\n\n".join([f"{doc.metadata.get('title', 'Untitled')}\n{doc.page_content}" for doc in docs])
@tool
def add_to_knowledge_base(content: str, title: str = "Untitled") -> str:
"""Add a new document (or chunk) to the knowledge base.
Parameters
----------
content: str
The text content to add.
title: str, optional
A humanreadable title for the document.
"""
doc = {
"page_content": content,
"metadata": {"title": title},
}
vector_store.add_documents([doc])
return f"Added document '{title}'."
# -----------------------------------------------------------------
# Agent setup using LangChain's create_agent with a custom system prompt.
# -----------------------------------------------------------------
SYSTEM_PROMPT = (
"You are a helpful assistant with access to a knowledge base. "
"Use the tools `search_knowledge_base` and `add_to_knowledge_base` "
"to answer user queries. If the user asks to add information, "
"use `add_to_knowledge_base`. If the user asks for information, "
"use `search_knowledge_base`. Do not fabricate facts."
)
agent = create_agent(
llm=llm,
tools=[search_knowledge_base, add_to_knowledge_base],
system_prompt=SYSTEM_PROMPT,
agent_type=AgentType.ZERO_SHOT_REACT_DESCRIPTION,
)
executor = AgentExecutor(agent=agent, tools=[search_knowledge_base, add_to_knowledge_base], verbose=True)
# -----------------------------------------------------------------
# Document ingestion split into chunks and add to Qdrant.
# -----------------------------------------------------------------
def ingest_directory(directory: str, chunk_size: int = 1000, chunk_overlap: int = 200):
"""Load all text files from *directory*, split into chunks, and store.
Parameters
----------
directory: str
Path to the directory containing documents.
chunk_size: int, optional
Size of each chunk in characters.
chunk_overlap: int, optional
Overlap between consecutive chunks.
"""
loader = DirectoryLoader(directory, glob="**/*.txt")
documents = loader.load()
splitter = RecursiveCharacterTextSplitter(chunk_size=chunk_size, chunk_overlap=chunk_overlap)
chunks = splitter.split_documents(documents)
# Convert LangChain Document objects to dicts expected by Qdrant
docs_to_add = []
for doc in chunks:
title = doc.metadata.get("source", "Untitled")
docs_to_add.append({
"page_content": doc.page_content,
"metadata": {"title": title, "source": doc.metadata.get("source", "")},
})
vector_store.add_documents(docs_to_add)
print(f"Ingested {len(docs_to_add)} chunks into the knowledge base.")
# -----------------------------------------------------------------
# CLI simple interactive loop.
# -----------------------------------------------------------------
async def main():
parser = argparse.ArgumentParser(description="RAG Agent CLI")
parser.add_argument("--ingest", type=str, help="Path to directory to ingest")
args = parser.parse_args()
if args.ingest:
ingest_directory(args.ingest)
return
print("RAG Agent ready. Type /quit to exit.")
while True:
user_input = input("You: ")
if user_input.strip() == "/quit":
print("Goodbye!")
break
if user_input.startswith("/add "):
# Expected format: /add <title> | <content>
try:
_, rest = user_input.split("/add ", 1)
title, content = rest.split("|", 1)
title = title.strip()
content = content.strip()
result = await executor.ainvoke({"messages": [HumanMessage(content=f"Add document {title}")], "configurable": {"thread_id": "session-1"}})
# Directly call tool to add content
add_to_knowledge_base(content, title)
print("Agent: Document added.")
except Exception as e:
print(f"Error parsing /add command: {e}")
continue
if user_input.startswith("/search "):
query = user_input[len("/search "):].strip()
result = await executor.ainvoke({"messages": [HumanMessage(content=f"Search for {query}")], "configurable": {"thread_id": "session-1"}})
print("Agent:", result["messages"][-1].content)
continue
# Default: normal chat
result = await executor.ainvoke({"messages": [HumanMessage(content=user_input)], "configurable": {"thread_id": "session-1"}})
print("Agent:", result["messages"][-1].content)
if __name__ == "__main__":
asyncio.run(main())
"""