complete working code
This commit is contained in:
parent
7e8375831a
commit
93331106d7
1 changed files with 99 additions and 136 deletions
|
|
@ -12,10 +12,14 @@ from langchain import hub
|
||||||
import tempfile
|
import tempfile
|
||||||
from langgraph.prebuilt import create_react_agent
|
from langgraph.prebuilt import create_react_agent
|
||||||
from langchain_community.tools import DuckDuckGoSearchRun
|
from langchain_community.tools import DuckDuckGoSearchRun
|
||||||
|
from typing import TypedDict, List
|
||||||
|
from langchain_core.language_models import BaseLanguageModel
|
||||||
|
from langchain_core.messages import AIMessage, HumanMessage, SystemMessage
|
||||||
|
from time import sleep
|
||||||
|
from tenacity import retry, wait_exponential, stop_after_attempt
|
||||||
|
|
||||||
|
|
||||||
def init_session_state():
|
def init_session_state():
|
||||||
"""Initialize session state variables."""
|
|
||||||
if 'api_keys_submitted' not in st.session_state:
|
if 'api_keys_submitted' not in st.session_state:
|
||||||
st.session_state.api_keys_submitted = False
|
st.session_state.api_keys_submitted = False
|
||||||
if 'chat_history' not in st.session_state:
|
if 'chat_history' not in st.session_state:
|
||||||
|
|
@ -28,11 +32,9 @@ def init_session_state():
|
||||||
st.session_state.qdrant_url = ""
|
st.session_state.qdrant_url = ""
|
||||||
|
|
||||||
def sidebar_api_form():
|
def sidebar_api_form():
|
||||||
"""Render API credentials form in sidebar."""
|
|
||||||
with st.sidebar:
|
with st.sidebar:
|
||||||
st.header("API Credentials")
|
st.header("API Credentials")
|
||||||
|
|
||||||
# Show current status
|
|
||||||
if st.session_state.api_keys_submitted:
|
if st.session_state.api_keys_submitted:
|
||||||
st.success("API credentials verified")
|
st.success("API credentials verified")
|
||||||
if st.button("Reset Credentials"):
|
if st.button("Reset Credentials"):
|
||||||
|
|
@ -40,32 +42,18 @@ def sidebar_api_form():
|
||||||
st.rerun()
|
st.rerun()
|
||||||
return True
|
return True
|
||||||
|
|
||||||
# Show API form
|
|
||||||
with st.form("api_credentials"):
|
with st.form("api_credentials"):
|
||||||
cohere_key = st.text_input("Cohere API Key", type="password")
|
cohere_key = st.text_input("Cohere API Key", type="password")
|
||||||
qdrant_key = st.text_input(
|
qdrant_key = st.text_input("Qdrant API Key", type="password", help="Enter your Qdrant API key")
|
||||||
"Qdrant API Key",
|
qdrant_url = st.text_input("Qdrant URL",
|
||||||
type="password",
|
placeholder="https://xyz-example.eu-central.aws.cloud.qdrant.io:6333",
|
||||||
help="Enter your Qdrant API key"
|
help="Enter your Qdrant instance URL")
|
||||||
)
|
|
||||||
qdrant_url = st.text_input(
|
|
||||||
"Qdrant URL",
|
|
||||||
placeholder="https://xyz-example.eu-central.aws.cloud.qdrant.io:6333",
|
|
||||||
help="Enter your Qdrant instance URL"
|
|
||||||
)
|
|
||||||
|
|
||||||
if st.form_submit_button("Submit Credentials"):
|
if st.form_submit_button("Submit Credentials"):
|
||||||
try:
|
try:
|
||||||
# First validate the credentials before saving to session state
|
client = QdrantClient(url=qdrant_url, api_key=qdrant_key, timeout=60)
|
||||||
client = QdrantClient(
|
|
||||||
url=qdrant_url,
|
|
||||||
api_key=qdrant_key,
|
|
||||||
timeout=60
|
|
||||||
)
|
|
||||||
# Test connection
|
|
||||||
client.get_collections()
|
client.get_collections()
|
||||||
|
|
||||||
# Only save to session state after successful validation
|
|
||||||
st.session_state.cohere_api_key = cohere_key
|
st.session_state.cohere_api_key = cohere_key
|
||||||
st.session_state.qdrant_api_key = qdrant_key
|
st.session_state.qdrant_api_key = qdrant_key
|
||||||
st.session_state.qdrant_url = qdrant_url
|
st.session_state.qdrant_url = qdrant_url
|
||||||
|
|
@ -78,59 +66,43 @@ def sidebar_api_form():
|
||||||
return False
|
return False
|
||||||
|
|
||||||
def init_qdrant() -> QdrantClient:
|
def init_qdrant() -> QdrantClient:
|
||||||
"""Initialize Qdrant vector database."""
|
|
||||||
if not st.session_state.get("qdrant_api_key"):
|
if not st.session_state.get("qdrant_api_key"):
|
||||||
raise ValueError("Qdrant API key not provided")
|
raise ValueError("Qdrant API key not provided")
|
||||||
if not st.session_state.get("qdrant_url"):
|
if not st.session_state.get("qdrant_url"):
|
||||||
raise ValueError("Qdrant URL not provided")
|
raise ValueError("Qdrant URL not provided")
|
||||||
|
|
||||||
return QdrantClient(
|
return QdrantClient(url=st.session_state.qdrant_url,
|
||||||
url=st.session_state.qdrant_url,
|
api_key=st.session_state.qdrant_api_key,
|
||||||
api_key=st.session_state.qdrant_api_key,
|
timeout=60)
|
||||||
timeout=60
|
|
||||||
)
|
|
||||||
|
|
||||||
# Initialize session state
|
|
||||||
init_session_state()
|
init_session_state()
|
||||||
|
|
||||||
# Main application logic
|
|
||||||
if not sidebar_api_form():
|
if not sidebar_api_form():
|
||||||
st.info("Please enter your API credentials in the sidebar to continue.")
|
st.info("Please enter your API credentials in the sidebar to continue.")
|
||||||
st.stop()
|
st.stop()
|
||||||
|
|
||||||
# Initialize services with verified credentials
|
embedding = CohereEmbeddings(model="embed-english-v3.0",
|
||||||
embedding = CohereEmbeddings(
|
cohere_api_key=st.session_state.cohere_api_key)
|
||||||
model="embed-english-v3.0",
|
|
||||||
cohere_api_key=st.session_state.cohere_api_key
|
|
||||||
)
|
|
||||||
|
|
||||||
chat_model = ChatCohere(
|
chat_model = ChatCohere(model="command-r7b-12-2024",
|
||||||
model="command-r7b-12-2024",
|
temperature=0.1,
|
||||||
temperature=0.1,
|
max_tokens=512,
|
||||||
max_tokens=512,
|
verbose=True,
|
||||||
verbose=True,
|
cohere_api_key=st.session_state.cohere_api_key)
|
||||||
cohere_api_key=st.session_state.cohere_api_key
|
|
||||||
)
|
|
||||||
|
|
||||||
client = init_qdrant()
|
client = init_qdrant()
|
||||||
|
|
||||||
#document preprocessing
|
|
||||||
|
|
||||||
def process_document(file):
|
def process_document(file):
|
||||||
"""Process uploaded PDF document using a temporary file."""
|
|
||||||
try:
|
try:
|
||||||
# Create a temporary file
|
|
||||||
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())
|
||||||
tmp_path = tmp_file.name
|
tmp_path = tmp_file.name
|
||||||
|
|
||||||
# Process the temporary file
|
|
||||||
loader = PyPDFLoader(tmp_path)
|
loader = PyPDFLoader(tmp_path)
|
||||||
documents = loader.load()
|
documents = loader.load()
|
||||||
text_splitter = RecursiveCharacterTextSplitter(chunk_size=1000, chunk_overlap=200)
|
text_splitter = RecursiveCharacterTextSplitter(chunk_size=1000, chunk_overlap=200)
|
||||||
texts = text_splitter.split_documents(documents)
|
texts = text_splitter.split_documents(documents)
|
||||||
|
|
||||||
# Clean up the temporary file
|
|
||||||
os.unlink(tmp_path)
|
os.unlink(tmp_path)
|
||||||
|
|
||||||
return texts
|
return texts
|
||||||
|
|
@ -143,26 +115,18 @@ COLLECTION_NAME = "cohere_rag"
|
||||||
def create_vector_stores(texts):
|
def create_vector_stores(texts):
|
||||||
"""Create and populate vector store with documents."""
|
"""Create and populate vector store with documents."""
|
||||||
try:
|
try:
|
||||||
# First, create the collection explicitly
|
|
||||||
try:
|
try:
|
||||||
client.create_collection(
|
client.create_collection(collection_name=COLLECTION_NAME,
|
||||||
collection_name=COLLECTION_NAME,
|
vectors_config=VectorParams(size=1024,
|
||||||
vectors_config=VectorParams(
|
distance=Distance.COSINE))
|
||||||
size=1024, # Dimension for Cohere embed-english-v3.0
|
|
||||||
distance=Distance.COSINE
|
|
||||||
)
|
|
||||||
)
|
|
||||||
st.success(f"Created new collection: {COLLECTION_NAME}")
|
st.success(f"Created new collection: {COLLECTION_NAME}")
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
if "already exists" not in str(e).lower():
|
if "already exists" not in str(e).lower():
|
||||||
raise e
|
raise e
|
||||||
|
|
||||||
# Then initialize the vector store
|
vector_store = QdrantVectorStore(client=client,
|
||||||
vector_store = QdrantVectorStore(
|
collection_name=COLLECTION_NAME,
|
||||||
client=client,
|
embedding=embedding)
|
||||||
collection_name=COLLECTION_NAME,
|
|
||||||
embedding=embedding,
|
|
||||||
)
|
|
||||||
|
|
||||||
with st.spinner('Storing documents in Qdrant...'):
|
with st.spinner('Storing documents in Qdrant...'):
|
||||||
vector_store.add_documents(texts)
|
vector_store.add_documents(texts)
|
||||||
|
|
@ -174,91 +138,99 @@ def create_vector_stores(texts):
|
||||||
st.error(f"Error in vector store creation: {str(e)}")
|
st.error(f"Error in vector store creation: {str(e)}")
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def create_fallback_agent():
|
# Define the state schema using TypedDict
|
||||||
"""Create a LangGraph agent with DuckDuckGo search tool."""
|
class AgentState(TypedDict):
|
||||||
|
"""State schema for the agent."""
|
||||||
|
messages: List[HumanMessage | AIMessage | SystemMessage]
|
||||||
|
is_last_step: bool
|
||||||
|
|
||||||
|
class RateLimitedDuckDuckGo(DuckDuckGoSearchRun):
|
||||||
|
@retry(wait=wait_exponential(multiplier=1, min=4, max=10),
|
||||||
|
stop=stop_after_attempt(3))
|
||||||
|
def run(self, query: str) -> str:
|
||||||
|
"""Run search with rate limiting."""
|
||||||
|
try:
|
||||||
|
sleep(2) # Add delay between requests
|
||||||
|
return super().run(query)
|
||||||
|
except Exception as e:
|
||||||
|
if "Ratelimit" in str(e):
|
||||||
|
sleep(5) # Longer delay on rate limit
|
||||||
|
return super().run(query)
|
||||||
|
raise e
|
||||||
|
|
||||||
|
def create_fallback_agent(chat_model: BaseLanguageModel):
|
||||||
|
"""Create a LangGraph agent for web research."""
|
||||||
|
|
||||||
def web_research(query: str) -> str:
|
def web_research(query: str) -> str:
|
||||||
"""Search the web for information about a query."""
|
"""Web search with result formatting."""
|
||||||
search = DuckDuckGoSearchRun()
|
try:
|
||||||
results = search.run(query)
|
search = DuckDuckGoSearchRun(num_results=5)
|
||||||
return f"Web search results: {results}"
|
results = search.run(query)
|
||||||
|
return results
|
||||||
|
except Exception as e:
|
||||||
|
return f"Search failed: {str(e)}. Providing answer based on general knowledge."
|
||||||
|
|
||||||
tools = [web_research]
|
tools = [web_research]
|
||||||
|
|
||||||
# Create agent with Cohere model
|
agent = create_react_agent(model=chat_model,
|
||||||
agent = create_react_agent(
|
tools=tools,
|
||||||
chat_model, # Using the already initialized Cohere model
|
debug=False)
|
||||||
tools=tools,
|
|
||||||
)
|
|
||||||
|
|
||||||
return agent
|
return agent
|
||||||
|
|
||||||
def process_query(vectorstore, query) -> tuple[str, list]:
|
def process_query(vectorstore, query) -> tuple[str, list]:
|
||||||
"""Process a query using RAG with fallback to web search."""
|
"""Process a query using RAG with fallback to web search."""
|
||||||
try:
|
try:
|
||||||
# First try vector store retrieval
|
|
||||||
retriever = vectorstore.as_retriever(
|
retriever = vectorstore.as_retriever(
|
||||||
search_type="similarity_score_threshold",
|
search_type="similarity_score_threshold",
|
||||||
search_kwargs={
|
search_kwargs={
|
||||||
"k": 10,
|
"k": 10,
|
||||||
"score_threshold": 0.7 # Only return relevant documents
|
"score_threshold": 0.7
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
|
||||||
# Get relevant documents
|
relevant_docs = retriever.get_relevant_documents(query)
|
||||||
with st.spinner('Searching document database...'):
|
|
||||||
relevant_docs = retriever.get_relevant_documents(query)
|
|
||||||
|
|
||||||
if relevant_docs:
|
if relevant_docs:
|
||||||
# Use RAG with document context
|
|
||||||
retrieval_qa_prompt = hub.pull("langchain-ai/retrieval-qa-chat")
|
retrieval_qa_prompt = hub.pull("langchain-ai/retrieval-qa-chat")
|
||||||
|
combine_docs_chain = create_stuff_documents_chain(chat_model, retrieval_qa_prompt)
|
||||||
|
retrieval_chain = create_retrieval_chain(retriever, combine_docs_chain)
|
||||||
|
response = retrieval_chain.invoke({"input": query})
|
||||||
|
return response['answer'], relevant_docs
|
||||||
|
|
||||||
combine_docs_chain = create_stuff_documents_chain(
|
|
||||||
chat_model,
|
|
||||||
retrieval_qa_prompt
|
|
||||||
)
|
|
||||||
|
|
||||||
retrieval_chain = create_retrieval_chain(
|
|
||||||
retriever,
|
|
||||||
combine_docs_chain
|
|
||||||
)
|
|
||||||
|
|
||||||
with st.spinner('Generating response from documents...'):
|
|
||||||
response = retrieval_chain.invoke({"input": query})
|
|
||||||
if not response or 'answer' not in response:
|
|
||||||
raise ValueError("No response generated")
|
|
||||||
|
|
||||||
return response['answer'], relevant_docs
|
|
||||||
else:
|
else:
|
||||||
# Fallback to web search using LangGraph agent
|
st.info("No relevant documents found. Searching web...")
|
||||||
st.info("No relevant documents found. Searching the web...")
|
fallback_agent = create_fallback_agent(chat_model)
|
||||||
|
|
||||||
fallback_agent = create_fallback_agent()
|
with st.spinner('Researching...'):
|
||||||
|
|
||||||
with st.spinner('Searching web and generating response...'):
|
|
||||||
# Prepare input for the agent
|
|
||||||
agent_input = {
|
agent_input = {
|
||||||
"messages": [
|
"messages": [
|
||||||
("user", f"Please search and answer this question: {query}")
|
HumanMessage(content=f"""Please thoroughly research the question: '{query}' and provide a detailed and comprehensive response. Make sure to gather the latest information from credible sources. Minimum 400 words.""")
|
||||||
]
|
],
|
||||||
|
"is_last_step": False
|
||||||
}
|
}
|
||||||
|
|
||||||
# Get agent response
|
config = {"recursion_limit": 100}
|
||||||
response = fallback_agent.invoke(agent_input)
|
|
||||||
last_message = response["messages"][-1]
|
|
||||||
|
|
||||||
if isinstance(last_message, tuple):
|
try:
|
||||||
answer = last_message[1]
|
response = fallback_agent.invoke(agent_input, config=config)
|
||||||
else:
|
|
||||||
answer = last_message.content
|
|
||||||
|
|
||||||
return f"Based on web search: {answer}", []
|
if isinstance(response, dict) and "messages" in response:
|
||||||
|
last_message = response["messages"][-1]
|
||||||
|
answer = last_message.content if hasattr(last_message, 'content') else str(last_message)
|
||||||
|
|
||||||
|
return f"""Comprehensive Research Results:
|
||||||
|
{answer}
|
||||||
|
""", []
|
||||||
|
|
||||||
|
except Exception as agent_error:
|
||||||
|
fallback_response = chat_model.invoke(f"Please provide a general answer to: {query}").content
|
||||||
|
return f"Web search unavailable. General response: {fallback_response}", []
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
st.error(f"Error processing query: {str(e)}")
|
st.error(f"Error: {str(e)}")
|
||||||
return "I encountered an error processing your query. Please try again.", []
|
return "I encountered an error. Please try rephrasing your question.", []
|
||||||
|
|
||||||
#post processing - strip, summarize along with formatted sources
|
|
||||||
def post_process(answer, sources):
|
def post_process(answer, sources):
|
||||||
"""Post-process the answer and format sources."""
|
"""Post-process the answer and format sources."""
|
||||||
answer = answer.strip()
|
answer = answer.strip()
|
||||||
|
|
@ -266,7 +238,7 @@ def post_process(answer, sources):
|
||||||
# Summarize long answers
|
# Summarize long answers
|
||||||
if len(answer) > 500:
|
if len(answer) > 500:
|
||||||
summary_prompt = f"Summarize the following answer in 2-3 sentences: {answer}"
|
summary_prompt = f"Summarize the following answer in 2-3 sentences: {answer}"
|
||||||
summary = chat_model.invoke(summary_prompt).content # Changed from predict to invoke
|
summary = chat_model.invoke(summary_prompt).content
|
||||||
answer = f"{summary}\n\nFull Answer: {answer}"
|
answer = f"{summary}\n\nFull Answer: {answer}"
|
||||||
|
|
||||||
formatted_sources = []
|
formatted_sources = []
|
||||||
|
|
@ -275,26 +247,25 @@ def post_process(answer, sources):
|
||||||
formatted_sources.append(formatted_source)
|
formatted_sources.append(formatted_source)
|
||||||
return answer, formatted_sources
|
return answer, formatted_sources
|
||||||
|
|
||||||
st.title("RAG Agent with Cohere 🤖") # New heading
|
st.title("RAG Agent with Cohere 🤖")
|
||||||
|
|
||||||
uploaded_file = st.file_uploader("Choose a PDF or Image File", type=["pdf", "jpg", "jpeg"])
|
uploaded_file = st.file_uploader("Choose a PDF or Image File", type=["pdf", "jpg", "jpeg"])
|
||||||
|
|
||||||
if uploaded_file is not None:
|
if uploaded_file is not None and 'processed_file' not in st.session_state:
|
||||||
with st.spinner('Processing file... This may take a while for images.'):
|
with st.spinner('Processing file... This may take a while for images.'):
|
||||||
texts = process_document(uploaded_file)
|
texts = process_document(uploaded_file)
|
||||||
vectorstore = create_vector_stores(texts)
|
vectorstore = create_vector_stores(texts)
|
||||||
if vectorstore:
|
if vectorstore:
|
||||||
st.session_state.vectorstore = vectorstore
|
st.session_state.vectorstore = vectorstore
|
||||||
|
st.session_state.processed_file = True
|
||||||
st.success('File uploaded and processed successfully!')
|
st.success('File uploaded and processed successfully!')
|
||||||
else:
|
else:
|
||||||
st.error('Failed to process file. Please try again.')
|
st.error('Failed to process file. Please try again.')
|
||||||
|
|
||||||
# Display chat history
|
|
||||||
for message in st.session_state.chat_history:
|
for message in st.session_state.chat_history:
|
||||||
with st.chat_message(message["role"]):
|
with st.chat_message(message["role"]):
|
||||||
st.markdown(message["content"])
|
st.markdown(message["content"])
|
||||||
|
|
||||||
# Chat input
|
|
||||||
if query := st.chat_input("Ask a question about the document:"):
|
if query := st.chat_input("Ask a question about the document:"):
|
||||||
st.session_state.chat_history.append({"role": "user", "content": query})
|
st.session_state.chat_history.append({"role": "user", "content": query})
|
||||||
with st.chat_message("user"):
|
with st.chat_message("user"):
|
||||||
|
|
@ -304,31 +275,24 @@ if query := st.chat_input("Ask a question about the document:"):
|
||||||
with st.chat_message("assistant"):
|
with st.chat_message("assistant"):
|
||||||
try:
|
try:
|
||||||
answer, sources = process_query(st.session_state.vectorstore, query)
|
answer, sources = process_query(st.session_state.vectorstore, query)
|
||||||
|
st.markdown(answer)
|
||||||
|
|
||||||
if sources: # Only post-process if we have sources
|
if sources:
|
||||||
processed_answer, formatted_sources = post_process(answer, sources)
|
|
||||||
else:
|
|
||||||
processed_answer, formatted_sources = answer, []
|
|
||||||
|
|
||||||
st.markdown(f"{processed_answer}")
|
|
||||||
|
|
||||||
if formatted_sources:
|
|
||||||
with st.expander("Sources"):
|
with st.expander("Sources"):
|
||||||
for source in formatted_sources:
|
for source in sources:
|
||||||
st.markdown(f"- {source}")
|
st.markdown(f"- {source.page_content[:200]}...")
|
||||||
|
|
||||||
st.session_state.chat_history.append({
|
st.session_state.chat_history.append({
|
||||||
"role": "assistant",
|
"role": "assistant",
|
||||||
"content": processed_answer
|
"content": answer
|
||||||
})
|
})
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
st.error(f"Error: {str(e)}")
|
st.error(f"Error: {str(e)}")
|
||||||
st.info("Please try asking your question again.")
|
st.info("Please try asking your question again.")
|
||||||
else:
|
else:
|
||||||
st.error("Please upload a document first.")
|
st.error("Please upload a document first.")
|
||||||
|
|
||||||
|
|
||||||
# Add to sidebar
|
|
||||||
with st.sidebar:
|
with st.sidebar:
|
||||||
st.divider()
|
st.divider()
|
||||||
col1, col2 = st.columns(2)
|
col1, col2 = st.columns(2)
|
||||||
|
|
@ -339,7 +303,6 @@ with st.sidebar:
|
||||||
with col2:
|
with col2:
|
||||||
if st.button('Clear All Data'):
|
if st.button('Clear All Data'):
|
||||||
try:
|
try:
|
||||||
# Check if collections exist before deleting
|
|
||||||
collections = client.get_collections().collections
|
collections = client.get_collections().collections
|
||||||
collection_names = [col.name for col in collections]
|
collection_names = [col.name for col in collections]
|
||||||
|
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue