Deepakraj02's picture
Update app.py
7148f99 verified
Raw History Blame Contribute Delete
5.12 kB
import os
import streamlit as st
from langchain.chains import create_history_aware_retriever, create_retrieval_chain
from langchain.chains.combine_documents import create_stuff_documents_chain
from langchain_chroma import Chroma
from langchain_community.chat_message_histories import ChatMessageHistory
from langchain_core.chat_history import BaseChatMessageHistory
from langchain_core.prompts import ChatPromptTemplate, MessagesPlaceholder
from langchain_groq import ChatGroq
from langchain_core.runnables.history import RunnableWithMessageHistory
from langchain_huggingface import HuggingFaceEmbeddings
from langchain_text_splitters import RecursiveCharacterTextSplitter
from langchain_community.document_loaders import PyPDFLoader, WebBaseLoader
from chromadb.config import Settings
# Load secrets from environment (set in Hugging Face Space Settings > Secrets)
groq_api_key = os.environ.get("GROQ_API_KEY")
hf_token = os.environ.get("HUGGINGFACE_API_KEY")
# Initialize LLM and embeddings
embeddings = HuggingFaceEmbeddings(model_name="all-MiniLM-L6-v2")
llm = ChatGroq(groq_api_key=groq_api_key, model_name="llama3-70b-8192")
# Streamlit UI
st.title("Conversational RAG | PDF + Website")
st.write("Upload PDFs or enter a website URL to chat with their content.")
session_id = st.text_input("Session ID:", value="default_session")
# Initialize session state
if 'store' not in st.session_state:
st.session_state.store = {}
if 'vectorstore' not in st.session_state:
st.session_state.vectorstore = None
if 'document_source' not in st.session_state:
st.session_state.document_source = ""
if 'documents' not in st.session_state:
st.session_state.documents = []
# Input Mode Selector
input_mode = st.radio("Select Input Type:", ("PDF Upload", "Website URL"))
# --- Load PDF ---
if input_mode == "PDF Upload":
uploaded_files = st.file_uploader("Upload PDF(s)", type="pdf", accept_multiple_files=True)
if uploaded_files:
for uploaded_file in uploaded_files:
with open("temp.pdf", "wb") as f:
f.write(uploaded_file.getvalue())
loader = PyPDFLoader("temp.pdf")
st.session_state.documents.extend(loader.load())
st.session_state.document_source = "PDF"
# --- Load Website ---
elif input_mode == "Website URL":
url = st.text_input("Enter Website URL")
if st.button("Load Website") and url:
loader = WebBaseLoader(web_paths=[url])
st.session_state.documents = loader.load()
st.session_state.document_source = url
# --- Build Vector Store and Chain ---
if st.session_state.documents:
splitter = RecursiveCharacterTextSplitter(chunk_size=1000, chunk_overlap=200)
splits = splitter.split_documents(st.session_state.documents)
vectorstore = Chroma.from_documents(
documents=splits,
embedding=embeddings,
collection_name="rag_collection",
persist_directory=None,
client_settings=Settings(anonymized_telemetry=False)
)
st.session_state.vectorstore = vectorstore
retriever = vectorstore.as_retriever()
# Prompts
contextualize_prompt = ChatPromptTemplate.from_messages([
("system", "Given a chat history and the latest user question which might reference context in the chat history, formulate a standalone question."),
MessagesPlaceholder("chat_history"),
("human", "{input}")
])
qa_prompt = ChatPromptTemplate.from_messages([
("system", "Use the following context to answer the question in 3 sentences or less.{context}"),
MessagesPlaceholder("chat_history"),
("human", "{input}")
])
# Chains
history_aware_retriever = create_history_aware_retriever(llm, retriever, contextualize_prompt)
qa_chain = create_stuff_documents_chain(llm, qa_prompt)
rag_chain = create_retrieval_chain(history_aware_retriever, qa_chain)
def get_session_history(session_id: str) -> BaseChatMessageHistory:
if session_id not in st.session_state.store:
st.session_state.store[session_id] = ChatMessageHistory()
return st.session_state.store[session_id]
conversational_rag_chain = RunnableWithMessageHistory(
rag_chain,
get_session_history,
input_messages_key="input",
history_messages_key="chat_history",
output_messages_key="answer"
)
st.divider()
st.write(f"You are chatting with: **{st.session_state.document_source}**")
user_input = st.text_input("Your question:")
if user_input:
session_history = get_session_history(session_id)
response = conversational_rag_chain.invoke(
{"input": user_input},
config={"configurable": {"session_id": session_id}}
)
st.write("**Assistant:**", response['answer'])
with st.expander("πŸ“œ Chat History"):
for msg in session_history.messages:
role = "You" if msg.type == "human" else "Assistant"
st.markdown(f"**{role}:** {msg.content}")
elif not groq_api_key:
st.warning("❗ Please set your GROQ_API_KEY in the Space secrets.")