rag document, url working
This commit is contained in:
parent
7aa9a761ff
commit
40ac9f865d
2 changed files with 119 additions and 50 deletions
|
|
@ -1,9 +1,11 @@
|
||||||
import os
|
import os
|
||||||
|
import tempfile
|
||||||
|
from datetime import datetime
|
||||||
|
from typing import List
|
||||||
|
|
||||||
import streamlit as st
|
import streamlit as st
|
||||||
import google.generativeai as genai
|
import google.generativeai as genai
|
||||||
import tempfile
|
|
||||||
import bs4
|
import bs4
|
||||||
from typing import List
|
|
||||||
from agno.agent import Agent
|
from agno.agent import Agent
|
||||||
from agno.models.google import Gemini
|
from agno.models.google import Gemini
|
||||||
from langchain_community.document_loaders import PyPDFLoader, WebBaseLoader
|
from langchain_community.document_loaders import PyPDFLoader, WebBaseLoader
|
||||||
|
|
@ -14,10 +16,10 @@ from qdrant_client.models import Distance, VectorParams
|
||||||
from langchain_core.embeddings import Embeddings
|
from langchain_core.embeddings import Embeddings
|
||||||
|
|
||||||
|
|
||||||
# Custom Gemini Embedder Class
|
# Custom Classes
|
||||||
class GeminiEmbedder(Embeddings):
|
class GeminiEmbedder(Embeddings):
|
||||||
def __init__(self, model_name="models/embedding-004"):
|
def __init__(self, model_name="models/text-embedding-004"):
|
||||||
genai.configure(api_key=os.environ["GOOGLE_API_KEY"])
|
genai.configure(api_key=st.session_state.google_api_key)
|
||||||
self.model = model_name
|
self.model = model_name
|
||||||
|
|
||||||
def embed_documents(self, texts: List[str]) -> List[List[float]]:
|
def embed_documents(self, texts: List[str]) -> List[List[float]]:
|
||||||
|
|
@ -31,37 +33,74 @@ class GeminiEmbedder(Embeddings):
|
||||||
)
|
)
|
||||||
return response['embedding']
|
return response['embedding']
|
||||||
|
|
||||||
# Initialize Streamlit App
|
|
||||||
|
# Constants
|
||||||
|
COLLECTION_NAME = "gemini-rag-agno"
|
||||||
|
|
||||||
|
|
||||||
|
# Streamlit App Initialization
|
||||||
st.title("🤖 AI Agent with Gemini & Qdrant RAG")
|
st.title("🤖 AI Agent with Gemini & Qdrant RAG")
|
||||||
|
|
||||||
|
# Session State Initialization
|
||||||
|
if 'google_api_key' not in st.session_state:
|
||||||
|
st.session_state.google_api_key = ""
|
||||||
|
if 'qdrant_api_key' not in st.session_state:
|
||||||
|
st.session_state.qdrant_api_key = ""
|
||||||
|
if 'qdrant_url' not in st.session_state:
|
||||||
|
st.session_state.qdrant_url = ""
|
||||||
|
if 'vector_store' not in st.session_state:
|
||||||
|
st.session_state.vector_store = None
|
||||||
|
if 'processed_documents' not in st.session_state:
|
||||||
|
st.session_state.processed_documents = []
|
||||||
|
|
||||||
|
|
||||||
# Sidebar Configuration
|
# Sidebar Configuration
|
||||||
st.sidebar.header("🔑 API Configuration")
|
st.sidebar.header("🔑 API Configuration")
|
||||||
google_api_key = st.sidebar.text_input("Google API Key", type="password")
|
google_api_key = st.sidebar.text_input("Google API Key", type="password", value=st.session_state.google_api_key)
|
||||||
qdrant_api_key = st.sidebar.text_input("Qdrant API Key", type="password")
|
qdrant_api_key = st.sidebar.text_input("Qdrant API Key", type="password", value=st.session_state.qdrant_api_key)
|
||||||
qdrant_url = st.sidebar.text_input("Qdrant URL",
|
qdrant_url = st.sidebar.text_input("Qdrant URL",
|
||||||
placeholder="https://your-cluster.cloud.qdrant.io:6333")
|
placeholder="https://your-cluster.cloud.qdrant.io:6333",
|
||||||
|
value=st.session_state.qdrant_url)
|
||||||
|
|
||||||
# Initialize Qdrant Client
|
# Update session state
|
||||||
|
st.session_state.google_api_key = google_api_key
|
||||||
|
st.session_state.qdrant_api_key = qdrant_api_key
|
||||||
|
st.session_state.qdrant_url = qdrant_url
|
||||||
|
|
||||||
|
|
||||||
|
# Utility Functions
|
||||||
def init_qdrant():
|
def init_qdrant():
|
||||||
if not all([qdrant_api_key, qdrant_url]):
|
"""Initialize Qdrant client with configured settings."""
|
||||||
|
if not all([st.session_state.qdrant_api_key, st.session_state.qdrant_url]):
|
||||||
return None
|
return None
|
||||||
try:
|
try:
|
||||||
return QdrantClient(
|
return QdrantClient(
|
||||||
url=qdrant_url,
|
url=st.session_state.qdrant_url,
|
||||||
api_key=qdrant_api_key,
|
api_key=st.session_state.qdrant_api_key,
|
||||||
timeout=60
|
timeout=60
|
||||||
)
|
)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
st.error(f"🔴 Qdrant connection failed: {str(e)}")
|
st.error(f"🔴 Qdrant connection failed: {str(e)}")
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
# Document Processing Functions
|
# Document Processing Functions
|
||||||
def process_pdf(file):
|
def process_pdf(file) -> List:
|
||||||
|
"""Process PDF file and add source metadata."""
|
||||||
try:
|
try:
|
||||||
with tempfile.NamedTemporaryFile(delete=False, suffix='.pdf') as tmp_file:
|
with tempfile.NamedTemporaryFile(delete=False, suffix='.pdf') as tmp_file:
|
||||||
tmp_file.write(file.getvalue())
|
tmp_file.write(file.getvalue())
|
||||||
loader = PyPDFLoader(tmp_file.name)
|
loader = PyPDFLoader(tmp_file.name)
|
||||||
documents = loader.load()
|
documents = loader.load()
|
||||||
|
|
||||||
|
# Add source metadata
|
||||||
|
for doc in documents:
|
||||||
|
doc.metadata.update({
|
||||||
|
"source_type": "pdf",
|
||||||
|
"file_name": file.name,
|
||||||
|
"timestamp": datetime.now().isoformat()
|
||||||
|
})
|
||||||
|
|
||||||
text_splitter = RecursiveCharacterTextSplitter(
|
text_splitter = RecursiveCharacterTextSplitter(
|
||||||
chunk_size=1000,
|
chunk_size=1000,
|
||||||
chunk_overlap=200
|
chunk_overlap=200
|
||||||
|
|
@ -71,7 +110,9 @@ def process_pdf(file):
|
||||||
st.error(f"📄 PDF processing error: {str(e)}")
|
st.error(f"📄 PDF processing error: {str(e)}")
|
||||||
return []
|
return []
|
||||||
|
|
||||||
def process_web(url):
|
|
||||||
|
def process_web(url: str) -> List:
|
||||||
|
"""Process web URL and add source metadata."""
|
||||||
try:
|
try:
|
||||||
loader = WebBaseLoader(
|
loader = WebBaseLoader(
|
||||||
web_paths=(url,),
|
web_paths=(url,),
|
||||||
|
|
@ -82,6 +123,15 @@ def process_web(url):
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
documents = loader.load()
|
documents = loader.load()
|
||||||
|
|
||||||
|
# Add source metadata
|
||||||
|
for doc in documents:
|
||||||
|
doc.metadata.update({
|
||||||
|
"source_type": "url",
|
||||||
|
"url": url,
|
||||||
|
"timestamp": datetime.now().isoformat()
|
||||||
|
})
|
||||||
|
|
||||||
text_splitter = RecursiveCharacterTextSplitter(
|
text_splitter = RecursiveCharacterTextSplitter(
|
||||||
chunk_size=1000,
|
chunk_size=1000,
|
||||||
chunk_overlap=200
|
chunk_overlap=200
|
||||||
|
|
@ -91,10 +141,10 @@ def process_web(url):
|
||||||
st.error(f"🌐 Web processing error: {str(e)}")
|
st.error(f"🌐 Web processing error: {str(e)}")
|
||||||
return []
|
return []
|
||||||
|
|
||||||
# Vector Store Management
|
|
||||||
COLLECTION_NAME = "agno_rag"
|
|
||||||
|
|
||||||
|
# Vector Store Management
|
||||||
def create_vector_store(client, texts):
|
def create_vector_store(client, texts):
|
||||||
|
"""Create and initialize vector store with documents."""
|
||||||
try:
|
try:
|
||||||
# Create collection if needed
|
# Create collection if needed
|
||||||
try:
|
try:
|
||||||
|
|
@ -127,10 +177,11 @@ def create_vector_store(client, texts):
|
||||||
st.error(f"🔴 Vector store error: {str(e)}")
|
st.error(f"🔴 Vector store error: {str(e)}")
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
# Main Application Flow
|
# Main Application Flow
|
||||||
if google_api_key:
|
if st.session_state.google_api_key:
|
||||||
os.environ["GOOGLE_API_KEY"] = google_api_key
|
os.environ["GOOGLE_API_KEY"] = st.session_state.google_api_key
|
||||||
genai.configure(api_key=google_api_key)
|
genai.configure(api_key=st.session_state.google_api_key)
|
||||||
|
|
||||||
qdrant_client = init_qdrant()
|
qdrant_client = init_qdrant()
|
||||||
|
|
||||||
|
|
@ -140,15 +191,39 @@ if google_api_key:
|
||||||
web_url = st.sidebar.text_input("Or enter URL")
|
web_url = st.sidebar.text_input("Or enter URL")
|
||||||
|
|
||||||
# Process documents
|
# Process documents
|
||||||
vector_store = None
|
|
||||||
if uploaded_file:
|
if uploaded_file:
|
||||||
|
file_name = uploaded_file.name
|
||||||
|
if file_name not in st.session_state.processed_documents:
|
||||||
|
with st.spinner('Processing PDF...'):
|
||||||
texts = process_pdf(uploaded_file)
|
texts = process_pdf(uploaded_file)
|
||||||
if texts and qdrant_client:
|
if texts and qdrant_client:
|
||||||
vector_store = create_vector_store(qdrant_client, texts)
|
if st.session_state.vector_store:
|
||||||
elif web_url:
|
st.session_state.vector_store.add_documents(texts)
|
||||||
|
else:
|
||||||
|
st.session_state.vector_store = create_vector_store(qdrant_client, texts)
|
||||||
|
st.session_state.processed_documents.append(file_name)
|
||||||
|
st.success(f"✅ Added PDF: {file_name}")
|
||||||
|
|
||||||
|
if web_url:
|
||||||
|
if web_url not in st.session_state.processed_documents:
|
||||||
|
with st.spinner('Processing URL...'):
|
||||||
texts = process_web(web_url)
|
texts = process_web(web_url)
|
||||||
if texts and qdrant_client:
|
if texts and qdrant_client:
|
||||||
vector_store = create_vector_store(qdrant_client, texts)
|
if st.session_state.vector_store:
|
||||||
|
st.session_state.vector_store.add_documents(texts)
|
||||||
|
else:
|
||||||
|
st.session_state.vector_store = create_vector_store(qdrant_client, texts)
|
||||||
|
st.session_state.processed_documents.append(web_url)
|
||||||
|
st.success(f"✅ Added URL: {web_url}")
|
||||||
|
|
||||||
|
# Display sources in sidebar
|
||||||
|
if st.session_state.processed_documents:
|
||||||
|
st.sidebar.header("📚 Processed Sources")
|
||||||
|
for source in st.session_state.processed_documents:
|
||||||
|
if source.endswith('.pdf'):
|
||||||
|
st.sidebar.text(f"📄 {source}")
|
||||||
|
else:
|
||||||
|
st.sidebar.text(f"🌐 {source}")
|
||||||
|
|
||||||
# Initialize Agent
|
# Initialize Agent
|
||||||
agent = Agent(
|
agent = Agent(
|
||||||
|
|
@ -159,7 +234,7 @@ if google_api_key:
|
||||||
markdown=True,
|
markdown=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Initialize chat history
|
# Chat Interface
|
||||||
if 'history' not in st.session_state:
|
if 'history' not in st.session_state:
|
||||||
st.session_state.history = []
|
st.session_state.history = []
|
||||||
|
|
||||||
|
|
@ -168,7 +243,7 @@ if google_api_key:
|
||||||
with st.chat_message(msg["role"]):
|
with st.chat_message(msg["role"]):
|
||||||
st.write(msg["content"])
|
st.write(msg["content"])
|
||||||
|
|
||||||
# User input
|
# Handle user input
|
||||||
if prompt := st.chat_input("Ask about your documents..."):
|
if prompt := st.chat_input("Ask about your documents..."):
|
||||||
# Add user message to history
|
# Add user message to history
|
||||||
st.session_state.history.append({"role": "user", "content": prompt})
|
st.session_state.history.append({"role": "user", "content": prompt})
|
||||||
|
|
@ -177,8 +252,8 @@ if google_api_key:
|
||||||
|
|
||||||
# Retrieve relevant documents
|
# Retrieve relevant documents
|
||||||
context = ""
|
context = ""
|
||||||
if vector_store:
|
if st.session_state.vector_store:
|
||||||
retriever = vector_store.as_retriever(
|
retriever = st.session_state.vector_store.as_retriever(
|
||||||
search_type="similarity_score_threshold",
|
search_type="similarity_score_threshold",
|
||||||
search_kwargs={"k": 5, "score_threshold": 0.7}
|
search_kwargs={"k": 5, "score_threshold": 0.7}
|
||||||
)
|
)
|
||||||
|
|
@ -200,10 +275,14 @@ if google_api_key:
|
||||||
with st.chat_message("assistant"):
|
with st.chat_message("assistant"):
|
||||||
st.write(response.content)
|
st.write(response.content)
|
||||||
|
|
||||||
if vector_store and docs:
|
if st.session_state.vector_store and docs:
|
||||||
with st.expander("🔍 See sources"):
|
with st.expander("🔍 See sources"):
|
||||||
for i, doc in enumerate(docs, 1):
|
for i, doc in enumerate(docs, 1):
|
||||||
st.write(f"Source {i}: {doc.page_content[:200]}...")
|
source_type = doc.metadata.get("source_type", "unknown")
|
||||||
|
source_icon = "📄" if source_type == "pdf" else "🌐"
|
||||||
|
source_name = doc.metadata.get("file_name" if source_type == "pdf" else "url", "unknown")
|
||||||
|
st.write(f"{source_icon} Source {i} from {source_name}:")
|
||||||
|
st.write(f"{doc.page_content[:200]}...")
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
st.error(f"❌ Error generating response: {str(e)}")
|
st.error(f"❌ Error generating response: {str(e)}")
|
||||||
|
|
|
||||||
|
|
@ -1,22 +1,12 @@
|
||||||
from google import genai
|
import google.generativeai as genai
|
||||||
import os
|
import os
|
||||||
from dotenv import load_dotenv
|
from dotenv import load_dotenv
|
||||||
load_dotenv()
|
load_dotenv()
|
||||||
|
|
||||||
client = genai.Client(api_key=os.getenv("GOOGLE_API_KEY"), http_options={'api_version':'v1alpha'})
|
genai.configure(api_key=os.getenv("GOOGLE_API_KEY"))
|
||||||
|
|
||||||
import asyncio
|
result = genai.embed_content(
|
||||||
|
model="models/text-embedding-004",
|
||||||
|
content="What is the meaning of life?")
|
||||||
|
|
||||||
config = {'thinking_config': {'include_thoughts': True}}
|
print(str(result['embedding']))
|
||||||
|
|
||||||
async def main():
|
|
||||||
chat = client.aio.chats.create(
|
|
||||||
model='gemini-2.0-flash-thinking-exp-01-21',
|
|
||||||
config=config
|
|
||||||
)
|
|
||||||
response = await chat.send_message('Explain Deep Q Networks from first principles')
|
|
||||||
print(response.text)
|
|
||||||
response = await chat.send_message('What did you just say before this?')
|
|
||||||
print(response.text)
|
|
||||||
|
|
||||||
asyncio.run(main())
|
|
||||||
Loading…
Reference in a new issue